Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
16 changes: 14 additions & 2 deletions backend/app/database/connection.py
Original file line number Diff line number Diff line change
@@ -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]:
Expand All @@ -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
Expand Down
109 changes: 35 additions & 74 deletions backend/app/database/image_embeddings.py
Original file line number Diff line number Diff line change
@@ -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(
"""
Expand All @@ -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)
Expand All @@ -64,42 +56,32 @@ 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
WHERE model_version = ?
""",
(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(
Expand All @@ -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.
Expand All @@ -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]:
Expand All @@ -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(
"""
Expand All @@ -160,19 +135,14 @@ 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(
model_version: str, signature: str, limit: int
) -> 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(
"""
Expand All @@ -183,33 +153,24 @@ 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 = ?",
(model_version,),
)
else:
cursor.execute("SELECT COUNT(*) FROM image_embeddings")

result = cursor.fetchone()
return result[0] if result else 0
finally:
if conn:
conn.close()
68 changes: 32 additions & 36 deletions backend/app/database/images.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Comment thread
AbiramiR-27 marked this conversation as resolved.
from app.logging.setup_logging import get_logger

DATABASE_PATH = connection_module.DATABASE_PATH

# Initialize logger
logger = get_logger(__name__)

Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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]:
Expand Down
Loading
Loading