questionnaire.py 2.2 KB

1234567891011121314151617181920212223242526272829303132333435363738394041424344454647484950515253545556575859606162636465666768697071
  1. import json
  2. from typing import TypedDict, Optional
  3. from langgraph.graph import StateGraph, START, END
  4. from langchain_core.messages import HumanMessage
  5. from ..llm.client import get_llm
  6. from ..prompts.questionnaire import PARENT_TEMPLATE, CHILD_TEMPLATE
  7. from ..schemas.questionnaire import Questionnaire, GenerateRequest
  8. class GraphState(TypedDict):
  9. request: GenerateRequest
  10. _prompt: Optional[str]
  11. raw_response: str
  12. questionnaire: Optional[dict]
  13. error: Optional[str]
  14. def build_prompt(state: GraphState) -> GraphState:
  15. req = state["request"]
  16. template = PARENT_TEMPLATE if req.relationship_type == "parent" else CHILD_TEMPLATE
  17. prompt = template.format(member_name=req.member_name)
  18. return {**state, "_prompt": prompt}
  19. def call_llm(state: GraphState) -> GraphState:
  20. llm = get_llm()
  21. messages = [HumanMessage(content=state["_prompt"])]
  22. response = llm.invoke(messages)
  23. return {**state, "raw_response": response.content}
  24. def validate(state: GraphState) -> GraphState:
  25. raw = state["raw_response"]
  26. # 提取 JSON 块
  27. text = raw.strip()
  28. if "```json" in text:
  29. text = text.split("```json")[1].split("```")[0].strip()
  30. elif "```" in text:
  31. text = text.split("```")[1].split("```")[0].strip()
  32. try:
  33. data = json.loads(text)
  34. q = Questionnaire(**data)
  35. if not q.validate_structure():
  36. return {**state, "error": "问卷结构校验失败:至少需要1题且每题有options或scale"}
  37. return {**state, "questionnaire": data, "error": None}
  38. except Exception as e:
  39. return {**state, "error": f"JSON 解析失败: {str(e)}"}
  40. def build_questionnaire_graph():
  41. graph = StateGraph(GraphState)
  42. graph.add_node("build_prompt", build_prompt)
  43. graph.add_node("call_llm", call_llm)
  44. graph.add_node("validate", validate)
  45. graph.add_edge(START, "build_prompt")
  46. graph.add_edge("build_prompt", "call_llm")
  47. graph.add_edge("call_llm", "validate")
  48. graph.add_edge("validate", END)
  49. return graph.compile()
  50. _questionnaire_graph = None
  51. def get_questionnaire_graph():
  52. global _questionnaire_graph
  53. if _questionnaire_graph is None:
  54. _questionnaire_graph = build_questionnaire_graph()
  55. return _questionnaire_graph