Агент с RAG-памятью: solution.py
This commit is contained in:
@@ -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)
|
||||||
Reference in New Issue
Block a user