193 lines
9.0 KiB
Python
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()
|