forked from zhouruizhe/gmTouringMiniApp
150 lines
5.0 KiB
Python
150 lines
5.0 KiB
Python
from __future__ import annotations
|
||||
|
|
|
|||
|
|
import time
|
|||
|
|
import uuid
|
|||
|
|
from collections import OrderedDict
|
|||
|
|
from dataclasses import dataclass, field
|
|||
|
|
from typing import Optional
|
|||
|
|
|
|||
|
|
from .knowledge import KnowledgeBase
|
|||
|
|
from .llm import LLMService
|
|||
|
|
from .schemas import (
|
|||
|
|
GeneratedItinerary,
|
|||
|
|
Itinerary,
|
|||
|
|
PlanResponse,
|
|||
|
|
Source,
|
|||
|
|
TravelPreferences,
|
|||
|
|
)
|
|||
|
|
|
|||
|
|
|
|||
|
|
@dataclass
|
|||
|
|
class Session:
|
|||
|
|
preferences: TravelPreferences
|
|||
|
|
itinerary: GeneratedItinerary
|
|||
|
|
created_at: float
|
|||
|
|
history: list[str] = field(default_factory=list)
|
|||
|
|
|
|||
|
|
|
|||
|
|
class SessionStore:
|
|||
|
|
def __init__(self, max_size: int = 100, ttl_seconds: int = 7200) -> None:
|
|||
|
|
self.max_size = max_size
|
|||
|
|
self.ttl_seconds = ttl_seconds
|
|||
|
|
self.sessions: OrderedDict[str, Session] = OrderedDict()
|
|||
|
|
|
|||
|
|
def put(self, session: Session) -> str:
|
|||
|
|
self.cleanup()
|
|||
|
|
while len(self.sessions) >= self.max_size:
|
|||
|
|
self.sessions.popitem(last=False)
|
|||
|
|
session_id = str(uuid.uuid4())
|
|||
|
|
self.sessions[session_id] = session
|
|||
|
|
return session_id
|
|||
|
|
|
|||
|
|
def get(self, session_id: str) -> Optional[Session]:
|
|||
|
|
self.cleanup()
|
|||
|
|
session = self.sessions.get(session_id)
|
|||
|
|
if session:
|
|||
|
|
self.sessions.move_to_end(session_id)
|
|||
|
|
return session
|
|||
|
|
|
|||
|
|
def cleanup(self) -> None:
|
|||
|
|
deadline = time.time() - self.ttl_seconds
|
|||
|
|
expired = [
|
|||
|
|
key for key, session in self.sessions.items()
|
|||
|
|
if session.created_at < deadline
|
|||
|
|
]
|
|||
|
|
for key in expired:
|
|||
|
|
self.sessions.pop(key, None)
|
|||
|
|
|
|||
|
|
|
|||
|
|
class PlannerService:
|
|||
|
|
def __init__(
|
|||
|
|
self,
|
|||
|
|
knowledge: KnowledgeBase,
|
|||
|
|
llm: LLMService,
|
|||
|
|
store: SessionStore,
|
|||
|
|
) -> None:
|
|||
|
|
self.knowledge = knowledge
|
|||
|
|
self.llm = llm
|
|||
|
|
self.store = store
|
|||
|
|
|
|||
|
|
async def create(self, preferences: TravelPreferences) -> PlanResponse:
|
|||
|
|
candidates = self.knowledge.retrieve(preferences)
|
|||
|
|
if not candidates:
|
|||
|
|
raise LookupError("没有找到符合当前条件的光明区地点")
|
|||
|
|
generated = await self.llm.generate(preferences, candidates)
|
|||
|
|
generated = self._canonicalize(generated, candidates)
|
|||
|
|
session_id = self.store.put(
|
|||
|
|
Session(preferences=preferences, itinerary=generated, created_at=time.time())
|
|||
|
|
)
|
|||
|
|
return PlanResponse(
|
|||
|
|
conversation_id=session_id,
|
|||
|
|
assistant_message="已根据你的偏好生成光明区行程。",
|
|||
|
|
itinerary=self._with_sources(generated),
|
|||
|
|
)
|
|||
|
|
|
|||
|
|
async def adjust(self, session_id: str, message: str) -> PlanResponse:
|
|||
|
|
session = self.store.get(session_id)
|
|||
|
|
if not session:
|
|||
|
|
raise LookupError("会话不存在或已过期,请重新生成行程")
|
|||
|
|
updated_preferences = session.preferences.model_copy(
|
|||
|
|
update={
|
|||
|
|
"extra_requirements": (
|
|||
|
|
session.preferences.extra_requirements + ";" + message
|
|||
|
|
).strip(";")
|
|||
|
|
}
|
|||
|
|
)
|
|||
|
|
candidates = self.knowledge.retrieve(updated_preferences)
|
|||
|
|
if not candidates:
|
|||
|
|
raise LookupError("没有找到符合调整条件的光明区地点")
|
|||
|
|
generated = await self.llm.generate(
|
|||
|
|
updated_preferences, candidates, session.itinerary, message
|
|||
|
|
)
|
|||
|
|
generated = self._canonicalize(generated, candidates)
|
|||
|
|
session.preferences = updated_preferences
|
|||
|
|
session.itinerary = generated
|
|||
|
|
session.history.append(message)
|
|||
|
|
return PlanResponse(
|
|||
|
|
conversation_id=session_id,
|
|||
|
|
assistant_message="已按你的补充要求调整行程。",
|
|||
|
|
change_summary=f"已根据“{message}”更新相关安排。",
|
|||
|
|
itinerary=self._with_sources(generated),
|
|||
|
|
)
|
|||
|
|
|
|||
|
|
def _canonicalize(
|
|||
|
|
self,
|
|||
|
|
itinerary: GeneratedItinerary,
|
|||
|
|
candidates: list[dict],
|
|||
|
|
) -> GeneratedItinerary:
|
|||
|
|
candidate_map = {place["id"]: place for place in candidates}
|
|||
|
|
for item in itinerary.items:
|
|||
|
|
place = candidate_map.get(item.place_id)
|
|||
|
|
if not place:
|
|||
|
|
raise ValueError(f"行程包含知识库外地点:{item.place_id}")
|
|||
|
|
item.place_name = place["name"]
|
|||
|
|
item.tips = [
|
|||
|
|
tip for tip in item.tips if tip in place.get("tips", [])
|
|||
|
|
] or place.get("tips", [])[:2]
|
|||
|
|
return itinerary
|
|||
|
|
|
|||
|
|
def _with_sources(self, generated: GeneratedItinerary) -> Itinerary:
|
|||
|
|
seen: set[str] = set()
|
|||
|
|
sources: list[Source] = []
|
|||
|
|
for item in generated.items:
|
|||
|
|
if item.place_id in seen:
|
|||
|
|
continue
|
|||
|
|
seen.add(item.place_id)
|
|||
|
|
place = self.knowledge.by_id(item.place_id)
|
|||
|
|
if place:
|
|||
|
|
sources.append(
|
|||
|
|
Source(
|
|||
|
|
place_id=place["id"],
|
|||
|
|
source_name=place["sourceName"],
|
|||
|
|
source_url=place["sourceUrl"],
|
|||
|
|
updated_at=place["updatedAt"],
|
|||
|
|
)
|
|||
|
|
)
|
|||
|
|
return Itinerary(
|
|||
|
|
**generated.model_dump(),
|
|||
|
|
sources=sources,
|
|||
|
|
)
|