fix: preserve countdown dates and restore ids
ci / gitleaks (push) Successful in 7s
ci / docker (push) Successful in 3m19s

This commit is contained in:
2026-09-09 11:45:26 +08:00
parent 010bfd3653
commit e163f6f368
2 changed files with 100 additions and 10 deletions
+23 -10
View File
@@ -5,7 +5,7 @@ import json
import re import re
from datetime import UTC, date, datetime, time, timedelta from datetime import UTC, date, datetime, time, timedelta
from pathlib import Path from pathlib import Path
from uuid import UUID from uuid import UUID, uuid5
from zoneinfo import ZoneInfo from zoneinfo import ZoneInfo
from fastapi import APIRouter, Depends, File, HTTPException, Query, Response, UploadFile from fastapi import APIRouter, Depends, File, HTTPException, Query, Response, UploadFile
@@ -461,6 +461,9 @@ async def edit_countdown(countdown_id: UUID, payload: CountdownUpdate, user: Use
"lunar_day": values.get("lunar_day", row.lunar_day), "lunar_day": values.get("lunar_day", row.lunar_day),
"ignore_year": values.get("ignore_year", row.ignore_year), "ignore_year": values.get("ignore_year", row.ignore_year),
} }
calendar_fields_changed = bool(
{"event_date", "calendar_mode", "lunar_month", "lunar_day"} & values.keys()
)
if combined["calendar_mode"] == "solar": if combined["calendar_mode"] == "solar":
if combined["lunar_month"] is not None or combined["lunar_day"] is not None: if combined["lunar_month"] is not None or combined["lunar_day"] is not None:
if "calendar_mode" not in values: if "calendar_mode" not in values:
@@ -469,12 +472,17 @@ async def edit_countdown(countdown_id: UUID, payload: CountdownUpdate, user: Use
else: else:
if combined["lunar_month"] is None or combined["lunar_day"] is None: if combined["lunar_month"] is None or combined["lunar_day"] is None:
raise HTTPException(422, "农历倒数日需要月份和日期") raise HTTPException(422, "农历倒数日需要月份和日期")
converted = lunar_to_solar_safe( if calendar_fields_changed:
combined["event_date"].year, combined["lunar_month"], combined["lunar_day"] if "event_date" in values:
) lunar_year = combined["event_date"].year
if converted is None: else:
raise HTTPException(422, "所选年份不存在该农历日期") lunar_year, _, _ = solar_to_lunar_parts(row.event_date)
combined["event_date"] = converted converted = lunar_to_solar_safe(
lunar_year, combined["lunar_month"], combined["lunar_day"]
)
if converted is None:
raise HTTPException(422, "所选年份不存在该农历日期")
combined["event_date"] = converted
values.update(combined) values.update(combined)
for key, value in values.items(): for key, value in values.items():
setattr(row, key, value) setattr(row, key, value)
@@ -982,13 +990,17 @@ async def restore_json(payload: dict, mode: str = Query("merge", pattern="^(merg
) )
db.add(row) db.add(row)
existing_countdown_ids = set((await db.scalars(select(Countdown.id).where(Countdown.user_id == user.id))).all()) existing_countdown_ids = set((await db.scalars(select(Countdown.id).where(Countdown.user_id == user.id))).all())
occupied_countdown_ids = dict((await db.execute(select(Countdown.id, Countdown.user_id))).all())
has_pinned_countdown = bool(await db.scalar(select(Countdown.id).where(Countdown.user_id == user.id, Countdown.pinned.is_(True)))) has_pinned_countdown = bool(await db.scalar(select(Countdown.id).where(Countdown.user_id == user.id, Countdown.pinned.is_(True))))
for raw in payload.get("countdowns", []): for raw in payload.get("countdowns", []):
try: try:
source_id = UUID(raw["id"]) source_id = UUID(raw["id"])
except (KeyError, TypeError, ValueError) as exc: except (KeyError, TypeError, ValueError) as exc:
raise HTTPException(422, "无效的倒数日备份 ID") from exc raise HTTPException(422, "无效的倒数日备份 ID") from exc
if mode == "merge" and source_id in existing_countdown_ids: row_id = source_id
while row_id in occupied_countdown_ids and occupied_countdown_ids[row_id] != user.id:
row_id = uuid5(user.id, str(row_id))
if mode == "merge" and row_id in existing_countdown_ids:
continue continue
try: try:
item = CountdownInput( item = CountdownInput(
@@ -1009,7 +1021,7 @@ async def restore_json(payload: dict, mode: str = Query("merge", pattern="^(merg
except (KeyError, TypeError, ValueError) as exc: except (KeyError, TypeError, ValueError) as exc:
raise HTTPException(422, "无效的倒数日备份数据") from exc raise HTTPException(422, "无效的倒数日备份数据") from exc
row = Countdown( row = Countdown(
id=source_id, id=row_id,
user_id=user.id, user_id=user.id,
**item.model_dump(), **item.model_dump(),
archived_at=datetime.fromisoformat(raw["archived_at"]) if raw.get("archived_at") else None, archived_at=datetime.fromisoformat(raw["archived_at"]) if raw.get("archived_at") else None,
@@ -1020,7 +1032,8 @@ async def restore_json(payload: dict, mode: str = Query("merge", pattern="^(merg
row.pinned = False row.pinned = False
has_pinned_countdown = has_pinned_countdown or row.pinned has_pinned_countdown = has_pinned_countdown or row.pinned
db.add(row) db.add(row)
existing_countdown_ids.add(source_id) existing_countdown_ids.add(row_id)
occupied_countdown_ids[row_id] = user.id
audit(db, user.id, "restore", "backup", count=restored, mode=mode) audit(db, user.id, "restore", "backup", count=restored, mode=mode)
await db.commit() await db.commit()
return {"restored": restored, "mode": mode} return {"restored": restored, "mode": mode}
+77
View File
@@ -84,6 +84,35 @@ def test_lunar_countdown_crud_validation_and_display(client):
assert create_countdown(client, calendar_mode="solar", lunar_month=8, lunar_day=4).status_code == 422 assert create_countdown(client, calendar_mode="solar", lunar_month=8, lunar_day=4).status_code == 422
def test_editing_only_lunar_countdown_title_preserves_original_lunar_year(client):
boot(client)
created = create_countdown(
client,
title="农历纪念日",
event_date="2001-01-01",
calendar_mode="lunar",
lunar_month=12,
lunar_day=20,
)
assert created.status_code == 201
before = created.json()
assert before["event_date"] == "2002-02-01"
assert before["lunar_year"] == 2001
updated = client.patch(
f"/api/v1/countdowns/{before['id']}",
json={"title": "只修改标题"},
)
assert updated.status_code == 200
after = updated.json()
assert after["title"] == "只修改标题"
assert after["event_date"] == before["event_date"]
assert after["lunar_year"] == before["lunar_year"]
assert after["lunar_month"] == before["lunar_month"]
assert after["lunar_day"] == before["lunar_day"]
def test_countdown_crud_single_pin_archive_restore_and_purge(client): def test_countdown_crud_single_pin_archive_restore_and_purge(client):
boot(client) boot(client)
first = create_countdown(client).json() first = create_countdown(client).json()
@@ -188,3 +217,51 @@ def test_countdowns_backup_replace_and_merge_round_trip(client):
assert [item["title"] for item in all_active].count("周年") == 1 assert [item["title"] for item in all_active].count("周年") == 1
assert sum(item["pinned"] for item in all_active) == 1 assert sum(item["pinned"] for item in all_active) == 1
assert active["title"] == "周年" assert active["title"] == "周年"
def test_countdown_backup_merge_remaps_ids_owned_by_another_user(client):
boot(client)
original = create_countdown(client, title="跨账号纪念日").json()
exported = client.get("/api/v1/export").json()
client.post("/api/v1/auth/logout")
from backend.auth import hash_password
from backend.db import get_db
from backend.models import TaskList, User
async def add_other_user():
db_gen = get_db()
db = await anext(db_gen)
try:
other = User(username="other", password_hash=hash_password("correct horse battery staple"))
db.add(other)
await db.flush()
db.add(TaskList(user_id=other.id, name="收集箱", is_inbox=True))
await db.commit()
finally:
await db_gen.aclose()
client.portal.call(add_other_user)
assert client.post(
"/api/v1/auth/login",
json={"username": "other", "password": "correct horse battery staple"},
).status_code == 200
restored = client.post("/api/v1/restore", params={"mode": "merge"}, json=exported)
assert restored.status_code == 200
items = client.get("/api/v1/countdowns").json()
assert len(items) == 1
assert items[0]["title"] == "跨账号纪念日"
assert items[0]["id"] != original["id"]
merged_again = client.post("/api/v1/restore", params={"mode": "merge"}, json=exported)
assert merged_again.status_code == 200
assert len(client.get("/api/v1/countdowns").json()) == 1
exported_by_other = client.get("/api/v1/export").json()
restored_again = client.post("/api/v1/restore", params={"mode": "replace"}, json=exported_by_other)
assert restored_again.status_code == 200
replaced = client.get("/api/v1/countdowns").json()
assert len(replaced) == 1
assert replaced[0]["id"] == items[0]["id"]