Files
templates/python/tests/test_firmware_database.py

193 lines
9.0 KiB
Python

from __future__ import annotations
import base64
import hashlib
import io
import json
import tempfile
import unittest
from dataclasses import replace
from pathlib import Path
from unittest.mock import Mock
from urllib.error import HTTPError
from setprotocol.firmware_catalog import FirmwareRelease
from setprotocol.firmware_database import (
Credentials, FirmwareDatabase, GiteaFirmwarePublisher, GiteaRepository, HttpsClient,
)
from setprotocol.firmware_publish import FirmwarePublication
class Response(io.BytesIO):
def __init__(self, data, headers=None):
super().__init__(data)
self.headers = headers or {}
class Server:
"""In-memory HTTP boundary; no real credentials, network or repository writes."""
def __init__(self, *, bad_image=False, conflict=False, missing_release=False):
self.manifest = {"windows": {"versionCode": 42}, "firmware": {"releases": []}}
self.assets = {}
self.calls = []
self.bad_image = bad_image
self.conflict = conflict
self.missing_release = missing_release
def open(self, url, *, method="GET", data=None, **kwargs):
self.calls.append((method, url))
if self.missing_release and "/releases/tags/" in url:
raise HTTPError(url, 404, "Not found", {}, None)
if "/releases/download/" in url:
return Response(b"wrong" if self.bad_image else self.assets[url.rsplit("/", 1)[1]])
if "/raw/branch/" in url:
result = self.manifest
elif "/contents/" in url:
if method == "PUT":
if self.conflict:
raise HTTPError(url, 409, "Conflict", {}, None)
payload = json.loads(data)
assert payload["sha"] == "revision-1"
self.manifest = json.loads(base64.b64decode(payload["content"]))
result = {"sha": "revision-1", "content": base64.b64encode(json.dumps(self.manifest).encode()).decode()}
elif "/assets" in url:
if method == "POST":
self.assets[url.split("?name=", 1)[1]] = data
result = {"id": 2}
else:
result = [{"name": name, "id": index} for index, name in enumerate(self.assets)]
else:
result = {"id": 1}
return Response(json.dumps(result).encode())
class FirmwareDatabaseTests(unittest.TestCase):
def setUp(self):
self.temporary = tempfile.TemporaryDirectory()
self.addCleanup(self.temporary.cleanup)
self.root = Path(self.temporary.name)
self.image = self.root / "image.bin"
self.image.write_bytes(b"test firmware")
self.repository = GiteaRepository("https://git.example", "team", "images")
self.publication = FirmwarePublication(self.image, "DEVICE", "1.2.3", 0x10203, "can")
def test_publish_read_and_download_round_trip_and_idempotency(self):
server = Server()
publisher = GiteaFirmwarePublisher(self.repository, client=server)
entry = publisher.publish(self.publication)
self.assertEqual(server.manifest["windows"], {"versionCode": 42})
self.assertEqual(server.manifest["firmware"]["catalogVersion"], 1)
first_commit = next(i for i, (method, _) in enumerate(server.calls) if method == "PUT")
first_verification = next(i for i, (_, url) in enumerate(server.calls) if "/releases/download/" in url)
self.assertLess(first_verification, first_commit)
publisher.publish(self.publication)
self.assertEqual(sum(method == "PUT" for method, _ in server.calls), 1)
self.assertEqual(len(server.assets), 1)
database = FirmwareDatabase(self.repository.manifest_url, self.root / "cache", client=server)
rows = database.read_catalog(product="device", transport="can")
self.assertEqual(len(rows), 1)
self.assertEqual(database.read_catalog(transport="rs485"), [])
progress = []
downloaded = database.download(rows[0], progress.append)
self.assertEqual(downloaded.read_bytes(), self.image.read_bytes())
self.assertEqual(progress[-1], 100)
count = len(server.calls)
self.assertEqual(database.download(rows[0]), downloaded)
self.assertEqual(len(server.calls), count)
self.assertEqual(entry["fileName"], "image.bin")
def test_failed_image_verification_never_writes_catalog(self):
server = Server(bad_image=True)
with self.assertRaisesRegex(ValueError, "SHA-256"):
GiteaFirmwarePublisher(self.repository, client=server).publish(self.publication)
self.assertFalse(any(method == "PUT" for method, _ in server.calls))
def test_rebuilt_image_keeps_old_asset_and_manifest_on_conflict(self):
server = Server()
publisher = GiteaFirmwarePublisher(self.repository, client=server)
publisher.publish(self.publication)
before = json.dumps(server.manifest)
self.image.write_bytes(b"new firmware")
server.conflict = True
with self.assertRaises(HTTPError) as caught:
publisher.publish(self.publication)
caught.exception.close()
self.assertEqual(json.dumps(server.manifest), before)
self.assertEqual(len(server.assets), 2)
self.assertFalse(any(method == "DELETE" for method, _ in server.calls))
def test_preflight_is_offline(self):
server = Server()
GiteaFirmwarePublisher(self.repository, client=server).preflight(self.publication)
self.assertEqual(server.calls, [])
def test_missing_release_is_created(self):
server = Server(missing_release=True)
GiteaFirmwarePublisher(self.repository, client=server).publish(self.publication)
self.assertTrue(any(method == "POST" and url.endswith("/releases")
for method, url in server.calls))
def test_download_size_limit_and_catalog_readback_failure(self):
from setprotocol.firmware_publish import MAX_FIRMWARE_BYTES
client = Mock()
client.open.return_value = Response(b"", {"Content-Length": str(MAX_FIRMWARE_BYTES + 1)})
database = FirmwareDatabase(self.repository.manifest_url, self.root / "cache", client=client)
release = FirmwareRelease("D", "1", 1, "https://git.example/a.bin", "ab" * 32, "a.bin")
with self.assertRaisesRegex(ValueError, "maximum size"):
database.download(release)
self.assertEqual(list((self.root / "cache").rglob("*.part")), [])
server = Server()
original_open = server.open
def ignore_manifest_write(url, **kwargs):
if kwargs.get("method") == "PUT":
return Response(b"{}")
return original_open(url, **kwargs)
server.open = ignore_manifest_write
with self.assertRaisesRegex(RuntimeError, "readback"):
GiteaFirmwarePublisher(self.repository, client=server).publish(self.publication)
def test_bad_download_cleans_staging_and_rejects_unsafe_filename(self):
release = FirmwareRelease("D", "1", 1, "https://git.example/a.bin",
hashlib.sha256(b"good").hexdigest(), "a.bin")
client = Mock()
client.open.return_value = Response(b"bad")
database = FirmwareDatabase(self.repository.manifest_url, self.root / "cache", client=client)
with self.assertRaisesRegex(ValueError, "SHA-256"):
database.download(release)
self.assertEqual(list((self.root / "cache").rglob("*.part")), [])
with self.assertRaises(ValueError):
database.download(replace(release, file_name="../escaped.bin"))
self.assertEqual(client.open.call_count, 1)
def test_https_credentials_do_not_follow_cross_origin_redirect(self):
client = HttpsClient("https://git.example", Credentials("test", "secret"))
client.opener = Mock()
client.opener.open.side_effect = [
HTTPError("https://git.example/a", 302, "Redirect", {"Location": "https://cdn.example/a"}, None),
Response(b"image"),
]
client.open("https://git.example/a").close()
first, second = [call.args[0] for call in client.opener.open.call_args_list]
self.assertIn("Authorization", first.headers)
self.assertNotIn("Authorization", second.headers)
def test_http_redirect_and_write_redirect_are_rejected(self):
for method, destination in (("GET", "http://git.example/a"), ("POST", "https://cdn.example/a")):
with self.subTest(method=method):
client = HttpsClient("https://git.example")
client.opener = Mock()
client.opener.open.side_effect = HTTPError(
"https://git.example/a", 302, "Redirect", {"Location": destination}, None)
with self.assertRaises((ValueError, HTTPError)) as caught:
client.open("https://git.example/a", method=method)
if isinstance(caught.exception, HTTPError):
caught.exception.close()
self.assertEqual(client.opener.open.call_count, 1)
if __name__ == "__main__":
unittest.main()