"""多用户部署的 SQLite 存储:邀请码门禁 + 会话进度(引擎快照)。

- invites:预生成的邀请码,一码一会话(首次激活即绑定一个 session,之后该码作废);
- sessions:每个用户的独立进度,存引擎快照 JSON——离开即存档,回来从此恢复续播。
线程安全:每线程一条 sqlite 连接(HTTP handler 多线程);写操作各自 commit。
"""

from __future__ import annotations

import secrets
import sqlite3
import threading
import time
from pathlib import Path


class SessionStore:
    def __init__(self, db_path: str | Path) -> None:
        self._db_path = str(db_path)
        Path(self._db_path).parent.mkdir(parents=True, exist_ok=True)
        self._local = threading.local()
        self._init()

    def _conn(self) -> sqlite3.Connection:
        if not hasattr(self._local, "conn"):
            c = sqlite3.connect(self._db_path, check_same_thread=False, timeout=10)
            c.row_factory = sqlite3.Row
            c.execute("PRAGMA journal_mode=WAL")        # 并发读写更稳
            self._local.conn = c
        return self._local.conn

    def _init(self) -> None:
        self._conn().executescript("""
        CREATE TABLE IF NOT EXISTS invites(
          code TEXT PRIMARY KEY,
          world TEXT,
          session_id TEXT,
          created_at REAL,
          activated_at REAL);
        CREATE TABLE IF NOT EXISTS sessions(
          session_id TEXT PRIMARY KEY,
          invite_code TEXT,
          world TEXT,
          created_at REAL,
          last_active REAL,
          step INTEGER DEFAULT 0,
          snapshot TEXT);
        """)
        for tbl in ("invites", "sessions"):             # 旧库迁移:补 world 列(幂等)
            try:
                self._conn().execute(f"ALTER TABLE {tbl} ADD COLUMN world TEXT")
            except sqlite3.OperationalError:
                pass
        self._conn().commit()

    # ── 邀请码(每世界隔离:每个码绑定一个 world)──
    def gen_invites(self, n: int, world: str) -> list[str]:
        """为指定世界预生成 n 个邀请码(8 位十六进制,大写好念)。"""
        c, now, codes = self._conn(), time.time(), []
        for _ in range(n):
            code = secrets.token_hex(4).upper()
            c.execute("INSERT OR IGNORE INTO invites(code, world, created_at) VALUES(?, ?, ?)", (code, world, now))
            codes.append(code)
        c.commit()
        return codes

    def list_invites(self, world: str | None = None) -> list[dict]:
        q = "SELECT code, world, session_id, created_at, activated_at FROM invites"
        args: tuple = ()
        if world:
            q += " WHERE world=?"; args = (world,)
        rows = self._conn().execute(q + " ORDER BY created_at", args).fetchall()
        return [dict(r) for r in rows]

    def activate_invite(self, code: str) -> tuple[str, str] | None:
        """激活邀请码:未用则创建新会话并绑定到该码的世界,返回 (session_id, world);无效/已用返回 None。"""
        code = (code or "").strip().upper()
        c = self._conn()
        row = c.execute("SELECT session_id, world FROM invites WHERE code=?", (code,)).fetchone()
        if row is None or row["session_id"]:           # 不存在 或 已被激活
            return None
        world = row["world"]
        sid, now = secrets.token_hex(16), time.time()
        c.execute("UPDATE invites SET session_id=?, activated_at=? WHERE code=? AND session_id IS NULL",
                  (sid, now, code))
        if c.total_changes == 0:                        # 并发竞争:刚被别人激活
            c.rollback()
            return None
        c.execute("INSERT INTO sessions(session_id, invite_code, world, created_at, last_active) VALUES(?,?,?,?,?)",
                  (sid, code, world, now, now))
        c.commit()
        return sid, world

    def session_world(self, sid: str) -> str | None:
        """会话绑定的世界 id(路由到对应世界的引擎池)。"""
        if not sid:
            return None
        row = self._conn().execute("SELECT world FROM sessions WHERE session_id=?", (sid,)).fetchone()
        return row["world"] if row else None

    # ── 会话进度 ──
    def session_exists(self, sid: str) -> bool:
        if not sid:
            return False
        return self._conn().execute(
            "SELECT 1 FROM sessions WHERE session_id=?", (sid,)).fetchone() is not None

    def save_snapshot(self, sid: str, step: int, snapshot_json: str) -> None:
        c = self._conn()
        c.execute("UPDATE sessions SET step=?, snapshot=?, last_active=? WHERE session_id=?",
                  (step, snapshot_json, time.time(), sid))
        c.commit()

    def load_snapshot(self, sid: str) -> tuple[int | None, str | None]:
        row = self._conn().execute(
            "SELECT step, snapshot FROM sessions WHERE session_id=?", (sid,)).fetchone()
        return (row["step"], row["snapshot"]) if row else (None, None)

    def touch(self, sid: str) -> None:
        c = self._conn()
        c.execute("UPDATE sessions SET last_active=? WHERE session_id=?", (time.time(), sid))
        c.commit()
