faisaltitu/Drift-Detection
0
1"""2Database Module - SQLite for Prediction Logging3"""4 5import logging6import sqlite37from datetime import datetime8from pathlib import Path9from typing import Dict, List, Optional10import json11 12logger = logging.getLogger(__name__)13 14# Database path15DB_PATH = Path(__file__).parent.parent / "data" / "predictions.db"16 17 18def get_connection() -> sqlite3.Connection:19 """Get SQLite database connection."""20 DB_PATH.parent.mkdir(parents=True, exist_ok=True)21 conn = sqlite3.connect(str(DB_PATH), check_same_thread=False)22 conn.row_factory = sqlite3.Row23 return conn24 25 26def init_db() -> None:27 """Initialize database tables."""28 conn = get_connection()29 cursor = conn.cursor()30 31 cursor.execute("""32 CREATE TABLE IF NOT EXISTS predictions (33 id INTEGER PRIMARY KEY AUTOINCREMENT,34 input_data TEXT NOT NULL,35 prediction REAL NOT NULL,36 model_version INTEGER NOT NULL,37 timestamp DATETIME DEFAULT CURRENT_TIMESTAMP38 )39 """)40 41 conn.commit()42 conn.close()43 logger.info("Database initialized: %s", DB_PATH)44 45 46def log_prediction(47 input_data: Dict,48 prediction: float,49 model_version: int50) -> int:51 """52 Log a prediction to the database.53 54 Args:55 input_data: Input features as dictionary56 prediction: Model prediction57 model_version: Version of model used58 59 Returns:60 Prediction ID61 """62 conn = get_connection()63 cursor = conn.cursor()64 65 cursor.execute("""66 INSERT INTO predictions (input_data, prediction, model_version, timestamp)67 VALUES (?, ?, ?, ?)68 """, (json.dumps(input_data), prediction, model_version, datetime.now().isoformat()))69 70 prediction_id = cursor.lastrowid71 conn.commit()72 conn.close()73 74 return prediction_id75 76 77def get_predictions(limit: int = 100) -> List[Dict]:78 """79 Get recent predictions.80 81 Args:82 limit: Maximum number of predictions to return83 84 Returns:85 List of prediction records86 """87 conn = get_connection()88 cursor = conn.cursor()89 90 cursor.execute("""91 SELECT id, input_data, prediction, model_version, timestamp92 FROM predictions93 ORDER BY id DESC94 LIMIT ?95 """, (limit,))96 97 rows = cursor.fetchall()98 conn.close()99 100 return [101 {102 "id": row["id"],103 "input_data": json.loads(row["input_data"]),104 "prediction": row["prediction"],105 "model_version": row["model_version"],106 "timestamp": row["timestamp"],107 }108 for row in rows109 ]110 111 112def get_prediction_count() -> int:113 """Get total number of predictions."""114 conn = get_connection()115 cursor = conn.cursor()116 117 cursor.execute("SELECT COUNT(*) FROM predictions")118 count = cursor.fetchone()[0]119 conn.close()120 121 return count122 123 124# Initialize on import125init_db()126 