pandeydigant31's picture
Batch C: #10 reset cache button, #4 enhanced Guide tab help
0bbc8ce verified
Raw
History Blame Contribute Delete
4.67 kB
"""SQLite-backed cache for simulation sweep results, keyed on SHA256 of params."""
from __future__ import annotations
import hashlib
import json
import logging
import sqlite3
from typing import Any, Dict, Optional, Tuple
logger = logging.getLogger(__name__)
SCHEMA_VERSION = 3
class SweepCache:
"""Stores simulation scalar results in a local SQLite database.
Keys are derived from a SHA256 hash of the canonicalised parameter dict,
so identical runs are served from cache instead of re-simulated.
Schema is versioned via ``PRAGMA user_version``. On open, if the stored
version differs from ``SCHEMA_VERSION``, all rows are dropped and the
version is written.
"""
def __init__(self, db_path: str = "sweep_cache.db") -> None:
# check_same_thread=False: Gradio serves requests from a worker pool,
# so the same SweepCache instance is used across threads. Access is
# serialized at the Python level by SQLite's own locking, so cross-
# thread use is safe in practice.
self._conn = sqlite3.connect(db_path, check_same_thread=False)
self._conn.execute("PRAGMA journal_mode=WAL")
self._conn.execute(
"CREATE TABLE IF NOT EXISTS cache ("
" key TEXT PRIMARY KEY,"
" scalars TEXT,"
" wall_time REAL,"
" ts TIMESTAMP DEFAULT CURRENT_TIMESTAMP"
")"
)
self._migrate_if_needed()
self._conn.commit()
self._session_hits = 0
self._session_misses = 0
def _migrate_if_needed(self) -> None:
"""If stored schema version != SCHEMA_VERSION, drop all rows."""
cur = self._conn.execute("PRAGMA user_version")
current = cur.fetchone()[0]
if current != SCHEMA_VERSION:
logger.warning(
"sweep_cache schema mismatch (stored=%s, expected=%s); "
"wiping cache",
current, SCHEMA_VERSION,
)
self._conn.execute("DELETE FROM cache")
self._conn.execute(f"PRAGMA user_version = {SCHEMA_VERSION}")
# ------------------------------------------------------------------
# Internal helpers
# ------------------------------------------------------------------
@staticmethod
def _make_key(params: Dict[str, Any]) -> str:
"""Canonical SHA256 key: round floats to 8 dp, sort keys, hash."""
cleaned: Dict[str, Any] = {}
for k in sorted(params.keys()):
v = params[k]
cleaned[k] = round(v, 8) if isinstance(v, float) else v
blob = json.dumps(cleaned, sort_keys=True).encode()
return hashlib.sha256(blob).hexdigest()[:16]
# ------------------------------------------------------------------
# Public API
# ------------------------------------------------------------------
def get(self, params: Dict[str, Any]) -> Optional[Tuple[Dict[str, Any], float]]:
"""Return ``(scalars_dict, wall_time)`` or ``None`` on cache miss."""
key = self._make_key(params)
row = self._conn.execute(
"SELECT scalars, wall_time FROM cache WHERE key = ?", (key,)
).fetchone()
if row is None:
self._session_misses += 1
return None
self._session_hits += 1
return json.loads(row[0]), row[1]
def put(
self,
params: Dict[str, Any],
scalars: Dict[str, Any],
wall_time: float,
) -> None:
"""Insert (or replace) a cache entry."""
key = self._make_key(params)
self._conn.execute(
"INSERT OR REPLACE INTO cache (key, scalars, wall_time) VALUES (?, ?, ?)",
(key, json.dumps(scalars), wall_time),
)
self._conn.commit()
def clear(self) -> int:
"""Delete all cached entries. Returns the number of rows deleted."""
count = self._conn.execute("SELECT COUNT(*) FROM cache").fetchone()[0]
self._conn.execute("DELETE FROM cache")
self._conn.commit()
self._session_hits = 0
self._session_misses = 0
return count
@property
def stats(self) -> Dict[str, Any]:
"""Return cache statistics for the current session."""
total = self._conn.execute("SELECT COUNT(*) FROM cache").fetchone()[0]
total_lookups = self._session_hits + self._session_misses
return {
"total_cached": total,
"session_hits": self._session_hits,
"session_misses": self._session_misses,
"hit_rate": (
self._session_hits / total_lookups if total_lookups > 0 else 0.0
),
}