| 1234567891011121314151617181920212223242526272829303132333435363738394041424344454647484950515253545556575859606162636465666768697071 |
- 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
|