detection.py 9.3 KB

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