forked from zhouruizhe/gmTouringMiniApp
Initial commit: gmTouringMiniApp project
This commit is contained in:
@@ -0,0 +1,155 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from enum import Enum
|
||||
from typing import Literal, Optional
|
||||
|
||||
from pydantic import BaseModel, ConfigDict, Field, model_validator
|
||||
|
||||
|
||||
def to_camel(value: str) -> str:
|
||||
parts = value.split("_")
|
||||
return parts[0] + "".join(part.capitalize() for part in parts[1:])
|
||||
|
||||
|
||||
class APIModel(BaseModel):
|
||||
model_config = ConfigDict(
|
||||
alias_generator=to_camel,
|
||||
populate_by_name=True,
|
||||
extra="forbid",
|
||||
str_strip_whitespace=True,
|
||||
)
|
||||
|
||||
|
||||
class PlanningMode(str, Enum):
|
||||
QUICK = "quick"
|
||||
CUSTOM = "custom"
|
||||
|
||||
|
||||
class Duration(str, Enum):
|
||||
HALF_DAY = "half_day"
|
||||
FULL_DAY = "full_day"
|
||||
|
||||
|
||||
class Pace(str, Enum):
|
||||
RELAXED = "relaxed"
|
||||
MODERATE = "moderate"
|
||||
COMPACT = "compact"
|
||||
|
||||
|
||||
class Transport(str, Enum):
|
||||
WALKING = "walking"
|
||||
DRIVING = "driving"
|
||||
PUBLIC_TRANSPORT = "public_transport"
|
||||
|
||||
|
||||
class BudgetLevel(str, Enum):
|
||||
ECONOMY = "economy"
|
||||
STANDARD = "standard"
|
||||
QUALITY = "quality"
|
||||
|
||||
|
||||
Theme = Literal["亲子", "情侣", "朋友", "银发", "研学"]
|
||||
Interest = Literal["自然风光", "文化场馆", "生态科普", "美食", "摄影"]
|
||||
|
||||
|
||||
class TravelPreferences(APIModel):
|
||||
destination: Literal["深圳市光明区"] = "深圳市光明区"
|
||||
themes: list[Theme] = Field(default_factory=list, max_length=5)
|
||||
duration: Duration
|
||||
pace: Pace
|
||||
interests: list[Interest] = Field(default_factory=list, max_length=5)
|
||||
adults: int = Field(default=2, ge=0, le=20)
|
||||
children: int = Field(default=0, ge=0, le=20)
|
||||
child_ages: list[int] = Field(default_factory=list, max_length=20)
|
||||
transport: Transport = Transport.PUBLIC_TRANSPORT
|
||||
budget_level: BudgetLevel = BudgetLevel.STANDARD
|
||||
extra_requirements: str = Field(default="", max_length=500)
|
||||
|
||||
@model_validator(mode="after")
|
||||
def validate_group(self) -> "TravelPreferences":
|
||||
if self.adults + self.children < 1:
|
||||
raise ValueError("出行总人数至少为 1")
|
||||
if self.child_ages and len(self.child_ages) != self.children:
|
||||
raise ValueError("儿童年龄数量必须与儿童人数一致")
|
||||
if any(age < 0 or age > 17 for age in self.child_ages):
|
||||
raise ValueError("儿童年龄须在 0 至 17 岁之间")
|
||||
if len(set(self.themes)) != len(self.themes):
|
||||
raise ValueError("出行主题不能重复")
|
||||
if len(set(self.interests)) != len(self.interests):
|
||||
raise ValueError("兴趣偏好不能重复")
|
||||
return self
|
||||
|
||||
|
||||
class PoiCandidate(APIModel):
|
||||
"""Coordinate-free, presentation-free POI data that may reach the model."""
|
||||
|
||||
id: str = Field(min_length=1, max_length=100, pattern=r"^[A-Za-z0-9][A-Za-z0-9_-]*$")
|
||||
name: str = Field(min_length=1, max_length=100)
|
||||
category_code: str = Field(min_length=1, max_length=64)
|
||||
tag_codes: list[str] = Field(default_factory=list, max_length=20)
|
||||
summary: str = Field(default="", max_length=300)
|
||||
recommendation_index: int = Field(ge=1, le=5)
|
||||
|
||||
@model_validator(mode="after")
|
||||
def validate_tags(self) -> "PoiCandidate":
|
||||
if len(set(self.tag_codes)) != len(self.tag_codes):
|
||||
raise ValueError("POI 标签不能重复")
|
||||
return self
|
||||
|
||||
|
||||
class PlanRecommendationRequest(APIModel):
|
||||
mode: PlanningMode
|
||||
preferences: TravelPreferences
|
||||
duration_minutes: int = Field(ge=120, le=600, multiple_of=30)
|
||||
selected_poi_ids: list[str] = Field(default_factory=list, max_length=8)
|
||||
candidates: list[PoiCandidate] = Field(min_length=1, max_length=100)
|
||||
|
||||
@model_validator(mode="after")
|
||||
def validate_planning_scope(self) -> "PlanRecommendationRequest":
|
||||
candidate_ids = [candidate.id for candidate in self.candidates]
|
||||
if len(set(candidate_ids)) != len(candidate_ids):
|
||||
raise ValueError("候选 POI ID 不能重复")
|
||||
if len(set(self.selected_poi_ids)) != len(self.selected_poi_ids):
|
||||
raise ValueError("自选 POI ID 不能重复")
|
||||
|
||||
if self.mode == PlanningMode.QUICK:
|
||||
if self.selected_poi_ids:
|
||||
raise ValueError("快速规划不能预选 POI")
|
||||
elif not 2 <= len(self.selected_poi_ids) <= 8:
|
||||
raise ValueError("自选规划必须选择 2 至 8 个 POI")
|
||||
|
||||
unknown = set(self.selected_poi_ids) - set(candidate_ids)
|
||||
if unknown:
|
||||
raise ValueError(f"自选规划包含候选白名单外的 POI:{sorted(unknown)}")
|
||||
return self
|
||||
|
||||
|
||||
class ModelRecommendation(APIModel):
|
||||
poi_id: str = Field(min_length=1, max_length=100)
|
||||
reason: str = Field(min_length=1, max_length=300)
|
||||
order: Optional[int] = Field(default=None, ge=1, le=8)
|
||||
|
||||
|
||||
class ModelRecommendationSet(APIModel):
|
||||
recommendations: list[ModelRecommendation] = Field(min_length=1, max_length=8)
|
||||
|
||||
|
||||
class RecommendationItem(APIModel):
|
||||
poi_id: str
|
||||
reason: str
|
||||
order: int = Field(ge=1, le=8)
|
||||
|
||||
|
||||
class RecommendationResponse(APIModel):
|
||||
request_id: str
|
||||
mode: PlanningMode
|
||||
assistant_message: str
|
||||
recommendations: list[RecommendationItem] = Field(min_length=1, max_length=8)
|
||||
generation_mode: Literal["mock", "real"]
|
||||
|
||||
|
||||
class HealthResponse(APIModel):
|
||||
status: Literal["ok", "degraded"]
|
||||
mode: Literal["mock", "real"]
|
||||
configured: bool
|
||||
ready: bool
|
||||
Reference in New Issue
Block a user