| 12345678910111213141516171819202122232425262728293031323334353637383940414243444546474849505152535455565758 |
- import uuid
- from fastapi import APIRouter
- from app.models.chat import ChatRequest, ChatResponse, SourceInfo
- from app.graphs.chat_graph import create_chat_graph
- router = APIRouter(prefix="/api/v1", tags=["chat"])
- _graph = None
- def get_graph():
- global _graph
- if _graph is None:
- _graph = create_chat_graph()
- return _graph
- @router.post("/chat", response_model=ChatResponse)
- async def chat(req: ChatRequest):
- """家庭聊天: 意图分类→上下文→LLM→记忆"""
- trace_id = str(uuid.uuid4())
- graph = get_graph()
- initial_state = {
- "query": req.query,
- "user_id": req.user_id,
- "conversation_id": req.conversation_id,
- "child_id": req.context.child_id if req.context else None,
- "intent": None,
- "context": None,
- "messages": None,
- "answer": None,
- "tasks": [],
- "sources": [],
- }
- config = {
- "configurable": {"thread_id": req.conversation_id or str(req.user_id)},
- }
- result = await graph.ainvoke(initial_state, config)
- sources = []
- for s in result.get("sources", []):
- sources.append(SourceInfo(
- type=s.get("type", "tool"),
- title=s.get("name", ""),
- ))
- conv_id = req.conversation_id or f"conv_{req.user_id}_{__import__('time').time()}"
- return ChatResponse(
- answer=result.get("answer", ""),
- conversation_id=conv_id,
- sources=sources,
- tasks=result.get("tasks", []),
- trace_id=trace_id,
- )
|