diff --git a/pipeline/scripts/sync_typesense.py b/pipeline/scripts/sync_typesense.py index ebe08a3..753c377 100644 --- a/pipeline/scripts/sync_typesense.py +++ b/pipeline/scripts/sync_typesense.py @@ -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(): diff --git a/pipeline/tests/test_sync_typesense.py b/pipeline/tests/test_sync_typesense.py new file mode 100644 index 0000000..57da930 --- /dev/null +++ b/pipeline/tests/test_sync_typesense.py @@ -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)'