| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212 |
- """调用 LLM 判断当天是否处于节日或预热期内。"""
- from __future__ import annotations
- import json
- import re
- import time
- from datetime import date
- 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.llm_json import extract_json_object
- from app.festival_demand.types import FestivalDemandConfig
- from app.festival_demand.validation import filter_festivals_by_query_date
- SYSTEM_PROMPT = """
- 你是一个中国节日与节点事件日历专家。
- # 任务
- 根据给定的 query_date 和 events 列表,判断 query_date 当天处于哪些事件的活跃期内,并返回这些事件。
- # 节日类型(events 中 kind 字段)
- 1. statutory(法定节假日):须查准当年国务院放假调休区间,festival_start 为放假首日,festival_end 为放假末日。
- 2. non_statutory(非法定节日):节日活跃期仅当天,festival_start = festival_end = 该事件当天。
- # 预热期(两类通用)
- - 预热只在节日开始之前,节日一到预热立即结束。
- - 设节日开始日 S(法定为 festival_start,非法定为当天),预热天数为 events 中的 prewarm_days(记为 N):
- - 预热区间 = [S-N, S-1](含首尾,不含 S)
- - 节日期 = [festival_start, festival_end](含首尾)
- - 仅当 query_date 落在预热区间或节日期内,才返回该事件。
- # 你的职责
- - 由你推算当年每个事件的 festival_start、festival_end、event_date(农历、节气、固定纪念日、放假安排等均由你判断)。
- - 严格遵守上述活跃期定义筛选,不要把已过节日的项返回进来。
- # 输出格式
- 严格只输出一个 JSON 对象,禁止在 JSON 前后或之后追加任何说明、注释、markdown 代码块标记。
- {
- "date": "YYYY-MM-DD",
- "festivals": [
- {
- "name": "与 events 中 name 完全一致",
- "kind": "statutory 或 non_statutory",
- "event_date": "YYYY-MM-DD",
- "festival_start": "YYYY-MM-DD",
- "festival_end": "YYYY-MM-DD",
- "phase": "prewarm 或 festival",
- "days_to_festival": 0,
- "reason": "简短中文说明"
- }
- ]
- }
- # 约束
- - festivals 为空数组表示当天无活跃事件。
- - name 只能来自 events,不得编造。
- - phase=prewarm 时 query_date 必须 < festival_start;phase=festival 时 query_date 必须在 [festival_start, festival_end]。
- """
- def _parse_date_field(item: dict[str, Any], *keys: str) -> str | None:
- for key in keys:
- value = str(item.get(key) or "").strip()
- if re.fullmatch(r"\d{4}-\d{2}-\d{2}", value):
- return value
- return None
- def _parse_festival_item(
- item: Any,
- allowed_names: set[str],
- allowed_kinds: dict[str, str],
- ) -> dict[str, Any] | None:
- if not isinstance(item, dict):
- return None
- name = str(item.get("name") or "").strip()
- if not name or name not in allowed_names:
- return None
- festival_start = _parse_date_field(item, "festival_start")
- festival_end = _parse_date_field(item, "festival_end")
- event_date = _parse_date_field(item, "event_date")
- if not festival_start and event_date:
- festival_start = event_date
- if not festival_end and festival_start:
- festival_end = festival_start
- if not festival_start or not festival_end:
- return None
- kind = str(item.get("kind") or allowed_kinds.get(name, "non_statutory")).strip()
- if kind not in {"statutory", "non_statutory"}:
- kind = allowed_kinds.get(name, "non_statutory")
- return {
- "name": name,
- "kind": kind,
- "event_date": event_date or festival_start,
- "festival_start": festival_start,
- "festival_end": festival_end,
- "phase": str(item.get("phase") or "").strip().lower(),
- "days_to_festival": item.get("days_to_festival"),
- "reason": str(item.get("reason") or "").strip(),
- }
- def _enrich_festivals_with_event_type(
- festivals: list[dict[str, Any]],
- event_type_by_name: dict[str, str],
- ) -> list[dict[str, Any]]:
- enriched: list[dict[str, Any]] = []
- for item in festivals:
- name = str(item.get("name") or "").strip()
- enriched.append(
- {
- **item,
- "event_type": event_type_by_name.get(name, ""),
- }
- )
- return enriched
- def _normalize_detection_result(
- parsed: dict[str, Any],
- *,
- query_date: date,
- allowed_names: set[str],
- allowed_kinds: dict[str, str],
- prewarm_by_name: dict[str, int],
- event_type_by_name: dict[str, str],
- ) -> dict[str, Any]:
- festivals_raw = parsed.get("festivals")
- if festivals_raw is None:
- festivals_raw = parsed.get("active_festivals") or []
- if not isinstance(festivals_raw, list):
- raise FestivalDemandError("llm output festivals must be a list")
- parsed_items: list[dict[str, Any]] = []
- seen_names: set[str] = set()
- for item in festivals_raw:
- if isinstance(item, str):
- continue
- normalized = _parse_festival_item(item, allowed_names, allowed_kinds)
- if not normalized or normalized["name"] in seen_names:
- continue
- parsed_items.append(normalized)
- seen_names.add(normalized["name"])
- festivals = filter_festivals_by_query_date(
- parsed_items,
- query_date,
- prewarm_by_name,
- )
- festivals = _enrich_festivals_with_event_type(festivals, event_type_by_name)
- return {
- "date": query_date.isoformat(),
- "festivals": festivals,
- "festival_names": [item["name"] for item in festivals],
- }
- 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}
- payload = {
- "query_date": query_date.isoformat(),
- "year": query_date.year,
- "events": events,
- }
- last_error: Exception | None = None
- for attempt in range(1, config.llm_max_attempts + 1):
- try:
- resp = create_chat_completion(
- [
- {"role": "system", "content": SYSTEM_PROMPT.strip()},
- {
- "role": "user",
- "content": json.dumps(payload, ensure_ascii=False),
- },
- ],
- model=config.llm_model,
- temperature=config.llm_temperature,
- max_tokens=config.llm_max_tokens,
- )
- parsed = extract_json_object(str(resp.get("content") or ""))
- return _normalize_detection_result(
- parsed,
- query_date=query_date,
- allowed_names=allowed_names,
- allowed_kinds=allowed_kinds,
- prewarm_by_name=prewarm_by_name,
- event_type_by_name=event_type_by_name,
- )
- except (OpenRouterCallError, FestivalDemandError, ValueError) as exc:
- last_error = exc
- if attempt < config.llm_max_attempts:
- time.sleep(config.llm_retry_sleep_seconds)
- raise FestivalDemandError(
- f"festival detection failed after {config.llm_max_attempts} attempts: {last_error}"
- ) from last_error
|