forked from zhouruizhe/gmTouringMiniApp
149 lines
5.2 KiB
Python
149 lines
5.2 KiB
Python
from __future__ import annotations
|
|
|
|
import json
|
|
import os
|
|
from typing import Any
|
|
|
|
from pydantic import ValidationError
|
|
|
|
from .knowledge import ranked_candidates
|
|
from .prompts import SYSTEM_PROMPT
|
|
from .schemas import (
|
|
ModelRecommendationSet,
|
|
PlanRecommendationRequest,
|
|
PlanningMode,
|
|
)
|
|
|
|
|
|
class LLMError(RuntimeError):
|
|
pass
|
|
|
|
|
|
class LLMTimeout(LLMError):
|
|
pass
|
|
|
|
|
|
class LLMNotConfigured(LLMError):
|
|
pass
|
|
|
|
|
|
class LLMService:
|
|
def __init__(self) -> None:
|
|
mode = os.getenv("LLM_MODE", "mock").strip().lower()
|
|
if mode not in {"mock", "real"}:
|
|
raise ValueError("LLM_MODE must be 'mock' or 'real'")
|
|
self.mode = mode
|
|
self.timeout = float(os.getenv("LLM_TIMEOUT_SECONDS", "45"))
|
|
|
|
@property
|
|
def configured(self) -> bool:
|
|
return bool(os.getenv("OPENAI_API_KEY") and os.getenv("OPENAI_MODEL"))
|
|
|
|
@property
|
|
def ready(self) -> bool:
|
|
return self.mode == "mock" or self.configured
|
|
|
|
async def recommend(
|
|
self, request: PlanRecommendationRequest
|
|
) -> ModelRecommendationSet:
|
|
if self.mode == "mock":
|
|
return self._mock_recommend(request)
|
|
return await self._real_recommend(request)
|
|
|
|
async def _real_recommend(
|
|
self, request: PlanRecommendationRequest
|
|
) -> ModelRecommendationSet:
|
|
if not self.configured:
|
|
raise LLMNotConfigured("真实模型未配置,请在服务端设置模型凭据")
|
|
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 = request.model_dump(by_alias=True, mode="json")
|
|
feedback = ""
|
|
allowed_ids = {candidate.id for candidate in request.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:
|
|
# Provider exceptions can include request metadata; do not echo them.
|
|
raise LLMError("模型调用失败") from exc
|
|
|
|
try:
|
|
content = response.choices[0].message.content or "{}"
|
|
generated = ModelRecommendationSet.model_validate_json(content)
|
|
self._validate_ids(generated, allowed_ids)
|
|
return generated
|
|
except (ValidationError, ValueError, json.JSONDecodeError) as exc:
|
|
if attempt == 1:
|
|
raise LLMError("模型输出未通过 POI 白名单校验") from exc
|
|
feedback = (
|
|
"\n上次输出不符合契约。请只使用候选白名单 ID,"
|
|
"并重新输出完整 JSON。"
|
|
)
|
|
|
|
raise LLMError("模型未返回有效推荐")
|
|
|
|
def _mock_recommend(
|
|
self, request: PlanRecommendationRequest
|
|
) -> ModelRecommendationSet:
|
|
candidates_by_id = {candidate.id: candidate for candidate in request.candidates}
|
|
if request.mode == PlanningMode.CUSTOM:
|
|
selected = [candidates_by_id[poi_id] for poi_id in request.selected_poi_ids]
|
|
else:
|
|
# This count is only an AI recommendation hint. The client remains
|
|
# responsible for calculating transfers and final itinerary time.
|
|
desired_count = max(1, min(8, request.duration_minutes // 75))
|
|
selected = ranked_candidates(request)[:desired_count]
|
|
|
|
return ModelRecommendationSet.model_validate(
|
|
{
|
|
"recommendations": [
|
|
{
|
|
"poiId": candidate.id,
|
|
"reason": self._mock_reason(candidate.name, candidate.tag_codes),
|
|
"order": index,
|
|
}
|
|
for index, candidate in enumerate(selected, start=1)
|
|
]
|
|
}
|
|
)
|
|
|
|
@staticmethod
|
|
def _mock_reason(name: str, tag_codes: list[str]) -> str:
|
|
if tag_codes:
|
|
return f"{name}与当前偏好标签较匹配,可作为路线候选。"
|
|
return f"{name}的推荐指数较高,可作为路线候选。"
|
|
|
|
@staticmethod
|
|
def _validate_ids(
|
|
generated: ModelRecommendationSet, allowed_ids: set[str]
|
|
) -> None:
|
|
ids = [item.poi_id for item in generated.recommendations]
|
|
unknown = set(ids) - allowed_ids
|
|
if unknown:
|
|
raise ValueError(f"unknown POI ids: {sorted(unknown)}")
|
|
if len(ids) != len(set(ids)):
|
|
raise ValueError("duplicate POI ids")
|
|
|