forked from zhouruizhe/gmTouringMiniApp
204 lines
7.6 KiB
Python
204 lines
7.6 KiB
Python
from __future__ import annotations
|
||
|
||
import json
|
||
import os
|
||
from datetime import datetime, timedelta
|
||
from typing import Any, Optional
|
||
|
||
from pydantic import ValidationError
|
||
|
||
from .knowledge import KnowledgeBase
|
||
from .prompts import SYSTEM_PROMPT
|
||
from .schemas import GeneratedItinerary, Pace, TravelPreferences
|
||
|
||
|
||
class LLMError(RuntimeError):
|
||
pass
|
||
|
||
|
||
class LLMTimeout(LLMError):
|
||
pass
|
||
|
||
|
||
class LLMService:
|
||
def __init__(self, knowledge: KnowledgeBase) -> None:
|
||
self.knowledge = knowledge
|
||
self.mode = os.getenv("LLM_MODE", "mock").lower()
|
||
self.timeout = float(os.getenv("LLM_TIMEOUT_SECONDS", "45"))
|
||
|
||
@property
|
||
def configured(self) -> bool:
|
||
if self.mode == "mock":
|
||
return True
|
||
return bool(os.getenv("OPENAI_API_KEY") and os.getenv("OPENAI_MODEL"))
|
||
|
||
async def generate(
|
||
self,
|
||
preferences: TravelPreferences,
|
||
candidates: list[dict[str, Any]],
|
||
previous: Optional[GeneratedItinerary] = None,
|
||
user_message: Optional[str] = None,
|
||
) -> GeneratedItinerary:
|
||
if self.mode == "mock":
|
||
return self._mock_generate(preferences, candidates, previous, user_message)
|
||
return await self._real_generate(preferences, candidates, previous, user_message)
|
||
|
||
async def _real_generate(
|
||
self,
|
||
preferences: TravelPreferences,
|
||
candidates: list[dict[str, Any]],
|
||
previous: Optional[GeneratedItinerary],
|
||
user_message: Optional[str],
|
||
) -> GeneratedItinerary:
|
||
if not self.configured:
|
||
raise LLMError("真实模型尚未配置")
|
||
try:
|
||
from openai import APITimeoutError, AsyncOpenAI
|
||
except ImportError as exc:
|
||
raise LLMError("缺少 openai 依赖") from exc
|
||
|
||
client = AsyncOpenAI(
|
||
api_key=os.environ["OPENAI_API_KEY"],
|
||
base_url=os.getenv("OPENAI_BASE_URL") or None,
|
||
timeout=self.timeout,
|
||
)
|
||
payload = {
|
||
"preferences": preferences.model_dump(by_alias=True),
|
||
"candidates": candidates,
|
||
"travelTimes": self.knowledge.travel_context(
|
||
candidates, preferences.transport.value
|
||
),
|
||
"previousItinerary": (
|
||
previous.model_dump(by_alias=True) if previous else None
|
||
),
|
||
"adjustmentRequest": user_message,
|
||
}
|
||
feedback = ""
|
||
allowed_ids = {place["id"] for place in candidates}
|
||
|
||
for attempt in range(2):
|
||
try:
|
||
response = await client.chat.completions.create(
|
||
model=os.environ["OPENAI_MODEL"],
|
||
temperature=0.2,
|
||
response_format={"type": "json_object"},
|
||
messages=[
|
||
{"role": "system", "content": SYSTEM_PROMPT},
|
||
{
|
||
"role": "user",
|
||
"content": json.dumps(payload, ensure_ascii=False)
|
||
+ feedback,
|
||
},
|
||
],
|
||
)
|
||
except APITimeoutError as exc:
|
||
raise LLMTimeout("模型调用超时") from exc
|
||
except Exception as exc:
|
||
raise LLMError(f"模型调用失败:{exc}") from exc
|
||
|
||
try:
|
||
content = response.choices[0].message.content or "{}"
|
||
itinerary = GeneratedItinerary.model_validate_json(content)
|
||
unknown = {
|
||
item.place_id for item in itinerary.items
|
||
} - allowed_ids
|
||
if unknown:
|
||
raise ValueError(f"包含未知地点ID:{sorted(unknown)}")
|
||
return itinerary
|
||
except (ValidationError, ValueError, json.JSONDecodeError) as exc:
|
||
if attempt == 1:
|
||
raise LLMError(f"模型输出无法通过结构校验:{exc}") from exc
|
||
feedback = f"\n上次输出校验失败:{exc}。请重新输出完整合法 JSON。"
|
||
|
||
raise LLMError("模型未返回有效结果")
|
||
|
||
def _mock_generate(
|
||
self,
|
||
preferences: TravelPreferences,
|
||
candidates: list[dict[str, Any]],
|
||
previous: Optional[GeneratedItinerary],
|
||
user_message: Optional[str],
|
||
) -> GeneratedItinerary:
|
||
if not candidates:
|
||
raise LLMError("没有符合条件的候选地点")
|
||
|
||
max_items = {
|
||
Pace.RELAXED: 3,
|
||
Pace.MODERATE: 4,
|
||
Pace.COMPACT: 5,
|
||
}[preferences.pace]
|
||
target = 240 if preferences.duration.value == "half_day" else 480
|
||
start = datetime(2026, 1, 1, 9, 0)
|
||
elapsed = 0
|
||
items = []
|
||
previous_place: Optional[dict[str, Any]] = None
|
||
|
||
for place in candidates:
|
||
if len(items) >= max_items:
|
||
break
|
||
transfer_value: Optional[int] = 0
|
||
if previous_place:
|
||
transfer_value = self.knowledge.travel_minutes(
|
||
previous_place["id"],
|
||
place["id"],
|
||
preferences.transport.value,
|
||
)
|
||
transfer_for_math = transfer_value or 0
|
||
duration = int(place.get("recommendedMinutes", 75))
|
||
if elapsed + transfer_for_math + duration > target + 30 and items:
|
||
continue
|
||
item_start = start + timedelta(minutes=transfer_for_math)
|
||
end = item_start + timedelta(minutes=duration)
|
||
items.append(
|
||
{
|
||
"startTime": item_start.strftime("%H:%M"),
|
||
"endTime": end.strftime("%H:%M"),
|
||
"placeId": place["id"],
|
||
"placeName": place["name"],
|
||
"activity": place["summary"],
|
||
"reason": self._reason(preferences, place),
|
||
"transferFromPreviousMinutes": (
|
||
transfer_value if previous_place else 0
|
||
),
|
||
"tips": place.get("tips", [])[:2],
|
||
}
|
||
)
|
||
elapsed += transfer_for_math + duration
|
||
start = end
|
||
previous_place = place
|
||
|
||
themes = "、".join(preferences.themes or preferences.interests[:1])
|
||
adjustment = f";已响应“{user_message}”" if user_message else ""
|
||
pace_label = {
|
||
Pace.RELAXED: "轻松",
|
||
Pace.MODERATE: "适中",
|
||
Pace.COMPACT: "紧凑",
|
||
}[preferences.pace]
|
||
return GeneratedItinerary.model_validate(
|
||
{
|
||
"title": f"光明区{themes or '精选'}{'半日' if preferences.duration.value == 'half_day' else '一日'}游",
|
||
"summary": f"以科学、人文与都市自然为线索,按{pace_label}节奏安排{adjustment}。",
|
||
"totalMinutes": max(elapsed, 1),
|
||
"estimatedCostText": "费用以场馆、景区及实际交通信息为准",
|
||
"items": items,
|
||
"notes": [
|
||
"开放时间、预约和票价请在出行前通过官方渠道再次确认。",
|
||
"交通耗时为POC估算,请以出发时的实际导航为准。",
|
||
],
|
||
}
|
||
)
|
||
|
||
@staticmethod
|
||
def _reason(
|
||
preferences: TravelPreferences, place: dict[str, Any]
|
||
) -> str:
|
||
matches = list(
|
||
(set(preferences.themes) & set(place.get("themes", [])))
|
||
| (set(preferences.interests) & set(place.get("interests", [])))
|
||
)
|
||
return (
|
||
f"符合你的{'、'.join(sorted(matches))}偏好"
|
||
if matches
|
||
else "作为光明区同路线备选,便于控制整体节奏"
|
||
)
|