Files
gmTouringMiniApp/files/归档/server/app/planner.py
T

150 lines
5.0 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
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,
)