detection.py 7.6 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212
  1. """调用 LLM 判断当天是否处于节日或预热期内。"""
  2. from __future__ import annotations
  3. import json
  4. import re
  5. import time
  6. from datetime import date
  7. from typing import Any
  8. from app.core.open_router_llm import OpenRouterCallError, create_chat_completion
  9. from app.festival_demand.calendar import festival_events_payload
  10. from app.festival_demand.exceptions import FestivalDemandError
  11. from app.festival_demand.llm_json import extract_json_object
  12. from app.festival_demand.types import FestivalDemandConfig
  13. from app.festival_demand.validation import filter_festivals_by_query_date
  14. SYSTEM_PROMPT = """
  15. 你是一个中国节日与节点事件日历专家。
  16. # 任务
  17. 根据给定的 query_date 和 events 列表,判断 query_date 当天处于哪些事件的活跃期内,并返回这些事件。
  18. # 节日类型(events 中 kind 字段)
  19. 1. statutory(法定节假日):须查准当年国务院放假调休区间,festival_start 为放假首日,festival_end 为放假末日。
  20. 2. non_statutory(非法定节日):节日活跃期仅当天,festival_start = festival_end = 该事件当天。
  21. # 预热期(两类通用)
  22. - 预热只在节日开始之前,节日一到预热立即结束。
  23. - 设节日开始日 S(法定为 festival_start,非法定为当天),预热天数为 events 中的 prewarm_days(记为 N):
  24. - 预热区间 = [S-N, S-1](含首尾,不含 S)
  25. - 节日期 = [festival_start, festival_end](含首尾)
  26. - 仅当 query_date 落在预热区间或节日期内,才返回该事件。
  27. # 你的职责
  28. - 由你推算当年每个事件的 festival_start、festival_end、event_date(农历、节气、固定纪念日、放假安排等均由你判断)。
  29. - 严格遵守上述活跃期定义筛选,不要把已过节日的项返回进来。
  30. # 输出格式
  31. 严格只输出一个 JSON 对象,禁止在 JSON 前后或之后追加任何说明、注释、markdown 代码块标记。
  32. {
  33. "date": "YYYY-MM-DD",
  34. "festivals": [
  35. {
  36. "name": "与 events 中 name 完全一致",
  37. "kind": "statutory 或 non_statutory",
  38. "event_date": "YYYY-MM-DD",
  39. "festival_start": "YYYY-MM-DD",
  40. "festival_end": "YYYY-MM-DD",
  41. "phase": "prewarm 或 festival",
  42. "days_to_festival": 0,
  43. "reason": "简短中文说明"
  44. }
  45. ]
  46. }
  47. # 约束
  48. - festivals 为空数组表示当天无活跃事件。
  49. - name 只能来自 events,不得编造。
  50. - phase=prewarm 时 query_date 必须 < festival_start;phase=festival 时 query_date 必须在 [festival_start, festival_end]。
  51. """
  52. def _parse_date_field(item: dict[str, Any], *keys: str) -> str | None:
  53. for key in keys:
  54. value = str(item.get(key) or "").strip()
  55. if re.fullmatch(r"\d{4}-\d{2}-\d{2}", value):
  56. return value
  57. return None
  58. def _parse_festival_item(
  59. item: Any,
  60. allowed_names: set[str],
  61. allowed_kinds: dict[str, str],
  62. ) -> dict[str, Any] | None:
  63. if not isinstance(item, dict):
  64. return None
  65. name = str(item.get("name") or "").strip()
  66. if not name or name not in allowed_names:
  67. return None
  68. festival_start = _parse_date_field(item, "festival_start")
  69. festival_end = _parse_date_field(item, "festival_end")
  70. event_date = _parse_date_field(item, "event_date")
  71. if not festival_start and event_date:
  72. festival_start = event_date
  73. if not festival_end and festival_start:
  74. festival_end = festival_start
  75. if not festival_start or not festival_end:
  76. return None
  77. kind = str(item.get("kind") or allowed_kinds.get(name, "non_statutory")).strip()
  78. if kind not in {"statutory", "non_statutory"}:
  79. kind = allowed_kinds.get(name, "non_statutory")
  80. return {
  81. "name": name,
  82. "kind": kind,
  83. "event_date": event_date or festival_start,
  84. "festival_start": festival_start,
  85. "festival_end": festival_end,
  86. "phase": str(item.get("phase") or "").strip().lower(),
  87. "days_to_festival": item.get("days_to_festival"),
  88. "reason": str(item.get("reason") or "").strip(),
  89. }
  90. def _enrich_festivals_with_event_type(
  91. festivals: list[dict[str, Any]],
  92. event_type_by_name: dict[str, str],
  93. ) -> list[dict[str, Any]]:
  94. enriched: list[dict[str, Any]] = []
  95. for item in festivals:
  96. name = str(item.get("name") or "").strip()
  97. enriched.append(
  98. {
  99. **item,
  100. "event_type": event_type_by_name.get(name, ""),
  101. }
  102. )
  103. return enriched
  104. def _normalize_detection_result(
  105. parsed: dict[str, Any],
  106. *,
  107. query_date: date,
  108. allowed_names: set[str],
  109. allowed_kinds: dict[str, str],
  110. prewarm_by_name: dict[str, int],
  111. event_type_by_name: dict[str, str],
  112. ) -> dict[str, Any]:
  113. festivals_raw = parsed.get("festivals")
  114. if festivals_raw is None:
  115. festivals_raw = parsed.get("active_festivals") or []
  116. if not isinstance(festivals_raw, list):
  117. raise FestivalDemandError("llm output festivals must be a list")
  118. parsed_items: list[dict[str, Any]] = []
  119. seen_names: set[str] = set()
  120. for item in festivals_raw:
  121. if isinstance(item, str):
  122. continue
  123. normalized = _parse_festival_item(item, allowed_names, allowed_kinds)
  124. if not normalized or normalized["name"] in seen_names:
  125. continue
  126. parsed_items.append(normalized)
  127. seen_names.add(normalized["name"])
  128. festivals = filter_festivals_by_query_date(
  129. parsed_items,
  130. query_date,
  131. prewarm_by_name,
  132. )
  133. festivals = _enrich_festivals_with_event_type(festivals, event_type_by_name)
  134. return {
  135. "date": query_date.isoformat(),
  136. "festivals": festivals,
  137. "festival_names": [item["name"] for item in festivals],
  138. }
  139. def detect_active_festivals(
  140. query_date: date,
  141. config: FestivalDemandConfig,
  142. ) -> dict[str, Any]:
  143. """调用 LLM 判断查询日期处于活跃期内的节日列表。"""
  144. events = festival_events_payload()
  145. allowed_names = {str(item["name"]) for item in events}
  146. allowed_kinds = {str(item["name"]): str(item["kind"]) for item in events}
  147. prewarm_by_name = {str(item["name"]): int(item["prewarm_days"]) for item in events}
  148. event_type_by_name = {str(item["name"]): str(item["event_type"]) for item in events}
  149. payload = {
  150. "query_date": query_date.isoformat(),
  151. "year": query_date.year,
  152. "events": events,
  153. }
  154. last_error: Exception | None = None
  155. for attempt in range(1, config.llm_max_attempts + 1):
  156. try:
  157. resp = create_chat_completion(
  158. [
  159. {"role": "system", "content": SYSTEM_PROMPT.strip()},
  160. {
  161. "role": "user",
  162. "content": json.dumps(payload, ensure_ascii=False),
  163. },
  164. ],
  165. model=config.llm_model,
  166. temperature=config.llm_temperature,
  167. max_tokens=config.llm_max_tokens,
  168. )
  169. parsed = extract_json_object(str(resp.get("content") or ""))
  170. return _normalize_detection_result(
  171. parsed,
  172. query_date=query_date,
  173. allowed_names=allowed_names,
  174. allowed_kinds=allowed_kinds,
  175. prewarm_by_name=prewarm_by_name,
  176. event_type_by_name=event_type_by_name,
  177. )
  178. except (OpenRouterCallError, FestivalDemandError, ValueError) as exc:
  179. last_error = exc
  180. if attempt < config.llm_max_attempts:
  181. time.sleep(config.llm_retry_sleep_seconds)
  182. raise FestivalDemandError(
  183. f"festival detection failed after {config.llm_max_attempts} attempts: {last_error}"
  184. ) from last_error