Merge pull request #77 from Amateur-Digital-Network/fix/alias-download-atomic-write

fix: write alias downloads atomically and validate checksum before persisting
master
ce5rpy 1 week ago committed by GitHub
commit 89d6cb19b9
No known key found for this signature in database
GPG Key ID: B5690EEEBB952194

@ -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",

@ -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.*")) == []

Loading…
Cancel
Save

Powered by TurnKey Linux.