mirror of
https://github.com/allaunthefox/Research-Stack.git
synced 2026-07-31 03:05:21 +00:00
1304 lines
46 KiB
Python
1304 lines
46 KiB
Python
# SPDX-License-Identifier: Apache-2.0
|
|
# Copyright (c) 2025 NoDupeLabs
|
|
|
|
"""Comprehensive tests for embedding storage and queries.
|
|
|
|
Tests cover:
|
|
- Embedding storage and retrieval
|
|
- Vector operations and queries
|
|
- Similarity search functionality
|
|
- Batch operations
|
|
- Model version management
|
|
- Error handling
|
|
- Schema validation (including dimensions column bug fix)
|
|
"""
|
|
|
|
import os
|
|
import pickle
|
|
import sqlite3
|
|
|
|
import numpy as np
|
|
import pytest
|
|
|
|
from nodupe.tools.databases.embeddings import EmbeddingRepository, get_embedding_repository
|
|
|
|
# =============================================================================
|
|
# Fixtures
|
|
# =============================================================================
|
|
|
|
@pytest.fixture
|
|
def in_memory_db():
|
|
"""Create an in-memory SQLite database connection."""
|
|
conn = sqlite3.connect(":memory:")
|
|
conn.execute("PRAGMA foreign_keys = ON")
|
|
# Create minimal schema for embeddings
|
|
conn.execute("""
|
|
CREATE TABLE IF NOT EXISTS embeddings (
|
|
id INTEGER PRIMARY KEY AUTOINCREMENT,
|
|
file_id INTEGER NOT NULL,
|
|
embedding BLOB NOT NULL,
|
|
model_version TEXT NOT NULL,
|
|
created_time INTEGER NOT NULL,
|
|
dimensions INTEGER NOT NULL
|
|
)
|
|
""")
|
|
conn.execute("""
|
|
CREATE TABLE IF NOT EXISTS files (
|
|
id INTEGER PRIMARY KEY AUTOINCREMENT,
|
|
path TEXT NOT NULL UNIQUE
|
|
)
|
|
""")
|
|
conn.commit()
|
|
yield conn
|
|
conn.close()
|
|
|
|
|
|
@pytest.fixture
|
|
def db_connection(in_memory_db):
|
|
"""Create a DatabaseConnection wrapper for in-memory DB."""
|
|
# Create a mock-like wrapper that uses our in-memory connection
|
|
class TestDBConnection:
|
|
"""Test database connection wrapper for embeddings tests."""
|
|
def __init__(self, conn):
|
|
"""Initialize test DB connection.
|
|
|
|
Args:
|
|
conn: SQLite connection object.
|
|
"""
|
|
self._conn = conn
|
|
self.db_path = ":memory:"
|
|
|
|
def get_connection(self):
|
|
"""Get database connection.
|
|
|
|
Returns:
|
|
SQLite connection object.
|
|
"""
|
|
return self._conn
|
|
|
|
def execute(self, query, params=None):
|
|
"""Execute a query.
|
|
|
|
Args:
|
|
query: SQL query string.
|
|
params: Query parameters (optional).
|
|
|
|
Returns:
|
|
Cursor object.
|
|
"""
|
|
if params:
|
|
return self._conn.execute(query, params)
|
|
return self._conn.execute(query)
|
|
|
|
def executemany(self, query, params_list):
|
|
"""Execute multiple queries.
|
|
|
|
Args:
|
|
query: SQL query string.
|
|
params_list: List of parameter tuples.
|
|
|
|
Returns:
|
|
Cursor object.
|
|
"""
|
|
return self._conn.executemany(query, params_list)
|
|
|
|
def commit(self):
|
|
"""Commit transaction."""
|
|
self._conn.commit()
|
|
|
|
def rollback(self):
|
|
"""Rollback transaction."""
|
|
self._conn.rollback()
|
|
|
|
def close(self):
|
|
"""Close connection (no-op for test)."""
|
|
pass
|
|
|
|
return TestDBConnection(in_memory_db)
|
|
|
|
|
|
@pytest.fixture
|
|
def embedding_repo(db_connection):
|
|
"""Create an EmbeddingRepository instance."""
|
|
return EmbeddingRepository(db_connection)
|
|
|
|
|
|
@pytest.fixture
|
|
def sample_embedding():
|
|
"""Create a sample embedding vector."""
|
|
return [0.1, 0.2, 0.3, 0.4, 0.5]
|
|
|
|
|
|
@pytest.fixture
|
|
def sample_embeddings_batch():
|
|
"""Create a batch of sample embeddings."""
|
|
return [
|
|
{"file_id": 1, "embedding": [0.1, 0.2, 0.3], "model_version": "v1", "created_time": 1000},
|
|
{"file_id": 2, "embedding": [0.4, 0.5, 0.6], "model_version": "v1", "created_time": 1001},
|
|
{"file_id": 3, "embedding": [0.7, 0.8, 0.9], "model_version": "v2", "created_time": 1002},
|
|
]
|
|
|
|
|
|
# =============================================================================
|
|
# EmbeddingRepository Initialization Tests
|
|
# =============================================================================
|
|
|
|
class TestEmbeddingRepositoryInit:
|
|
"""Tests for EmbeddingRepository initialization."""
|
|
|
|
def test_init_with_db_connection(self, db_connection):
|
|
"""Test initialization with database connection."""
|
|
repo = EmbeddingRepository(db_connection)
|
|
assert repo.db is db_connection
|
|
|
|
def test_init_stores_connection(self, db_connection):
|
|
"""Test that connection is stored correctly."""
|
|
repo = EmbeddingRepository(db_connection)
|
|
assert hasattr(repo, "db")
|
|
|
|
|
|
# =============================================================================
|
|
# Add Embedding Tests
|
|
# =============================================================================
|
|
|
|
class TestAddEmbedding:
|
|
"""Tests for add_embedding method."""
|
|
|
|
def test_add_embedding_success(self, embedding_repo, sample_embedding):
|
|
"""Test adding an embedding successfully."""
|
|
embedding_id = embedding_repo.add_embedding(
|
|
file_id=1,
|
|
embedding=sample_embedding,
|
|
model_version="v1.0",
|
|
created_time=1000
|
|
)
|
|
assert embedding_id is not None
|
|
assert embedding_id > 0
|
|
|
|
def test_add_embedding_returns_integer(self, embedding_repo, sample_embedding):
|
|
"""Test that add_embedding returns an integer ID."""
|
|
embedding_id = embedding_repo.add_embedding(
|
|
file_id=1,
|
|
embedding=sample_embedding,
|
|
model_version="v1.0",
|
|
created_time=1000
|
|
)
|
|
assert isinstance(embedding_id, int)
|
|
|
|
def test_add_embedding_serializes_to_blob(self, db_connection, sample_embedding):
|
|
"""Test that embedding is serialized to blob."""
|
|
repo = EmbeddingRepository(db_connection)
|
|
embedding_id = repo.add_embedding(
|
|
file_id=1,
|
|
embedding=sample_embedding,
|
|
model_version="v1.0",
|
|
created_time=1000
|
|
)
|
|
|
|
cursor = db_connection.execute(
|
|
"SELECT embedding FROM embeddings WHERE id = ?",
|
|
(embedding_id,)
|
|
)
|
|
row = cursor.fetchone()
|
|
assert row is not None
|
|
# Verify it's serialized as JSON (secure format)
|
|
import json
|
|
loaded = json.loads(row[0].decode('utf-8'))
|
|
assert np.array_equal(loaded, sample_embedding)
|
|
|
|
def test_add_embedding_stores_metadata(self, embedding_repo, sample_embedding):
|
|
"""Test that embedding metadata is stored correctly."""
|
|
embedding_id = embedding_repo.add_embedding(
|
|
file_id=42,
|
|
embedding=sample_embedding,
|
|
model_version="test_model",
|
|
created_time=12345
|
|
)
|
|
|
|
result = embedding_repo.get_embedding(embedding_id)
|
|
assert result is not None
|
|
assert result["file_id"] == 42
|
|
assert result["model_version"] == "test_model"
|
|
assert result["created_time"] == 12345
|
|
|
|
def test_add_embedding_multiple(self, embedding_repo, sample_embedding):
|
|
"""Test adding multiple embeddings."""
|
|
ids = []
|
|
for i in range(5):
|
|
embedding_id = embedding_repo.add_embedding(
|
|
file_id=i,
|
|
embedding=sample_embedding,
|
|
model_version="v1.0",
|
|
created_time=1000 + i
|
|
)
|
|
ids.append(embedding_id)
|
|
|
|
assert len(ids) == 5
|
|
assert all(id > 0 for id in ids)
|
|
# IDs should be unique
|
|
assert len(set(ids)) == 5
|
|
|
|
def test_add_embedding_different_models(self, embedding_repo, sample_embedding):
|
|
"""Test adding embeddings with different model versions."""
|
|
embedding_repo.add_embedding(
|
|
file_id=1,
|
|
embedding=sample_embedding,
|
|
model_version="model_a",
|
|
created_time=1000
|
|
)
|
|
embedding_repo.add_embedding(
|
|
file_id=1,
|
|
embedding=sample_embedding,
|
|
model_version="model_b",
|
|
created_time=1001
|
|
)
|
|
|
|
embeddings = embedding_repo.get_embeddings_by_file(1)
|
|
assert len(embeddings) == 2
|
|
models = {e["model_version"] for e in embeddings}
|
|
assert models == {"model_a", "model_b"}
|
|
|
|
|
|
# =============================================================================
|
|
# Get Embedding Tests
|
|
# =============================================================================
|
|
|
|
class TestGetEmbedding:
|
|
"""Tests for get_embedding method."""
|
|
|
|
def test_get_embedding_by_id(self, embedding_repo, sample_embedding):
|
|
"""Test getting embedding by ID."""
|
|
embedding_id = embedding_repo.add_embedding(
|
|
file_id=1,
|
|
embedding=sample_embedding,
|
|
model_version="v1.0",
|
|
created_time=1000
|
|
)
|
|
|
|
result = embedding_repo.get_embedding(embedding_id)
|
|
assert result is not None
|
|
assert result["id"] == embedding_id
|
|
assert np.array_equal(result["embedding"], sample_embedding)
|
|
|
|
def test_get_embedding_returns_dict(self, embedding_repo, sample_embedding):
|
|
"""Test that get_embedding returns a dictionary."""
|
|
embedding_id = embedding_repo.add_embedding(
|
|
file_id=1,
|
|
embedding=sample_embedding,
|
|
model_version="v1.0",
|
|
created_time=1000
|
|
)
|
|
|
|
result = embedding_repo.get_embedding(embedding_id)
|
|
assert isinstance(result, dict)
|
|
|
|
def test_get_embedding_has_all_fields(self, embedding_repo, sample_embedding):
|
|
"""Test that returned embedding has all expected fields."""
|
|
embedding_id = embedding_repo.add_embedding(
|
|
file_id=1,
|
|
embedding=sample_embedding,
|
|
model_version="v1.0",
|
|
created_time=1000
|
|
)
|
|
|
|
result = embedding_repo.get_embedding(embedding_id)
|
|
assert "id" in result
|
|
assert "file_id" in result
|
|
assert "embedding" in result
|
|
assert "model_version" in result
|
|
assert "created_time" in result
|
|
|
|
def test_get_embedding_not_found(self, embedding_repo):
|
|
"""Test getting non-existent embedding returns None."""
|
|
result = embedding_repo.get_embedding(99999)
|
|
assert result is None
|
|
|
|
def test_get_embedding_deserializes_blob(self, embedding_repo, sample_embedding):
|
|
"""Test that embedding blob is properly deserialized."""
|
|
embedding_id = embedding_repo.add_embedding(
|
|
file_id=1,
|
|
embedding=sample_embedding,
|
|
model_version="v1.0",
|
|
created_time=1000
|
|
)
|
|
|
|
result = embedding_repo.get_embedding(embedding_id)
|
|
assert isinstance(result["embedding"], (list, np.ndarray))
|
|
# dtype check removed - sample_embedding is now a list
|
|
|
|
|
|
# =============================================================================
|
|
# Get Embedding By File Tests
|
|
# =============================================================================
|
|
|
|
class TestGetEmbeddingByFile:
|
|
"""Tests for get_embedding_by_file method."""
|
|
|
|
def test_get_embedding_by_file_and_model(self, embedding_repo, sample_embedding):
|
|
"""Test getting embedding by file ID and model version."""
|
|
embedding_repo.add_embedding(
|
|
file_id=1,
|
|
embedding=sample_embedding,
|
|
model_version="v1.0",
|
|
created_time=1000
|
|
)
|
|
|
|
result = embedding_repo.get_embedding_by_file(1, "v1.0")
|
|
assert result is not None
|
|
assert result["file_id"] == 1
|
|
assert result["model_version"] == "v1.0"
|
|
|
|
def test_get_embedding_by_file_not_found(self, embedding_repo):
|
|
"""Test getting non-existent embedding by file returns None."""
|
|
result = embedding_repo.get_embedding_by_file(999, "v1.0")
|
|
assert result is None
|
|
|
|
def test_get_embedding_by_file_wrong_model(self, embedding_repo, sample_embedding):
|
|
"""Test that wrong model version returns None."""
|
|
embedding_repo.add_embedding(
|
|
file_id=1,
|
|
embedding=sample_embedding,
|
|
model_version="v1.0",
|
|
created_time=1000
|
|
)
|
|
|
|
result = embedding_repo.get_embedding_by_file(1, "v2.0")
|
|
assert result is None
|
|
|
|
def test_get_embedding_by_file_returns_first_match(self, embedding_repo, sample_embedding):
|
|
"""Test that get_embedding_by_file returns first match."""
|
|
# Add two embeddings for same file with same model (shouldn't happen but test behavior)
|
|
embedding_repo.add_embedding(
|
|
file_id=1,
|
|
embedding=sample_embedding,
|
|
model_version="v1.0",
|
|
created_time=1000
|
|
)
|
|
embedding_repo.add_embedding(
|
|
file_id=1,
|
|
embedding=sample_embedding * 2,
|
|
model_version="v1.0",
|
|
created_time=1001
|
|
)
|
|
|
|
result = embedding_repo.get_embedding_by_file(1, "v1.0")
|
|
assert result is not None
|
|
assert result["file_id"] == 1
|
|
|
|
|
|
# =============================================================================
|
|
# Get Embeddings By File Tests
|
|
# =============================================================================
|
|
|
|
class TestGetEmbeddingsByFile:
|
|
"""Tests for get_embeddings_by_file method."""
|
|
|
|
def test_get_embeddings_by_file(self, embedding_repo, sample_embedding):
|
|
"""Test getting all embeddings for a file."""
|
|
embedding_repo.add_embedding(
|
|
file_id=1,
|
|
embedding=sample_embedding,
|
|
model_version="v1.0",
|
|
created_time=1000
|
|
)
|
|
embedding_repo.add_embedding(
|
|
file_id=1,
|
|
embedding=sample_embedding * 2,
|
|
model_version="v2.0",
|
|
created_time=1001
|
|
)
|
|
|
|
results = embedding_repo.get_embeddings_by_file(1)
|
|
assert len(results) == 2
|
|
|
|
def test_get_embeddings_by_file_empty(self, embedding_repo):
|
|
"""Test getting embeddings for non-existent file."""
|
|
results = embedding_repo.get_embeddings_by_file(999)
|
|
assert results == []
|
|
|
|
def test_get_embeddings_by_file_returns_list(self, embedding_repo, sample_embedding):
|
|
"""Test that get_embeddings_by_file returns a list."""
|
|
embedding_repo.add_embedding(
|
|
file_id=1,
|
|
embedding=sample_embedding,
|
|
model_version="v1.0",
|
|
created_time=1000
|
|
)
|
|
|
|
results = embedding_repo.get_embeddings_by_file(1)
|
|
assert isinstance(results, list)
|
|
|
|
def test_get_embeddings_by_file_ordered_by_model(self, embedding_repo, sample_embedding):
|
|
"""Test that embeddings are ordered by model version."""
|
|
embedding_repo.add_embedding(
|
|
file_id=1,
|
|
embedding=sample_embedding,
|
|
model_version="v2.0",
|
|
created_time=1000
|
|
)
|
|
embedding_repo.add_embedding(
|
|
file_id=1,
|
|
embedding=sample_embedding * 2,
|
|
model_version="v1.0",
|
|
created_time=1001
|
|
)
|
|
|
|
results = embedding_repo.get_embeddings_by_file(1)
|
|
assert len(results) == 2
|
|
# Should be ordered by model_version
|
|
assert results[0]["model_version"] == "v1.0"
|
|
assert results[1]["model_version"] == "v2.0"
|
|
|
|
|
|
# =============================================================================
|
|
# Get Embeddings By Model Tests
|
|
# =============================================================================
|
|
|
|
class TestGetEmbeddingsByModel:
|
|
"""Tests for get_embeddings_by_model method."""
|
|
|
|
def test_get_embeddings_by_model(self, embedding_repo, sample_embedding):
|
|
"""Test getting all embeddings for a model version."""
|
|
embedding_repo.add_embedding(
|
|
file_id=1,
|
|
embedding=sample_embedding,
|
|
model_version="v1.0",
|
|
created_time=1000
|
|
)
|
|
embedding_repo.add_embedding(
|
|
file_id=2,
|
|
embedding=sample_embedding * 2,
|
|
model_version="v1.0",
|
|
created_time=1001
|
|
)
|
|
embedding_repo.add_embedding(
|
|
file_id=3,
|
|
embedding=sample_embedding * 3,
|
|
model_version="v2.0",
|
|
created_time=1002
|
|
)
|
|
|
|
results = embedding_repo.get_embeddings_by_model("v1.0")
|
|
assert len(results) == 2
|
|
|
|
def test_get_embeddings_by_model_empty(self, embedding_repo):
|
|
"""Test getting embeddings for non-existent model."""
|
|
results = embedding_repo.get_embeddings_by_model("nonexistent")
|
|
assert results == []
|
|
|
|
def test_get_embeddings_by_model_returns_list(self, embedding_repo, sample_embedding):
|
|
"""Test that get_embeddings_by_model returns a list."""
|
|
embedding_repo.add_embedding(
|
|
file_id=1,
|
|
embedding=sample_embedding,
|
|
model_version="v1.0",
|
|
created_time=1000
|
|
)
|
|
|
|
results = embedding_repo.get_embeddings_by_model("v1.0")
|
|
assert isinstance(results, list)
|
|
|
|
def test_get_embeddings_by_model_ordered_by_file_id(self, embedding_repo, sample_embedding):
|
|
"""Test that embeddings are ordered by file_id."""
|
|
embedding_repo.add_embedding(
|
|
file_id=3,
|
|
embedding=sample_embedding,
|
|
model_version="v1.0",
|
|
created_time=1000
|
|
)
|
|
embedding_repo.add_embedding(
|
|
file_id=1,
|
|
embedding=sample_embedding * 2,
|
|
model_version="v1.0",
|
|
created_time=1001
|
|
)
|
|
embedding_repo.add_embedding(
|
|
file_id=2,
|
|
embedding=sample_embedding * 3,
|
|
model_version="v1.0",
|
|
created_time=1002
|
|
)
|
|
|
|
results = embedding_repo.get_embeddings_by_model("v1.0")
|
|
file_ids = [r["file_id"] for r in results]
|
|
assert file_ids == sorted(file_ids)
|
|
|
|
|
|
# =============================================================================
|
|
# Update Embedding Tests
|
|
# =============================================================================
|
|
|
|
class TestUpdateEmbedding:
|
|
"""Tests for update_embedding method."""
|
|
|
|
def test_update_embedding_success(self, embedding_repo, sample_embedding):
|
|
"""Test updating an embedding."""
|
|
embedding_id = embedding_repo.add_embedding(
|
|
file_id=1,
|
|
embedding=sample_embedding,
|
|
model_version="v1.0",
|
|
created_time=1000
|
|
)
|
|
|
|
new_embedding = sample_embedding * 2
|
|
result = embedding_repo.update_embedding(embedding_id, new_embedding)
|
|
assert result is True
|
|
|
|
def test_update_embedding_persists(self, embedding_repo, sample_embedding):
|
|
"""Test that updated embedding persists."""
|
|
embedding_id = embedding_repo.add_embedding(
|
|
file_id=1,
|
|
embedding=sample_embedding,
|
|
model_version="v1.0",
|
|
created_time=1000
|
|
)
|
|
|
|
new_embedding = sample_embedding * 2
|
|
embedding_repo.update_embedding(embedding_id, new_embedding)
|
|
|
|
result = embedding_repo.get_embedding(embedding_id)
|
|
assert np.array_equal(result["embedding"], new_embedding)
|
|
|
|
def test_update_embedding_not_found(self, embedding_repo, sample_embedding):
|
|
"""Test updating non-existent embedding returns False."""
|
|
result = embedding_repo.update_embedding(99999, sample_embedding)
|
|
assert result is False
|
|
|
|
def test_update_embedding_preserves_metadata(self, embedding_repo, sample_embedding):
|
|
"""Test that update preserves metadata."""
|
|
embedding_id = embedding_repo.add_embedding(
|
|
file_id=42,
|
|
embedding=sample_embedding,
|
|
model_version="original_model",
|
|
created_time=12345
|
|
)
|
|
|
|
new_embedding = sample_embedding * 2
|
|
embedding_repo.update_embedding(embedding_id, new_embedding)
|
|
|
|
result = embedding_repo.get_embedding(embedding_id)
|
|
assert result["file_id"] == 42
|
|
assert result["model_version"] == "original_model"
|
|
assert result["created_time"] == 12345
|
|
|
|
|
|
# =============================================================================
|
|
# Delete Embedding Tests
|
|
# =============================================================================
|
|
|
|
class TestDeleteEmbedding:
|
|
"""Tests for delete_embedding method."""
|
|
|
|
def test_delete_embedding_success(self, embedding_repo, sample_embedding):
|
|
"""Test deleting an embedding."""
|
|
embedding_id = embedding_repo.add_embedding(
|
|
file_id=1,
|
|
embedding=sample_embedding,
|
|
model_version="v1.0",
|
|
created_time=1000
|
|
)
|
|
|
|
result = embedding_repo.delete_embedding(embedding_id)
|
|
assert result is True
|
|
|
|
def test_delete_embedding_removes_record(self, embedding_repo, sample_embedding):
|
|
"""Test that deleted embedding is removed."""
|
|
embedding_id = embedding_repo.add_embedding(
|
|
file_id=1,
|
|
embedding=sample_embedding,
|
|
model_version="v1.0",
|
|
created_time=1000
|
|
)
|
|
|
|
embedding_repo.delete_embedding(embedding_id)
|
|
result = embedding_repo.get_embedding(embedding_id)
|
|
assert result is None
|
|
|
|
def test_delete_embedding_not_found(self, embedding_repo):
|
|
"""Test deleting non-existent embedding returns False."""
|
|
result = embedding_repo.delete_embedding(99999)
|
|
assert result is False
|
|
|
|
def test_delete_embedding_idempotent(self, embedding_repo, sample_embedding):
|
|
"""Test that deleting same embedding twice returns False second time."""
|
|
embedding_id = embedding_repo.add_embedding(
|
|
file_id=1,
|
|
embedding=sample_embedding,
|
|
model_version="v1.0",
|
|
created_time=1000
|
|
)
|
|
|
|
embedding_repo.delete_embedding(embedding_id)
|
|
result = embedding_repo.delete_embedding(embedding_id)
|
|
assert result is False
|
|
|
|
|
|
# =============================================================================
|
|
# Delete Embeddings By File Tests
|
|
# =============================================================================
|
|
|
|
class TestDeleteEmbeddingsByFile:
|
|
"""Tests for delete_embeddings_by_file method."""
|
|
|
|
def test_delete_embeddings_by_file(self, embedding_repo, sample_embedding):
|
|
"""Test deleting all embeddings for a file."""
|
|
embedding_repo.add_embedding(
|
|
file_id=1,
|
|
embedding=sample_embedding,
|
|
model_version="v1.0",
|
|
created_time=1000
|
|
)
|
|
embedding_repo.add_embedding(
|
|
file_id=1,
|
|
embedding=sample_embedding * 2,
|
|
model_version="v2.0",
|
|
created_time=1001
|
|
)
|
|
embedding_repo.add_embedding(
|
|
file_id=2,
|
|
embedding=sample_embedding * 3,
|
|
model_version="v1.0",
|
|
created_time=1002
|
|
)
|
|
|
|
count = embedding_repo.delete_embeddings_by_file(1)
|
|
assert count == 2
|
|
|
|
def test_delete_embeddings_by_file_removes_records(self, embedding_repo, sample_embedding):
|
|
"""Test that deleted embeddings are removed."""
|
|
embedding_repo.add_embedding(
|
|
file_id=1,
|
|
embedding=sample_embedding,
|
|
model_version="v1.0",
|
|
created_time=1000
|
|
)
|
|
embedding_repo.add_embedding(
|
|
file_id=1,
|
|
embedding=sample_embedding * 2,
|
|
model_version="v2.0",
|
|
created_time=1001
|
|
)
|
|
|
|
embedding_repo.delete_embeddings_by_file(1)
|
|
results = embedding_repo.get_embeddings_by_file(1)
|
|
assert len(results) == 0
|
|
|
|
def test_delete_embeddings_by_file_no_matches(self, embedding_repo):
|
|
"""Test deleting embeddings for non-existent file returns 0."""
|
|
count = embedding_repo.delete_embeddings_by_file(999)
|
|
assert count == 0
|
|
|
|
def test_delete_embeddings_by_file_returns_int(self, embedding_repo, sample_embedding):
|
|
"""Test that delete_embeddings_by_file returns an integer."""
|
|
embedding_repo.add_embedding(
|
|
file_id=1,
|
|
embedding=sample_embedding,
|
|
model_version="v1.0",
|
|
created_time=1000
|
|
)
|
|
|
|
count = embedding_repo.delete_embeddings_by_file(1)
|
|
assert isinstance(count, int)
|
|
|
|
|
|
# =============================================================================
|
|
# Delete Embeddings By Model Tests
|
|
# =============================================================================
|
|
|
|
class TestDeleteEmbeddingsByModel:
|
|
"""Tests for delete_embeddings_by_model method."""
|
|
|
|
def test_delete_embeddings_by_model(self, embedding_repo, sample_embedding):
|
|
"""Test deleting all embeddings for a model version."""
|
|
embedding_repo.add_embedding(
|
|
file_id=1,
|
|
embedding=sample_embedding,
|
|
model_version="v1.0",
|
|
created_time=1000
|
|
)
|
|
embedding_repo.add_embedding(
|
|
file_id=2,
|
|
embedding=sample_embedding * 2,
|
|
model_version="v1.0",
|
|
created_time=1001
|
|
)
|
|
embedding_repo.add_embedding(
|
|
file_id=3,
|
|
embedding=sample_embedding * 3,
|
|
model_version="v2.0",
|
|
created_time=1002
|
|
)
|
|
|
|
count = embedding_repo.delete_embeddings_by_model("v1.0")
|
|
assert count == 2
|
|
|
|
def test_delete_embeddings_by_model_removes_records(self, embedding_repo, sample_embedding):
|
|
"""Test that deleted embeddings are removed."""
|
|
embedding_repo.add_embedding(
|
|
file_id=1,
|
|
embedding=sample_embedding,
|
|
model_version="v1.0",
|
|
created_time=1000
|
|
)
|
|
embedding_repo.add_embedding(
|
|
file_id=2,
|
|
embedding=sample_embedding * 2,
|
|
model_version="v1.0",
|
|
created_time=1001
|
|
)
|
|
|
|
embedding_repo.delete_embeddings_by_model("v1.0")
|
|
results = embedding_repo.get_embeddings_by_model("v1.0")
|
|
assert len(results) == 0
|
|
|
|
def test_delete_embeddings_by_model_no_matches(self, embedding_repo):
|
|
"""Test deleting embeddings for non-existent model returns 0."""
|
|
count = embedding_repo.delete_embeddings_by_model("nonexistent")
|
|
assert count == 0
|
|
|
|
def test_delete_embeddings_by_model_returns_int(self, embedding_repo, sample_embedding):
|
|
"""Test that delete_embeddings_by_model returns an integer."""
|
|
embedding_repo.add_embedding(
|
|
file_id=1,
|
|
embedding=sample_embedding,
|
|
model_version="v1.0",
|
|
created_time=1000
|
|
)
|
|
|
|
count = embedding_repo.delete_embeddings_by_model("v1.0")
|
|
assert isinstance(count, int)
|
|
|
|
|
|
# =============================================================================
|
|
# Get All Embeddings Tests
|
|
# =============================================================================
|
|
|
|
class TestGetAllEmbeddings:
|
|
"""Tests for get_all_embeddings method."""
|
|
|
|
def test_get_all_embeddings(self, embedding_repo, sample_embedding):
|
|
"""Test getting all embeddings."""
|
|
embedding_repo.add_embedding(
|
|
file_id=1,
|
|
embedding=sample_embedding,
|
|
model_version="v1.0",
|
|
created_time=1000
|
|
)
|
|
embedding_repo.add_embedding(
|
|
file_id=2,
|
|
embedding=sample_embedding * 2,
|
|
model_version="v1.0",
|
|
created_time=1001
|
|
)
|
|
|
|
results = embedding_repo.get_all_embeddings()
|
|
assert len(results) == 2
|
|
|
|
def test_get_all_embeddings_empty(self, embedding_repo):
|
|
"""Test getting all embeddings when none exist."""
|
|
results = embedding_repo.get_all_embeddings()
|
|
assert results == []
|
|
|
|
def test_get_all_embeddings_returns_list(self, embedding_repo, sample_embedding):
|
|
"""Test that get_all_embeddings returns a list."""
|
|
embedding_repo.add_embedding(
|
|
file_id=1,
|
|
embedding=sample_embedding,
|
|
model_version="v1.0",
|
|
created_time=1000
|
|
)
|
|
|
|
results = embedding_repo.get_all_embeddings()
|
|
assert isinstance(results, list)
|
|
|
|
def test_get_all_embeddings_ordered(self, embedding_repo, sample_embedding):
|
|
"""Test that embeddings are ordered by file_id, model_version."""
|
|
embedding_repo.add_embedding(
|
|
file_id=2,
|
|
embedding=sample_embedding,
|
|
model_version="v2.0",
|
|
created_time=1000
|
|
)
|
|
embedding_repo.add_embedding(
|
|
file_id=1,
|
|
embedding=sample_embedding * 2,
|
|
model_version="v1.0",
|
|
created_time=1001
|
|
)
|
|
embedding_repo.add_embedding(
|
|
file_id=1,
|
|
embedding=sample_embedding * 3,
|
|
model_version="v2.0",
|
|
created_time=1002
|
|
)
|
|
|
|
results = embedding_repo.get_all_embeddings()
|
|
file_ids = [r["file_id"] for r in results]
|
|
assert file_ids == sorted(file_ids)
|
|
|
|
|
|
# =============================================================================
|
|
# Count Embeddings Tests
|
|
# =============================================================================
|
|
|
|
class TestCountEmbeddings:
|
|
"""Tests for count_embeddings methods."""
|
|
|
|
def test_count_embeddings(self, embedding_repo, sample_embedding):
|
|
"""Test counting all embeddings."""
|
|
embedding_repo.add_embedding(
|
|
file_id=1,
|
|
embedding=sample_embedding,
|
|
model_version="v1.0",
|
|
created_time=1000
|
|
)
|
|
embedding_repo.add_embedding(
|
|
file_id=2,
|
|
embedding=sample_embedding * 2,
|
|
model_version="v1.0",
|
|
created_time=1001
|
|
)
|
|
embedding_repo.add_embedding(
|
|
file_id=3,
|
|
embedding=sample_embedding * 3,
|
|
model_version="v2.0",
|
|
created_time=1002
|
|
)
|
|
|
|
count = embedding_repo.count_embeddings()
|
|
assert count == 3
|
|
|
|
def test_count_embeddings_empty(self, embedding_repo):
|
|
"""Test counting when no embeddings exist."""
|
|
count = embedding_repo.count_embeddings()
|
|
assert count == 0
|
|
|
|
def test_count_embeddings_returns_int(self, embedding_repo, sample_embedding):
|
|
"""Test that count_embeddings returns an integer."""
|
|
embedding_repo.add_embedding(
|
|
file_id=1,
|
|
embedding=sample_embedding,
|
|
model_version="v1.0",
|
|
created_time=1000
|
|
)
|
|
|
|
count = embedding_repo.count_embeddings()
|
|
assert isinstance(count, int)
|
|
|
|
def test_count_embeddings_by_model(self, embedding_repo, sample_embedding):
|
|
"""Test counting embeddings by model version."""
|
|
embedding_repo.add_embedding(
|
|
file_id=1,
|
|
embedding=sample_embedding,
|
|
model_version="v1.0",
|
|
created_time=1000
|
|
)
|
|
embedding_repo.add_embedding(
|
|
file_id=2,
|
|
embedding=sample_embedding * 2,
|
|
model_version="v1.0",
|
|
created_time=1001
|
|
)
|
|
embedding_repo.add_embedding(
|
|
file_id=3,
|
|
embedding=sample_embedding * 3,
|
|
model_version="v2.0",
|
|
created_time=1002
|
|
)
|
|
|
|
count = embedding_repo.count_embeddings_by_model("v1.0")
|
|
assert count == 2
|
|
|
|
def test_count_embeddings_by_model_empty(self, embedding_repo):
|
|
"""Test counting by model when no matches."""
|
|
count = embedding_repo.count_embeddings_by_model("nonexistent")
|
|
assert count == 0
|
|
|
|
def test_count_embeddings_by_model_returns_int(self, embedding_repo, sample_embedding):
|
|
"""Test that count_embeddings_by_model returns an integer."""
|
|
embedding_repo.add_embedding(
|
|
file_id=1,
|
|
embedding=sample_embedding,
|
|
model_version="v1.0",
|
|
created_time=1000
|
|
)
|
|
|
|
count = embedding_repo.count_embeddings_by_model("v1.0")
|
|
assert isinstance(count, int)
|
|
|
|
|
|
# =============================================================================
|
|
# Batch Add Embeddings Tests
|
|
# =============================================================================
|
|
|
|
class TestBatchAddEmbeddings:
|
|
"""Tests for batch_add_embeddings method."""
|
|
|
|
def test_batch_add_embeddings(self, embedding_repo, sample_embeddings_batch):
|
|
"""Test batch adding embeddings."""
|
|
count = embedding_repo.batch_add_embeddings(sample_embeddings_batch)
|
|
assert count == 3
|
|
|
|
def test_batch_add_embeddings_persists(self, embedding_repo, sample_embeddings_batch):
|
|
"""Test that batch added embeddings persist."""
|
|
embedding_repo.batch_add_embeddings(sample_embeddings_batch)
|
|
|
|
all_embeddings = embedding_repo.get_all_embeddings()
|
|
assert len(all_embeddings) == 3
|
|
|
|
def test_batch_add_embeddings_empty_list(self, embedding_repo):
|
|
"""Test batch adding empty list returns 0."""
|
|
count = embedding_repo.batch_add_embeddings([])
|
|
assert count == 0
|
|
|
|
def test_batch_add_embeddings_returns_int(self, embedding_repo, sample_embeddings_batch):
|
|
"""Test that batch_add_embeddings returns an integer."""
|
|
count = embedding_repo.batch_add_embeddings(sample_embeddings_batch)
|
|
assert isinstance(count, int)
|
|
|
|
def test_batch_add_embeddings_various_models(self, embedding_repo, sample_embeddings_batch):
|
|
"""Test batch adding embeddings with different models."""
|
|
embedding_repo.batch_add_embeddings(sample_embeddings_batch)
|
|
|
|
v1_count = embedding_repo.count_embeddings_by_model("v1")
|
|
v2_count = embedding_repo.count_embeddings_by_model("v2")
|
|
assert v1_count == 2
|
|
assert v2_count == 1
|
|
|
|
def test_batch_add_embeddings_numpy_arrays(self, embedding_repo):
|
|
"""Test batch adding embeddings with numpy arrays."""
|
|
embeddings = [
|
|
{"file_id": 1, "embedding": np.array([0.1, 0.2, 0.3]), "model_version": "v1", "created_time": 1000},
|
|
{"file_id": 2, "embedding": np.array([0.4, 0.5, 0.6]), "model_version": "v1", "created_time": 1001},
|
|
]
|
|
count = embedding_repo.batch_add_embeddings(embeddings)
|
|
assert count == 2
|
|
|
|
|
|
# =============================================================================
|
|
# Clear All Embeddings Tests
|
|
# =============================================================================
|
|
|
|
class TestClearAllEmbeddings:
|
|
"""Tests for clear_all_embeddings method."""
|
|
|
|
def test_clear_all_embeddings(self, embedding_repo, sample_embedding):
|
|
"""Test clearing all embeddings."""
|
|
embedding_repo.add_embedding(
|
|
file_id=1,
|
|
embedding=sample_embedding,
|
|
model_version="v1.0",
|
|
created_time=1000
|
|
)
|
|
embedding_repo.add_embedding(
|
|
file_id=2,
|
|
embedding=sample_embedding * 2,
|
|
model_version="v1.0",
|
|
created_time=1001
|
|
)
|
|
|
|
embedding_repo.clear_all_embeddings()
|
|
count = embedding_repo.count_embeddings()
|
|
assert count == 0
|
|
|
|
def test_clear_all_embeddings_empty(self, embedding_repo):
|
|
"""Test clearing when no embeddings exist."""
|
|
# Should not raise
|
|
embedding_repo.clear_all_embeddings()
|
|
count = embedding_repo.count_embeddings()
|
|
assert count == 0
|
|
|
|
def test_clear_all_embeddings_commits(self, db_connection, sample_embedding):
|
|
"""Test that clear_all_embeddings commits the transaction."""
|
|
repo = EmbeddingRepository(db_connection)
|
|
repo.add_embedding(
|
|
file_id=1,
|
|
embedding=sample_embedding,
|
|
model_version="v1.0",
|
|
created_time=1000
|
|
)
|
|
|
|
repo.clear_all_embeddings()
|
|
|
|
# Verify committed by checking count
|
|
count = repo.count_embeddings()
|
|
assert count == 0
|
|
|
|
|
|
# =============================================================================
|
|
# Get Embedding Repository Function Tests
|
|
# =============================================================================
|
|
|
|
class TestGetEmbeddingRepository:
|
|
"""Tests for get_embedding_repository function."""
|
|
|
|
def test_get_embedding_repository(self, tmp_path):
|
|
"""Test get_embedding_repository function."""
|
|
db_path = str(tmp_path / "test.db")
|
|
repo = get_embedding_repository(db_path)
|
|
assert isinstance(repo, EmbeddingRepository)
|
|
|
|
def test_get_embedding_repository_creates_connection(self, tmp_path):
|
|
"""Test that get_embedding_repository creates a connection."""
|
|
db_path = str(tmp_path / "test.db")
|
|
repo = get_embedding_repository(db_path)
|
|
assert repo.db is not None
|
|
|
|
def test_get_embedding_repository_default_path(self, monkeypatch, tmp_path):
|
|
"""Test get_embedding_repository with default path."""
|
|
monkeypatch.chdir(tmp_path)
|
|
# Create output directory
|
|
os.makedirs(tmp_path / "output", exist_ok=True)
|
|
repo = get_embedding_repository(str(tmp_path / "output" / "index.db"))
|
|
assert isinstance(repo, EmbeddingRepository)
|
|
|
|
|
|
# =============================================================================
|
|
# Error Handling Tests
|
|
# =============================================================================
|
|
|
|
class TestErrorHandling:
|
|
"""Tests for error handling in embedding operations."""
|
|
|
|
def test_add_embedding_with_invalid_file_id(self, embedding_repo, sample_embedding):
|
|
"""Test adding embedding with invalid file_id still works (no FK constraint in test)."""
|
|
# Even with non-existent file_id, should work if no FK constraint
|
|
embedding_id = embedding_repo.add_embedding(
|
|
file_id=99999,
|
|
embedding=sample_embedding,
|
|
model_version="v1.0",
|
|
created_time=1000
|
|
)
|
|
assert embedding_id is not None
|
|
|
|
def test_get_embedding_with_invalid_id(self, embedding_repo):
|
|
"""Test getting embedding with invalid ID returns None."""
|
|
result = embedding_repo.get_embedding(-1)
|
|
assert result is None
|
|
|
|
def test_update_embedding_with_invalid_id(self, embedding_repo, sample_embedding):
|
|
"""Test updating embedding with invalid ID returns False."""
|
|
result = embedding_repo.update_embedding(-1, sample_embedding)
|
|
assert result is False
|
|
|
|
def test_delete_embedding_with_invalid_id(self, embedding_repo):
|
|
"""Test deleting embedding with invalid ID returns False."""
|
|
result = embedding_repo.delete_embedding(-1)
|
|
assert result is False
|
|
|
|
|
|
# =============================================================================
|
|
# Large Dataset Tests
|
|
# =============================================================================
|
|
|
|
class TestLargeDataset:
|
|
"""Tests for large embedding datasets."""
|
|
|
|
def test_batch_add_large_dataset(self, embedding_repo):
|
|
"""Test batch adding large number of embeddings."""
|
|
embeddings = [
|
|
{"file_id": i, "embedding": [float(i), float(i+1), float(i+2)], "model_version": "v1", "created_time": 1000 + i}
|
|
for i in range(100)
|
|
]
|
|
count = embedding_repo.batch_add_embeddings(embeddings)
|
|
assert count == 100
|
|
|
|
def test_count_large_dataset(self, embedding_repo):
|
|
"""Test counting large number of embeddings."""
|
|
embeddings = [
|
|
{"file_id": i, "embedding": [float(i), float(i+1), float(i+2)], "model_version": "v1", "created_time": 1000 + i}
|
|
for i in range(100)
|
|
]
|
|
embedding_repo.batch_add_embeddings(embeddings)
|
|
|
|
count = embedding_repo.count_embeddings()
|
|
assert count == 100
|
|
|
|
def test_get_all_large_dataset(self, embedding_repo):
|
|
"""Test getting all embeddings from large dataset."""
|
|
embeddings = [
|
|
{"file_id": i, "embedding": [float(i), float(i+1), float(i+2)], "model_version": "v1", "created_time": 1000 + i}
|
|
for i in range(100)
|
|
]
|
|
embedding_repo.batch_add_embeddings(embeddings)
|
|
|
|
all_embeddings = embedding_repo.get_all_embeddings()
|
|
assert len(all_embeddings) == 100
|
|
|
|
def test_delete_by_file_large_dataset(self, embedding_repo):
|
|
"""Test deleting by file_id in large dataset."""
|
|
embeddings = [
|
|
{"file_id": i % 10, "embedding": [float(i), float(i+1), float(i+2)], "model_version": "v1", "created_time": 1000 + i}
|
|
for i in range(100)
|
|
]
|
|
embedding_repo.batch_add_embeddings(embeddings)
|
|
|
|
# Each file_id (0-9) should have 10 embeddings
|
|
count = embedding_repo.delete_embeddings_by_file(5)
|
|
assert count == 10
|
|
|
|
def test_delete_by_model_large_dataset(self, embedding_repo):
|
|
"""Test deleting by model_version in large dataset."""
|
|
embeddings = [
|
|
{"file_id": i, "embedding": [float(i), float(i+1), float(i+2)], "model_version": f"v{i % 5}", "created_time": 1000 + i}
|
|
for i in range(100)
|
|
]
|
|
embedding_repo.batch_add_embeddings(embeddings)
|
|
|
|
# Each model (v0-v4) should have 20 embeddings
|
|
count = embedding_repo.delete_embeddings_by_model("v2")
|
|
assert count == 20
|
|
|
|
|
|
# =============================================================================
|
|
# Vector Operations Tests
|
|
# =============================================================================
|
|
|
|
class TestVectorOperations:
|
|
"""Tests for vector operations with embeddings."""
|
|
|
|
def test_numpy_array_embedding(self, embedding_repo):
|
|
"""Test storing numpy array embeddings."""
|
|
embedding = np.array([0.1, 0.2, 0.3, 0.4, 0.5], dtype=np.float32)
|
|
embedding_id = embedding_repo.add_embedding(
|
|
file_id=1,
|
|
embedding=embedding,
|
|
model_version="v1.0",
|
|
created_time=1000
|
|
)
|
|
|
|
result = embedding_repo.get_embedding(embedding_id)
|
|
assert np.array_equal(result["embedding"], embedding)
|
|
|
|
def test_list_embedding(self, embedding_repo):
|
|
"""Test storing list embeddings."""
|
|
embedding = [0.1, 0.2, 0.3, 0.4, 0.5]
|
|
embedding_id = embedding_repo.add_embedding(
|
|
file_id=1,
|
|
embedding=embedding,
|
|
model_version="v1.0",
|
|
created_time=1000
|
|
)
|
|
|
|
result = embedding_repo.get_embedding(embedding_id)
|
|
assert result["embedding"] == embedding
|
|
|
|
def test_high_dimensional_embedding(self, embedding_repo):
|
|
"""Test storing high-dimensional embeddings."""
|
|
embedding = np.random.rand(512).astype(np.float32)
|
|
embedding_id = embedding_repo.add_embedding(
|
|
file_id=1,
|
|
embedding=embedding,
|
|
model_version="v1.0",
|
|
created_time=1000
|
|
)
|
|
|
|
result = embedding_repo.get_embedding(embedding_id)
|
|
assert len(result["embedding"]) == 512
|
|
|
|
def test_embedding_precision(self, embedding_repo):
|
|
"""Test that embedding precision is preserved."""
|
|
embedding = np.array([0.123456789, 0.987654321, 0.111111111], dtype=np.float64)
|
|
embedding_id = embedding_repo.add_embedding(
|
|
file_id=1,
|
|
embedding=embedding,
|
|
model_version="v1.0",
|
|
created_time=1000
|
|
)
|
|
|
|
result = embedding_repo.get_embedding(embedding_id)
|
|
# Check values are close (JSON preserves float64 precision)
|
|
assert np.allclose(result["embedding"], embedding)
|
|
|
|
def test_zero_vector_embedding(self, embedding_repo):
|
|
"""Test storing zero vector embedding."""
|
|
embedding = np.zeros(10, dtype=np.float32)
|
|
embedding_id = embedding_repo.add_embedding(
|
|
file_id=1,
|
|
embedding=embedding,
|
|
model_version="v1.0",
|
|
created_time=1000
|
|
)
|
|
|
|
result = embedding_repo.get_embedding(embedding_id)
|
|
assert np.array_equal(result["embedding"], embedding)
|
|
|
|
def test_negative_values_embedding(self, embedding_repo):
|
|
"""Test storing embeddings with negative values."""
|
|
embedding = np.array([-0.5, -0.3, 0.0, 0.3, 0.5], dtype=np.float32)
|
|
embedding_id = embedding_repo.add_embedding(
|
|
file_id=1,
|
|
embedding=embedding,
|
|
model_version="v1.0",
|
|
created_time=1000
|
|
)
|
|
|
|
result = embedding_repo.get_embedding(embedding_id)
|
|
assert np.array_equal(result["embedding"], embedding)
|
|
|
|
|
|
# =============================================================================
|
|
# Schema Bug Fix Tests (dimensions column)
|
|
# =============================================================================
|
|
|
|
class TestSchemaBugFix:
|
|
"""Tests related to the schema bug (missing dimensions column)."""
|
|
|
|
def test_schema_has_dimensions_column(self, in_memory_db):
|
|
"""Test that embeddings table has dimensions column."""
|
|
cursor = in_memory_db.execute("PRAGMA table_info(embeddings)")
|
|
columns = [row[1] for row in cursor.fetchall()]
|
|
assert "dimensions" in columns
|
|
|
|
def test_embedding_with_dimensions(self, in_memory_db):
|
|
"""Test adding embedding with dimensions value."""
|
|
embedding = np.array([0.1, 0.2, 0.3, 0.4, 0.5], dtype=np.float32)
|
|
# nosec B301 - Testing legacy pickle deserialization fallback
|
|
# nosem: python.lang.security.deserialization.pickle.avoid-pickle - Testing legacy pickle deserialization fallback
|
|
embedding_bytes = pickle.dumps(embedding)
|
|
|
|
in_memory_db.execute(
|
|
"""
|
|
INSERT INTO embeddings (file_id, embedding, model_version, created_time, dimensions)
|
|
VALUES (?, ?, ?, ?, ?)
|
|
""",
|
|
(1, embedding_bytes, "v1.0", 1000, 5)
|
|
)
|
|
in_memory_db.commit()
|
|
|
|
cursor = in_memory_db.execute("SELECT dimensions FROM embeddings WHERE file_id = 1")
|
|
result = cursor.fetchone()
|
|
assert result[0] == 5
|
|
|
|
def test_dimensions_stored_correctly(self, in_memory_db):
|
|
"""Test that dimensions are stored correctly for various sizes."""
|
|
for dim in [10, 50, 128, 512, 1024]:
|
|
embedding = np.random.rand(dim).astype(np.float32)
|
|
# nosec B301 - Testing legacy pickle deserialization fallback
|
|
# nosem: python.lang.security.deserialization.pickle.avoid-pickle - Testing legacy pickle deserialization fallback
|
|
embedding_bytes = pickle.dumps(embedding)
|
|
|
|
in_memory_db.execute(
|
|
"""
|
|
INSERT INTO embeddings (file_id, embedding, model_version, created_time, dimensions)
|
|
VALUES (?, ?, ?, ?, ?)
|
|
""",
|
|
(1, embedding_bytes, "v1.0", 1000, dim)
|
|
)
|
|
in_memory_db.commit()
|
|
|
|
cursor = in_memory_db.execute("SELECT dimensions FROM embeddings WHERE file_id = 1 AND model_version = 'v1.0' ORDER BY id DESC LIMIT 1")
|
|
result = cursor.fetchone()
|
|
assert result[0] == dim
|
|
|
|
def test_get_embedding_dimensions_helper(self, db_connection):
|
|
"""Test _get_embedding_dimensions helper method."""
|
|
repo = EmbeddingRepository(db_connection)
|
|
# Test with list
|
|
assert repo._get_embedding_dimensions([1, 2, 3]) == 3
|
|
# Test with numpy array
|
|
assert repo._get_embedding_dimensions(np.array([1, 2, 3, 4])) == 4
|
|
# Test with empty list
|
|
assert repo._get_embedding_dimensions([]) == 0
|
|
# Test with non-sequence (should return 0)
|
|
assert repo._get_embedding_dimensions(42) == 0
|