"""Accès DuckDB en lecture seule + garde-fous pour le bac à sable SQL.""" from __future__ import annotations import os import re import threading import time from pathlib import Path import duckdb import numpy as np ROOT = Path(__file__).resolve().parent.parent DB_PATH = os.environ.get("SPID_DB", str(ROOT / "data" / "spid.duckdb")) QUERY_TIMEOUT_S = float(os.environ.get("SQL_TIMEOUT_S", "20")) MAX_ROWS = int(os.environ.get("SQL_MAX_ROWS", "5000")) TABLES = ["projects", "project_mentions", "project_timeline", "kg_nodes", "kg_edges", "project_embeddings", "spid_sections"] _ALIAS = {"yr": "year", "loc": "location", "lbl": "label"} _lock = threading.Lock() _conn: duckdb.DuckDBPyConnection | None = None _schema_cache: list[dict] | None = None _emb_cache: dict | None = None def connect() -> duckdb.DuckDBPyConnection: global _conn with _lock: if _conn is None: if not os.path.exists(DB_PATH): raise RuntimeError(f"Base introuvable : {DB_PATH}") _conn = duckdb.connect( DB_PATH, read_only=True, config={ "enable_external_access": "false", # pas de read_parquet/read_csv sur le disque "threads": os.environ.get("DUCKDB_THREADS", "4"), "memory_limit": os.environ.get("DUCKDB_MEM", "4GB"), }, ) try: _conn.execute("SET lock_configuration = true") except Exception: pass return _conn def cursor() -> duckdb.DuckDBPyConnection: """Connexion dupliquée (thread-safe) sur la connexion principale.""" return connect().cursor() def rows(sql: str, params: list | tuple = ()) -> list[dict]: cur = cursor() try: res = cur.execute(sql, list(params)) cols = [_ALIAS.get(d[0], d[0]) for d in res.description] # alias courts : year/location/label sont réservés en DuckDB return [dict(zip(cols, _jsonable(r))) for r in res.fetchall()] finally: cur.close() def one(sql: str, params: list | tuple = ()) -> dict | None: r = rows(sql, params) return r[0] if r else None def scalar(sql: str, params: list | tuple = ()): cur = cursor() try: return cur.execute(sql, list(params)).fetchone()[0] finally: cur.close() def _jsonable(row): out = [] for v in row: if hasattr(v, "isoformat"): v = v.isoformat() elif isinstance(v, (np.floating,)): v = float(v) elif isinstance(v, (np.integer,)): v = int(v) elif isinstance(v, float) and (v != v): # NaN v = None out.append(v) return out # ---------------------------------------------------------------- schéma def schema() -> list[dict]: global _schema_cache if _schema_cache is None: cur = cursor() try: out = [] for t in TABLES: cols = cur.execute( "select column_name, data_type from information_schema.columns where table_name=? order by ordinal_position", [t] ).fetchall() n = cur.execute(f'select count(*) from "{t}"').fetchone()[0] out.append({"table": t, "rows": n, "columns": [{"name": c, "type": d} for c, d in cols]}) _schema_cache = out finally: cur.close() return _schema_cache # ---------------------------------------------------------------- bac à sable SQL _FORBIDDEN = re.compile( r"\b(attach|detach|copy|export|import|install|load|pragma|create|insert|update|delete|drop|alter|call|set|reset|" r"checkpoint|vacuum|force|begin|commit|rollback|grant|revoke|use|read_\w+|glob|getenv|current_setting|" r"sqlite_\w+|postgres_\w+|http\w*|write_\w+|list_files|duckdb_secrets|secrets?|parquet_\w+|iceberg_\w+|delta_\w+|" r"json_extract_path_text|read_json\w*|read_text|read_blob)\b", re.IGNORECASE, ) def _strip_comments(sql: str) -> str: sql = re.sub(r"/\*.*?\*/", " ", sql, flags=re.S) sql = re.sub(r"--[^\n]*", " ", sql) return sql.strip() def guard_sql(sql: str) -> str: s = _strip_comments(sql).rstrip(";").strip() if not s: raise ValueError("Requête vide.") if ";" in s: raise ValueError("Une seule instruction à la fois.") if not re.match(r"^(select|with|from|describe|show|summarize|explain)\b", s, re.IGNORECASE): raise ValueError("Seules les requêtes de lecture (SELECT / WITH / DESCRIBE / SUMMARIZE / EXPLAIN) sont autorisées.") m = _FORBIDDEN.search(s) if m: raise ValueError(f"Mot-clé interdit dans le bac à sable : {m.group(0)}") return s def run_sql(sql: str, limit: int = 500) -> dict: """Exécute une requête de lecture avec délai maximal et plafond de lignes.""" s = guard_sql(sql) limit = max(1, min(int(limit), MAX_ROWS)) wrapped = s if re.match(r"^(describe|show|summarize|explain)\b", s, re.IGNORECASE) else f"select * from ({s}) as _q limit {limit + 1}" cur = cursor() result: dict = {} err: list[Exception] = [] def work(): try: t0 = time.perf_counter() res = cur.execute(wrapped) cols = [d[0] for d in res.description] data = res.fetchall() result["columns"] = cols result["truncated"] = len(data) > limit result["rows"] = [_jsonable(r) for r in data[:limit]] result["elapsed_ms"] = round((time.perf_counter() - t0) * 1000, 1) except Exception as e: # noqa: BLE001 err.append(e) th = threading.Thread(target=work, daemon=True) th.start() th.join(QUERY_TIMEOUT_S) if th.is_alive(): try: cur.interrupt() except Exception: pass th.join(2) cur.close() raise TimeoutError(f"Requête interrompue après {QUERY_TIMEOUT_S:.0f} s.") cur.close() if err: raise ValueError(str(err[0]).split("\n")[0][:400]) result["sql"] = s result["limit"] = limit return result # ---------------------------------------------------------------- embeddings (similarité entre projets) def embeddings() -> dict: global _emb_cache if _emb_cache is None: cur = cursor() try: data = cur.execute("select project_id, vector from project_embeddings").fetchall() finally: cur.close() ids = [d[0] for d in data] mat = np.asarray([d[1] for d in data], dtype=np.float32) norms = np.linalg.norm(mat, axis=1, keepdims=True) norms[norms == 0] = 1 mat = mat / norms _emb_cache = {"ids": ids, "index": {pid: i for i, pid in enumerate(ids)}, "mat": mat} return _emb_cache def similar_ids(project_id: str, k: int = 10) -> list[tuple[str, float]]: e = embeddings() i = e["index"].get(project_id) if i is None: return [] sims = e["mat"] @ e["mat"][i] order = np.argsort(-sims) out = [] for j in order: if j == i: continue out.append((e["ids"][j], float(sims[j]))) if len(out) >= k: break return out def nearest_to_vector(vec: np.ndarray, k: int = 20) -> list[tuple[str, float]]: e = embeddings() v = np.asarray(vec, dtype=np.float32) v = v / (np.linalg.norm(v) or 1) sims = e["mat"] @ v order = np.argsort(-sims)[:k] return [(e["ids"][j], float(sims[j])) for j in order]