Files

107 lines
5.2 KiB
Python
Raw Permalink 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.
from langchain_openai import ChatOpenAI
from pydantic import SecretStr
from langchain.agents import create_agent
from langchain.tools import tool
from langgraph.checkpoint.memory import MemorySaver
from rich.console import Console
# Инициализация LLM с использованием плейсхолдеров
llm = ChatOpenAI(
model="google/gemma-4-26b-a4b",
base_url="http://192.168.0.120:1234/v1",
api_key=SecretStr("lm-studio"),
temperature=0.7,
)
console = Console()
# Определение инструмента
@tool
def get_price(city: str, date: str):
"""Возвращает прогноз погоды (цены/состояние) для указанного города и даты."""
# Имитация логики
return f"В городе {city} на дату {date} ожидается солнечная погода, +20°C."
tools = [get_price]
# Настройка памяти и агента с механизмом прерывания (interrupt)
memory = MemorySaver()
agent = create_agent(
model=llm,
tools=tools,
system_prompt="Ты полезный помощмущник. Если пользователь спрашивает о погоде, используй инструмент get_price.",
checkpointer=memory,
interrupt_before=['tools'],
)
# Конфигурация потока (thread_id обеспечивает память разговора)
config = {"configurable": {"thread_id": "chat-session-123"}}
def ask_and_run(user_input, config):
"""Основная функция обработки сообщений и управления циклом подтверждения."""
# Если user_input is None, мы просто продолжаем выполнение (возобновление после паузы)
# В LangGraph для возобновления через stream передается None или пустой список сообщений
stream_input = user_input if user_input is not None else []
# Используем stream_mode=['messages', 'updates'] согласно заданию
try:
for chunk in agent.stream(stream_input, config=config, stream_mode=['messages', 'updates']):
state = agent.get_state(config)
chunk_type, chunk_data = chunk
if chunk_type == 'messages':
# Потоковый вывод токенов (сообщений)
message, metadata = chunk_data
if hasattr(message, "content") and message.content:
print(message.append if hasattr(message, 'append') else message.content, end="", flush=True)
elif chunk_type == 'updates':
# Вывод информации об обновлениях (например, вызовы инструментов)
pass
# Проверка на прерывание перед вызовом инструмента
if '__interrupt__' in chunk_data and state.next == ('tools',):
console.print("\n" + "---" * 15)
# Извлекаем информацию о том, какой инструмент хочет вызвать агент
last_message = state.values['messages'][-1]
if hasattr(last_message, 'tool_calls') and last_message.tool_calls:
tool_call = last_message.tool_calls[0]
console.print(f"{tool_call['name']}({tool_call['args']})")
console.print(f"Агент хочет вызвать утилиту {tool_call['name']}({tool_call['args']})")
answer = input("Разрешить? (Y/n): ")
if answer.lower().strip() == 'y':
# Рекурсивный вызов для продолжения без нового сообщения пользователя
ask_and_run(None, config)
else:
console.print("Отменено")
return # Выходим из текущей итерации стрима
if user_input: # Печать переноса строки только если был новый ввод
print()
except Exception as e:
# Если произошла ошибка в процессе стрима (например, при рекурсии), пробрасываем её выше
raise e
if __name__ == "__main__":
console.print("[bold blue]Чат запущен. Напишите 'exit' для выхода.[/bold blue]")
while True:
try:
user_text = input("\nВы: ")
if user_text.lower().strip() == 'exit':
break
# Формируем входные данные для агента
input_payload = {"messages": [{"role": "human", "content": user_text}]}
ask_and_run(input_payload, config)
except KeyboardInterrupt:
break
except Exception as e:
console.print(f"[bold red]Ошибка: {e}[/bold red]")