Агент с 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