Stream-режим AI-агента: agent.py
This commit is contained in:
@@ -1,13 +1,12 @@
|
||||
import os
|
||||
from typing import Dict, Any
|
||||
|
||||
from langgraph.graph import StateGraph, START
|
||||
from langgraph.prebuilt.tool_executor import ToolExecutorNode
|
||||
from langgraph.checkpoint.memory import MemorySaver
|
||||
from langchain_ollama import ChatOllama
|
||||
from nomic_embed_text import NomicEmbedText
|
||||
from langchain.tools import BaseTool
|
||||
from langchain.schema import HumanMessage, AIMessage, SystemMessage
|
||||
from langchain.agents import ToolExecutor, AgentExecutor, create_openai_tools_agent
|
||||
from langchain.prompts import ChatPromptTemplate
|
||||
|
||||
# ---------- Настройки LLM и эмбеддингов ----------
|
||||
LLM_MODEL = os.getenv("OLLAMA_MODEL", "llama3")
|
||||
@@ -30,58 +29,18 @@ class DummyTool(BaseTool):
|
||||
|
||||
tool = DummyTool()
|
||||
tools = [tool]
|
||||
tools_dict = {t.name: t for t in tools}
|
||||
|
||||
# ---------- Создание агента ----------
|
||||
def create_agent() -> StateGraph:
|
||||
graph = StateGraph()
|
||||
prompt_template = ChatPromptTemplate.from_messages(
|
||||
[
|
||||
SystemMessage(content="Ты помощник. Используй инструменты при необходимости."),
|
||||
("human", "{input}"),
|
||||
]
|
||||
)
|
||||
|
||||
def agent(state: Dict[str, Any]) -> Dict[str, Any]:
|
||||
messages = state.get("messages", [])
|
||||
# Добавляем системное сообщение с инструкцией
|
||||
system_msg = SystemMessage(
|
||||
content="Ты помощник. Используй инструменты при необходимости."
|
||||
)
|
||||
all_messages = [system_msg] + messages
|
||||
|
||||
# Запускаем LLM в режиме stream
|
||||
response = llm.invoke(all_messages, stream=True)
|
||||
return {"messages": [response]}
|
||||
|
||||
graph.add_node("agent", agent)
|
||||
|
||||
def tool_executor(state: Dict[str, Any]) -> Dict[str, Any]:
|
||||
last_msg = state["messages"][-1]
|
||||
if hasattr(last_msg, "tool_calls") and last_msg.tool_calls:
|
||||
# Выполняем первый вызов инструмента
|
||||
call = last_msg.tool_calls[0]
|
||||
name = call["name"]
|
||||
args = eval(call["args"])
|
||||
result = tools_dict[name].invoke(args)
|
||||
# Создаём сообщение с результатом
|
||||
tool_msg = AIMessage(
|
||||
content=result,
|
||||
tool_calls=[call],
|
||||
)
|
||||
return {"messages": [tool_msg]}
|
||||
return {}
|
||||
|
||||
graph.add_node("tool_executor", ToolExecutorNode(tools))
|
||||
graph.set_entry_point("agent")
|
||||
graph.add_edge(START, "agent")
|
||||
graph.add_conditional_edges(
|
||||
"agent",
|
||||
lambda x: "messages" in x and hasattr(x["messages"][-1], "tool_calls"),
|
||||
{
|
||||
True: "tool_executor",
|
||||
False: START,
|
||||
},
|
||||
)
|
||||
graph.set_finish_node("agent")
|
||||
|
||||
return graph.compile(checkpointer=MemorySaver())
|
||||
|
||||
tools_dict = {t.name: t for t in tools}
|
||||
graph = create_agent()
|
||||
agent = create_openai_tools_agent(llm, tools, prompt=prompt_template)
|
||||
executor = AgentExecutor(agent=agent, tools=tools, verbose=False)
|
||||
|
||||
# ---------- Функции форматирования ----------
|
||||
def format_message(message) -> str:
|
||||
@@ -108,7 +67,7 @@ def run_agent(user_input: str):
|
||||
step = 1
|
||||
# Инициализируем состояние с пользовательским сообщением
|
||||
init_state = {"messages": [HumanMessage(content=user_input)]}
|
||||
stream = graph.stream(init_state, stream_mode=["messages", "updates"])
|
||||
stream = executor.stream(init_state, stream_mode=["messages", "updates"])
|
||||
|
||||
for chunk_type, chunk_data in stream:
|
||||
if chunk_type == "messages":
|
||||
|
||||
Reference in New Issue
Block a user