diff --git a/app/api/similarity.py b/app/api/similarity.py
index ca15b23e..26bb7dc6 100644
--- a/app/api/similarity.py
+++ b/app/api/similarity.py
@@ -82,6 +82,19 @@ def get_similar_documents(
"message": "No OCR text available for similarity comparison",
}
+ # Check whether an embedding has been computed yet
+ if not file_record.embedding:
+ return {
+ "file_id": file_id,
+ "similar_documents": [],
+ "count": 0,
+ "message": (
+ "Embedding not yet computed for this file. "
+ "It will be generated automatically during processing or via the backfill task. "
+ "You can also trigger it manually with POST /api/files/{file_id}/compute-embedding."
+ ),
+ }
+
try:
from app.utils.similarity import find_similar_documents
@@ -331,3 +344,120 @@ def trigger_compute_all_embeddings(
"status": "queued",
"files_queued": queued,
}
+
+
+@router.get("/similarity/pairs")
+@require_login
+def get_similarity_pairs(
+ request: Request,
+ db: DbSession,
+ threshold: float = Query(0.7, ge=0.0, le=1.0, description="Minimum similarity score for a pair"),
+ limit: int = Query(50, ge=1, le=200, description="Maximum number of pairs to return"),
+ page: int = Query(1, ge=1, description="Page number"),
+):
+ """Return pairs of documents with high similarity across the entire corpus.
+
+ Unlike the per-file ``/files/{id}/similar`` endpoint, this scans every
+ document that has a pre-computed embedding and returns **all** pairs
+ whose cosine similarity exceeds ``threshold``, sorted by descending
+ score.
+
+ To keep memory bounded the query loads only the columns needed for
+ scoring and streams results in chunks.
+
+ Response:
+ ```json
+ {
+ "pairs": [
+ {
+ "file_a": {"file_id": 1, "original_filename": "invoice_jan.pdf", ...},
+ "file_b": {"file_id": 5, "original_filename": "invoice_feb.pdf", ...},
+ "similarity_score": 0.94
+ }
+ ],
+ "total_pairs": 12,
+ "threshold": 0.7,
+ "page": 1,
+ "pages": 1,
+ "embedding_coverage": {"total_files": 120, "files_with_embedding": 95}
+ }
+ ```
+ """
+ from app.utils.similarity import cosine_similarity
+
+ # Load all files that have embeddings (columns only for efficiency)
+ rows = (
+ db.query(
+ FileRecord.id,
+ FileRecord.original_filename,
+ FileRecord.document_title,
+ FileRecord.mime_type,
+ FileRecord.created_at,
+ FileRecord.embedding,
+ )
+ .filter(
+ FileRecord.embedding.isnot(None),
+ FileRecord.embedding != "",
+ )
+ .order_by(FileRecord.id)
+ .all()
+ )
+
+ # Parse embeddings upfront
+ parsed: list[tuple] = []
+ for row in rows:
+ try:
+ vec = json.loads(row.embedding)
+ parsed.append((row, vec))
+ except (json.JSONDecodeError, TypeError):
+ continue
+
+ # Pairwise comparison (triangle: i < j avoids duplicating A↔B / B↔A)
+ all_pairs: list[dict] = []
+ for i in range(len(parsed)):
+ row_a, vec_a = parsed[i]
+ for j in range(i + 1, len(parsed)):
+ row_b, vec_b = parsed[j]
+ score = cosine_similarity(vec_a, vec_b)
+ if score >= threshold:
+ all_pairs.append(
+ {
+ "file_a": _row_to_dict(row_a),
+ "file_b": _row_to_dict(row_b),
+ "similarity_score": round(score, 4),
+ }
+ )
+
+ # Sort by score descending
+ all_pairs.sort(key=lambda p: p["similarity_score"], reverse=True)
+
+ total_pairs = len(all_pairs)
+ total_pages = max(1, (total_pairs + limit - 1) // limit)
+ offset = (page - 1) * limit
+ page_pairs = all_pairs[offset : offset + limit]
+
+ total_files = db.query(FileRecord).count()
+
+ return {
+ "pairs": page_pairs,
+ "total_pairs": total_pairs,
+ "threshold": threshold,
+ "page": page,
+ "pages": total_pages,
+ "per_page": limit,
+ "embedding_coverage": {
+ "total_files": total_files,
+ "files_with_embedding": len(parsed),
+ },
+ }
+
+
+def _row_to_dict(row) -> dict:
+ """Serialise a column-only query row to a dict for JSON responses."""
+ return {
+ "file_id": row.id,
+ "original_filename": row.original_filename,
+ "document_title": row.document_title,
+ "mime_type": row.mime_type,
+ "created_at": row.created_at.isoformat() if row.created_at else None,
+ }
diff --git a/app/tasks/compute_embedding.py b/app/tasks/compute_embedding.py
index 0df05175..fedc6585 100644
--- a/app/tasks/compute_embedding.py
+++ b/app/tasks/compute_embedding.py
@@ -6,12 +6,14 @@ access.
"""
import logging
+from datetime import datetime, timezone
from app.celery_app import celery
from app.database import SessionLocal
from app.models import FileRecord
from app.tasks.retry_config import BaseTaskWithRetry
from app.utils import log_task_progress
+from app.utils.step_manager import update_step_status
logger = logging.getLogger(__name__)
@@ -47,6 +49,9 @@ def compute_document_embedding(self, file_id: int) -> dict:
logger.warning("[%s] File %s not found, skipping embedding", task_id, file_id)
return {"status": "skipped", "detail": "File not found"}
+ now = datetime.now(timezone.utc)
+ update_step_status(db, file_id, "compute_embedding", "in_progress", started_at=now)
+
# Already has a cached embedding – nothing to do
if file_record.embedding:
logger.info("[%s] File %s already has a cached embedding", task_id, file_id)
@@ -57,6 +62,7 @@ def compute_document_embedding(self, file_id: int) -> dict:
"Embedding already cached",
file_id=file_id,
)
+ update_step_status(db, file_id, "compute_embedding", "success", completed_at=now)
return {"status": "skipped", "detail": "Embedding already cached"}
if not file_record.ocr_text or not file_record.ocr_text.strip():
@@ -68,12 +74,14 @@ def compute_document_embedding(self, file_id: int) -> dict:
"No OCR text available",
file_id=file_id,
)
+ update_step_status(db, file_id, "compute_embedding", "skipped", completed_at=now)
return {"status": "skipped", "detail": "No OCR text available"}
try:
from app.utils.similarity import compute_and_store_embedding
embedding = compute_and_store_embedding(db, file_record)
+ completed = datetime.now(timezone.utc)
if embedding:
log_task_progress(
task_id,
@@ -82,6 +90,7 @@ def compute_document_embedding(self, file_id: int) -> dict:
f"Embedding computed ({len(embedding)} dimensions)",
file_id=file_id,
)
+ update_step_status(db, file_id, "compute_embedding", "success", completed_at=completed)
return {
"status": "success",
"detail": f"Embedding computed ({len(embedding)} dimensions)",
@@ -94,6 +103,14 @@ def compute_document_embedding(self, file_id: int) -> dict:
"Embedding computation returned None",
file_id=file_id,
)
+ update_step_status(
+ db,
+ file_id,
+ "compute_embedding",
+ "failure",
+ error_message="Embedding computation returned None",
+ completed_at=completed,
+ )
return {"status": "error", "detail": "Embedding computation returned None"}
except Exception as exc:
logger.exception("[%s] Embedding computation failed for file %s: %s", task_id, file_id, exc)
@@ -104,6 +121,14 @@ def compute_document_embedding(self, file_id: int) -> dict:
f"Exception: {exc}",
file_id=file_id,
)
+ update_step_status(
+ db,
+ file_id,
+ "compute_embedding",
+ "failure",
+ error_message=str(exc),
+ completed_at=datetime.now(timezone.utc),
+ )
return {"status": "error", "detail": str(exc)}
diff --git a/app/utils/step_manager.py b/app/utils/step_manager.py
index 9e4449a3..ca094847 100644
--- a/app/utils/step_manager.py
+++ b/app/utils/step_manager.py
@@ -23,6 +23,7 @@ BASE_MAIN_PROCESSING_STEPS = [
"embed_metadata_into_pdf",
"finalize_document_storage",
"send_to_all_destinations",
+ "compute_embedding",
]
OPTIONAL_PROCESSING_STEPS = {
diff --git a/app/views/files.py b/app/views/files.py
index dad67680..c1a2c7da 100644
--- a/app/views/files.py
+++ b/app/views/files.py
@@ -389,8 +389,12 @@ def _compute_processing_flow(logs):
},
"extract_metadata_with_gpt": {"label": "Extract Metadata (GPT)", "next": ["embed_metadata_into_pdf"]},
"embed_metadata_into_pdf": {"label": "Embed Metadata into PDF", "next": ["finalize_document_storage"]},
- "finalize_document_storage": {"label": "Finalize & Queue Distribution", "next": ["send_to_all_destinations"]},
+ "finalize_document_storage": {
+ "label": "Finalize & Queue Distribution",
+ "next": ["send_to_all_destinations", "compute_embedding"],
+ },
"send_to_all_destinations": {"label": "Upload to Destinations", "next": [], "has_branches": True},
+ "compute_embedding": {"label": "Compute Embedding", "next": []},
}
# Filter out deduplication step if not enabled or if not showing it
@@ -815,3 +819,54 @@ def duplicates_page(
"error": str(e),
},
)
+
+
+@router.get("/similarity")
+@require_login
+def similarity_dashboard_page(
+ request: Request,
+ db: Session = Depends(get_db),
+):
+ """Render the corpus-wide similarity dashboard.
+
+ Passes the configured threshold and embedding coverage stats so the
+ template can display them immediately while the JS fetches the actual
+ pairs from the API asynchronously.
+ """
+ from app.config import settings
+ from app.models import FileRecord
+
+ try:
+ total_files = db.query(FileRecord).count()
+ files_with_embedding = (
+ db.query(FileRecord).filter(FileRecord.embedding.isnot(None), FileRecord.embedding != "").count()
+ )
+ files_with_ocr = db.query(FileRecord).filter(FileRecord.ocr_text.isnot(None), FileRecord.ocr_text != "").count()
+
+ return templates.TemplateResponse(
+ "similarity_dashboard.html",
+ {
+ "request": request,
+ "default_threshold": settings.near_duplicate_threshold,
+ "embedding_model": settings.embedding_model,
+ "total_files": total_files,
+ "files_with_embedding": files_with_embedding,
+ "files_with_ocr": files_with_ocr,
+ "files_missing_embedding": files_with_ocr - files_with_embedding,
+ },
+ )
+ except Exception as e:
+ logger.error(f"Error rendering similarity dashboard: {e}")
+ return templates.TemplateResponse(
+ "similarity_dashboard.html",
+ {
+ "request": request,
+ "default_threshold": 0.85,
+ "embedding_model": "text-embedding-3-small",
+ "total_files": 0,
+ "files_with_embedding": 0,
+ "files_with_ocr": 0,
+ "files_missing_embedding": 0,
+ "error": str(e),
+ },
+ )
diff --git a/frontend/templates/base.html b/frontend/templates/base.html
index ef42f9cc..44f537fe 100644
--- a/frontend/templates/base.html
+++ b/frontend/templates/base.html
@@ -104,6 +104,9 @@
Duplicates
+
+ Similarity
+
Queue Monitor
@@ -184,6 +187,9 @@
Duplicates
+
+ Similarity
+
Queue Monitor
diff --git a/frontend/templates/similarity_dashboard.html b/frontend/templates/similarity_dashboard.html
new file mode 100644
index 00000000..1fdf2c9a
--- /dev/null
+++ b/frontend/templates/similarity_dashboard.html
@@ -0,0 +1,368 @@
+{% extends "base.html" %}
+{% block title %}Document Similarity - DocuElevate{% endblock %}
+
+{% block head_extra %}
+
+{% endblock %}
+
+{% block content %}
+
+
+
+ Document Similarity
+
+
+ Pairs of documents with high semantic similarity, ranked by score.
+ Embeddings are computed during document ingestion; a background task
+ also backfills any files that were processed before this feature was enabled.
+
+
+
+
+
+
{{ total_files }}
+
Total Files
+
+
+
{{ files_with_embedding }}
+
With Embedding
+
+
+
+ {{ files_missing_embedding }}
+
+
Missing Embedding
+
+
+
{{ embedding_model }}
+
Embedding Model
+
+
+
+ {% if files_missing_embedding > 0 %}
+
+
+ {{ files_missing_embedding }} file(s) have OCR text but no embedding yet.
+ The background task will compute them automatically every 5 minutes, or you can
+ .
+
+ {% endif %}
+
+
+
+
+
+
+
+
+
+
Scanning for similar document pairs…
+
+
+
+
+
No similar pairs found
+
+ No document pairs exceed the similarity threshold.
+
+
+
+
+
Failed to load pairs.
+
+
+
+
+
+{% endblock %}
+
+{% block scripts %}
+
+{% endblock %}
diff --git a/tests/test_similarity.py b/tests/test_similarity.py
index c790b51b..45357747 100644
--- a/tests/test_similarity.py
+++ b/tests/test_similarity.py
@@ -108,18 +108,23 @@ class TestFindSimilarDocuments:
assert result == []
@pytest.mark.unit
- @patch("app.utils.similarity.generate_embedding")
- def test_finds_similar_documents(self, mock_embed, db_session):
- """Should find similar documents based on embedding similarity."""
- # Create a target file with OCR text
+ def test_finds_similar_documents(self, db_session):
+ """Should find similar documents based on pre-computed embedding similarity."""
+ # Pre-computed embeddings that reflect similarity
+ target_embedding = [1.0, 0.0, 0.0]
+ similar_embedding = [0.95, 0.05, 0.0]
+ different_embedding = [0.0, 0.0, 1.0]
+
+ # Create a target file with pre-computed embedding
target = FileRecord(
filehash="hash1",
local_filename="/tmp/target.pdf",
file_size=1024,
original_filename="target.pdf",
ocr_text="This is an invoice from Amazon for January 2026",
+ embedding=json.dumps(target_embedding),
)
- # Create a similar file
+ # Create a similar file with pre-computed embedding
similar = FileRecord(
filehash="hash2",
local_filename="/tmp/similar.pdf",
@@ -128,8 +133,9 @@ class TestFindSimilarDocuments:
ocr_text="This is an invoice from Amazon for February 2026",
document_title="Amazon Invoice Feb",
mime_type="application/pdf",
+ embedding=json.dumps(similar_embedding),
)
- # Create a different file
+ # Create a different file with pre-computed embedding
different = FileRecord(
filehash="hash3",
local_filename="/tmp/different.pdf",
@@ -138,26 +144,12 @@ class TestFindSimilarDocuments:
ocr_text="Recipe for chocolate cake with detailed instructions",
document_title="Chocolate Cake Recipe",
mime_type="application/pdf",
+ embedding=json.dumps(different_embedding),
)
db_session.add_all([target, similar, different])
db_session.commit()
- # Mock embeddings that reflect similarity
- target_embedding = [1.0, 0.0, 0.0]
- similar_embedding = [0.95, 0.05, 0.0]
- different_embedding = [0.0, 0.0, 1.0]
-
- def mock_embed_side_effect(text):
- if "January" in text or "invoice" in text.lower()[:30]:
- return target_embedding
- elif "February" in text:
- return similar_embedding
- else:
- return different_embedding
-
- mock_embed.side_effect = mock_embed_side_effect
-
result = find_similar_documents(db_session, file_id=target.id, threshold=0.3)
assert len(result) == 1
@@ -166,8 +158,7 @@ class TestFindSimilarDocuments:
assert result[0]["original_filename"] == "similar.pdf"
@pytest.mark.unit
- @patch("app.utils.similarity.generate_embedding")
- def test_respects_threshold(self, mock_embed, db_session):
+ def test_respects_threshold(self, db_session):
"""Should filter out documents below the threshold."""
target = FileRecord(
filehash="hash1",
@@ -175,6 +166,7 @@ class TestFindSimilarDocuments:
file_size=100,
original_filename="target.pdf",
ocr_text="target text",
+ embedding=json.dumps([1.0, 0.0]),
)
candidate = FileRecord(
filehash="hash2",
@@ -182,26 +174,25 @@ class TestFindSimilarDocuments:
file_size=100,
original_filename="candidate.pdf",
ocr_text="different text",
+ embedding=json.dumps([0.1, 0.99]),
)
db_session.add_all([target, candidate])
db_session.commit()
- # Return nearly orthogonal vectors -> low similarity
- mock_embed.side_effect = lambda text: [1.0, 0.0] if "target" in text else [0.1, 0.99]
-
result = find_similar_documents(db_session, file_id=target.id, threshold=0.9)
assert len(result) == 0
@pytest.mark.unit
- @patch("app.utils.similarity.generate_embedding")
- def test_respects_limit(self, mock_embed, db_session):
+ def test_respects_limit(self, db_session):
"""Should respect the limit parameter."""
+ embedding = [1.0, 0.0, 0.0]
target = FileRecord(
filehash="hash0",
local_filename="/tmp/t.pdf",
file_size=100,
original_filename="target.pdf",
ocr_text="target text",
+ embedding=json.dumps(embedding),
)
db_session.add(target)
@@ -212,12 +203,11 @@ class TestFindSimilarDocuments:
file_size=100,
original_filename=f"candidate_{i}.pdf",
ocr_text=f"similar text {i}",
+ embedding=json.dumps(embedding),
)
db_session.add(f)
db_session.commit()
- mock_embed.return_value = [1.0, 0.0, 0.0]
-
result = find_similar_documents(db_session, file_id=target.id, limit=2, threshold=0.0)
assert len(result) <= 2
@@ -286,15 +276,16 @@ class TestSimilarDocumentsAPI:
assert "message" in data
@pytest.mark.integration
- @patch("app.utils.similarity.generate_embedding")
- def test_returns_similar_documents(self, mock_embed, client: TestClient, db_session):
+ def test_returns_similar_documents(self, client: TestClient, db_session):
"""Should return similar documents with scores."""
+ embedding = [1.0, 0.0, 0.0]
target = FileRecord(
filehash="hash1",
local_filename="/tmp/target.pdf",
file_size=1024,
original_filename="target.pdf",
ocr_text="Invoice from Amazon January 2026",
+ embedding=json.dumps(embedding),
)
similar = FileRecord(
filehash="hash2",
@@ -304,12 +295,11 @@ class TestSimilarDocumentsAPI:
ocr_text="Invoice from Amazon February 2026",
document_title="Amazon Invoice",
mime_type="application/pdf",
+ embedding=json.dumps(embedding),
)
db_session.add_all([target, similar])
db_session.commit()
- mock_embed.return_value = [1.0, 0.0, 0.0]
-
response = client.get(f"/api/files/{target.id}/similar")
assert response.status_code == 200
data = response.json()
@@ -324,15 +314,16 @@ class TestSimilarDocumentsAPI:
assert "original_filename" in doc
@pytest.mark.integration
- @patch("app.utils.similarity.generate_embedding")
- def test_query_parameters(self, mock_embed, client: TestClient, db_session):
+ def test_query_parameters(self, client: TestClient, db_session):
"""Should respect limit and threshold query parameters."""
+ embedding = [1.0, 0.0]
target = FileRecord(
filehash="hash1",
local_filename="/tmp/t.pdf",
file_size=100,
original_filename="t.pdf",
ocr_text="test",
+ embedding=json.dumps(embedding),
)
db_session.add(target)
@@ -343,12 +334,11 @@ class TestSimilarDocumentsAPI:
file_size=100,
original_filename=f"c{i}.pdf",
ocr_text=f"text {i}",
+ embedding=json.dumps(embedding),
)
db_session.add(f)
db_session.commit()
- mock_embed.return_value = [1.0, 0.0]
-
response = client.get(f"/api/files/{target.id}/similar?limit=2&threshold=0.0")
assert response.status_code == 200
data = response.json()
@@ -403,15 +393,36 @@ class TestSimilarDocumentsAPI:
assert data["count"] == 0
@pytest.mark.integration
- @patch("app.utils.similarity.generate_embedding")
- def test_response_structure(self, mock_embed, client: TestClient, db_session):
+ def test_embedding_not_computed_message(self, client: TestClient, db_session):
+ """Should return a message when OCR text exists but no embedding yet."""
+ file_record = FileRecord(
+ filehash="noembhash",
+ local_filename="/tmp/noemb.pdf",
+ file_size=100,
+ original_filename="noemb.pdf",
+ ocr_text="Some OCR text content",
+ embedding=None,
+ )
+ db_session.add(file_record)
+ db_session.commit()
+
+ response = client.get(f"/api/files/{file_record.id}/similar")
+ assert response.status_code == 200
+ data = response.json()
+ assert data["count"] == 0
+ assert "message" in data
+ assert "not yet computed" in data["message"].lower()
+
+ def test_response_structure(self, client: TestClient, db_session):
"""Should return proper response structure for each similar document."""
+ embedding = [1.0, 0.0]
target = FileRecord(
filehash="h1",
local_filename="/tmp/t.pdf",
file_size=100,
original_filename="target.pdf",
ocr_text="Some text content here",
+ embedding=json.dumps(embedding),
)
other = FileRecord(
filehash="h2",
@@ -421,12 +432,11 @@ class TestSimilarDocumentsAPI:
ocr_text="Some similar text content",
document_title="Other Doc",
mime_type="application/pdf",
+ embedding=json.dumps(embedding),
)
db_session.add_all([target, other])
db_session.commit()
- mock_embed.return_value = [1.0, 0.0]
-
response = client.get(f"/api/files/{target.id}/similar")
assert response.status_code == 200
data = response.json()
@@ -829,3 +839,176 @@ class TestComputeDocumentEmbeddingTask:
assert result["status"] == "skipped"
assert "already cached" in result["detail"]
+
+
+# ---------------------------------------------------------------------------
+# Tests for similarity pairs endpoint
+# ---------------------------------------------------------------------------
+
+
+class TestSimilarityPairsAPI:
+ """Integration tests for GET /api/similarity/pairs."""
+
+ @pytest.mark.integration
+ def test_empty_database(self, client: TestClient):
+ """Should return zero pairs on empty database."""
+ response = client.get("/api/similarity/pairs")
+ assert response.status_code == 200
+ data = response.json()
+ assert data["total_pairs"] == 0
+ assert data["pairs"] == []
+ assert "embedding_coverage" in data
+
+ @pytest.mark.integration
+ def test_finds_similar_pairs(self, client: TestClient, db_session):
+ """Should find and return pairs of similar files."""
+ emb_a = [1.0, 0.0, 0.0]
+ emb_b = [0.98, 0.02, 0.0] # Very similar to A
+ emb_c = [0.0, 0.0, 1.0] # Different from A and B
+
+ f1 = FileRecord(
+ filehash="pairA",
+ local_filename="/tmp/pA.pdf",
+ file_size=100,
+ original_filename="fileA.pdf",
+ ocr_text="text A",
+ embedding=json.dumps(emb_a),
+ )
+ f2 = FileRecord(
+ filehash="pairB",
+ local_filename="/tmp/pB.pdf",
+ file_size=100,
+ original_filename="fileB.pdf",
+ ocr_text="text B",
+ embedding=json.dumps(emb_b),
+ )
+ f3 = FileRecord(
+ filehash="pairC",
+ local_filename="/tmp/pC.pdf",
+ file_size=100,
+ original_filename="fileC.pdf",
+ ocr_text="text C",
+ embedding=json.dumps(emb_c),
+ )
+ db_session.add_all([f1, f2, f3])
+ db_session.commit()
+
+ response = client.get("/api/similarity/pairs?threshold=0.9")
+ assert response.status_code == 200
+ data = response.json()
+
+ # Only A-B pair should be above 0.9
+ assert data["total_pairs"] == 1
+ pair = data["pairs"][0]
+ assert pair["similarity_score"] > 0.9
+ pair_ids = {pair["file_a"]["file_id"], pair["file_b"]["file_id"]}
+ assert pair_ids == {f1.id, f2.id}
+
+ @pytest.mark.integration
+ def test_respects_threshold(self, client: TestClient, db_session):
+ """Should filter pairs below threshold."""
+ emb = [1.0, 0.0]
+ different_emb = [0.0, 1.0]
+
+ f1 = FileRecord(
+ filehash="thA",
+ local_filename="/tmp/thA.pdf",
+ file_size=100,
+ original_filename="thA.pdf",
+ ocr_text="a",
+ embedding=json.dumps(emb),
+ )
+ f2 = FileRecord(
+ filehash="thB",
+ local_filename="/tmp/thB.pdf",
+ file_size=100,
+ original_filename="thB.pdf",
+ ocr_text="b",
+ embedding=json.dumps(different_emb),
+ )
+ db_session.add_all([f1, f2])
+ db_session.commit()
+
+ response = client.get("/api/similarity/pairs?threshold=0.9")
+ assert response.status_code == 200
+ data = response.json()
+ assert data["total_pairs"] == 0
+
+ @pytest.mark.integration
+ def test_pagination(self, client: TestClient, db_session):
+ """Should respect pagination parameters."""
+ emb = [1.0, 0.0, 0.0]
+ for i in range(5):
+ f = FileRecord(
+ filehash=f"pg{i}",
+ local_filename=f"/tmp/pg{i}.pdf",
+ file_size=100,
+ original_filename=f"pg{i}.pdf",
+ ocr_text=f"text {i}",
+ embedding=json.dumps(emb),
+ )
+ db_session.add(f)
+ db_session.commit()
+
+ response = client.get("/api/similarity/pairs?threshold=0.0&limit=2&page=1")
+ assert response.status_code == 200
+ data = response.json()
+ assert len(data["pairs"]) <= 2
+ assert data["per_page"] == 2
+
+
+# ---------------------------------------------------------------------------
+# Tests for backfill_missing_embeddings task
+# ---------------------------------------------------------------------------
+
+
+class TestBackfillMissingEmbeddingsTask:
+ """Unit tests for the backfill_missing_embeddings Celery task."""
+
+ @pytest.mark.unit
+ @patch("app.tasks.compute_embedding.compute_document_embedding.delay")
+ def test_queues_files_without_embeddings(self, mock_delay, db_session):
+ """Should queue tasks for files with OCR text but no embedding."""
+ f1 = FileRecord(
+ filehash="bf1",
+ local_filename="/tmp/bf1.pdf",
+ file_size=100,
+ original_filename="bf1.pdf",
+ ocr_text="Some text",
+ )
+ f2 = FileRecord(
+ filehash="bf2",
+ local_filename="/tmp/bf2.pdf",
+ file_size=100,
+ original_filename="bf2.pdf",
+ ocr_text="More text",
+ embedding=json.dumps([0.1]),
+ )
+ db_session.add_all([f1, f2])
+ db_session.commit()
+
+ from app.tasks.compute_embedding import backfill_missing_embeddings
+
+ with patch("app.tasks.compute_embedding.SessionLocal") as mock_session_local:
+ mock_session_local.return_value.__enter__ = lambda self: db_session
+ mock_session_local.return_value.__exit__ = lambda self, *args: None
+
+ result = backfill_missing_embeddings()
+
+ assert result["queued"] == 1
+ mock_delay.assert_called_once_with(f1.id)
+
+ @pytest.mark.unit
+ @patch("app.tasks.compute_embedding.compute_document_embedding.delay")
+ def test_empty_database(self, mock_delay, db_session):
+ """Should queue nothing when no files need embeddings."""
+ from app.tasks.compute_embedding import backfill_missing_embeddings
+
+ with patch("app.tasks.compute_embedding.SessionLocal") as mock_session_local:
+ mock_session_local.return_value.__enter__ = lambda self: db_session
+ mock_session_local.return_value.__exit__ = lambda self, *args: None
+
+ result = backfill_missing_embeddings()
+
+ assert result["queued"] == 0
+ mock_delay.assert_not_called()