intent_classifier.py 2.0 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657
  1. from enum import Enum
  2. from langchain_openai import ChatOpenAI
  3. from app.config import settings
  4. import logging
  5. logger = logging.getLogger(__name__)
  6. class Intent(str, Enum):
  7. CHAT = "chat" # 日常聊天 / 育儿交流
  8. ANALYSIS = "analysis" # 报告解读 / 数据分析
  9. RECOMMEND = "recommend" # 商品 / 活动推荐
  10. TASK = "task" # 任务创建 / 进度查询
  11. MIND = "mind" # 情绪疏导 / 心理支持
  12. HEALTH = "health" # 健康咨询 / 舌诊
  13. CLASSIFY_PROMPT = """从用户消息中识别意图, 只返回意图代码, 不要解释:
  14. - chat: 日常聊天、育儿交流、询问建议(非具体商品/报告)
  15. - analysis: 报告解读\数据分析\趋势查看(提到"报告"/"分析"/"评分")
  16. - recommend: 商品\活动\文章推荐(提到"推荐"/"买"/"吃什么"/"适合")
  17. - task: 任务相关(提到"任务"/"打卡"/"完成"/"奖励")
  18. - mind: 情绪问题\心理支持(提到"心情"/"难过"/"焦虑"/"不开心")
  19. - health: 健康咨询\舌诊(提到"舌"/"健康"/"体质")
  20. 用户消息: {query}
  21. 历史上下文: {context}
  22. 意图代码:
  23. """
  24. class IntentClassifier:
  25. def __init__(self):
  26. self.llm = ChatOpenAI(
  27. model=settings.llm_model,
  28. api_key=settings.llm_api_key,
  29. base_url=settings.llm_base_url,
  30. temperature=settings.llm_temperature,
  31. max_tokens=20,
  32. )
  33. async def classify(self, query: str, context: str = "") -> Intent:
  34. prompt = CLASSIFY_PROMPT.format(query=query[:200], context=context[:300])
  35. try:
  36. resp = await self.llm.ainvoke(prompt)
  37. intent_str = resp.content.strip().lower()
  38. for intent in Intent:
  39. if intent.value in intent_str:
  40. return intent
  41. logger.warning("无法识别的意图: %s, 默认 chat", intent_str)
  42. return Intent.CHAT
  43. except Exception as e:
  44. logger.warning("意图分类失败: %s, 默认 chat", e)
  45. return Intent.CHAT