Browse Source

增加保底代码校验

xueyiming 6 days ago
parent
commit
7a7921046c

+ 67 - 7
app/festival_demand/detection.py

@@ -11,6 +11,10 @@ from typing import Any
 from app.core.open_router_llm import OpenRouterCallError, create_chat_completion
 from app.festival_demand.calendar import festival_events_payload
 from app.festival_demand.exceptions import FestivalDemandError
+from app.festival_demand.fixed_dates import (
+    detect_deterministic_active_festivals,
+    merge_festival_detections,
+)
 from app.festival_demand.llm_json import extract_json_object
 from app.festival_demand.types import FestivalDemandConfig
 from app.festival_demand.validation import filter_festivals_by_query_date
@@ -162,16 +166,28 @@ def _normalize_detection_result(
     }
 
 
-def detect_active_festivals(
+def _build_detection_response(
+    *,
+    query_date: date,
+    festivals: list[dict[str, Any]],
+) -> dict[str, Any]:
+    return {
+        "date": query_date.isoformat(),
+        "festivals": festivals,
+        "festival_names": [item["name"] for item in festivals],
+    }
+
+
+def _detect_active_festivals_via_llm(
     query_date: date,
     config: FestivalDemandConfig,
+    *,
+    allowed_names: set[str],
+    allowed_kinds: dict[str, str],
+    prewarm_by_name: dict[str, int],
+    event_type_by_name: dict[str, str],
+    events: list[dict[str, int | str]],
 ) -> dict[str, Any]:
-    """调用 LLM 判断查询日期处于活跃期内的节日列表。"""
-    events = festival_events_payload()
-    allowed_names = {str(item["name"]) for item in events}
-    allowed_kinds = {str(item["name"]): str(item["kind"]) for item in events}
-    prewarm_by_name = {str(item["name"]): int(item["prewarm_days"]) for item in events}
-    event_type_by_name = {str(item["name"]): str(item["event_type"]) for item in events}
     payload = {
         "query_date": query_date.isoformat(),
         "year": query_date.year,
@@ -210,3 +226,47 @@ def detect_active_festivals(
     raise FestivalDemandError(
         f"festival detection failed after {config.llm_max_attempts} attempts: {last_error}"
     ) from last_error
+
+
+def detect_active_festivals(
+    query_date: date,
+    config: FestivalDemandConfig,
+) -> dict[str, Any]:
+    """判断查询日期处于活跃期内的节日列表(LLM + 确定性日期代码兜底)。"""
+    events = festival_events_payload()
+    allowed_names = {str(item["name"]) for item in events}
+    allowed_kinds = {str(item["name"]): str(item["kind"]) for item in events}
+    prewarm_by_name = {str(item["name"]): int(item["prewarm_days"]) for item in events}
+    event_type_by_name = {str(item["name"]): str(item["event_type"]) for item in events}
+
+    code_festivals = detect_deterministic_active_festivals(
+        query_date,
+        allowed_kinds=allowed_kinds,
+        prewarm_by_name=prewarm_by_name,
+        event_type_by_name=event_type_by_name,
+    )
+
+    try:
+        llm_result = _detect_active_festivals_via_llm(
+            query_date,
+            config,
+            allowed_names=allowed_names,
+            allowed_kinds=allowed_kinds,
+            prewarm_by_name=prewarm_by_name,
+            event_type_by_name=event_type_by_name,
+            events=events,
+        )
+        llm_festivals = llm_result.get("festivals") or []
+    except FestivalDemandError:
+        if code_festivals:
+            return _build_detection_response(
+                query_date=query_date,
+                festivals=code_festivals,
+            )
+        raise
+
+    merged_festivals = merge_festival_detections(llm_festivals, code_festivals)
+    return _build_detection_response(
+        query_date=query_date,
+        festivals=merged_festivals,
+    )

+ 216 - 0
app/festival_demand/fixed_dates.py

@@ -0,0 +1,216 @@
+"""确定性节日/节气/纪念日的代码日期计算与活跃检测兜底。"""
+
+from __future__ import annotations
+
+from datetime import date, timedelta
+from typing import Any
+
+from app.festival_demand.validation import filter_festivals_by_query_date
+
+# 21 世纪节气常数 C(索引 0=小寒 … 23=冬至)
+_CENTURY_21_C = (
+    5.4055,
+    20.12,
+    3.87,
+    18.73,
+    5.63,
+    20.646,
+    4.81,
+    20.1,
+    5.52,
+    21.04,
+    5.678,
+    21.37,
+    7.108,
+    22.83,
+    7.5,
+    23.13,
+    7.646,
+    23.042,
+    8.318,
+    23.438,
+    7.438,
+    22.36,
+    7.18,
+    21.94,
+)
+
+_SOLAR_TERM_INDEX: dict[str, int] = {
+    "小寒": 0,
+    "大寒": 1,
+    "立春": 2,
+    "雨水": 3,
+    "惊蛰": 4,
+    "春分": 5,
+    "清明": 6,
+    "谷雨": 7,
+    "立夏": 8,
+    "小满": 9,
+    "芒种": 10,
+    "夏至": 11,
+    "小暑": 12,
+    "大暑": 13,
+    "立秋": 14,
+    "处暑": 15,
+    "白露": 16,
+    "秋分": 17,
+    "寒露": 18,
+    "霜降": 19,
+    "立冬": 20,
+    "小雪": 21,
+    "大雪": 22,
+    "冬至": 23,
+}
+
+# 固定公历月日(非法定单日事件:festival_start = festival_end = 当天)
+FIXED_MONTH_DAY: dict[str, tuple[int, int]] = {
+    "妇女节": (3, 8),
+    "儿童节": (6, 1),
+    "建党节": (7, 1),
+    "建军节": (8, 1),
+    "教师节": (9, 10),
+    "周总理逝世": (1, 8),
+    "周总理诞辰": (3, 5),
+    "七七事变纪念日": (7, 7),
+    "日本投降": (8, 15),
+    "毛主席逝世": (9, 9),
+    "918纪念日": (9, 18),
+    "毛主席诞辰": (12, 26),
+    "公祭日": (12, 13),
+    "台湾光复纪念日": (10, 25),
+    "315": (3, 15),
+}
+
+# 固定公历锚点日(法定节日兜底:仅锚点日当天为节日期,放假区间仍由 LLM 补充)
+FIXED_STATUTORY_ANCHOR: dict[str, tuple[int, int]] = {
+    "元旦": (1, 1),
+    "劳动节": (5, 1),
+    "国庆节": (10, 1),
+}
+
+# 第 N 个星期 X(month 1-12, weekday 0=周一 … 6=周日)
+_NTH_WEEKDAY_RULES: dict[str, tuple[int, int, int]] = {
+    "母亲节": (5, 6, 2),  # 5 月第 2 个周日
+    "父亲节": (6, 6, 3),  # 6 月第 3 个周日
+}
+
+DETERMINISTIC_EVENT_NAMES: frozenset[str] = frozenset(
+    {
+        *_SOLAR_TERM_INDEX.keys(),
+        *FIXED_MONTH_DAY.keys(),
+        *FIXED_STATUTORY_ANCHOR.keys(),
+        *_NTH_WEEKDAY_RULES.keys(),
+    }
+)
+
+
+def is_deterministic_event(name: str) -> bool:
+    return name in DETERMINISTIC_EVENT_NAMES
+
+
+def solar_term_date(year: int, term_name: str) -> date:
+    """计算指定年份的二十四节气公历日期。"""
+    if term_name not in _SOLAR_TERM_INDEX:
+        raise ValueError(f"unknown solar term: {term_name}")
+    term_index = _SOLAR_TERM_INDEX[term_name]
+    century_c = _CENTURY_21_C[term_index]
+    y = year % 100
+    leap_years = y // 4
+    day = int(y * 0.2422 + century_c) - leap_years
+    month = term_index // 2 + 1
+    return date(year, month, day)
+
+
+def nth_weekday_of_month(year: int, month: int, weekday: int, nth: int) -> date:
+    """返回某月第 nth 个指定星期几的日期(weekday: 0=周一 … 6=周日)。"""
+    if nth <= 0:
+        raise ValueError("nth must be positive")
+    first = date(year, month, 1)
+    offset = (weekday - first.weekday()) % 7
+    candidate = first + timedelta(days=offset)
+    candidate += timedelta(weeks=nth - 1)
+    if candidate.month != month:
+        raise ValueError("nth weekday does not exist in month")
+    return candidate
+
+
+def resolve_event_date(year: int, name: str) -> date | None:
+    """解析确定性事件的当年公历日期。"""
+    if name in _SOLAR_TERM_INDEX:
+        return solar_term_date(year, name)
+    if name in FIXED_MONTH_DAY:
+        month, day = FIXED_MONTH_DAY[name]
+        return date(year, month, day)
+    if name in FIXED_STATUTORY_ANCHOR:
+        month, day = FIXED_STATUTORY_ANCHOR[name]
+        return date(year, month, day)
+    if name in _NTH_WEEKDAY_RULES:
+        month, weekday, nth = _NTH_WEEKDAY_RULES[name]
+        return nth_weekday_of_month(year, month, weekday, nth)
+    return None
+
+
+def detect_deterministic_active_festivals(
+    query_date: date,
+    *,
+    allowed_kinds: dict[str, str],
+    prewarm_by_name: dict[str, int],
+    event_type_by_name: dict[str, str],
+) -> list[dict[str, Any]]:
+    """用代码计算确定性事件在 query_date 是否处于预热/节日期。"""
+    candidates: list[dict[str, Any]] = []
+    year = query_date.year
+
+    for name in DETERMINISTIC_EVENT_NAMES:
+        if name not in prewarm_by_name:
+            continue
+        try:
+            event_date_value = resolve_event_date(year, name)
+        except (ValueError, OverflowError):
+            continue
+        if event_date_value is None:
+            continue
+
+        kind = allowed_kinds.get(name, "non_statutory")
+        event_iso = event_date_value.isoformat()
+        candidates.append(
+            {
+                "name": name,
+                "kind": kind,
+                "event_date": event_iso,
+                "festival_start": event_iso,
+                "festival_end": event_iso,
+            }
+        )
+
+    filtered = filter_festivals_by_query_date(
+        candidates,
+        query_date,
+        prewarm_by_name,
+    )
+    enriched: list[dict[str, Any]] = []
+    for item in filtered:
+        name = str(item.get("name") or "").strip()
+        enriched.append(
+            {
+                **item,
+                "event_type": event_type_by_name.get(name, ""),
+                "detection_source": "code",
+                "reason": "代码兜底:确定性日期",
+            }
+        )
+    return enriched
+
+
+def merge_festival_detections(
+    *detection_lists: list[dict[str, Any]],
+) -> list[dict[str, Any]]:
+    """按名称合并多路检测结果,后出现的结果覆盖先前的同名项。"""
+    merged: dict[str, dict[str, Any]] = {}
+    for detection_list in detection_lists:
+        for item in detection_list:
+            name = str(item.get("name") or "").strip()
+            if not name:
+                continue
+            merged[name] = item
+    return list(merged.values())

+ 107 - 0
tests/test_festival_detection.py

@@ -0,0 +1,107 @@
+"""节日检测流程集成测试(含代码兜底)。"""
+
+from __future__ import annotations
+
+import unittest
+from datetime import date
+from unittest.mock import patch
+
+from app.festival_demand.demand_generation import generate_demands_from_matches
+from app.festival_demand.detection import detect_active_festivals
+from app.festival_demand.exceptions import FestivalDemandError
+from app.festival_demand.fixed_dates import merge_festival_detections
+from app.festival_demand.types import FestivalDemandConfig
+
+
+def _test_config() -> FestivalDemandConfig:
+    return FestivalDemandConfig(
+        cron_hours="7",
+        cron_minute=0,
+        llm_model="test-model",
+        llm_max_attempts=1,
+        llm_retry_sleep_seconds=0.0,
+        llm_max_tokens=100,
+        llm_temperature=0.0,
+        demand_pool_source_table="dwd_multi_demand_pool_di",
+        demand_pool_strategy="去年同期阳历",
+        output_strategy="去年同期阳历-节点事件",
+        match_batch_size=200,
+    )
+
+
+class FestivalDetectionIntegrationTest(unittest.TestCase):
+    def test_code_fallback_when_llm_returns_empty(self) -> None:
+        llm_result = {
+            "date": "2026-08-03",
+            "festivals": [],
+            "festival_names": [],
+        }
+        with patch(
+            "app.festival_demand.detection._detect_active_festivals_via_llm",
+            return_value=llm_result,
+        ):
+            result = detect_active_festivals(date(2026, 8, 3), _test_config())
+
+        self.assertIn("立秋", result["festival_names"])
+        liqiu = next(item for item in result["festivals"] if item["name"] == "立秋")
+        self.assertEqual(liqiu["detection_source"], "code")
+        self.assertEqual(liqiu["event_type"], "节日/节气")
+        self.assertEqual(liqiu["phase"], "prewarm")
+
+    def test_code_fallback_when_llm_fails(self) -> None:
+        with patch(
+            "app.festival_demand.detection._detect_active_festivals_via_llm",
+            side_effect=FestivalDemandError("llm unavailable"),
+        ):
+            result = detect_active_festivals(date(2026, 8, 3), _test_config())
+
+        self.assertIn("立秋", result["festival_names"])
+
+    def test_code_overrides_llm_for_same_event(self) -> None:
+        llm_festivals = [
+            {
+                "name": "立秋",
+                "kind": "non_statutory",
+                "event_date": "2026-08-08",
+                "festival_start": "2026-08-08",
+                "festival_end": "2026-08-08",
+                "phase": "prewarm",
+                "days_to_festival": 5,
+                "event_type": "节日/节气",
+                "detection_source": "llm",
+            }
+        ]
+        code_festivals = [
+            {
+                "name": "立秋",
+                "kind": "non_statutory",
+                "event_date": "2026-08-07",
+                "festival_start": "2026-08-07",
+                "festival_end": "2026-08-07",
+                "phase": "prewarm",
+                "days_to_festival": 4,
+                "event_type": "节日/节气",
+                "detection_source": "code",
+            }
+        ]
+        merged = merge_festival_detections(llm_festivals, code_festivals)
+        liqiu = next(item for item in merged if item["name"] == "立秋")
+        self.assertEqual(liqiu["detection_source"], "code")
+        self.assertEqual(liqiu["event_date"], "2026-08-07")
+
+    def test_downstream_demand_generation_accepts_code_festival(self) -> None:
+        matches = [
+            {
+                "demand_name": "贴秋膘",
+                "festival_name": "立秋",
+                "event_type": "节日/节气",
+            }
+        ]
+        generated = generate_demands_from_matches(matches, query_date=date(2026, 8, 3))
+        self.assertIn("立秋", generated["generated_demand_names"])
+        self.assertIn("2026 立秋", generated["generated_demand_names"])
+        self.assertIn("2026 立秋 贴秋膘", generated["generated_demand_names"])
+
+
+if __name__ == "__main__":
+    unittest.main()

+ 48 - 0
tests/test_festival_fixed_dates.py

@@ -0,0 +1,48 @@
+"""确定性节日日期兜底测试。"""
+
+from __future__ import annotations
+
+import unittest
+from datetime import date
+
+from app.festival_demand.fixed_dates import (
+    detect_deterministic_active_festivals,
+    is_deterministic_event,
+    resolve_event_date,
+    solar_term_date,
+)
+from app.festival_demand.calendar import DEFAULT_FESTIVAL_EVENTS
+
+
+class FestivalFixedDatesTest(unittest.TestCase):
+    def test_solar_term_liqiu_2026(self) -> None:
+        self.assertEqual(solar_term_date(2026, "立秋"), date(2026, 8, 7))
+
+    def test_fixed_month_day(self) -> None:
+        self.assertEqual(resolve_event_date(2026, "建党节"), date(2026, 7, 1))
+        self.assertEqual(resolve_event_date(2026, "918纪念日"), date(2026, 9, 18))
+
+    def test_liqiu_prewarm_detected_on_2026_08_03(self) -> None:
+        allowed_kinds = {event.name: event.kind for event in DEFAULT_FESTIVAL_EVENTS}
+        prewarm_by_name = {event.name: event.prewarm_days for event in DEFAULT_FESTIVAL_EVENTS}
+        event_type_by_name = {event.name: event.event_type for event in DEFAULT_FESTIVAL_EVENTS}
+
+        festivals = detect_deterministic_active_festivals(
+            date(2026, 8, 3),
+            allowed_kinds=allowed_kinds,
+            prewarm_by_name=prewarm_by_name,
+            event_type_by_name=event_type_by_name,
+        )
+        names = {item["name"] for item in festivals}
+        self.assertIn("立秋", names)
+        liqiu = next(item for item in festivals if item["name"] == "立秋")
+        self.assertEqual(liqiu["phase"], "prewarm")
+        self.assertEqual(liqiu["days_to_festival"], 4)
+
+    def test_lunar_festival_not_deterministic(self) -> None:
+        self.assertFalse(is_deterministic_event("春节"))
+        self.assertTrue(is_deterministic_event("立秋"))
+
+
+if __name__ == "__main__":
+    unittest.main()