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()