From aad30baebab0bf5195b9a5313fa5c7a6e82f23e6 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Rodrigo=20P=C3=A9rez?= Date: Sat, 19 Sep 2026 01:38:05 -0300 Subject: [PATCH] fix: write alias downloads atomically and validate checksum before persisting --- .../persistence/alias_loader.py | 51 +++++++++--- .../test_alias_reload_resilience.py | 77 +++++++++++++++++++ 2 files changed, 116 insertions(+), 12 deletions(-) diff --git a/src/adn_server/infrastructure/persistence/alias_loader.py b/src/adn_server/infrastructure/persistence/alias_loader.py index 2952330..1e0250e 100644 --- a/src/adn_server/infrastructure/persistence/alias_loader.py +++ b/src/adn_server/infrastructure/persistence/alias_loader.py @@ -29,8 +29,10 @@ import csv import hashlib import json import logging +import os import shutil import ssl +import threading import time from pathlib import Path from typing import Any @@ -51,6 +53,7 @@ def try_download( first_timeout: float = 30, retry_timeout: float = 10, retry_delay_sec: float = 3, + expected_checksum: str | None = None, ) -> str: """Legacy try_download: download file from url if missing or older than stale_sec. @@ -73,7 +76,7 @@ def try_download( return f"ID ALIAS MAPPER: '{file_name}' is current, not downloaded" data: bytes | None = None - last_error: OSError | None = None + last_error: Exception | None = None for attempt in range(1, max_attempts + 1): timeout = first_timeout if attempt == 1 else retry_timeout try: @@ -82,10 +85,13 @@ def try_download( ctx.verify_mode = ssl.CERT_NONE with urlopen(url, context=ctx, timeout=timeout) as response: data = response.read() + if expected_checksum and hashlib.blake2b(data).hexdigest() != expected_checksum: + raise ValueError("downloaded data does not match expected checksum") last_error = None break - except OSError as e: + except (OSError, ValueError) as e: last_error = e + data = None if attempt < max_attempts: logger.warning( "(ALIAS) ID ALIAS MAPPER: '%s' download attempt %d/%d failed (%s), retrying in %gs", @@ -93,13 +99,20 @@ def try_download( ) if retry_delay_sec: time.sleep(retry_delay_sec) + if isinstance(last_error, ValueError): + return f"ID ALIAS MAPPER: '{file_name}' could not be downloaded, checksum mismatch after {max_attempts} attempts" if last_error is not None: return f"ID ALIAS MAPPER: '{file_name}' could not be downloaded due to an IOError: {last_error}" if not data or data == b"{}": return f"ID ALIAS MAPPER: '{file_name}' file not written because downloaded data is empty" try: full.parent.mkdir(parents=True, exist_ok=True) - full.write_bytes(data) + tmp = full.with_name(f"{full.name}.tmp.{os.getpid()}.{threading.get_ident()}") + try: + tmp.write_bytes(data) + tmp.replace(full) + finally: + tmp.unlink(missing_ok=True) except OSError as e: return f"ID ALIAS mapper '{file_name}' file could not be written: {e}" return f"ID ALIAS MAPPER: '{file_name}' successfully downloaded" @@ -129,6 +142,15 @@ def _blake2bsum(file_path: Path) -> str: return h.hexdigest() +def _atomic_copy(src: Path, dst: Path) -> None: + tmp = dst.with_name(f"{dst.name}.tmp.{os.getpid()}.{threading.get_ident()}") + try: + shutil.copy(src, tmp) + tmp.replace(dst) + finally: + tmp.unlink(missing_ok=True) + + class DefaultAliasLoader(AliasLoader): """Load aliases from JSON files and optional downloads. Legacy mk_aliases.""" @@ -151,17 +173,22 @@ class DefaultAliasLoader(AliasLoader): if aliases.get("CHECKSUM_FILE") and aliases.get("CHECKSUM_URL"): result = try_download(path, aliases["CHECKSUM_FILE"], aliases.get("CHECKSUM_URL", ""), stale_sec) _log_download_result(result) - for key, url_key in [ - ("PEER_FILE", "PEER_URL"), - ("SUBSCRIBER_FILE", "SUBSCRIBER_URL"), - ("TGID_FILE", "TGID_URL"), - ("SERVER_ID_FILE", "SERVER_ID_URL"), + checksums = self._load_checksums(path, aliases.get("CHECKSUM_FILE")) + for key, url_key, checksum_key in [ + ("PEER_FILE", "PEER_URL", "peer_ids"), + ("SUBSCRIBER_FILE", "SUBSCRIBER_URL", "subscriber_ids"), + ("TGID_FILE", "TGID_URL", "talkgroup_ids"), + ("SERVER_ID_FILE", "SERVER_ID_URL", "server_ids"), ]: url = aliases.get(url_key) if url and aliases.get(key): - result = try_download(path, aliases[key], url, stale_sec) + result = try_download( + path, aliases[key], url, stale_sec, + expected_checksum=checksums.get(checksum_key), + ) _log_download_result(result) - checksums = self._load_checksums(path, aliases.get("CHECKSUM_FILE")) + else: + checksums = self._load_checksums(path, aliases.get("CHECKSUM_FILE")) peer_file = aliases.get("PEER_FILE", "peer_ids.json") sub_file = aliases.get("SUBSCRIBER_FILE", "subscriber_ids.json") tgid_file = aliases.get("TGID_FILE", "talkgroup_ids.json") @@ -314,7 +341,7 @@ class DefaultAliasLoader(AliasLoader): logger.warning("(ALIAS) ID ALIAS MAPPER: %s dictionary is empty", name) if loaded_from_primary and full.is_file(): try: - shutil.copy(full, bak) + _atomic_copy(full, bak) except OSError as g: logger.info( "(ALIAS) ID ALIAS MAPPER: couldn't make backup copy of %s file %s", @@ -390,7 +417,7 @@ class DefaultAliasLoader(AliasLoader): logger.warning("(ALIAS) ID ALIAS MAPPER: server_ids dictionary is empty") if loaded_from_primary and full.is_file(): try: - shutil.copy(full, bak) + _atomic_copy(full, bak) except OSError as g: logger.info( "(ALIAS) ID ALIAS MAPPER: couldn't make backup copy of server_ids file %s", diff --git a/tests/infrastructure/test_alias_reload_resilience.py b/tests/infrastructure/test_alias_reload_resilience.py index f7c1fd7..6f17160 100644 --- a/tests/infrastructure/test_alias_reload_resilience.py +++ b/tests/infrastructure/test_alias_reload_resilience.py @@ -2,7 +2,9 @@ from __future__ import annotations +import hashlib import json +import threading from pathlib import Path from unittest.mock import patch @@ -55,6 +57,49 @@ def test_try_download_retries_and_recovers_from_transient_failure(tmp_path: Path assert calls[1] == 10 +def test_try_download_rejects_data_that_fails_checksum(tmp_path: Path) -> None: + file_name = "server_ids.tsv" + full = tmp_path / file_name + full.write_bytes(b"old-good-content") + + with patch( + "adn_server.infrastructure.persistence.alias_loader.urlopen", + return_value=_FakeResponse(b"corrupted-content"), + ): + result = try_download( + tmp_path, + file_name, + "https://example.invalid/server_ids.tsv", + stale_sec=0, + max_attempts=1, + expected_checksum="deadbeef", + ) + assert "checksum mismatch" in result + assert full.read_bytes() == b"old-good-content" + assert list(tmp_path.glob(f"{file_name}.tmp.*")) == [] + + +def test_try_download_accepts_data_matching_checksum(tmp_path: Path) -> None: + file_name = "server_ids.tsv" + payload = b"good-content" + digest = hashlib.blake2b(payload).hexdigest() + + with patch( + "adn_server.infrastructure.persistence.alias_loader.urlopen", + return_value=_FakeResponse(payload), + ): + result = try_download( + tmp_path, + file_name, + "https://example.invalid/server_ids.tsv", + stale_sec=0, + max_attempts=1, + expected_checksum=digest, + ) + assert "successfully downloaded" in result + assert (tmp_path / file_name).read_bytes() == payload + + def test_try_download_gives_up_after_max_attempts(tmp_path: Path) -> None: file_name = "peer_ids.json" calls: list[float] = [] @@ -150,3 +195,35 @@ def test_load_server_tsv_with_backup_uses_bak_when_primary_missing(tmp_path: Pat ) loaded = loader._load_server_tsv_with_backup(tmp_path, file_name, None) assert loaded.get("1234") == "Chile" + + +def test_concurrent_downloads_of_same_file_never_corrupt_it(tmp_path: Path) -> None: + file_name = "peer_ids.json" + payload_a = b"A" * 500_000 + payload_b = b"B" * 700_000 + + def _urlopen_for(payload): + def _fake(url, context=None, timeout=None): + return _FakeResponse(payload) + return _fake + + barrier = threading.Barrier(2) + + def _run(payload): + with patch( + "adn_server.infrastructure.persistence.alias_loader.urlopen", + side_effect=_urlopen_for(payload), + ): + barrier.wait() + try_download(tmp_path, file_name, "https://example.invalid/peer_ids.json", stale_sec=0) + + t1 = threading.Thread(target=_run, args=(payload_a,)) + t2 = threading.Thread(target=_run, args=(payload_b,)) + t1.start() + t2.start() + t1.join() + t2.join() + + final = (tmp_path / file_name).read_bytes() + assert final == payload_a or final == payload_b + assert list(tmp_path.glob(f"{file_name}.tmp.*")) == []