from __future__ import annotations import hashlib import io import json from pathlib import Path import shutil import tarfile import zipfile import pytest from fastapi.testclient import TestClient from app.main import create_app from app.ota.diagnostics import clear_failure_log, persist_failure_log from app.ota.package import OtaPackageError, build_package, extract_payload, inspect_package from app.ota.state import write_last_result from app.ota.versioning import read_feature_updated_at, SoftwareVersion, read_software_version SOURCE_ROOT = Path(__file__).resolve().parents[1] CURRENT_VERSION = read_software_version(SOURCE_ROOT) FEATURE_UPDATED_AT = read_feature_updated_at(SOURCE_ROOT) def following_version(version: SoftwareVersion) -> str: return str(SoftwareVersion(version.major, version.minor, version.patch + 1)) def make_source(tmp_path: Path, version: str = "1.0.2") -> tuple[Path, Path]: source = tmp_path / "source" source.mkdir(parents=True) (source / "VERSION").write_text(version + "\n", encoding="utf-8") (source / "FEATURE_UPDATED_AT").write_text("2026-09-04T12:40+08:00\n", encoding="utf-8") (source / "requirements.txt").write_text("fastapi==0.116.1\n", encoding="utf-8") (source / "module.py").write_text("VALUE = 1\n", encoding="utf-8") (source / "data").mkdir() (source / "data" / "must-not-ship.json").write_text("{}", encoding="utf-8") (source / "__pycache__").mkdir() (source / "__pycache__" / "bad.pyc").write_bytes(b"bad") wheelhouse = tmp_path / "wheelhouse" wheelhouse.mkdir() wheel = wheelhouse / "example.whl" wheel.write_bytes(b"wheel") digest = hashlib.sha256(wheel.read_bytes()).hexdigest() (wheelhouse / "SHA256SUMS").write_text(f"{digest} {wheel.name}\n", encoding="utf-8") return source, wheelhouse def make_package(tmp_path: Path, version: str = "1.0.2") -> Path: source, wheelhouse = make_source(tmp_path, version) package = tmp_path / f"matrix-screen-controller-{version}.ota" build_package( source, wheelhouse, package, version=SoftwareVersion.parse(version), release_notes="test release", ) return package def rewrite_zip(package: Path, replacements: dict[str, bytes]) -> None: with zipfile.ZipFile(package, "r") as archive: values = {name: archive.read(name) for name in archive.namelist()} values.update(replacements) package.unlink() with zipfile.ZipFile(package, "w", compression=zipfile.ZIP_STORED) as archive: for name, value in values.items(): archive.writestr(name, value) def test_strict_software_version_ordering(): assert SoftwareVersion.parse("1.0.1") > SoftwareVersion.parse("1.0.0") assert SoftwareVersion.parse("1.10.0") > SoftwareVersion.parse("1.2.99") for invalid in ("v1.0.0", "1.0", "01.0.0", "1.0.0-beta"): with pytest.raises(ValueError): SoftwareVersion.parse(invalid) def test_feature_updated_at_requires_valid_beijing_minute(tmp_path): (tmp_path / "FEATURE_UPDATED_AT").write_text("2026-09-04T12:40+08:00\n", encoding="utf-8") assert read_feature_updated_at(tmp_path) == "2026-09-04T12:40+08:00" for invalid in ( "2026-09-04T12:40:00+08:00", "2026-09-04T12:40Z", "2026-09-04T04:40+00:00", "2026-02-30T12:40+08:00", ): (tmp_path / "FEATURE_UPDATED_AT").write_text(invalid + "\n", encoding="utf-8") with pytest.raises(ValueError): read_feature_updated_at(tmp_path) (tmp_path / "FEATURE_UPDATED_AT").unlink() with pytest.raises(RuntimeError, match="timestamp file is unreadable"): read_feature_updated_at(tmp_path) def test_full_package_build_inspect_extract_and_exclusions(tmp_path): package = make_package(tmp_path) info = inspect_package(package, current_version=SoftwareVersion.parse("1.0.0")) assert str(info.target_version) == "1.0.2" destination = tmp_path / "unpacked" extract_payload(info, destination) assert (destination / "software" / "module.py").is_file() assert (destination / "software" / "FEATURE_UPDATED_AT").read_text(encoding="utf-8").strip() == "2026-09-04T12:40+08:00" assert (destination / "wheelhouse" / "SHA256SUMS").is_file() assert not (destination / "software" / "data").exists() assert not (destination / "software" / "__pycache__").exists() def test_package_carries_verified_frpc_system_dependency(tmp_path): source, wheelhouse = make_source(tmp_path) system_dependencies = tmp_path / "frpc" system_dependencies.mkdir() binary = system_dependencies / "frpc" binary.write_bytes(b"official-aarch64-frpc") digest = hashlib.sha256(binary.read_bytes()).hexdigest() (system_dependencies / "SHA256SUMS").write_text(f"{digest} frpc\n", encoding="utf-8") package = tmp_path / "with-frpc.ota" build_package( source, wheelhouse, package, version=SoftwareVersion.parse("1.0.2"), release_notes="FRP payload", system_dependencies=system_dependencies, ) destination = tmp_path / "with-frpc-unpacked" extract_payload(inspect_package(package), destination) assert (destination / "software" / "system-dependencies" / "frpc" / "frpc").read_bytes() == binary.read_bytes() assert (destination / "software" / "system-dependencies" / "frpc" / "SHA256SUMS").is_file() def test_same_or_older_version_is_rejected(tmp_path): package = make_package(tmp_path) for current in ("1.0.2", "1.0.3", "2.0.0"): with pytest.raises(OtaPackageError, match="must be newer"): inspect_package(package, current_version=SoftwareVersion.parse(current)) def test_payload_checksum_tampering_is_rejected(tmp_path): package = make_package(tmp_path) with zipfile.ZipFile(package, "r") as archive: payload = archive.read("payload.tar.gz") + b"tampered" rewrite_zip(package, {"payload.tar.gz": payload}) with pytest.raises(OtaPackageError, match="checksum or size"): inspect_package(package) def test_unsafe_tar_path_is_rejected_during_extract(tmp_path): package = make_package(tmp_path) with zipfile.ZipFile(package, "r") as archive: manifest = json.loads(archive.read("manifest.json")) buffer = io.BytesIO() with tarfile.open(fileobj=buffer, mode="w:gz") as tar: content = b"escape" member = tarfile.TarInfo("software/../../escape") member.size = len(content) tar.addfile(member, io.BytesIO(content)) payload = buffer.getvalue() manifest["payload_sha256"] = hashlib.sha256(payload).hexdigest() manifest["payload_bytes"] = len(payload) manifest["expanded_bytes"] = len(b"escape") rewrite_zip(package, { "manifest.json": (json.dumps(manifest) + "\n").encode(), "payload.tar.gz": payload, }) info = inspect_package(package) with pytest.raises(OtaPackageError, match="unsafe OTA payload path"): extract_payload(info, tmp_path / "unsafe") assert not (tmp_path / "escape").exists() def test_ota_api_exposes_software_version_and_accepts_one_new_package(tmp_path): project = tmp_path / "project" project.mkdir() target = following_version(CURRENT_VERSION) package = make_package(tmp_path / "package", target) app = create_app( project_root=project, driver_kind="mock", ota_worker_starter=lambda: None, ota_staging_root=tmp_path / "ota-staging", ) with TestClient(app) as client: status = client.get("/api/status").json() assert status["service"]["software_version"] == str(CURRENT_VERSION) assert status["service"]["feature_updated_at"] == FEATURE_UPDATED_AT initial = client.get("/api/ota/status").json() assert initial["current_version"] == str(CURRENT_VERSION) assert initial["feature_updated_at"] == FEATURE_UPDATED_AT assert initial["failure_log_available"] is False assert initial["failure_log_bytes"] == 0 assert client.get("/api/ota/failure-log").status_code == 404 response = client.post( f"/api/ota/update?filename=matrix-screen-controller-{target}.ota", content=package.read_bytes(), headers={"Content-Type": "application/octet-stream"}, ) assert response.status_code == 202 assert response.json()["job"]["target_version"] == target assert client.get("/api/ota/status").json()["active"] is True locked = client.put("/api/config", json={"brightness": 30}) assert locked.status_code == 423 concurrent = client.post( f"/api/ota/update?filename=matrix-screen-controller-{target}.ota", content=package.read_bytes(), headers={"Content-Type": "application/octet-stream"}, ) assert concurrent.status_code == 409 def test_ota_api_rejects_same_version_before_starting_worker(tmp_path): project = tmp_path / "project" project.mkdir() package = make_package(tmp_path / "package", str(CURRENT_VERSION)) starts = [] app = create_app( project_root=project, driver_kind="mock", ota_worker_starter=lambda: starts.append(True), ota_staging_root=tmp_path / "ota-staging", ) with TestClient(app) as client: response = client.post( f"/api/ota/update?filename=matrix-screen-controller-{CURRENT_VERSION}.ota", content=package.read_bytes(), headers={"Content-Type": "application/octet-stream"}, ) assert response.status_code == 409 assert starts == [] def test_ota_failure_log_api_requires_failed_result_and_disables_caching(tmp_path): project = tmp_path / "project" project.mkdir() app = create_app(project_root=project, driver_kind="mock") data_root = app.state.ota_manager.data_root runtime_log = tmp_path / "ota-worker.log" runtime_log.write_text("\npytest failed\n", encoding="utf-8") persist_failure_log(data_root, runtime_log) with TestClient(app) as client: assert client.get("/api/ota/failure-log").status_code == 404 write_last_result(data_root, {"status": "failed", "target_version": "9.9.9"}) status = client.get("/api/ota/status").json() assert status["failure_log_available"] is True assert status["failure_log_bytes"] == len(runtime_log.read_bytes()) response = client.get("/api/ota/failure-log") assert response.status_code == 200 assert response.headers["cache-control"] == "no-store" assert response.headers["content-type"].startswith("text/plain") assert response.content == runtime_log.read_bytes() write_last_result(data_root, {"status": "success", "target_version": "9.9.9"}) assert client.get("/api/ota/status").json()["failure_log_available"] is False assert client.get("/api/ota/failure-log").status_code == 404 clear_failure_log(data_root)