Files

306 lines
12 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.
import sys
import logging
from typing import List
# ------------------------------------------------------------------
# 1️⃣ Основные импорты LangChain / LangGraph
# ------------------------------------------------------------------
from langchain_ollama import OllamaEmbeddings, OllamaLLM
from langchain_qdrant import QdrantVectorStore
from langchain.schema import Document, HumanMessage, AIMessage
from langchain.agents import create_agent, Tool
from langchain.agents.middleware import HumanInTheLoopMiddleware
from langgraph.checkpoint.memory import MemorySaver
from langgraph.types import Command
# ------------------------------------------------------------------
# 2️⃣ Настройка логирования
# ------------------------------------------------------------------
logging.basicConfig(
level=logging.INFO,
format="%(asctime)s [%(levelname)s] %(message)s",
handlers=[logging.StreamHandler(sys.stdout)],
)
logger = logging.getLogger(__name__)
# ------------------------------------------------------------------
# 3️⃣ Конфигурация Ollama и Qdrant
# ------------------------------------------------------------------
EMBEDDING_MODEL = "nomic-embed-text"
LLM_MODEL = "llama3"
embeddings = OllamaEmbeddings(model=EMBEDDING_MODEL)
vector_store = QdrantVectorStore(
embedding_function=embeddings,
url="http://localhost:6333",
collection_name="documents",
)
# ------------------------------------------------------------------
# 4️⃣ Создание LLM и вспомогательных инструментов
# ------------------------------------------------------------------
llm = OllamaLLM(model=LLM_MODEL, temperature=0.7)
def get_weather(city: str, date: str) -> str:
"""Пример простого инструмента получения погоды."""
# В реальном проекте здесь будет запрос к API погоды.
return f"Погода в {city} на {date}: солнечно, 25°C."
weather_tool = Tool(
name="get_weather",
func=get_weather,
description=(
"Получает прогноз погоды для указанного города и даты. "
"Аргументы: city (строка), date (строка)."
),
)
# ------------------------------------------------------------------
# 5️⃣ Создание агента с HumanInTheLoopMiddleware
# ------------------------------------------------------------------
memory = MemorySaver()
agent = create_agent(
model=llm,
tools=[weather_tool],
system_prompt="Ты полезный ассистент, который может использовать инструменты.",
middleware=[
HumanInTheLoopMiddleware(
interrupt_on={
"get_weather": True, # разрешаем все решения: approve, edit, reject
},
description_prefix="Подтвердите вызов инструмента",
),
],
checkpointer=memory,
)
# ------------------------------------------------------------------
# 6️⃣ Декоратор для обработки ошибок и логирования
# ------------------------------------------------------------------
def middleware(func):
"""Декоратор для обработки ошибок и логирования."""
def wrapper(*args, **kwargs):
try:
logger.info(f"Выполняется команда: {func.__name__}")
result = func(*args, **kwargs)
logger.info("Команда выполнена успешно")
return result
except Exception as e:
logger.exception(f"Ошибка в команде {func.__name__}: {e}")
print(f"[ERROR] {e}")
return wrapper
# ------------------------------------------------------------------
# 7️⃣ Функции работы с Qdrant (добавление и поиск)
# ------------------------------------------------------------------
@middleware
def add_document(text: str) -> None:
"""Добавляет документ в Qdrant."""
doc = Document(page_content=text)
vector_store.add_documents([doc])
print("Документ добавлен.")
@middleware
def search(query: str, k: int = 3) -> List[Document]:
"""Выполняет поиск по запросу и выводит результаты."""
results = vector_store.similarity_search_with_score(query, k=k)
if not results:
print("Ничего не найдено.")
return []
for idx, (doc, score) in enumerate(results, start=1):
print(f"\nРезультат {idx} (score: {score:.4f})")
snippet = doc.page_content[:500]
if len(doc.page_content) > 500:
snippet += "..."
print(snippet)
return [doc for doc, _ in results]
# ------------------------------------------------------------------
# 8️⃣ Парсинг команд CLI
# ------------------------------------------------------------------
def parse_command(line: str):
"""Разбирает строку команды и вызывает соответствующую функцию."""
if not line.strip():
return
parts = line.split(maxsplit=1)
cmd = parts[0].lower()
if cmd == "/add":
if len(parts) < 2:
print("[ERROR] Необходимо указать текст для добавления.")
return
add_document(parts[1])
elif cmd == "/search":
if len(parts) < 2:
print("[ERROR] Необходимо указать запрос для поиска.")
return
search(parts[1])
elif cmd == "/ask":
if len(parts) < 2:
print("[ERROR] Необходимо задать вопрос.")
return
handle_agent_query(parts[1])
elif cmd == "/quit":
print("Выход из программы.")
sys.exit(0)
else:
print(f"[ERROR] Нераспознанная команда: {cmd}")
# ------------------------------------------------------------------
# 9️⃣ Обработка запроса к агенту с Human‑intheLoop
# ------------------------------------------------------------------
def handle_agent_query(user_text: str):
"""Запускает агента и обрабатывает возможные паузы."""
config = {"configurable": {"thread_id": "сессия-1"}}
# Первый вызов
result = agent.invoke(
{"messages": [HumanMessage(content=user_text)]},
config=config,
)
while "__interrupt__" in result:
interrupt_value = result["__interrupt__"][0].value
action_requests = interrupt_value["action_requests"]
review_configs = interrupt_value["review_configs"]
decisions = []
print("\n--- Подтверждение ---")
for idx, (req, cfg) in enumerate(zip(action_requests, review_configs), start=1):
name = req.get("name", "unknown")
args = req.get("args", {})
description = req.get("description", "")
print(f"\n{idx}. Инструмент: {name}")
print(f" Аргументы: {args}")
if description:
print(f" Описание: {description}")
allowed = cfg.get("allowed_decisions", ["approve", "edit", "reject"])
prompt = f"a = approve, r = reject"
if "edit" in allowed:
prompt += ", e = edit"
while True:
choice = input(prompt + ": ").strip().lower()
if choice == "a":
decisions.append({"type": "approve"})
break
elif choice == "r":
msg = input("Сообщение для агента (причина отказа): ")
decisions.append({"type": "reject", "message": msg})
break
elif choice == "e" and "edit" in allowed:
# простое редактирование аргументов в формате JSON
new_args = input("Введите новые аргументы в формате JSON: ")
try:
import json
edited = json.loads(new_args)
decisions.append(
{
"type": "edit",
"edited_action": {"name": name, "args": edited},
}
)
break
except Exception as e:
print(f"[ERROR] Неверный JSON: {e}")
else:
print("Неверный ввод. Попробуйте снова.")
# Возобновляем агент с решениями
result = agent.invoke(
Command(resume={"decisions": decisions}),
config=config,
)
# После завершения выводим финальный ответ
if result.get("messages"):
final_msg = result["messages"][-1]
if isinstance(final_msg, AIMessage):
print(f"\nАгент: {final_msg.content}")
else:
print("\nАгент не вернул сообщение.")
else:
print("\nАгент завершил работу без ответа.")
# ------------------------------------------------------------------
# 🔄 Основной цикл CLI
# ------------------------------------------------------------------
def main():
"""Основной цикл CLI."""
logger.info("Запуск клиента Human-in-the-Loop")
print(
"Добро пожаловать! Используйте команды /add, /search, /ask и /quit."
)
while True:
try:
line = input("> ")
parse_command(line)
except KeyboardInterrupt:
print("\nПрервано пользователем. Выход.")
break
except EOFError:
print("\nEOF. Выход.")
break
# ------------------------------------------------------------------
# 10️⃣ UNIT TESTS (не меняем)
# ------------------------------------------------------------------
import unittest
from unittest.mock import patch, MagicMock
class TestClient(unittest.TestCase):
def setUp(self):
# Подменяем vector_store для тестов
self.original_vector_store = vector_store
mock_vs = MagicMock()
mock_vs.add_documents.return_value = None
mock_vs.similarity_search_with_score.return_value = [
(Document(page_content="Test content 1"), 0.9),
(Document(page_content="Test content 2"), 0.8),
]
global vector_store
vector_store = mock_vs
def tearDown(self):
global vector_store
vector_store = self.original_vector_store
@patch("builtins.print")
def test_add_document_success(self, mock_print):
add_document("Sample text")
mock_print.assert_called_with("Документ добавлен.")
@patch("builtins.print")
def test_search_results(self, mock_print):
results = search("query")
self.assertEqual(len(results), 2)
# Проверяем, что print был вызван хотя бы один раз
self.assertTrue(mock_print.call_count > 0)
@patch("builtins.print")
def test_parse_unknown_command(self, mock_print):
parse_command("/unknown")
mock_print.assert_called_with("[ERROR] Нераспознанная команда: /unknown")
if __name__ == "__main__":
if len(sys.argv) > 1 and sys.argv[1] == "test":
# Запуск тестов
unittest.main(argv=[sys.argv[0]])
else:
main()