forked from zhouruizhe/gmTouringMiniApp
Initial commit: gmTouringMiniApp project
This commit is contained in:
@@ -0,0 +1,148 @@
|
||||
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")
|
||||
|
||||
Reference in New Issue
Block a user