Files

499 lines
19 KiB
Python

import base64
import json
from io import BytesIO
import pytest
from fastapi.testclient import TestClient
from PIL import Image
import app.persistence as persistence
from app.main import create_app
from app.demo_library import DEMO_ANIMATION_ID, DEMO_STATIC_ID, DEMO_RESTRICTED_MESSAGE
from app.display.startup_indicator import render_startup_smile
from app.templates.store import (
TEMPLATE_SCHEMA_VERSION,
TemplateError,
TemplateStorageFullError,
TemplateStore,
TemplateValidationError,
)
def scene(background=(1, 2, 3), text="TOP"):
pixels = bytes(background) * (64 * 64)
return {
"version": 1,
"width": 64,
"height": 64,
"pixelRgb": base64.b64encode(pixels).decode("ascii"),
"elements": [{
"id": "text-1",
"type": "text",
"text": text,
"font": "default",
"size": 12,
"x": 4,
"y": 24,
"align": "left",
"color": "#FFFFFF",
}],
}
def test_deleting_default_template_immediately_falls_back_to_demo_animation(tmp_path):
app = create_app(project_root=tmp_path, driver_kind="mock")
with TestClient(app) as client:
created = client.post(
"/api/templates",
json={"name": "即将删除的默认图", "scene": scene(text="DEFAULT")},
).json()
selected = client.put(
"/api/display/default-content",
json={"type": "template", "id": created["id"]},
)
assert selected.status_code == 200
deleted = client.delete(
f"/api/templates/{created['id']}",
headers={"If-Match": f'"{created["revision"]}"'},
)
assert deleted.status_code == 200
assert client.get("/api/display/default-content").json()["id"] == DEMO_ANIMATION_ID
state = client.get("/api/status").json()["state"]
assert state["animation_active"] is True
assert state["current_content"]["id"] == DEMO_ANIMATION_ID
def test_template_play_is_temporary_revision_checked_and_supports_demo(tmp_path):
app = create_app(project_root=tmp_path, driver_kind="mock")
with TestClient(app) as client:
created = client.post(
"/api/templates",
json={"name": "临时播放", "scene": scene(background=(12, 34, 56), text="PLAY")},
).json()
default_before = client.get("/api/display/default-content").json()
played = client.post(
f"/api/templates/{created['id']}/play",
headers={"If-Match": f'"{created["revision"]}"'},
json={},
)
assert played.status_code == 200
assert played.json() == {
"ok": True,
"template_id": created["id"],
"revision": created["revision"],
}
state = client.get("/api/status").json()["state"]
assert state["mode"] == "template"
assert state["animation_active"] is False
assert state["current_content"] == {
"category": "template",
"id": created["id"],
"name": "临时播放",
}
assert client.get("/api/display/default-content").json() == default_before
current_frame = Image.open(BytesIO(client.get("/api/display/current-frame").content)).convert("RGB")
thumbnail = Image.open(BytesIO(client.get(created["thumbnail_url"]).content)).convert("RGB")
assert current_frame.tobytes() == thumbnail.tobytes()
before_rejected = current_frame.tobytes()
assert client.post(f"/api/templates/{created['id']}/play", json={}).status_code == 428
assert client.post(
f"/api/templates/{created['id']}/play",
headers={"If-Match": '"stale"'},
json={},
).status_code == 409
assert Image.open(BytesIO(client.get("/api/display/current-frame").content)).convert("RGB").tobytes() == before_rejected
assert client.post(
"/api/templates/00000000-0000-4000-8000-000000000099/play",
headers={"If-Match": '"missing"'},
json={},
).status_code == 404
assert client.post(
"/api/templates/not-a-uuid/play",
headers={"If-Match": '"missing"'},
json={},
).status_code == 404
demo = client.get("/api/templates?include_demo=true").json()["templates"][0]
demo_play = client.post(
f"/api/templates/{DEMO_STATIC_ID}/play",
headers={"If-Match": f'"{demo["revision"]}"'},
json={},
)
assert demo_play.status_code == 200
demo_state = client.get("/api/status").json()["state"]
assert demo_state["current_content"]["id"] == DEMO_STATIC_ID
assert client.get("/api/display/default-content").json() == default_before
def test_template_crud_copy_persistence_and_thumbnail_cleanup(tmp_path):
app = create_app(project_root=tmp_path, driver_kind="mock")
client = TestClient(app)
created_response = client.post("/api/templates", json={"name": "中文模板", "scene": scene()})
assert created_response.status_code == 201
created = created_response.json()
assert created["name"] == "中文模板"
assert created["size_bytes"] > 0
assert "schema_version" not in created
old_url = created["thumbnail_url"]
old_thumbnail_name = f"{created['id']}-{created['digest']}.png"
persisted_record = json.loads(
(
tmp_path
/ "data"
/ "templates"
/ "records"
/ f"{created['id']}.json"
).read_text(encoding="utf-8")
)
assert persisted_record["schema_version"] == TEMPLATE_SCHEMA_VERSION
duplicate_name = client.post("/api/templates", json={"name": "中文模板", "scene": scene()})
assert duplicate_name.status_code == 409
duplicate_case = client.post("/api/templates", json={"name": "中文模板 ", "scene": scene()})
assert duplicate_case.status_code == 409
thumbnail_response = client.get(old_url)
assert thumbnail_response.status_code == 200
image = Image.open(BytesIO(thumbnail_response.content))
assert image.size == (64, 64)
assert image.mode == "RGB"
assert image.getpixel((0, 0)) == (1, 2, 3)
assert len(set(image.getdata())) > 1
updated_response = client.put(
f"/api/templates/{created['id']}",
headers={"If-Match": f'"{created["revision"]}"'},
json={"scene": scene(background=(9, 8, 7), text="NEW")},
)
assert updated_response.status_code == 200
updated = updated_response.json()
assert updated["digest"] != created["digest"]
assert updated["thumbnail_url"] != old_url
assert client.get(old_url).status_code == 404
assert not (tmp_path / "data" / "templates" / "thumbnails" / old_thumbnail_name).exists()
updated_image = Image.open(BytesIO(client.get(updated["thumbnail_url"]).content))
assert updated_image.getpixel((0, 0)) == (9, 8, 7)
renamed = client.patch(
f"/api/templates/{created['id']}",
headers={"If-Match": f'"{updated["revision"]}"'},
json={"name": "新名字"},
)
assert renamed.status_code == 200
renamed_template = renamed.json()
assert renamed_template["digest"] == updated["digest"]
copied = client.post(f"/api/templates/{created['id']}/copy", json={})
assert copied.status_code == 201
assert copied.json()["name"] == "新名字 - 副本"
copied_again = client.post(f"/api/templates/{created['id']}/copy", json={})
assert copied_again.json()["name"] == "新名字 - 副本 2"
restarted = TestClient(create_app(project_root=tmp_path, driver_kind="mock"))
listing = restarted.get("/api/templates").json()
assert [item["name"] for item in listing["templates"]] == [
"新名字 - 副本 2", "新名字 - 副本", "新名字",
]
assert listing["storage"]["templates_bytes"] == sum(item["size_bytes"] for item in listing["templates"])
assert listing["storage"]["free_bytes"] > 0
delete_headers = {"If-Match": f'"{renamed_template["revision"]}"'}
assert client.delete(f"/api/templates/{created['id']}", headers=delete_headers).json() == {"ok": True}
assert client.get(f"/api/templates/{created['id']}").status_code == 404
assert client.delete(f"/api/templates/{created['id']}", headers=delete_headers).status_code == 404
def test_demo_static_template_is_first_read_only_and_copies_to_normal_content(tmp_path):
client = TestClient(create_app(project_root=tmp_path, driver_kind="mock"))
assert client.get("/api/templates").json()["templates"] == []
listing = client.get("/api/templates?include_demo=true").json()
demo = listing["templates"][0]
assert demo["id"] == DEMO_STATIC_ID
assert demo["name"] == "演示静态图"
assert demo["demo_order"] == 1
assert demo["read_only"] is True
assert demo["size_bytes"] == 0
assert listing["storage"]["templates_bytes"] == 0
detail = client.get(f"/api/templates/{DEMO_STATIC_ID}").json()
assert base64.b64decode(detail["scene"]["pixelRgb"]) == render_startup_smile().tobytes()
thumbnail = client.get(demo["thumbnail_url"])
assert thumbnail.status_code == 200
assert Image.open(BytesIO(thumbnail.content)).tobytes() == render_startup_smile().tobytes()
for method, payload in (("put", {"scene": scene()}), ("patch", {"name": "不能改"})):
response = getattr(client, method)(
f"/api/templates/{DEMO_STATIC_ID}",
headers={"If-Match": f'"{demo["revision"]}"'},
json=payload,
)
assert response.status_code == 409
assert response.json()["detail"] == DEMO_RESTRICTED_MESSAGE
removed = client.delete(
f"/api/templates/{DEMO_STATIC_ID}",
headers={"If-Match": f'"{demo["revision"]}"'},
)
assert removed.status_code == 409
copied = client.post("/api/library/copy", json={
"source_type": "template",
"source_id": DEMO_STATIC_ID,
"source_revision": demo["revision"],
"destination_type": "static",
})
assert copied.status_code == 201
copy_item = copied.json()["item"]
assert copy_item["id"] != DEMO_STATIC_ID
assert copy_item["name"] == "演示静态图 - 副本"
assert "read_only" not in copy_item
renamed = client.patch(
f"/api/templates/{copy_item['id']}",
headers={"If-Match": f'"{copy_item["revision"]}"'},
json={"name": "我的笑脸"},
)
assert renamed.status_code == 200
assert len(list((tmp_path / "data" / "templates" / "records").glob("*.json"))) == 1
def test_template_preserves_multilingual_text(tmp_path):
client = TestClient(create_app(project_root=tmp_path, driver_kind="mock"))
multilingual = "你好 25°C"
created = client.post(
"/api/templates",
json={"name": "多语言", "scene": scene(text=multilingual)},
)
assert created.status_code == 201
stored = client.get(f"/api/templates/{created.json()['id']}")
assert stored.status_code == 200
assert stored.json()["scene"]["layers"][1]["elements"][0]["text"] == multilingual
def test_template_mutations_require_current_revision(tmp_path):
client = TestClient(create_app(project_root=tmp_path, driver_kind="mock"))
created = client.post("/api/templates", json={"name": "shared", "scene": scene()}).json()
path = f"/api/templates/{created['id']}"
original_revision = created["revision"]
assert client.put(path, json={"scene": scene(text="missing")}).status_code == 428
first = client.put(
path,
headers={"If-Match": f'"{original_revision}"'},
json={"scene": scene(text="first")},
)
assert first.status_code == 200
current = first.json()
stale_headers = {"If-Match": f'"{original_revision}"'}
assert client.put(path, headers=stale_headers, json={"scene": scene(text="stale")}).status_code == 409
assert client.patch(path, headers=stale_headers, json={"name": "stale name"}).status_code == 409
assert client.delete(path, headers=stale_headers).status_code == 409
preserved = client.get(path).json()
assert preserved["name"] == "shared"
assert preserved["scene"]["layers"][1]["elements"][0]["text"] == "first"
assert preserved["revision"] == current["revision"]
def test_template_validation_and_status_codes(tmp_path, monkeypatch):
app = create_app(project_root=tmp_path, driver_kind="mock")
client = TestClient(app)
broken = scene()
broken["pixelRgb"] = "bad"
assert client.post("/api/templates", json={"name": "bad", "scene": broken}).status_code == 422
duplicate_ids = scene()
duplicate_ids["elements"].append(dict(duplicate_ids["elements"][0]))
assert client.post("/api/templates", json={"name": "bad ids", "scene": duplicate_ids}).status_code == 422
assert client.get("/api/templates/not-a-uuid").status_code == 404
def full(*_args, **_kwargs):
raise TemplateStorageFullError("not enough disk space to save template")
monkeypatch.setattr(app.state.template_store, "create", full)
assert client.post("/api/templates", json={"name": "full", "scene": scene()}).status_code == 507
def test_reconcile_removes_only_strict_thumbnail_artifacts(tmp_path):
store = TemplateStore(tmp_path / "data")
created = store.create("keep", scene())
thumbnail_dir = tmp_path / "data" / "templates" / "thumbnails"
orphan = thumbnail_dir / "11111111-1111-1111-1111-111111111111-0123456789abcdef.png"
unrelated = thumbnail_dir / "user-picture.png"
temporary = thumbnail_dir / ".partial.png.abc.tmp"
orphan.write_bytes(b"orphan")
unrelated.write_bytes(b"unrelated")
temporary.write_bytes(b"temporary")
TemplateStore(tmp_path / "data")
assert not orphan.exists()
assert not temporary.exists()
assert unrelated.read_bytes() == b"unrelated"
assert store.thumbnail_path(created["id"], created["digest"]).exists()
def test_legacy_template_migrates_once_without_changing_revision_or_digest(tmp_path):
data_dir = tmp_path / "data"
original_store = TemplateStore(data_dir)
created = original_store.create("legacy", scene())
record_path = (
data_dir
/ "templates"
/ "records"
/ f"{created['id']}.json"
)
legacy = json.loads(record_path.read_text(encoding="utf-8"))
legacy.pop("schema_version")
record_path.write_text(
json.dumps(legacy, ensure_ascii=False, indent=2) + "\n",
encoding="utf-8",
)
migrated_store = TemplateStore(data_dir)
migrated = migrated_store.get(created["id"])
migrated_bytes = record_path.read_bytes()
assert migrated["revision"] == created["revision"]
assert migrated["digest"] == created["digest"]
assert "schema_version" not in migrated
assert json.loads(migrated_bytes)["schema_version"] == TEMPLATE_SCHEMA_VERSION
TemplateStore(data_dir)
assert record_path.read_bytes() == migrated_bytes
@pytest.mark.parametrize("failure_kind", ["future", "corrupt", "unknown"])
def test_invalid_template_record_blocks_startup_without_deleting_any_thumbnail(
tmp_path,
failure_kind,
):
data_dir = tmp_path / "data"
initial = TemplateStore(data_dir)
created = initial.create("keep", scene())
record_path = (
data_dir
/ "templates"
/ "records"
/ f"{created['id']}.json"
)
thumbnail_path = initial.thumbnail_path(created["id"], created["digest"])
thumbnail_bytes = thumbnail_path.read_bytes()
orphan = (
data_dir
/ "templates"
/ "thumbnails"
/ "11111111-1111-1111-1111-111111111111-0123456789abcdef.png"
)
orphan.write_bytes(b"orphan")
if failure_kind == "corrupt":
original = b"{bad json"
else:
document = json.loads(record_path.read_text(encoding="utf-8"))
if failure_kind == "future":
document["schema_version"] = TEMPLATE_SCHEMA_VERSION + 1
else:
document["unexpected"] = True
original = (json.dumps(document, separators=(",", ":")) + "\n").encode("utf-8")
record_path.write_bytes(original)
with pytest.raises(TemplateValidationError, match="invalid template record"):
TemplateStore(data_dir)
assert record_path.read_bytes() == original
assert thumbnail_path.read_bytes() == thumbnail_bytes
assert orphan.read_bytes() == b"orphan"
def test_template_preflight_validates_every_record_before_migrating_any(tmp_path):
data_dir = tmp_path / "data"
initial = TemplateStore(data_dir)
initial.create("one", scene(text="one"))
initial.create("two", scene(text="two"))
record_paths = sorted((data_dir / "templates" / "records").glob("*.json"))
legacy_path, invalid_path = record_paths
legacy_document = json.loads(legacy_path.read_text(encoding="utf-8"))
legacy_document.pop("schema_version")
legacy_bytes = (
json.dumps(legacy_document, ensure_ascii=False, indent=2) + "\n"
).encode("utf-8")
legacy_path.write_bytes(legacy_bytes)
invalid_path.write_bytes(b"{bad json")
with pytest.raises(TemplateValidationError):
TemplateStore(data_dir)
assert legacy_path.read_bytes() == legacy_bytes
def test_template_migration_write_failure_rolls_back_every_record(tmp_path, monkeypatch):
data_dir = tmp_path / "data"
initial = TemplateStore(data_dir)
initial.create("one", scene(text="one"))
initial.create("two", scene(text="two"))
record_paths = sorted((data_dir / "templates" / "records").glob("*.json"))
legacy_bytes = {}
for path in record_paths:
document = json.loads(path.read_text(encoding="utf-8"))
document.pop("schema_version")
content = (json.dumps(document, ensure_ascii=False, indent=2) + "\n").encode("utf-8")
path.write_bytes(content)
legacy_bytes[path] = content
real_replace = persistence.os.replace
replace_calls = 0
def fail_second_replace(source, target):
nonlocal replace_calls
replace_calls += 1
if replace_calls == 2:
raise OSError("second migration replace failed")
return real_replace(source, target)
monkeypatch.setattr(persistence.os, "replace", fail_second_replace)
with pytest.raises(TemplateError, match="second migration replace failed"):
TemplateStore(data_dir)
assert {path: path.read_bytes() for path in record_paths} == legacy_bytes
assert list((data_dir / "templates" / "records").glob(".*.tmp")) == []
@pytest.mark.parametrize(
("path_parts", "value"),
[
(("scene", "version"), True),
(("scene", "width"), 64.0),
(("scene", "layers", 1, "elements", 0, "size"), True),
],
)
def test_current_template_schema_rejects_noncanonical_json_types(
tmp_path,
path_parts,
value,
):
data_dir = tmp_path / "data"
initial = TemplateStore(data_dir)
created = initial.create("strict", scene())
record_path = (
data_dir
/ "templates"
/ "records"
/ f"{created['id']}.json"
)
document = json.loads(record_path.read_text(encoding="utf-8"))
target = document
for part in path_parts[:-1]:
target = target[part]
target[path_parts[-1]] = value
original = (json.dumps(document, separators=(",", ":")) + "\n").encode("utf-8")
record_path.write_bytes(original)
with pytest.raises(TemplateValidationError, match="integer|fields or dimensions"):
TemplateStore(data_dir)
assert record_path.read_bytes() == original