diff --git a/src/bank_importer/reminders.py b/src/bank_importer/reminders.py index ee15f83..7f4e3ff 100644 --- a/src/bank_importer/reminders.py +++ b/src/bank_importer/reminders.py @@ -73,26 +73,68 @@ def get_settings(connection: sqlite3.Connection) -> dict[str, str]: return settings +def _valid_scan_time(value: str) -> bool: + parts = value.strip().split(":") + if len(parts) != 2 or not parts[0].isdigit() or not parts[1].isdigit(): + return False + if len(parts[1]) != 2: + return False + hour, minute = int(parts[0]), int(parts[1]) + return 0 <= hour <= 23 and 0 <= minute <= 59 + + +def validate_setting(key: str, value: str) -> str: + raw = str(value).strip() + if key == SETTING_MONTHLY_START_DAY: + try: + day = int(raw) + except ValueError as exc: + raise ValueError("每月起始日须为 1 到 28 的整数。") from exc + if day < 1 or day > 28: + raise ValueError("每月起始日须为 1 到 28 的整数。") + return str(day) + if key == SETTING_GAP_DAYS: + try: + days = int(raw) + except ValueError as exc: + raise ValueError("断档天数须为大于等于 1 的整数。") from exc + if days < 1: + raise ValueError("断档天数须为大于等于 1 的整数。") + return str(days) + if key == SETTING_SCAN_TIME: + if not _valid_scan_time(raw): + raise ValueError("扫描时间须为合法的 HH:MM。") + hour, minute = raw.split(":") + return f"{int(hour):02d}:{minute}" + raise ValueError(f"unknown setting: {key}") + + def update_settings(connection: sqlite3.Connection, updates: dict[str, str]) -> dict[str, str]: allowed = set(DEFAULT_SETTINGS) + cleaned: dict[str, str] = {} + for key, value in updates.items(): + if key not in allowed: + raise ValueError(f"unknown setting: {key}") + cleaned[key] = validate_setting(key, value) now = utc_now() with connection: - for key, value in updates.items(): - if key not in allowed: - raise ValueError(f"unknown setting: {key}") + for key, value in cleaned.items(): connection.execute( """ INSERT INTO reminder_settings (key, value, updated_at) VALUES (?, ?, ?) ON CONFLICT(key) DO UPDATE SET value = excluded.value, updated_at = excluded.updated_at """, - (key, str(value), now), + (key, value, now), ) return get_settings(connection) def _setting_int(settings: dict[str, str], key: str) -> int: - return int(settings.get(key, DEFAULT_SETTINGS[key])) + try: + return int(validate_setting(key, settings.get(key, DEFAULT_SETTINGS[key]))) + except (TypeError, ValueError): + return int(DEFAULT_SETTINGS[key]) def _month_period(day: date | None = None) -> str: @@ -400,21 +442,33 @@ def _deliver( ).fetchone() with connection: - if existing is not None and existing["status"] != "resolved": + if existing is not None: reminder_id = int(existing["id"]) send_count = int(existing["send_count"]) + 1 connection.execute( """ UPDATE reminders - SET send_count = ?, last_sent_at = ?, status = 'open' + SET send_count = ?, last_sent_at = ?, status = 'open', + title = ?, content = ?, rule_params = ?, action_link = ? WHERE id = ? """, - (send_count, now, reminder_id), + ( + send_count, + now, + finding.title, + finding.content, + json.dumps(finding.rule_params, ensure_ascii=False), + finding.action_link, + reminder_id, + ), + ) + event_type = "sent" if existing["status"] == "resolved" else ( + "escalated" if send_count > 1 else "sent" ) _append_event( connection, reminder_id, - "escalated" if send_count > 1 else "sent", + event_type, actor_ref, f"第 {send_count} 次催办", ) @@ -465,7 +519,10 @@ def deliver_many( ) -> list[int]: sent: list[int] = [] for key in dedupe_keys: - reminder_id = deliver_finding(connection, key, actor=actor, ip=ip) + try: + reminder_id = deliver_finding(connection, key, actor=actor, ip=ip) + except Exception: + continue if reminder_id is not None: sent.append(reminder_id) return sent diff --git a/tests/test_reminders.py b/tests/test_reminders.py index 09714cc..a014584 100644 --- a/tests/test_reminders.py +++ b/tests/test_reminders.py @@ -3,8 +3,11 @@ from __future__ import annotations from datetime import date, timedelta +from pathlib import Path import json +import sqlite3 import unittest +from unittest.mock import patch from bank_importer import auth, reminders from bank_importer.db import connect, migrate, utc_now @@ -187,6 +190,63 @@ class DeliveryAndDedupTests(ReminderTestCase): ).fetchone()[0] self.assertEqual(2, events) + def test_resolved_reminder_reopens_same_row(self) -> None: + self._insert_pending_account(self.company_a) + finding = next( + item + for item in reminders.scan_findings(self.connection) + if item.rule_key == reminders.RULE_PENDING and item.company_id == self.company_a + ) + first = reminders.deliver_finding(self.connection, finding.dedupe_key, actor=self.admin) + self.assertIsNotNone(first) + reminders.update_reminder_status( + self.connection, first, "resolved", company_id=self.company_a, actor=self.actor_a + ) + second = reminders.deliver_finding(self.connection, finding.dedupe_key, actor=self.admin) + self.assertEqual(first, second) + row = self.connection.execute( + "SELECT send_count, status, dedupe_key FROM reminders WHERE id = ?", + (first,), + ).fetchone() + self.assertEqual(2, row["send_count"]) + self.assertEqual("open", row["status"]) + count = self.connection.execute( + "SELECT COUNT(*) FROM reminders WHERE dedupe_key = ?", (finding.dedupe_key,) + ).fetchone()[0] + self.assertEqual(1, count) + events = [ + item["event_type"] + for item in self.connection.execute( + "SELECT event_type FROM reminder_events WHERE reminder_id = ? ORDER BY id", + (first,), + ).fetchall() + ] + self.assertEqual(["sent", "resolved", "sent"], events) + + def test_deliver_many_continues_after_one_failure(self) -> None: + self._insert_pending_account(self.company_a) + finding = next( + item + for item in reminders.scan_findings(self.connection) + if item.rule_key == reminders.RULE_PENDING and item.company_id == self.company_a + ) + original = reminders.deliver_finding + + def flaky(connection, key, **kwargs): + if key == "boom": + raise sqlite3.IntegrityError("UNIQUE constraint failed: reminders.dedupe_key") + return original(connection, key, **kwargs) + + with patch.object(reminders, "deliver_finding", side_effect=flaky): + sent = reminders.deliver_many( + self.connection, ["boom", finding.dedupe_key], actor=self.admin + ) + self.assertEqual(1, len(sent)) + self.assertEqual( + 1, + self.connection.execute("SELECT COUNT(*) FROM reminders").fetchone()[0], + ) + def test_manual_reminder_each_send_is_separate(self) -> None: first = reminders.send_manual( self.connection, @@ -335,6 +395,79 @@ class SettingsTests(ReminderTestCase): with self.assertRaises(ValueError): reminders.update_settings(self.connection, {"unknown_key": "1"}) + def test_invalid_monthly_start_day_rejected(self) -> None: + with self.assertRaises(ValueError): + reminders.update_settings(self.connection, {"monthly_start_day": "0"}) + with self.assertRaises(ValueError): + reminders.update_settings(self.connection, {"monthly_start_day": "29"}) + self.assertEqual("5", reminders.get_settings(self.connection)["monthly_start_day"]) + + def test_invalid_gap_days_rejected(self) -> None: + with self.assertRaises(ValueError): + reminders.update_settings(self.connection, {"gap_days": "0"}) + with self.assertRaises(ValueError): + reminders.update_settings(self.connection, {"gap_days": "abc"}) + self.assertEqual("5", reminders.get_settings(self.connection)["gap_days"]) + + def test_invalid_scan_time_rejected(self) -> None: + with self.assertRaises(ValueError): + reminders.update_settings(self.connection, {"scan_time": "25:00"}) + with self.assertRaises(ValueError): + reminders.update_settings(self.connection, {"scan_time": "8:0"}) + self.assertEqual("08:00", reminders.get_settings(self.connection)["scan_time"]) + + def test_corrupt_settings_do_not_break_scan(self) -> None: + with self.connection: + self.connection.execute( + "UPDATE reminder_settings SET value = 'not-a-number' WHERE key = 'monthly_start_day'" + ) + self.connection.execute( + "UPDATE reminder_settings SET value = '0' WHERE key = 'gap_days'" + ) + findings = reminders.scan_findings(self.connection) + self.assertIsInstance(findings, list) + + +class NoticeListDelegationTests(unittest.TestCase): + def test_async_notice_go_handle_uses_event_delegation(self) -> None: + source = (Path(__file__).resolve().parents[1] / "web" / "app.js").read_text( + encoding="utf-8" + ) + init_body = source.split("function initNotifications()", 1)[1].split("\nfunction ", 1)[0] + self.assertIn('list.addEventListener("click"', init_body) + self.assertIn('closest("[data-view-link]")', init_body) + self.assertIn("showView(viewLink.dataset.viewLink)", init_body) + self.assertIn('document.addEventListener("click"', source) + self.assertNotIn( + '$$("[data-view-link]").forEach((button) => button.addEventListener("click"', + source, + ) + + def test_delegated_lookup_finds_button_inserted_after_init(self) -> None: + """Simulate #notice-list after async replaceChildren: click target is the new button.""" + list_root = {"id": "notice-list", "parent": None, "attrs": {}} + side = {"id": "lr-side", "parent": list_root, "attrs": {}} + button = { + "id": "go", + "parent": side, + "attrs": {"data-view-link": "reconcile"}, + } + list_root["children"] = [side] + side["children"] = [button] + + def closest(node, attr): + current = node + while current is not None: + if attr in current.get("attrs", {}): + return current + current = current.get("parent") + return None + + clicked = closest(button, "data-view-link") + self.assertIsNotNone(clicked) + self.assertEqual("reconcile", clicked["attrs"]["data-view-link"]) + self.assertIs(list_root, clicked["parent"]["parent"]) + class MigrationTests(unittest.TestCase): def test_v6_migration_applies_and_rolls_back(self) -> None: diff --git a/tests/test_server_auth.py b/tests/test_server_auth.py index 36039fe..9478c0d 100644 --- a/tests/test_server_auth.py +++ b/tests/test_server_auth.py @@ -287,6 +287,8 @@ class ServerAuthMatrixTests(unittest.TestCase): lambda: anon.get("/api/admin/users"), lambda: anon.get("/api/admin/companies"), lambda: anon.get("/api/admin/audit-log"), + lambda: anon.get("/api/admin/reminders"), + lambda: anon.get("/api/company/reminders"), lambda: anon.get("/api/me"), ): status, _, data = method_check() @@ -387,6 +389,10 @@ class ServerAuthMatrixTests(unittest.TestCase): lambda: self.cashier_a.request("POST", "/api/admin/users/1/enable"), lambda: self.cashier_a.request("POST", "/api/admin/users/1/reset-password"), lambda: self.cashier_a.get("/api/admin/audit-log"), + lambda: self.cashier_a.get("/api/admin/reminders"), + lambda: self.cashier_a.get("/api/admin/reminders/pending"), + lambda: self.cashier_a.post_json("/api/admin/reminder-settings", {"gap_days": "3"}), + lambda: self.cashier_a.post_json("/api/admin/reminders/send", {"dedupe_keys": ["x"]}), ) for call in calls: status, _, data = call() @@ -557,6 +563,28 @@ class ServerAuthMatrixTests(unittest.TestCase): self.assertNotIn(password, row["detail"] or "") self.assertNotIn(password, row["target"] or "") + def test_invalid_reminder_settings_return_400_and_do_not_persist(self) -> None: + cases = ( + {"monthly_start_day": "0"}, + {"gap_days": "0"}, + {"scan_time": "25:99"}, + ) + for payload in cases: + with self.subTest(payload=payload): + status, _, data = self.admin.post_json( + "/api/admin/reminder-settings", {"settings": payload} + ) + self.assertEqual(400, status, data) + self.assertEqual("error", as_json(data)["status"]) + status, _, data = self.admin.get("/api/admin/reminder-settings") + self.assertEqual(200, status) + settings = as_json(data)["settings"] + self.assertEqual("5", settings["monthly_start_day"]) + self.assertEqual("5", settings["gap_days"]) + self.assertEqual("08:00", settings["scan_time"]) + status, _, data = self.admin.get("/api/admin/reminders/pending") + self.assertEqual(200, status, data) + if __name__ == "__main__": unittest.main() diff --git a/web/app.js b/web/app.js index c4f1c9c..d75ad74 100644 --- a/web/app.js +++ b/web/app.js @@ -446,7 +446,12 @@ function initShell() { } }); $$("[data-view]").forEach((button) => button.addEventListener("click", (event) => { event.preventDefault(); showView(button.dataset.view); })); - $$("[data-view-link]").forEach((button) => button.addEventListener("click", (event) => { event.preventDefault(); showView(button.dataset.viewLink); })); + document.addEventListener("click", (event) => { + const button = event.target.closest("[data-view-link]"); + if (!button) return; + event.preventDefault(); + showView(button.dataset.viewLink); + }); $$(".logout").forEach((control) => control.addEventListener("click", async (event) => { event.preventDefault(); try { @@ -2373,6 +2378,12 @@ function initNotifications() { }); list.addEventListener("click", async (event) => { + const viewLink = event.target.closest("[data-view-link]"); + if (viewLink && list.contains(viewLink)) { + event.preventDefault(); + showView(viewLink.dataset.viewLink); + return; + } const readBtn = event.target.closest(".btn-mark-read"); const doneBtn = event.target.closest(".btn-mark-done"); const btn = readBtn || doneBtn;