106 lines
4.5 KiB
Python
106 lines
4.5 KiB
Python
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() # завершаем вывод новой строкой |