add agent_with_memory.py
This commit is contained in:
@@ -0,0 +1,173 @@
|
|||||||
|
#!/usr/bin/env python3
|
||||||
|
# -*- coding: utf-8 -*-
|
||||||
|
|
||||||
|
"""
|
||||||
|
Практическое задание №3 – Memory + Confirmation (Rich UI).
|
||||||
|
Запуск:
|
||||||
|
python agent_with_memory.py
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import json
|
||||||
|
import uuid
|
||||||
|
from pathlib import Path
|
||||||
|
from typing import Any, Dict, Iterable, List, Optional, Tuple
|
||||||
|
|
||||||
|
# ────────────────────── 1. Библиотеки ───────────────────────
|
||||||
|
from langchain_openai import ChatOpenAI # пример LLM‑модели
|
||||||
|
from langgraph.checkpoint.memory import MemorySaver
|
||||||
|
from langgraph.constants import START
|
||||||
|
from langgraph.graph import StateGraph
|
||||||
|
from langgraph.types import Command, interrupt
|
||||||
|
from rich.console import Console
|
||||||
|
|
||||||
|
# ────────────────────── 2. Настройки и утилиты ───────────────
|
||||||
|
console = Console()
|
||||||
|
OPENAI_API_KEY = "sk-..." # <-- ваш ключ OpenAI (или используйте переменную окружения)
|
||||||
|
llm = ChatOpenAI(temperature=0, model="gpt-4o-mini", openai_api_key=OPENAI_API_KEY)
|
||||||
|
|
||||||
|
# Пример простого инструмента
|
||||||
|
def get_price(args: Dict[str, Any]) -> str:
|
||||||
|
"""Возвращает цену (фиктивный ответ)."""
|
||||||
|
city = args.get("city")
|
||||||
|
date = args.get("date")
|
||||||
|
return f"Цена в {city} на {date}: 123₽"
|
||||||
|
|
||||||
|
# ────────────────────── 3. Структура состояния ───────────────
|
||||||
|
class State(dict):
|
||||||
|
"""Состояние агента – список сообщений."""
|
||||||
|
messages: List[Dict[str, Any]]
|
||||||
|
|
||||||
|
|
||||||
|
# ────────────────────── 4. Создание графа (агента) ────────
|
||||||
|
def build_agent() -> Tuple[StateGraph, MemorySaver]:
|
||||||
|
# ├─ хранилище памяти
|
||||||
|
memory = MemorySaver()
|
||||||
|
|
||||||
|
# └─ граф
|
||||||
|
builder = StateGraph(State)
|
||||||
|
|
||||||
|
# Узел генерации ответа LLM
|
||||||
|
@builder.node()
|
||||||
|
def llm_node(state: State) -> dict:
|
||||||
|
user_msg = state["messages"][-1]
|
||||||
|
response = llm.invoke(
|
||||||
|
[
|
||||||
|
{"role": "system", "content": SYSTEM_PROMPT},
|
||||||
|
*state["messages"],
|
||||||
|
]
|
||||||
|
)
|
||||||
|
# Добавляем системный ответ
|
||||||
|
new_message = {
|
||||||
|
"role": "assistant",
|
||||||
|
"content": response.content,
|
||||||
|
"tool_calls": response.tool_calls or [],
|
||||||
|
}
|
||||||
|
return {"messages": state["messages"] + [new_message]}
|
||||||
|
|
||||||
|
# Узел вызова инструмента
|
||||||
|
@builder.node()
|
||||||
|
def tool_node(state: State) -> dict:
|
||||||
|
# Последнее сообщение содержит tool_calls
|
||||||
|
last_msg = state["messages"][-1]
|
||||||
|
tool_call = last_msg["tool_calls"][0]
|
||||||
|
name, args = tool_call["name"], json.loads(tool_call["arguments"])
|
||||||
|
if name == "get_price":
|
||||||
|
result = get_price(args)
|
||||||
|
else:
|
||||||
|
result = f"Unknown tool {name}"
|
||||||
|
# Добавляем результат как новое сообщение
|
||||||
|
return {
|
||||||
|
"messages": state["messages"]
|
||||||
|
+ [
|
||||||
|
{
|
||||||
|
"role": "tool",
|
||||||
|
"content": result,
|
||||||
|
"tool_call_id": tool_call["id"],
|
||||||
|
}
|
||||||
|
]
|
||||||
|
}
|
||||||
|
|
||||||
|
builder.add_node("llm", llm_node)
|
||||||
|
builder.add_node("tool", tool_node)
|
||||||
|
|
||||||
|
# Переходы
|
||||||
|
builder.set_entry_point("llm")
|
||||||
|
builder.add_conditional_edges(
|
||||||
|
"llm",
|
||||||
|
lambda x: "tool" if any(m.get("tool_calls") for m in x["messages"]) else END,
|
||||||
|
{"tool": "tool", END: END},
|
||||||
|
)
|
||||||
|
builder.add_edge("tool", END)
|
||||||
|
|
||||||
|
# Компилируем с чекпоинтером
|
||||||
|
graph = builder.compile(checkpointer=memory, interrupt_before=["tools"])
|
||||||
|
return graph, memory
|
||||||
|
|
||||||
|
|
||||||
|
SYSTEM_PROMPT = """
|
||||||
|
Ты — помощник. В процессе разговора можешь вызывать инструмент get_price(city, date).
|
||||||
|
Перед каждым вызовом инструмента агент должен остановиться и запросить подтверждение у пользователя.
|
||||||
|
"""
|
||||||
|
|
||||||
|
|
||||||
|
# ────────────────────── 5. Функция «разговаривать» ───────────
|
||||||
|
def ask_and_run(user_input: Optional[Dict[str, Any]], config: dict):
|
||||||
|
"""Обрабатывает поток от агента, ловит паузы и спрашивает подтверждение."""
|
||||||
|
for chunk_type, chunk_data in agent.stream(
|
||||||
|
user_input,
|
||||||
|
config=config,
|
||||||
|
stream_mode=["messages", "updates"],
|
||||||
|
):
|
||||||
|
# ├─ 1. Вывод токенов
|
||||||
|
if chunk_type == "messages":
|
||||||
|
console.print(chunk_data["content"], end="", style="green")
|
||||||
|
continue
|
||||||
|
|
||||||
|
# ├─ 2. Информация о вызове инструмента (если есть)
|
||||||
|
if chunk_type == "updates" and "tool_calls" in chunk_data:
|
||||||
|
tool_call = chunk_data["tool_calls"][0]
|
||||||
|
name, args = tool_call["name"], json.loads(tool_call["arguments"])
|
||||||
|
console.print(f"\n[bold cyan]Агент хочет вызвать инструмент:[/bold cyan]")
|
||||||
|
console.print(f"[yellow]{name}({args})[/yellow]\n")
|
||||||
|
|
||||||
|
# ├─ 3. Пауза перед вызовом инструмента
|
||||||
|
if "__interrupt__" in chunk_data and state.next == ("tools",):
|
||||||
|
# Получаем состояние (поскольку после pause state ещё не обновился)
|
||||||
|
state = agent.get_state(config)
|
||||||
|
last_msg = state["messages"][-1]
|
||||||
|
tool_call = last_msg["tool_calls"][0]
|
||||||
|
name, args = tool_call["name"], json.loads(tool_call["arguments"])
|
||||||
|
console.print(f"[bold red]Пауза: запрос на вызов инструмента {name}[/bold red]")
|
||||||
|
answer = console.input("[green]Разрешить? (Y/n): [/green]").strip().lower()
|
||||||
|
if answer in ("y", ""):
|
||||||
|
# Возобновляем выполнение с тем же config
|
||||||
|
ask_and_run(None, config)
|
||||||
|
else:
|
||||||
|
console.print("[red]Отменено[/red]")
|
||||||
|
break
|
||||||
|
|
||||||
|
# └─ 4. Любой другой чанк – просто выводим
|
||||||
|
if chunk_type == "updates":
|
||||||
|
console.print(json.dumps(chunk_data, ensure_ascii=False), style="magenta")
|
||||||
|
|
||||||
|
|
||||||
|
# ────────────────────── 6. Основной цикл ─────────────────────
|
||||||
|
if __name__ == "__main__":
|
||||||
|
agent, memory = build_agent()
|
||||||
|
|
||||||
|
thread_id = str(uuid.uuid4()) # можно задать свой id (например "chat-1")
|
||||||
|
config = {"configurable": {"thread_id": thread_id}}
|
||||||
|
|
||||||
|
console.print("[bold blue]Добро пожаловать![/bold blue]")
|
||||||
|
while True:
|
||||||
|
user_msg = console.input("\n[bold]Вы:[/bold] ")
|
||||||
|
if user_msg.lower() in ("exit", "quit"):
|
||||||
|
break
|
||||||
|
|
||||||
|
# Формируем сообщение от пользователя
|
||||||
|
payload = {"messages": [{"role": "human", "content": user_msg}]}
|
||||||
|
ask_and_run(payload, config)
|
||||||
|
|
||||||
|
console.print("[bold blue]До свидания![/bold blue]")
|
||||||
Reference in New Issue
Block a user