Агент с RAG-памятью: solution.py

This commit is contained in:
2026-05-26 22:36:36 +00:00
parent ada78812a9
commit 5503acc7a0
@@ -0,0 +1,155 @@
import os
import logging
from typing import List, Optional
from fastapi import FastAPI, HTTPException
from fastapi.responses import JSONResponse
from pydantic import BaseModel
from langchain.embeddings import OpenAIEmbeddings
from langchain.vectorstores import Qdrant
from langchain.llms import OpenAI
from langchain.chains import RetrievalQA
from langchain.prompts import PromptTemplate
from langchain.text_splitter import RecursiveCharacterTextSplitter
from rich import print as rprint
from qdrant_client import QdrantClient
from qdrant_client.http import models as qdrant_models
# =======================
# Конфигурация и логирование
# =======================
logging.basicConfig(
level=logging.INFO,
format="%(asctime)s [%(levelname)s] %(name)s: %(message)s",
)
logger = logging.getLogger("rag_agent")
# =======================
# Параметры
# =======================
QDRANT_HOST = os.getenv("QDRANT_HOST", "localhost")
QDRANT_PORT = int(os.getenv("QDRANT_PORT", "6333"))
QDRANT_COLLECTION = os.getenv("QDRANT_COLLECTION", "rag_collection")
OPENAI_API_KEY = os.getenv("OPENAI_API_KEY")
if not OPENAI_API_KEY:
raise RuntimeError("OPENAI_API_KEY не задана в переменных окружения")
# =======================
# Инициализация Qdrant
# =======================
client = QdrantClient(host=QDRANT_HOST, port=QDRANT_PORT)
try:
client.retreive_collection(name=QDRANT_COLLECTION)
except Exception:
client.retreive_collection(name=QDRANT_COLLECTION, create=True)
# =======================
# Векторное хранилище
# =======================
embeddings = OpenAIEmbeddings(openai_api_key=OPENAI_API_KEY)
vectorstore = Qdrant(
client=client,
collection_name=QDRANT_COLLECTION,
embeddings=embeddings,
)
# =======================
# Модель и цепочка
# =======================
llm = OpenAI(openai_api_key=OPENAI_API_KEY, temperature=0.2)
prompt_template = PromptTemplate(
input_variables=["context", "question"],
template=(
"Ты эксперт по теме. Используй следующий контекст для ответа на вопрос.\n"
"Контекст:\n{context}\n\nВопрос: {question}\nОтвет:"
),
)
qa_chain = RetrievalQA.from_chain_type(
llm=llm,
chain_type="stuff",
retriever=vectorstore.as_retriever(search_kwargs={"k": 4}),
return_source_documents=True,
chain_type_kwargs={"prompt": prompt_template},
)
# =======================
# FastAPI приложение
# =======================
app = FastAPI(title="RAG Agent")
class DocumentIn(BaseModel):
text: str
chunk_size: Optional[int] = 1000
chunk_overlap: Optional[int] = 200
class QueryIn(BaseModel):
question: str
@app.post("/ingest")
async def ingest(doc: DocumentIn):
"""
Разбивает документ на чанки и сохраняет в Qdrant.
"""
try:
splitter = RecursiveCharacterTextSplitter(
chunk_size=doc.chunk_size,
chunk_overlap=doc.chunk_overlap,
separators=["\n\n", "\n", " ", ""],
)
chunks = splitter.split_text(doc.text)
if not chunks:
raise ValueError("Документ не содержит текста для разбивки")
ids = [f"{hash(chunk)}" for chunk in chunks]
vectorstore.add_texts(chunks, ids=ids)
logger.info(f"Добавлено {len(chunks)} чанков")
return JSONResponse({"status": "ok", "chunks": len(chunks)})
except Exception as e:
logger.exception("Ошибка при загрузке документа")
raise HTTPException(status_code=400, detail=str(e))
@app.post("/query")
async def query(q: QueryIn):
"""
Выполняет запрос к RAG-агенту.
"""
try:
result = qa_chain({"question": q.question})
answer = result["result"]
sources = [doc.metadata.get("source", "unknown") for doc in result["source_documents"]]
logger.info(f"Ответ на вопрос: {q.question}")
return {"answer": answer, "sources": sources}
except Exception as e:
logger.exception("Ошибка при обработке запроса")
raise HTTPException(status_code=500, detail=str(e))
# =======================
# Тесты (используются при запуске как скрипт)
# =======================
if __name__ == "__main__":
import uvicorn
import pytest
import sys
@pytest.fixture(scope="module")
def client_app():
from fastapi.testclient import TestClient
return TestClient(app)
def test_ingest_and_query(client_app):
doc = {
"text": "Python – это язык программирования высокого уровня, созданный Гвидо ван Россумом. Он известен своей простотой и читаемостью.",
"chunk_size": 50,
"chunk_overlap": 10,
}
r = client_app.post("/ingest", json=doc)
assert r.status_code == 200
r = client_app.post("/query", json={"question": "Кто создал Python?"})
assert r.status_code == 200
data = r.json()
assert "Python" in data["answer"]
# Запуск тестов
if "test" in sys.argv:
sys.exit(pytest.main([__file__]))
else:
uvicorn.run(app, host="0.0.0.0", port=8000)