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)