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,
|
||
)
|