import sqlite3 import json import os import uuid import zipfile import numpy as np from datetime import datetime from typing import List, Dict, Tuple, Optional from backend.config import SQLITE_DB_PATH, OBJECT_IMAGES_DIR from backend.database.models import TaughtObject, ObjectEmbedding, DetectionLog class DatabaseManager: def __init__(self, db_path: str = SQLITE_DB_PATH): self.db_path = db_path self.init_db() def get_connection(self): conn = sqlite3.connect(self.db_path) conn.row_factory = sqlite3.Row return conn def init_db(self): with self.get_connection() as conn: cursor = conn.cursor() cursor.execute(""" CREATE TABLE IF NOT EXISTS users ( id TEXT PRIMARY KEY, name TEXT NOT NULL, created_at TEXT NOT NULL ); """) cursor.execute(""" CREATE TABLE IF NOT EXISTS objects ( id TEXT PRIMARY KEY, user_id TEXT, name TEXT UNIQUE NOT NULL, category TEXT NOT NULL, description TEXT, status TEXT DEFAULT 'active', created_at TEXT NOT NULL, updated_at TEXT NOT NULL ); """) cursor.execute(""" CREATE TABLE IF NOT EXISTS object_images ( id TEXT PRIMARY KEY, object_id TEXT NOT NULL, file_path TEXT NOT NULL, image_hash TEXT, created_at TEXT NOT NULL, FOREIGN KEY (object_id) REFERENCES objects(id) ON DELETE CASCADE ); """) cursor.execute(""" CREATE TABLE IF NOT EXISTS object_embeddings ( id TEXT PRIMARY KEY, object_id TEXT NOT NULL, model_name TEXT NOT NULL, model_version TEXT NOT NULL, embedding TEXT NOT NULL, created_at TEXT NOT NULL, FOREIGN KEY (object_id) REFERENCES objects(id) ON DELETE CASCADE ); """) cursor.execute(""" CREATE TABLE IF NOT EXISTS detection_logs ( id TEXT PRIMARY KEY, user_id TEXT, object_name TEXT NOT NULL, detection_type TEXT NOT NULL, confidence REAL NOT NULL, timestamp TEXT NOT NULL ); """) conn.commit() def create_object(self, name: str, category: str = "general", description: str = "") -> TaughtObject: now = datetime.utcnow().isoformat() obj_id = f"obj_{uuid.uuid4().hex[:10]}" with self.get_connection() as conn: cursor = conn.cursor() cursor.execute( "INSERT INTO objects (id, name, category, description, created_at, updated_at) VALUES (?, ?, ?, ?, ?, ?)", (obj_id, name.strip(), category.strip(), description.strip(), now, now) ) conn.commit() return TaughtObject(id=obj_id, name=name, category=category, description=description, created_at=now) def get_object_by_name(self, name: str) -> Optional[TaughtObject]: with self.get_connection() as conn: cursor = conn.cursor() cursor.execute("SELECT * FROM objects WHERE LOWER(name) = LOWER(?)", (name,)) row = cursor.fetchone() if row: return TaughtObject( id=row["id"], name=row["name"], category=row["category"], description=row["description"], status=row["status"], created_at=row["created_at"] ) return None def get_all_objects(self) -> List[TaughtObject]: objects = [] with self.get_connection() as conn: cursor = conn.cursor() cursor.execute(""" SELECT o.*, (SELECT COUNT(*) FROM object_images WHERE object_id = o.id) as img_cnt, (SELECT COUNT(*) FROM object_embeddings WHERE object_id = o.id) as emb_cnt FROM objects o ORDER BY o.created_at DESC """) rows = cursor.fetchall() for r in rows: objects.append(TaughtObject( id=r["id"], name=r["name"], category=r["category"], description=r["description"] or "", status=r["status"], created_at=r["created_at"], image_count=r["img_cnt"], embedding_count=r["emb_cnt"] )) return objects def delete_object(self, object_id: str) -> bool: with self.get_connection() as conn: cursor = conn.cursor() cursor.execute("DELETE FROM object_embeddings WHERE object_id = ?", (object_id,)) cursor.execute("DELETE FROM object_images WHERE object_id = ?", (object_id,)) cursor.execute("DELETE FROM objects WHERE id = ?", (object_id,)) conn.commit() return cursor.rowcount > 0 def add_embedding(self, object_id: str, embedding: List[float], model_name: str = "resnet18", model_version: str = "1.0") -> str: emb_id = f"emb_{uuid.uuid4().hex[:10]}" now = datetime.utcnow().isoformat() emb_json = json.dumps(embedding) with self.get_connection() as conn: cursor = conn.cursor() cursor.execute( "INSERT INTO object_embeddings (id, object_id, model_name, model_version, embedding, created_at) VALUES (?, ?, ?, ?, ?, ?)", (emb_id, object_id, model_name, model_version, emb_json, now) ) conn.commit() return emb_id def add_image_record(self, object_id: str, file_path: str) -> str: img_id = f"img_{uuid.uuid4().hex[:10]}" now = datetime.utcnow().isoformat() with self.get_connection() as conn: cursor = conn.cursor() cursor.execute( "INSERT INTO object_images (id, object_id, file_path, created_at) VALUES (?, ?, ?, ?)", (img_id, object_id, file_path, now) ) conn.commit() return img_id def get_all_embeddings(self) -> List[Tuple[str, str, np.ndarray]]: """Returns list of (object_id, object_name, embedding_array)""" results = [] with self.get_connection() as conn: cursor = conn.cursor() cursor.execute(""" SELECT e.embedding, o.id, o.name FROM object_embeddings e JOIN objects o ON e.object_id = o.id WHERE o.status = 'active' """) rows = cursor.fetchall() for r in rows: emb_vec = np.array(json.loads(r["embedding"]), dtype=np.float32) results.append((r["id"], r["name"], emb_vec)) return results def find_best_matching_object(self, query_embedding: np.ndarray, min_threshold: float = 0.65) -> Tuple[Optional[str], float]: """Performs cosine similarity search over stored object embeddings""" all_embs = self.get_all_embeddings() if not all_embs: return None, 0.0 query_norm = query_embedding / (np.linalg.norm(query_embedding) + 1e-7) best_name = None best_similarity = -1.0 for obj_id, obj_name, emb_vec in all_embs: emb_norm = emb_vec / (np.linalg.norm(emb_vec) + 1e-7) sim = float(np.dot(query_norm, emb_norm)) if sim > best_similarity: best_similarity = sim best_name = obj_name if best_similarity >= min_threshold: return best_name, best_similarity return None, best_similarity def log_detection(self, object_name: str, detection_type: str, confidence: float): now = datetime.utcnow().isoformat() log_id = f"log_{uuid.uuid4().hex[:10]}" with self.get_connection() as conn: cursor = conn.cursor() cursor.execute( "INSERT INTO detection_logs (id, object_name, detection_type, confidence, timestamp) VALUES (?, ?, ?, ?, ?)", (log_id, object_name, detection_type, confidence, now) ) conn.commit() def export_database_zip(self, export_path: str): with zipfile.ZipFile(export_path, 'w') as zipf: if os.path.exists(self.db_path): zipf.write(self.db_path, arcname="database.db") manifest = { "version": "1.0", "exported_at": datetime.utcnow().isoformat(), "object_count": len(self.get_all_objects()) } zipf.writestr("manifest.json", json.dumps(manifest, indent=2)) db = DatabaseManager()