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