tongue.py 4.8 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139
  1. """
  2. 舌诊 LangGraph — 视觉 LLM 舌象分析
  3. 流程:
  4. START → load_image → analyze_tongue → END
  5. 输出:整体评估 + 7 类结构化舌诊指标
  6. """
  7. import base64
  8. from typing import TypedDict, Optional
  9. from langgraph.graph import StateGraph, START, END
  10. from langchain_core.messages import HumanMessage
  11. from ..llm.client import get_vision_llm
  12. TONGUE_INDICATOR_CODES = [
  13. "tongue_color", "coating_color", "coating_texture",
  14. "fissure", "teeth_mark", "sublingual_vein", "constitution",
  15. ]
  16. # LLM 可能自创的 code 别名 → 规范 code(容错归一化)
  17. INDICATOR_ALIASES = {
  18. "tongue_shape": "tongue_color",
  19. "tongue_coating_color": "coating_color",
  20. "tongue_coating_texture": "coating_texture",
  21. "coating": "coating_color",
  22. "sublingual_veins": "sublingual_vein",
  23. "tooth_mark": "teeth_mark",
  24. "teeth_marks": "teeth_mark",
  25. "body_constitution": "constitution",
  26. }
  27. SYSTEM_PROMPT = (
  28. "你是资深中医舌诊专家。根据用户上传的舌象图片,输出结构化 JSON,"
  29. "不要输出任何 JSON 之外的文字。"
  30. "indicators 数组的 code 字段必须严格从以下 7 个值中选择,禁止自创或改写:"
  31. "tongue_color(舌色)、coating_color(苔色)、coating_texture(苔质)、"
  32. "fissure(裂纹)、teeth_mark(齿痕)、sublingual_vein(舌下络脉)、constitution(体质)。"
  33. "JSON 格式:"
  34. '{"overall_assessment": "整体舌象评估结论", "indicators": ['
  35. '{"code": "tongue_color", "value": "淡红"}, '
  36. '{"code": "coating_color", "value": "薄白"}]}'
  37. "指标值用简短中文描述,如舌色「淡红」、苔色「薄白」、裂纹「无」、齿痕「轻」。"
  38. )
  39. class GraphState(TypedDict):
  40. request: dict
  41. image_base64: Optional[str]
  42. raw_response: str
  43. overall_assessment: str
  44. indicators: list
  45. error: Optional[str]
  46. def load_image(state: GraphState) -> GraphState:
  47. req = state["request"]
  48. if req.get("image_base64"):
  49. return {**state, "image_base64": req["image_base64"]}
  50. # 支持 data URL 格式 (data:image/jpeg;base64,xxx) — TongueDiagnosisAgent 调用场景
  51. url = req.get("image_url", "")
  52. if url and url.startswith("data:image"):
  53. payload = url.split(",", 1)[1]
  54. return {**state, "image_base64": payload}
  55. if url:
  56. return {**state, "error": "舌诊暂不支持外部 URL,请传 image_base64 或 data URL"}
  57. return {**state, "error": "image_url 或 image_base64 至少提供一项"}
  58. def analyze_tongue(state: GraphState) -> GraphState:
  59. if state.get("error"):
  60. return state
  61. llm = get_vision_llm()
  62. # 优先使用 Java 后端传入的 prompt 模板(可配置),缺省回退内置 SYSTEM_PROMPT
  63. prompt_text = state["request"].get("prompt_template") or SYSTEM_PROMPT
  64. content = [
  65. {"type": "text", "text": prompt_text},
  66. {
  67. "type": "image_url",
  68. "image_url": {"url": f"data:image/jpeg;base64,{state['image_base64']}"},
  69. },
  70. ]
  71. response = llm.invoke([HumanMessage(content=content)])
  72. return {**state, "raw_response": response.content}
  73. def parse_result(state: GraphState) -> GraphState:
  74. raw = state.get("raw_response", "").strip()
  75. text = raw
  76. if "```json" in text:
  77. text = text.split("```json")[1].split("```")[0].strip()
  78. elif "```" in text:
  79. text = text.split("```")[1].split("```")[0].strip()
  80. import json
  81. try:
  82. data = json.loads(text)
  83. assessment = data.get("overall_assessment", "")
  84. indicators = []
  85. for it in data.get("indicators", []):
  86. code = it.get("code")
  87. if code in INDICATOR_ALIASES:
  88. code = INDICATOR_ALIASES[code]
  89. if code not in TONGUE_INDICATOR_CODES:
  90. continue
  91. indicators.append({"code": code, "value": it.get("value")})
  92. if not assessment or not indicators:
  93. return {**state, "error": "舌诊结果缺少评估或指标"}
  94. return {
  95. **state,
  96. "overall_assessment": assessment,
  97. "indicators": indicators,
  98. "error": None,
  99. }
  100. except Exception as e:
  101. return {**state, "error": f"舌诊 JSON 解析失败: {e}"}
  102. def build_tongue_graph():
  103. graph = StateGraph(GraphState)
  104. graph.add_node("load_image", load_image)
  105. graph.add_node("analyze_tongue", analyze_tongue)
  106. graph.add_node("parse_result", parse_result)
  107. graph.add_edge(START, "load_image")
  108. graph.add_edge("load_image", "analyze_tongue")
  109. graph.add_edge("analyze_tongue", "parse_result")
  110. graph.add_edge("parse_result", END)
  111. return graph.compile()
  112. _tongue_graph = None
  113. def get_tongue_graph():
  114. global _tongue_graph
  115. if _tongue_graph is None:
  116. _tongue_graph = build_tongue_graph()
  117. return _tongue_graph