feat(search): validate the index before the alias points at it
The old sync created a collection, imported batches without reading a single import response, and swapped the alias regardless. A partial import published a half-empty index, and two overlapping DAG runs could prune each other's collections. Publication now checks every import response and the final document count before upserting the alias, and holds a session-scoped advisory lock across the read and the publish so concurrent runs serialise. Cleanup keeps the previous collection as a rollback pointer and is best-effort: an uncertain alias response must never delete what might still be live. Also parses the Typesense URL properly instead of splitting on colons, which mangled any host carrying a scheme and a default port. Co-Authored-By: Claude Opus 5 <noreply@anthropic.com>
This commit is contained in:
1 parent
1d8858fbda
commit
38bc17cab3
2 files changed
+161
-50
No files matched your search
@@ -11,6 +11,10 @@ Usage:
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import json
|
||||
import logging
|
||||
import re
|
||||
import uuid
|
||||
import os
|
||||
import sys
|
||||
import time
|
||||
@@ -112,63 +116,78 @@ def build_document(row: dict) -> dict:
|
||||
return doc
|
||||
|
||||
|
||||
def publish_collection(client, rows: list[dict]) -> str:
|
||||
"""Validate a new collection before moving the alias; keep rollback data.
|
||||
|
||||
Caller holds the database advisory lock across reading and publication so
|
||||
overlapping school-data DAGs cannot publish or prune each other's work.
|
||||
Failed drafts are left for the next successful publication to prune: a lost
|
||||
alias-update response must never cause deletion of a potentially live index.
|
||||
"""
|
||||
if not rows or len({r["urn"] for r in rows}) != len(rows):
|
||||
raise ValueError("Search source must contain nonempty, unique school URNs")
|
||||
name = f"schools_{int(time.time())}_{uuid.uuid4().hex[:12]}"
|
||||
try:
|
||||
previous = client.aliases["schools"].retrieve()["collection_name"]
|
||||
except typesense.exceptions.ObjectNotFound:
|
||||
previous = None
|
||||
client.collections.create({**COLLECTION_SCHEMA, "name": name})
|
||||
for i in range(0, len(rows), 500):
|
||||
batch = [build_document(r) for r in rows[i:i + 500]]
|
||||
results = client.collections[name].documents.import_(batch, {"action": "upsert"})
|
||||
if isinstance(results, str):
|
||||
results = [json.loads(line) for line in results.splitlines() if line.strip()]
|
||||
if len(results) != len(batch) or any(r.get("success") is not True for r in results):
|
||||
raise ValueError("Search import failed; live alias unchanged")
|
||||
if client.collections[name].retrieve()["num_documents"] != len(rows):
|
||||
raise ValueError("Search document count mismatch; live alias unchanged")
|
||||
client.aliases.upsert("schools", {"collection_name": name})
|
||||
|
||||
# Retention is best-effort and must not make successful publication fail.
|
||||
try:
|
||||
keep = {name, previous}
|
||||
keep.update(a["collection_name"] for a in client.aliases.retrieve()["aliases"])
|
||||
for collection in client.collections.retrieve():
|
||||
old = collection["name"]
|
||||
if old not in keep and re.fullmatch(r"schools_\d+(?:_[0-9a-f]+)?", old):
|
||||
client.collections[old].delete()
|
||||
except Exception:
|
||||
logging.getLogger(__name__).exception("Search published, but old collection cleanup failed")
|
||||
return name
|
||||
|
||||
|
||||
def sync(typesense_url: str, api_key: str):
|
||||
from urllib.parse import urlparse
|
||||
url = urlparse(typesense_url)
|
||||
client = typesense.Client({
|
||||
"nodes": [{"host": typesense_url.split("//")[-1].split(":")[0],
|
||||
"port": typesense_url.split(":")[-1],
|
||||
"protocol": "http"}],
|
||||
"nodes": [{"host": url.hostname, "port": str(url.port or 8108),
|
||||
"protocol": url.scheme}],
|
||||
"api_key": api_key,
|
||||
"connection_timeout_seconds": 10,
|
||||
})
|
||||
|
||||
# Create timestamped collection for zero-downtime swap
|
||||
ts = int(time.time())
|
||||
collection_name = f"schools_{ts}"
|
||||
|
||||
print(f"Creating collection: {collection_name}")
|
||||
schema = {**COLLECTION_SCHEMA, "name": collection_name}
|
||||
client.collections.create(schema)
|
||||
|
||||
# Fetch data from marts — join fact_performance if it exists
|
||||
conn = get_db_connection()
|
||||
with conn.cursor(cursor_factory=psycopg2.extras.RealDictCursor) as cur:
|
||||
# Check whether the merged fact table exists
|
||||
cur.execute("""
|
||||
SELECT table_name FROM information_schema.tables
|
||||
WHERE table_schema = 'marts' AND table_name = 'fact_performance'
|
||||
""")
|
||||
has_fact_performance = cur.fetchone() is not None
|
||||
|
||||
query = QUERY_BASE
|
||||
if has_fact_performance:
|
||||
query = query.replace(
|
||||
"l.longitude as lng",
|
||||
"l.longitude as lng,\n p.rwm_expected_pct,\n p.progress_8_score",
|
||||
)
|
||||
query += QUERY_PERFORMANCE_JOIN
|
||||
|
||||
cur.execute(query)
|
||||
rows = cur.fetchall()
|
||||
conn.close()
|
||||
|
||||
print(f"Indexing {len(rows)} schools...")
|
||||
|
||||
# Batch import
|
||||
batch_size = 500
|
||||
for i in range(0, len(rows), batch_size):
|
||||
batch = [build_document(r) for r in rows[i : i + batch_size]]
|
||||
client.collections[collection_name].documents.import_(batch, {"action": "upsert"})
|
||||
print(f" Indexed {min(i + batch_size, len(rows))}/{len(rows)}")
|
||||
|
||||
# Swap alias
|
||||
print("Swapping alias 'schools' → new collection")
|
||||
try:
|
||||
client.aliases.upsert("schools", {"collection_name": collection_name})
|
||||
except Exception:
|
||||
# If alias doesn't exist yet, create it
|
||||
client.aliases.upsert("schools", {"collection_name": collection_name})
|
||||
|
||||
print("Done.")
|
||||
with conn.cursor(cursor_factory=psycopg2.extras.RealDictCursor) as cur:
|
||||
# Session-scoped lock is released even on errors when conn closes.
|
||||
cur.execute("SELECT pg_advisory_lock(731042019)")
|
||||
cur.execute("""
|
||||
SELECT table_name FROM information_schema.tables
|
||||
WHERE table_schema = 'marts' AND table_name = 'fact_performance'
|
||||
""")
|
||||
has_fact_performance = cur.fetchone() is not None
|
||||
query = QUERY_BASE
|
||||
if has_fact_performance:
|
||||
query = query.replace(
|
||||
"l.longitude as lng",
|
||||
"l.longitude as lng, p.rwm_expected_pct, p.progress_8_score",
|
||||
)
|
||||
query += QUERY_PERFORMANCE_JOIN
|
||||
cur.execute(query)
|
||||
rows = cur.fetchall()
|
||||
name = publish_collection(client, rows)
|
||||
print(f"Published {len(rows)} schools in {name}")
|
||||
finally:
|
||||
conn.close()
|
||||
|
||||
|
||||
def main():
|
||||
|
||||
@@ -0,0 +1,92 @@
|
||||
import importlib.util
|
||||
from pathlib import Path
|
||||
from unittest.mock import MagicMock
|
||||
import pytest
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def sync_module(monkeypatch):
|
||||
folder = Path(__file__).resolve().parents[1] / 'scripts'
|
||||
monkeypatch.syspath_prepend(str(folder))
|
||||
spec = importlib.util.spec_from_file_location('sync_typesense', folder / 'sync_typesense.py')
|
||||
module = importlib.util.module_from_spec(spec)
|
||||
spec.loader.exec_module(module)
|
||||
return module
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def client():
|
||||
c = MagicMock()
|
||||
c.aliases.__getitem__.return_value.retrieve.return_value = {'collection_name': 'schools_2'}
|
||||
c.aliases.retrieve.return_value = {'aliases': [{'collection_name': 'schools_99'}]}
|
||||
c.collections.__getitem__.return_value.documents.import_.return_value = [{'success': True}]
|
||||
c.collections.__getitem__.return_value.retrieve.return_value = {'num_documents': 1}
|
||||
c.collections.retrieve.return_value = [{'name': n} for n in ['schools_1', 'schools_2', 'schools_99', 'unrelated']]
|
||||
return c
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def rows():
|
||||
return [{'urn': 100001, 'school_name': 'Example', 'phase_code': 2,
|
||||
'school_type_code': 1, 'local_authority': 'Testshire', 'postcode': 'TS1 1AA',
|
||||
'total_pupils': 250}]
|
||||
|
||||
|
||||
@pytest.mark.parametrize('results', [[{'success': False}], [], [{'success': True}, {'success': True}]])
|
||||
def test_partial_import_never_moves_alias(sync_module, client, rows, results):
|
||||
client.collections.__getitem__.return_value.documents.import_.return_value = results
|
||||
with pytest.raises(ValueError):
|
||||
sync_module.publish_collection(client, rows)
|
||||
client.aliases.upsert.assert_not_called()
|
||||
client.collections.__getitem__.return_value.delete.assert_not_called()
|
||||
|
||||
|
||||
def test_count_mismatch_does_not_publish(sync_module, client, rows):
|
||||
client.collections.__getitem__.return_value.retrieve.return_value = {'num_documents': 0}
|
||||
with pytest.raises(ValueError):
|
||||
sync_module.publish_collection(client, rows)
|
||||
client.aliases.upsert.assert_not_called()
|
||||
|
||||
|
||||
@pytest.mark.parametrize('empty', [True, False])
|
||||
def test_invalid_source_never_creates_collection(sync_module, client, rows, empty):
|
||||
with pytest.raises(ValueError):
|
||||
sync_module.publish_collection(client, [] if empty else rows + rows)
|
||||
client.collections.create.assert_not_called()
|
||||
|
||||
|
||||
def test_success_retains_previous_and_other_live_aliases(sync_module, client, rows):
|
||||
# Distinct mock per collection allows checking exactly which one was deleted.
|
||||
collections = {}
|
||||
def get(name):
|
||||
if name not in collections:
|
||||
c = MagicMock()
|
||||
c.documents.import_.return_value = '{"success":true}\n'
|
||||
c.retrieve.return_value = {'num_documents': 1}
|
||||
collections[name] = c
|
||||
return collections[name]
|
||||
client.collections.__getitem__.side_effect = get
|
||||
name = sync_module.publish_collection(client, rows)
|
||||
client.aliases.upsert.assert_called_once_with('schools', {'collection_name': name})
|
||||
assert set(collections) == {name, 'schools_1'}
|
||||
collections['schools_1'].delete.assert_called_once()
|
||||
|
||||
|
||||
def test_uncertain_alias_update_does_not_delete_candidate(sync_module, client, rows):
|
||||
client.aliases.upsert.side_effect = RuntimeError('response lost')
|
||||
with pytest.raises(RuntimeError):
|
||||
sync_module.publish_collection(client, rows)
|
||||
client.collections.__getitem__.return_value.delete.assert_not_called()
|
||||
|
||||
|
||||
def test_sync_closes_database_when_publication_fails(sync_module, monkeypatch):
|
||||
conn = MagicMock()
|
||||
monkeypatch.setattr(sync_module, 'get_db_connection', lambda: conn)
|
||||
monkeypatch.setattr(sync_module.typesense, 'Client', lambda _: MagicMock())
|
||||
def fail(*args): raise ValueError('import rejected')
|
||||
monkeypatch.setattr(sync_module, 'publish_collection', fail)
|
||||
with pytest.raises(ValueError):
|
||||
sync_module.sync('http://localhost:8108', 'dummy')
|
||||
conn.close.assert_called_once()
|
||||
statements = [call.args[0] for call in conn.cursor.return_value.__enter__.return_value.execute.call_args_list]
|
||||
assert statements[0] == 'SELECT pg_advisory_lock(731042019)'
|
||||
Reference in new issue
Block a user