Files
dz/solutions/699cc158d6d3a5544a3ed35b_Stream_режим_AI_агента/stream_agent.py
T

106 lines
4.5 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
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 GPT4o-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() # завершаем вывод новой строкой