94 lines
3.8 KiB
Python
94 lines
3.8 KiB
Python
# agent.py
|
|
|
|
from langchain.agents import create_agent
|
|
from langchain.tools import tool
|
|
from langgraph.checkpoint.memory import MemorySaver
|
|
from langchain_ollama import ChatOllama
|
|
from rich import print as rprint
|
|
|
|
|
|
# ---------- Настройка LLM ----------
|
|
llm = ChatOllama(model="llama3") # используем Ollama, как указано в условии
|
|
|
|
# ---------- Определение инструмента ----------
|
|
@tool
|
|
def get_price(city: str, date: str) -> str:
|
|
"""
|
|
Возвращает цену на указанную дату для города.
|
|
В реальном проекте здесь будет запрос к API или БД.
|
|
Для демонстрации возвращаем фиктивный результат.
|
|
"""
|
|
return f"Цена в {city} на {date}: 100₽"
|
|
|
|
# ---------- Память ----------
|
|
memory = MemorySaver()
|
|
|
|
# ---------- Создание агента с паузой перед инструментом ----------
|
|
agent = create_agent(
|
|
model=llm,
|
|
tools=[get_price],
|
|
system_prompt="""
|
|
Ты помощник, который может использовать инструменты.
|
|
Перед каждым вызовом инструмента агент должен остановиться и дождаться подтверждения пользователя.
|
|
""",
|
|
checkpointer=memory,
|
|
interrupt_before=["tools"], # пауза перед инструментом
|
|
)
|
|
|
|
# ---------- Конфигурация разговора ----------
|
|
config = {"configurable": {"thread_id": "разговор-1"}}
|
|
|
|
|
|
def ask_and_run(user_input, config):
|
|
"""
|
|
Обрабатывает вход пользователя и выводит потоковый ответ агента.
|
|
При необходимости запрашивает подтверждение перед вызовом инструмента.
|
|
"""
|
|
# Запускаем потоковую генерацию
|
|
for chunk in agent.stream(
|
|
user_input,
|
|
config=config,
|
|
stream_mode=["messages", "updates"],
|
|
):
|
|
state = agent.get_state(config)
|
|
chunk_type, chunk_data = chunk
|
|
|
|
# Печатаем токены сообщения
|
|
if chunk_type == "messages":
|
|
rprint(chunk_data["content"], end="")
|
|
|
|
# Печатаем вызовы инструментов (если они уже выполнены)
|
|
elif chunk_type == "updates" and "tool_calls" in chunk_data:
|
|
for call in chunk_data["tool_calls"]:
|
|
name = call["name"]
|
|
args = call.get("args", {})
|
|
rprint(f"\n{name}({args})")
|
|
|
|
# Обрабатываем паузу перед инструментом
|
|
if "__interrupt__" in chunk_data and state.next == ("tools",):
|
|
# Получаем информацию о предстоящем вызове инструмента
|
|
tool_call = state.values["messages"][-1].tool_calls[0]
|
|
name = tool_call["name"]
|
|
args = tool_call.get("args", {})
|
|
rprint(f"\nАгент хочет вызвать утилиту {name}({args})")
|
|
answer = input("Разрешить? (Y/n): ")
|
|
|
|
if answer.lower().strip() == "y":
|
|
# Возобновляем работу агента с того места, где остановились
|
|
ask_and_run(None, config)
|
|
else:
|
|
rprint("\nОтменено")
|
|
break
|
|
|
|
|
|
# ---------- Чат‑цикл ----------
|
|
if __name__ == "__main__":
|
|
while True:
|
|
user_input = input("\nВы: ")
|
|
if user_input.lower() in {"exit", "quit"}:
|
|
break
|
|
|
|
ask_and_run(
|
|
{"messages": [{"role": "human", "content": user_input}]},
|
|
config,
|
|
) |