import sqlite3 import uuid from datetime import datetime, timezone from typing import Optional, List from .models import CorrectionRecord DB_PATH = "./chroma_db/correction_store.db" _CREATE_SQL = """ CREATE TABLE IF NOT EXISTS corrections ( correction_id TEXT PRIMARY KEY, session_id TEXT NOT NULL, timestamp TEXT NOT NULL, section TEXT NOT NULL, correction_text TEXT NOT NULL, corrected_snippet TEXT, root_cause TEXT NOT NULL, action_taken TEXT NOT NULL, affected_chunk_ids TEXT NOT NULL, original_report TEXT NOT NULL, features_json TEXT NOT NULL, retrieved_context_json TEXT NOT NULL, verified INTEGER DEFAULT 0, verification_result TEXT ); CREATE TABLE IF NOT EXISTS chunk_edits ( edit_id TEXT PRIMARY KEY, correction_id TEXT NOT NULL REFERENCES corrections(correction_id), chroma_id TEXT NOT NULL, edit_type TEXT NOT NULL, original_content TEXT, new_content TEXT, timestamp TEXT NOT NULL ); """ class CorrectionStore: def __init__(self, db_path: str = DB_PATH): self.db_path = db_path self._init_db() def _init_db(self) -> None: with sqlite3.connect(self.db_path) as conn: conn.executescript(_CREATE_SQL) def save_correction(self, record: CorrectionRecord) -> None: with sqlite3.connect(self.db_path) as conn: conn.execute( """INSERT INTO corrections VALUES (?,?,?,?,?,?,?,?,?,?,?,?,?,?)""", ( record.correction_id, record.session_id, record.timestamp, record.section, record.correction_text, record.corrected_snippet, record.root_cause, record.action_taken, record.affected_chunk_ids, record.original_report, record.features_json, record.retrieved_context_json, 1 if record.verified else 0, record.verification_result, ), ) def save_chunk_edit( self, correction_id: str, chroma_id: str, edit_type: str, original: Optional[str], new_content: Optional[str], ) -> None: with sqlite3.connect(self.db_path) as conn: conn.execute( "INSERT INTO chunk_edits VALUES (?,?,?,?,?,?,?)", ( str(uuid.uuid4()), correction_id, chroma_id, edit_type, original, new_content, datetime.now(timezone.utc).isoformat(), ), ) def get_correction(self, correction_id: str) -> Optional[dict]: with sqlite3.connect(self.db_path) as conn: conn.row_factory = sqlite3.Row row = conn.execute( "SELECT * FROM corrections WHERE correction_id=?", (correction_id,), ).fetchone() return dict(row) if row else None def list_corrections(self, limit: int = 50) -> List[dict]: with sqlite3.connect(self.db_path) as conn: conn.row_factory = sqlite3.Row rows = conn.execute( "SELECT * FROM corrections ORDER BY timestamp DESC LIMIT ?", (limit,), ).fetchall() return [dict(r) for r in rows] def mark_verified(self, correction_id: str, result: str) -> None: with sqlite3.connect(self.db_path) as conn: conn.execute( "UPDATE corrections SET verified=1, verification_result=? WHERE correction_id=?", (result, correction_id), )