72 lines
3.0 KiB
Python
72 lines
3.0 KiB
Python
import os
|
||
import sqlite3
|
||
from typing import cast, Optional
|
||
|
||
from langchain.agents import create_agent, AgentState
|
||
from langchain.agents.middleware import SummarizationMiddleware
|
||
from langchain.agents.middleware.summarization import ContextMessages
|
||
from langchain.chat_models import init_chat_model
|
||
from langchain_core.language_models import BaseChatModel
|
||
from langchain_core.messages import HumanMessage
|
||
from langchain_core.runnables import RunnableConfig
|
||
from langgraph.checkpoint.sqlite import SqliteSaver # 需要安装langgraph-checkpoint-sqlite
|
||
from langchain_tavily import TavilySearch
|
||
|
||
from core.config.settings import settings
|
||
|
||
# 初始化模型 需要多模态大模型 deepseek已支持
|
||
llm_model = init_chat_model(
|
||
model="deepseek-v4-flash",
|
||
model_provider="deepseek",
|
||
api_key=settings.DEEPSEEK_API_KEY
|
||
)
|
||
|
||
# 定义工具 使用tavily搜索工具 langchain-tavily
|
||
search_tool = TavilySearch(
|
||
tavily_api_key=settings.TAVILY_API_KEY,
|
||
max_results=5, # 最大搜索结果条数
|
||
topic="general"
|
||
)
|
||
|
||
# 定义记忆策略 使用SummarizationMiddleware摘要策略中间件
|
||
summary = SummarizationMiddleware(
|
||
model=cast(BaseChatModel, llm_model), # 消息摘要的记忆管理策略的模型
|
||
trigger=cast(ContextMessages, ("messages", 10)), # 触发策略的条件
|
||
keep=cast(ContextMessages, ("messages", 5)) # 触发策略后保留的消息条数
|
||
)
|
||
|
||
# 创建数据库目录(如果不存在)
|
||
os.makedirs("sqlite", exist_ok=True)
|
||
# 创建 SQLite 数据库连接 check_same_thread=False 是为了确保在多线程环境下的安全性
|
||
conn = sqlite3.connect("sqlite/checkpoints.db", check_same_thread=False)
|
||
# 初始化 checkpointer
|
||
checkpointer = SqliteSaver(conn)
|
||
# 自动建表
|
||
checkpointer.setup()
|
||
|
||
# agent提示词
|
||
system_prompt = """
|
||
你是一名有名的国宴厨师。收到用户提供的食材照片或清单后,按照以下流程步骤操作:
|
||
1.识别和评估食材:若用户提供照片,首先辨别所有可见食材,基于食材的外观状态,评估其新鲜度和可用量,整理出一份“可用食材清单”。
|
||
2.智能食谱检索:优先调用search_tool工具,以“可用食材清单”为核心关键词,查找可行菜谱。
|
||
3.多维度评估与排序:从营养与烹饪难度这2个维度对检索到的候选食谱进行量化打分,并根据得分进行排序,制作简单且营养丰富的排名靠前。
|
||
4.结构化方案输出:把排序后的食谱整理成一份结构清晰的建议报告,要包括食谱信息、得分、推荐理由、食谱的参考图片,帮助用户快速做出决策。
|
||
"""
|
||
|
||
# 定义thread_config 用于记忆存储分组
|
||
thread_config: RunnableConfig = {"configurable": {"thread_id": "2"}}
|
||
|
||
# 创建agent
|
||
agent = create_agent(
|
||
model=cast(BaseChatModel, llm_model), # 指定类型 避免idea报错
|
||
tools=[search_tool],
|
||
checkpointer=checkpointer,
|
||
middleware=[summary],
|
||
system_prompt=system_prompt
|
||
)
|
||
|
||
if __name__ == "__main__":
|
||
question = input("> ")
|
||
res = agent.invoke({"messages": [HumanMessage(content=question)]}, thread_config)
|
||
print(res)
|