306 lines
12 KiB
Python
306 lines
12 KiB
Python
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‑in‑the‑Loop
|
||
# ------------------------------------------------------------------
|
||
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() |