chat.py 1.5 KB

12345678910111213141516171819202122232425262728293031323334353637383940414243444546474849505152535455565758
  1. import uuid
  2. from fastapi import APIRouter
  3. from app.models.chat import ChatRequest, ChatResponse, SourceInfo
  4. from app.graphs.chat_graph import create_chat_graph
  5. router = APIRouter(prefix="/api/v1", tags=["chat"])
  6. _graph = None
  7. def get_graph():
  8. global _graph
  9. if _graph is None:
  10. _graph = create_chat_graph()
  11. return _graph
  12. @router.post("/chat", response_model=ChatResponse)
  13. async def chat(req: ChatRequest):
  14. """家庭聊天: 意图分类→上下文→LLM→记忆"""
  15. trace_id = str(uuid.uuid4())
  16. graph = get_graph()
  17. initial_state = {
  18. "query": req.query,
  19. "user_id": req.user_id,
  20. "conversation_id": req.conversation_id,
  21. "child_id": req.context.child_id if req.context else None,
  22. "intent": None,
  23. "context": None,
  24. "messages": None,
  25. "answer": None,
  26. "tasks": [],
  27. "sources": [],
  28. }
  29. config = {
  30. "configurable": {"thread_id": req.conversation_id or str(req.user_id)},
  31. }
  32. result = await graph.ainvoke(initial_state, config)
  33. sources = []
  34. for s in result.get("sources", []):
  35. sources.append(SourceInfo(
  36. type=s.get("type", "tool"),
  37. title=s.get("name", ""),
  38. ))
  39. conv_id = req.conversation_id or f"conv_{req.user_id}_{__import__('time').time()}"
  40. return ChatResponse(
  41. answer=result.get("answer", ""),
  42. conversation_id=conv_id,
  43. sources=sources,
  44. tasks=result.get("tasks", []),
  45. trace_id=trace_id,
  46. )