import json from typing import TypedDict, Optional from langgraph.graph import StateGraph, START, END from langchain_core.messages import HumanMessage from ..llm.client import get_llm from ..prompts.questionnaire import PARENT_TEMPLATE, CHILD_TEMPLATE from ..schemas.questionnaire import Questionnaire, GenerateRequest class GraphState(TypedDict): request: GenerateRequest _prompt: Optional[str] raw_response: str questionnaire: Optional[dict] error: Optional[str] def build_prompt(state: GraphState) -> GraphState: req = state["request"] template = PARENT_TEMPLATE if req.relationship_type == "parent" else CHILD_TEMPLATE prompt = template.format(member_name=req.member_name) return {**state, "_prompt": prompt} def call_llm(state: GraphState) -> GraphState: llm = get_llm() messages = [HumanMessage(content=state["_prompt"])] response = llm.invoke(messages) return {**state, "raw_response": response.content} def validate(state: GraphState) -> GraphState: raw = state["raw_response"] # 提取 JSON 块 text = raw.strip() if "```json" in text: text = text.split("```json")[1].split("```")[0].strip() elif "```" in text: text = text.split("```")[1].split("```")[0].strip() try: data = json.loads(text) q = Questionnaire(**data) if not q.validate_structure(): return {**state, "error": "问卷结构校验失败:至少需要1题且每题有options或scale"} return {**state, "questionnaire": data, "error": None} except Exception as e: return {**state, "error": f"JSON 解析失败: {str(e)}"} def build_questionnaire_graph(): graph = StateGraph(GraphState) graph.add_node("build_prompt", build_prompt) graph.add_node("call_llm", call_llm) graph.add_node("validate", validate) graph.add_edge(START, "build_prompt") graph.add_edge("build_prompt", "call_llm") graph.add_edge("call_llm", "validate") graph.add_edge("validate", END) return graph.compile() _questionnaire_graph = None def get_questionnaire_graph(): global _questionnaire_graph if _questionnaire_graph is None: _questionnaire_graph = build_questionnaire_graph() return _questionnaire_graph