diff --git a/backend/app/database/connection.py b/backend/app/database/connection.py index d82ad6950..8eb04e8aa 100644 --- a/backend/app/database/connection.py +++ b/backend/app/database/connection.py @@ -1,11 +1,23 @@ import sqlite3 from contextlib import contextmanager from typing import Generator -from app.config.settings import DATABASE_PATH +import app.config.settings as settings from app.logging.setup_logging import get_logger logger = get_logger(__name__) +ORIGINAL_DATABASE_PATH = settings.DATABASE_PATH +DATABASE_PATH = settings.DATABASE_PATH + + +def get_database_path() -> str: + """Resolve the active database path dynamically, supporting + various test patching styles. + """ + if DATABASE_PATH != ORIGINAL_DATABASE_PATH: + return DATABASE_PATH + return settings.DATABASE_PATH + @contextmanager def get_db_connection() -> Generator[sqlite3.Connection, None, None]: @@ -16,7 +28,7 @@ def get_db_connection() -> Generator[sqlite3.Connection, None, None]: - Works for both single and multi-step transactions - Automatically commits on success or rolls back on failure """ - conn = sqlite3.connect(DATABASE_PATH) + conn = sqlite3.connect(get_database_path()) # --- Strict enforcement of all relational and logical rules --- conn.execute("PRAGMA foreign_keys = ON;") # Enforce FK constraints diff --git a/backend/app/database/image_embeddings.py b/backend/app/database/image_embeddings.py index 5a524d1c0..0eddb2d95 100644 --- a/backend/app/database/image_embeddings.py +++ b/backend/app/database/image_embeddings.py @@ -1,12 +1,11 @@ from typing import Dict, List, Tuple import numpy as np -from app.database.images import SQLITE_ID_CHUNK, _connect +from app.database.connection import get_db_connection +from app.database.images import SQLITE_ID_CHUNK def db_create_image_embeddings_table(): - conn = None - try: - conn = _connect() + with get_db_connection() as conn: cursor = conn.cursor() cursor.execute( """ @@ -31,28 +30,21 @@ def db_create_image_embeddings_table(): cursor.execute( "ALTER TABLE image_embeddings ADD COLUMN scored_signature TEXT" ) - conn.commit() - finally: - if conn: - conn.close() def db_upsert_image_embeddings(rows: List[Tuple[str, str, np.ndarray]]): - conn = None - try: - conn = _connect() - cursor = conn.cursor() - - # Convert each embedding - db_rows = [ - ( - image_id, - model_version, - np.ascontiguousarray(embedding, dtype=np.float32).tobytes(), - ) - for image_id, model_version, embedding in rows - ] + # Convert each embedding + db_rows = [ + ( + image_id, + model_version, + np.ascontiguousarray(embedding, dtype=np.float32).tobytes(), + ) + for image_id, model_version, embedding in rows + ] + with get_db_connection() as conn: + cursor = conn.cursor() cursor.executemany( """ INSERT INTO image_embeddings (image_id, model_version, embedding) @@ -64,18 +56,11 @@ def db_upsert_image_embeddings(rows: List[Tuple[str, str, np.ndarray]]): """, db_rows, ) - conn.commit() - finally: - if conn: - conn.close() def db_get_all_embeddings(model_version: str) -> Tuple[List[str], np.ndarray]: - conn = None - try: - conn = _connect() + with get_db_connection() as conn: cursor = conn.cursor() - cursor.execute( """ SELECT image_id, embedding FROM image_embeddings @@ -83,23 +68,20 @@ def db_get_all_embeddings(model_version: str) -> Tuple[List[str], np.ndarray]: """, (model_version,), ) - rows = cursor.fetchall() - if not rows: - return [], np.empty((0, 0), dtype=np.float32) - image_ids = [] - embeddings_list = [] + if not rows: + return [], np.empty((0, 0), dtype=np.float32) + + image_ids = [] + embeddings_list = [] - for image_id, blob in rows: - image_ids.append(image_id) - embeddings_list.append(np.frombuffer(blob, dtype=np.float32)) + for image_id, blob in rows: + image_ids.append(image_id) + embeddings_list.append(np.frombuffer(blob, dtype=np.float32)) - matrix = np.vstack(embeddings_list) - return image_ids, matrix - finally: - if conn: - conn.close() + matrix = np.vstack(embeddings_list) + return image_ids, matrix def db_get_embeddings_for_image_ids( @@ -115,9 +97,7 @@ def db_get_embeddings_for_image_ids( if not image_ids: return {} - conn = None - try: - conn = _connect() + with get_db_connection() as conn: cursor = conn.cursor() found: Dict[str, np.ndarray] = {} # Chunked to stay under SQLite's variable limit, as elsewhere. @@ -134,9 +114,6 @@ def db_get_embeddings_for_image_ids( for image_id, blob in cursor.fetchall(): found[image_id] = np.frombuffer(blob, dtype=np.float32) return found - finally: - if conn: - conn.close() def db_get_embedding_sample(model_version: str, limit: int) -> List[np.ndarray]: @@ -146,9 +123,7 @@ def db_get_embedding_sample(model_version: str, limit: int) -> List[np.ndarray]: Ordered by image id rather than sampled randomly so a run is reproducible; ids are UUIDs, so that is already an arbitrary slice. """ - conn = None - try: - conn = _connect() + with get_db_connection() as conn: cursor = conn.cursor() cursor.execute( """ @@ -160,9 +135,6 @@ def db_get_embedding_sample(model_version: str, limit: int) -> List[np.ndarray]: (model_version, limit), ) return [np.frombuffer(row[0], dtype=np.float32) for row in cursor.fetchall()] - finally: - if conn: - conn.close() def db_get_embeddings_needing_scoring( @@ -170,9 +142,7 @@ def db_get_embeddings_needing_scoring( ) -> Tuple[List[str], np.ndarray]: """Embeddings whose semantic scores are missing or from another vocabulary/label state. Returns up to `limit` (image_ids, matrix).""" - conn = None - try: - conn = _connect() + with get_db_connection() as conn: cursor = conn.cursor() cursor.execute( """ @@ -183,23 +153,18 @@ def db_get_embeddings_needing_scoring( (model_version, signature, limit), ) rows = cursor.fetchall() - if not rows: - return [], np.empty((0, 0), dtype=np.float32) - image_ids = [image_id for image_id, _ in rows] - matrix = np.vstack([np.frombuffer(blob, dtype=np.float32) for _, blob in rows]) - return image_ids, matrix - finally: - if conn: - conn.close() + if not rows: + return [], np.empty((0, 0), dtype=np.float32) + + image_ids = [image_id for image_id, _ in rows] + matrix = np.vstack([np.frombuffer(blob, dtype=np.float32) for _, blob in rows]) + return image_ids, matrix def db_count_embeddings(model_version: str | None = None) -> int: - conn = None - try: - conn = _connect() + with get_db_connection() as conn: cursor = conn.cursor() - if model_version is not None: cursor.execute( "SELECT COUNT(*) FROM image_embeddings WHERE model_version = ?", @@ -207,9 +172,5 @@ def db_count_embeddings(model_version: str | None = None) -> int: ) else: cursor.execute("SELECT COUNT(*) FROM image_embeddings") - result = cursor.fetchone() return result[0] if result else 0 - finally: - if conn: - conn.close() diff --git a/backend/app/database/images.py b/backend/app/database/images.py index 2eb1ab021..02584d762 100644 --- a/backend/app/database/images.py +++ b/backend/app/database/images.py @@ -6,11 +6,12 @@ from datetime import datetime # App-specific imports -from app.config.settings import ( - DATABASE_PATH, -) +import app.database.connection as connection_module +from app.database.connection import get_db_connection from app.logging.setup_logging import get_logger +DATABASE_PATH = connection_module.DATABASE_PATH + # Initialize logger logger = get_logger(__name__) @@ -56,7 +57,7 @@ class UntaggedImageRecord(TypedDict): def _connect() -> sqlite3.Connection: - conn = sqlite3.connect(DATABASE_PATH) + conn = sqlite3.connect(connection_module.get_database_path()) # Ensure ON DELETE CASCADE and other FKs are enforced conn.execute("PRAGMA foreign_keys = ON") return conn @@ -480,50 +481,45 @@ def db_delete_images_by_ids(image_ids: List[ImageId]) -> bool: if not image_ids: return True - conn = _connect() - cursor = conn.cursor() - try: - # Create placeholders for the IN clause - placeholders = ",".join("?" for _ in image_ids) - cursor.execute( - f"DELETE FROM images WHERE id IN ({placeholders})", - image_ids, - ) - conn.commit() - logger.info(f"Deleted {cursor.rowcount} obsolete image(s) from database") + total_deleted = 0 + with get_db_connection() as conn: + cursor = conn.cursor() + for start in range(0, len(image_ids), SQLITE_ID_CHUNK): + chunk = image_ids[start : start + SQLITE_ID_CHUNK] + placeholders = ",".join("?" for _ in chunk) + cursor.execute( + f"DELETE FROM images WHERE id IN ({placeholders})", + chunk, + ) + total_deleted += cursor.rowcount + logger.info(f"Deleted {total_deleted} obsolete image(s) from database") return True except sqlite3.Error as e: logger.error(f"Error deleting images: {e}") - conn.rollback() return False - finally: - conn.close() def db_toggle_image_favourite_status(image_id: str) -> bool: - conn = sqlite3.connect(DATABASE_PATH) - cursor = conn.cursor() try: - cursor.execute("SELECT id FROM images WHERE id = ?", (image_id,)) - if not cursor.fetchone(): - return False - cursor.execute( - """ - UPDATE images - SET isFavourite = CASE WHEN isFavourite = 1 THEN 0 ELSE 1 END - WHERE id = ? - """, - (image_id,), - ) - conn.commit() - return cursor.rowcount > 0 + with get_db_connection() as conn: + cursor = conn.cursor() + cursor.execute("SELECT id FROM images WHERE id = ?", (image_id,)) + if not cursor.fetchone(): + return False + cursor.execute( + """ + UPDATE images + SET isFavourite = CASE WHEN isFavourite = 1 THEN 0 ELSE 1 END + WHERE id = ? + """, + (image_id,), + ) + row_count = cursor.rowcount + return row_count > 0 except sqlite3.Error as e: logger.error(f"Database error: {e}") - conn.rollback() return False - finally: - conn.close() def db_get_image_by_id(image_id: str) -> Optional[dict]: diff --git a/backend/app/database/semantic_labels.py b/backend/app/database/semantic_labels.py index 18d1c0c24..e9e6bffc6 100644 --- a/backend/app/database/semantic_labels.py +++ b/backend/app/database/semantic_labels.py @@ -1,9 +1,7 @@ import json from typing import List, Tuple - import numpy as np - -from app.database.images import _connect +from app.database.connection import get_db_connection from app.logging.setup_logging import get_logger logger = get_logger(__name__) @@ -13,9 +11,7 @@ def db_create_semantic_labels_table(): - conn = None - try: - conn = _connect() + with get_db_connection() as conn: cursor = conn.cursor() # Migrate the pre-vocabulary shell schema. It shipped with no writer, @@ -71,11 +67,6 @@ def db_create_semantic_labels_table(): """ ) - conn.commit() - finally: - if conn: - conn.close() - def db_upsert_semantic_vocabulary(labels: List[dict]) -> None: """Idempotently sync the seed vocabulary into mappings + semantic_labels. @@ -85,9 +76,7 @@ def db_upsert_semantic_vocabulary(labels: List[dict]) -> None: missing from the seed are deactivated (rows kept -- image_classes may reference them). """ - conn = None - try: - conn = _connect() + with get_db_connection() as conn: cursor = conn.cursor() cursor.execute("SELECT class_id, name FROM mappings") @@ -186,15 +175,11 @@ def db_upsert_semantic_vocabulary(labels: List[dict]) -> None: ) deactivated += 1 - conn.commit() if added or updated or skipped or deactivated: logger.info( f"Semantic vocabulary sync: {added} added, {updated} updated, " f"{deactivated} deactivated, {skipped} skipped" ) - finally: - if conn: - conn.close() def db_get_labels_needing_embeddings( @@ -202,9 +187,7 @@ def db_get_labels_needing_embeddings( ) -> List[Tuple[int, List[str]]]: """Active labels whose cached embedding is missing or belongs to a different checkpoint. Returns (class_id, descriptions) pairs.""" - conn = None - try: - conn = _connect() + with get_db_connection() as conn: cursor = conn.cursor() cursor.execute( """ @@ -219,9 +202,6 @@ def db_get_labels_needing_embeddings( (class_id, json.loads(descriptions)) for class_id, descriptions in cursor.fetchall() ] - finally: - if conn: - conn.close() def db_update_label_embeddings( @@ -231,9 +211,7 @@ def db_update_label_embeddings( Same raw-float32 blob format as image_embeddings.embedding. """ - conn = None - try: - conn = _connect() + with get_db_connection() as conn: cursor = conn.cursor() cursor.executemany( """ @@ -250,10 +228,6 @@ def db_update_label_embeddings( for class_id, embedding, model_version in rows ], ) - conn.commit() - finally: - if conn: - conn.close() def db_get_active_label_embeddings( @@ -266,9 +240,7 @@ def db_get_active_label_embeddings( matrix, so row order must be deterministic); threshold is None where the label has no per-label override. """ - conn = None - try: - conn = _connect() + with get_db_connection() as conn: cursor = conn.cursor() cursor.execute( """ @@ -292,9 +264,6 @@ def db_get_active_label_embeddings( [np.frombuffer(blob, dtype=np.float32) for _, _, _, blob in rows] ) return meta, matrix - finally: - if conn: - conn.close() def db_write_image_semantic_scores( @@ -303,9 +272,7 @@ def db_write_image_semantic_scores( """Replace each image's semantic tag rows with the given (class_id, score) pairs and stamp its scored_signature. YOLO rows (class_id below the offset) are never touched.""" - conn = None - try: - conn = _connect() + with get_db_connection() as conn: cursor = conn.cursor() for image_id, pairs in batch: cursor.execute( @@ -322,7 +289,3 @@ def db_write_image_semantic_scores( "WHERE image_id = ?", (signature, image_id), ) - conn.commit() - finally: - if conn: - conn.close() diff --git a/backend/tests/test_image_embeddings.py b/backend/tests/test_image_embeddings.py index 0871953b2..08df8d043 100644 --- a/backend/tests/test_image_embeddings.py +++ b/backend/tests/test_image_embeddings.py @@ -4,7 +4,9 @@ import app.database.images as images_module import app.database.folders as folders_module import app.database.yolo_mapping as yolo_mapping_module +import app.database.connection as connection_module from app.database.images import _connect, db_create_images_table +from app.database.connection import get_db_connection from app.database.folders import db_create_folders_table from app.database.yolo_mapping import db_create_YOLO_classes_table from app.database.image_embeddings import ( @@ -36,6 +38,7 @@ def _isolated_db(tmp_path, monkeypatch): monkeypatch.setattr(images_module, "DATABASE_PATH", db_path) monkeypatch.setattr(folders_module, "DATABASE_PATH", db_path) monkeypatch.setattr(yolo_mapping_module, "DATABASE_PATH", db_path) + monkeypatch.setattr(connection_module, "DATABASE_PATH", db_path) # images' schema FK-references folders/mappings; SQLite validates that # the referenced tables exist at INSERT time even for a NULL FK value, @@ -162,10 +165,55 @@ def test_deleting_image_cascades_to_its_embedding(self): ) assert db_count_embeddings("siglip2-base-patch16-224") == 1 - conn = _connect() - conn.execute("DELETE FROM images WHERE id = ?", ("img7",)) - conn.commit() - conn.close() + with get_db_connection() as conn: + conn.execute("DELETE FROM images WHERE id = ?", ("img7",)) ids, _ = db_get_all_embeddings("siglip2-base-patch16-224") assert "img7" not in ids + + def test_deleting_image_cascades_to_embeddings_and_classes_regression(self): + # 1. Insert two dummy images + _insert_dummy_image("img7") + _insert_dummy_image("img8") + + # 2. Insert embeddings for both + db_upsert_image_embeddings( + [ + ("img7", "siglip2-base-patch16-224", np.ones(3, dtype=np.float32)), + ("img8", "siglip2-base-patch16-224", np.ones(3, dtype=np.float32)), + ] + ) + + # 3. Insert image classes (semantic scores) for both + with get_db_connection() as conn: + conn.execute( + "INSERT INTO image_classes (image_id, class_id, score) VALUES (?, 1, 0.85)", + ("img7",), + ) + conn.execute( + "INSERT INTO image_classes (image_id, class_id, score) VALUES (?, 1, 0.95)", + ("img8",), + ) + + # Verify initial state + assert db_count_embeddings("siglip2-base-patch16-224") == 2 + with get_db_connection() as conn: + res = conn.execute("SELECT COUNT(*) FROM image_classes").fetchone() + assert res[0] == 2 + + # 4. Call production delete logic for img7 + from app.database.images import db_delete_images_by_ids + + db_delete_images_by_ids(["img7"]) + + # 5. Assert img7 cascades deleted, but img8 remains intact + ids, _ = db_get_all_embeddings("siglip2-base-patch16-224") + assert "img7" not in ids + assert "img8" in ids + + with get_db_connection() as conn: + remaining_classes = conn.execute( + "SELECT image_id FROM image_classes" + ).fetchall() + assert len(remaining_classes) == 1 + assert remaining_classes[0][0] == "img8" diff --git a/backend/tests/test_semantic_labels.py b/backend/tests/test_semantic_labels.py index 965a515af..a8ae7bd72 100644 --- a/backend/tests/test_semantic_labels.py +++ b/backend/tests/test_semantic_labels.py @@ -4,6 +4,7 @@ import app.database.images as images_module import app.database.folders as folders_module import app.database.yolo_mapping as yolo_mapping_module +import app.database.connection as connection_module from app.database.images import _connect, db_create_images_table from app.database.folders import db_create_folders_table from app.database.yolo_mapping import db_create_YOLO_classes_table @@ -32,6 +33,7 @@ def _isolated_db(tmp_path, monkeypatch): monkeypatch.setattr(images_module, "DATABASE_PATH", db_path) monkeypatch.setattr(folders_module, "DATABASE_PATH", db_path) monkeypatch.setattr(yolo_mapping_module, "DATABASE_PATH", db_path) + monkeypatch.setattr(connection_module, "DATABASE_PATH", db_path) db_create_YOLO_classes_table() db_create_folders_table()