Stream‑режим AI‑агента: stream_agent.py
This commit is contained in:
@@ -0,0 +1,106 @@
|
|||||||
|
import os
|
||||||
|
from typing import Dict, Tuple
|
||||||
|
|
||||||
|
from langchain_community.tools.tavily_search import TavilySearchResults
|
||||||
|
from langchain_core.messages import HumanMessage, AIMessage
|
||||||
|
from langgraph.graph import StateGraph
|
||||||
|
from langgraph.prebuilt import create_agent_executor
|
||||||
|
from langgraph.schema import MessagesState
|
||||||
|
from langgraph.utils import format_messages
|
||||||
|
from openai import OpenAI
|
||||||
|
|
||||||
|
# ------------------------------------------------------------------
|
||||||
|
# 1. Подключаемся к LLM (OpenAI GPT‑4o-mini)
|
||||||
|
# ------------------------------------------------------------------
|
||||||
|
client = OpenAI(api_key=os.getenv("OPENAI_API_KEY"))
|
||||||
|
llm = client.chat.completions.create
|
||||||
|
|
||||||
|
# ------------------------------------------------------------------
|
||||||
|
# 2. Определяем инструмент поиска
|
||||||
|
# ------------------------------------------------------------------
|
||||||
|
search_tool = TavilySearchResults(max_results=3)
|
||||||
|
|
||||||
|
# ------------------------------------------------------------------
|
||||||
|
# 3. Создаём агент (используем готовый LangGraph‑агент)
|
||||||
|
# ------------------------------------------------------------------
|
||||||
|
def agent_executor(messages: list[HumanMessage]) -> AIMessage:
|
||||||
|
"""Вызов агента с использованием LangGraph."""
|
||||||
|
# Создаём простую схему, где агент может вызвать инструмент поиска
|
||||||
|
graph = StateGraph(MessagesState)
|
||||||
|
|
||||||
|
@graph.node
|
||||||
|
def start(state: MessagesState) -> MessagesState:
|
||||||
|
return state
|
||||||
|
|
||||||
|
@graph.node
|
||||||
|
def tool(state: MessagesState) -> MessagesState:
|
||||||
|
last_msg = state.messages[-1]
|
||||||
|
if isinstance(last_msg, AIMessage) and last_msg.tool_calls:
|
||||||
|
# вызываем инструмент
|
||||||
|
tool_name = last_msg.tool_calls[0]["name"]
|
||||||
|
args = last_msg.tool_calls[0]["args"]
|
||||||
|
result = search_tool.run(args)
|
||||||
|
new_message = AIMessage(content=result)
|
||||||
|
state.messages.append(new_message)
|
||||||
|
return state
|
||||||
|
|
||||||
|
graph.add_edge("start", "tool")
|
||||||
|
graph.set_entry_point("start")
|
||||||
|
|
||||||
|
# Запускаем граф
|
||||||
|
final_state = graph.invoke({"messages": messages})
|
||||||
|
return final_state["messages"][-1]
|
||||||
|
|
||||||
|
# ------------------------------------------------------------------
|
||||||
|
# 4. Функции форматирования сообщений
|
||||||
|
# ------------------------------------------------------------------
|
||||||
|
def format_message(message: AIMessage) -> str:
|
||||||
|
"""Возвращает строку для печати из сообщения."""
|
||||||
|
if message.content:
|
||||||
|
return message.content
|
||||||
|
# Если сообщение содержит вызов инструмента
|
||||||
|
tool_call = message.tool_calls[0]
|
||||||
|
name = tool_call["name"]
|
||||||
|
args = tool_call["args"]
|
||||||
|
return f"{name}({args})"
|
||||||
|
|
||||||
|
def format_chunk_message(chunk: Tuple) -> None:
|
||||||
|
"""Обрабатывает чанк типа 'messages'."""
|
||||||
|
global current_step
|
||||||
|
message, meta = chunk # type: ignore[assignment]
|
||||||
|
step_num = meta.get("langgraph_step", 0)
|
||||||
|
if step_num != current_step:
|
||||||
|
current_step = step_num
|
||||||
|
print("\n --- --- --- \n")
|
||||||
|
if message.content:
|
||||||
|
print(message.content, end="", flush=True)
|
||||||
|
|
||||||
|
# ------------------------------------------------------------------
|
||||||
|
# 5. Запускаем потоковый вывод
|
||||||
|
# ------------------------------------------------------------------
|
||||||
|
if __name__ == "__main__":
|
||||||
|
# Вводим запрос от пользователя
|
||||||
|
user_query = input("Введите ваш вопрос: ")
|
||||||
|
|
||||||
|
# Инициализируем состояние сообщений
|
||||||
|
messages = [HumanMessage(content=user_query)]
|
||||||
|
|
||||||
|
# Создаём итератор потока (используем готовый LangGraph executor)
|
||||||
|
stream = create_agent_executor(
|
||||||
|
llm=llm,
|
||||||
|
tools=[search_tool],
|
||||||
|
agent_name="stream_agent",
|
||||||
|
stream_mode=["messages", "updates"],
|
||||||
|
).stream({"messages": messages})
|
||||||
|
|
||||||
|
current_step = 1
|
||||||
|
|
||||||
|
for chunk_type, chunk_data in stream:
|
||||||
|
if chunk_type == "messages":
|
||||||
|
format_chunk_message(chunk_data)
|
||||||
|
elif chunk_type == "updates":
|
||||||
|
# Обрабатываем события обновления (например, завершение шага)
|
||||||
|
if isinstance(chunk_data, dict) and "model" in chunk_data:
|
||||||
|
last_msg = chunk_data["model"]["messages"][-1]
|
||||||
|
print("\n" + format_message(last_msg))
|
||||||
|
print() # завершаем вывод новой строкой
|
||||||
Reference in New Issue
Block a user