262 lines
11 KiB
Python
262 lines
11 KiB
Python
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("<script>alert('safe')</script>\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)
|