CoolFace
Apppublic

sahilfarib/Legal-Document-Intelligence

sourceHugging Faceupdated 4mo agoView on Hugging Face
0likes
pattern_store.py189 linesDownload Raw Back to feedback
1"""2feedback/pattern_store.py — SQLite CRUD for learned patterns and draft history3 4Uses context manager for safe connection handling. All schema is5created on first init. Supports draft history and feedback logging.6"""7 8import sqlite39from contextlib import contextmanager10from datetime import datetime11from uuid import uuid412from pathlib import Path13 14from loguru import logger15from config import PATTERNS_DB_PATH16 17 18SCHEMA = """\19CREATE TABLE IF NOT EXISTS learned_patterns (20    id          TEXT PRIMARY KEY,21    pattern     TEXT NOT NULL,22    doc_type    TEXT,23    source_diff TEXT,24    created_at  TEXT NOT NULL,25    use_count   INTEGER DEFAULT 0,26    active      INTEGER DEFAULT 127);28 29CREATE TABLE IF NOT EXISTS draft_history (30    draft_id        TEXT PRIMARY KEY,31    draft_type      TEXT NOT NULL,32    doc_ids         TEXT NOT NULL,33    draft_text      TEXT NOT NULL,34    validation_json TEXT,35    created_at      TEXT NOT NULL,36    was_edited      INTEGER DEFAULT 037);38 39CREATE TABLE IF NOT EXISTS feedback_log (40    id            TEXT PRIMARY KEY,41    draft_id      TEXT NOT NULL,42    original_text TEXT NOT NULL,43    edited_text   TEXT NOT NULL,44    change_ratio  REAL,45    pattern_id    TEXT,46    created_at    TEXT NOT NULL47);48 49CREATE INDEX IF NOT EXISTS idx_patterns_active50    ON learned_patterns (active, doc_type);51"""52 53 54class PatternStore:55    """SQLite-backed store for learned patterns and draft history."""56 57    def __init__(self, db_path: str | Path = PATTERNS_DB_PATH):58        self.db_path = str(db_path)59        Path(self.db_path).parent.mkdir(parents=True, exist_ok=True)60        self._init_db()61 62    @contextmanager63    def _conn(self):64        conn = sqlite3.connect(self.db_path)65        conn.row_factory = sqlite3.Row66        try:67            yield conn68            conn.commit()69        except Exception:70            conn.rollback()71            raise72        finally:73            conn.close()74 75    def _init_db(self):76        with self._conn() as conn:77            conn.executescript(SCHEMA)78        logger.debug(f"PatternStore initialised: {self.db_path}")79 80    # ----- Pattern CRUD -----81 82    def save_pattern(self, pattern: str, doc_type: str, source_diff: str = "") -> str:83        pid = str(uuid4())84        with self._conn() as conn:85            conn.execute(86                "INSERT INTO learned_patterns VALUES (?,?,?,?,?,?,?)",87                (pid, pattern, doc_type, source_diff,88                 datetime.utcnow().isoformat(), 0, 1),89            )90        logger.info(f"Pattern saved: [{pid[:8]}] {pattern}")91        return pid92 93    def update_pattern(self, pattern_id: str, new_text: str):94        with self._conn() as conn:95            conn.execute(96                "UPDATE learned_patterns SET pattern=? WHERE id=?",97                (new_text, pattern_id),98            )99        logger.info(f"Pattern updated: [{pattern_id[:8]}]")100 101    def deactivate_pattern(self, pattern_id: str):102        with self._conn() as conn:103            conn.execute(104                "UPDATE learned_patterns SET active=0 WHERE id=?",105                (pattern_id,),106            )107 108    def delete_pattern(self, pattern_id: str):109        with self._conn() as conn:110            conn.execute("DELETE FROM learned_patterns WHERE id=?", (pattern_id,))111 112    def get_active_patterns(113        self, doc_type: str | None = None, limit: int = 8114    ) -> list[tuple[str, str]]:115        with self._conn() as conn:116            if doc_type:117                rows = conn.execute(118                    "SELECT id, pattern FROM learned_patterns "119                    "WHERE active=1 AND (doc_type=? OR doc_type IS NULL) "120                    "ORDER BY use_count DESC LIMIT ?",121                    (doc_type, limit),122                ).fetchall()123            else:124                rows = conn.execute(125                    "SELECT id, pattern FROM learned_patterns "126                    "WHERE active=1 ORDER BY use_count DESC LIMIT ?",127                    (limit,),128                ).fetchall()129        return [(r["id"], r["pattern"]) for r in rows]130 131    def get_all_patterns(self) -> list[dict]:132        with self._conn() as conn:133            rows = conn.execute(134                "SELECT * FROM learned_patterns ORDER BY created_at DESC"135            ).fetchall()136        return [dict(r) for r in rows]137 138    def increment_use_count(self, pattern_ids: list[str]):139        with self._conn() as conn:140            for pid in pattern_ids:141                conn.execute(142                    "UPDATE learned_patterns SET use_count = use_count + 1 WHERE id=?",143                    (pid,),144                )145 146    def pattern_count(self, active_only: bool = True) -> int:147        with self._conn() as conn:148            q = "SELECT COUNT(*) FROM learned_patterns"149            if active_only:150                q += " WHERE active=1"151            return conn.execute(q).fetchone()[0]152 153    # ----- Draft History -----154 155    def save_draft(self, draft_id: str, draft_type: str, doc_ids: str,156                   draft_text: str, validation_json: str = ""):157        with self._conn() as conn:158            conn.execute(159                "INSERT OR REPLACE INTO draft_history VALUES (?,?,?,?,?,?,?)",160                (draft_id, draft_type, doc_ids, draft_text, validation_json,161                 datetime.utcnow().isoformat(), 0),162            )163 164    def get_draft(self, draft_id: str) -> dict | None:165        with self._conn() as conn:166            row = conn.execute(167                "SELECT * FROM draft_history WHERE draft_id=?", (draft_id,)168            ).fetchone()169        return dict(row) if row else None170 171    def mark_draft_edited(self, draft_id: str):172        with self._conn() as conn:173            conn.execute(174                "UPDATE draft_history SET was_edited=1 WHERE draft_id=?",175                (draft_id,),176            )177 178    # ----- Feedback Log -----179 180    def log_feedback(self, draft_id: str, original: str, edited: str,181                     change_ratio: float, pattern_id: str = ""):182        fid = str(uuid4())183        with self._conn() as conn:184            conn.execute(185                "INSERT INTO feedback_log VALUES (?,?,?,?,?,?,?)",186                (fid, draft_id, original, edited, change_ratio,187                 pattern_id, datetime.utcnow().isoformat()),188            )189