71 lines
2.4 KiB
Python
71 lines
2.4 KiB
Python
import os
|
|
from typing import List
|
|
|
|
from fastapi import FastAPI, HTTPException
|
|
from pydantic import BaseModel
|
|
from langchain.embeddings.openai import OpenAIEmbeddings
|
|
from langchain.vectorstores.qdrant import Qdrant
|
|
from langchain.chains.question_answering import load_qa_chain
|
|
from langchain.llms.openai import ChatOpenAI
|
|
from rich.console import Console
|
|
|
|
# Конфигурация
|
|
QDRANT_URL = os.getenv("QDRANT_URL", "http://localhost:6333")
|
|
QDRANT_COLLECTION = os.getenv("QDRANT_COLLECTION", "rag_memory")
|
|
OPENAI_API_KEY = os.getenv("OPENAI_API_KEY")
|
|
if not OPENAI_API_KEY:
|
|
raise RuntimeError("Не задана переменная окружения OPENAI_API_KEY")
|
|
|
|
console = Console()
|
|
|
|
# Инициализация компонентов
|
|
embeddings = OpenAIEmbeddings(openai_api_key=OPENAI_API_KEY)
|
|
llm = ChatOpenAI(temperature=0, openai_api_key=OPENAI_API_KEY)
|
|
|
|
# Создание или подключение к коллекции Qdrant
|
|
vectorstore = Qdrant(
|
|
client=None,
|
|
collection_name=QDRANT_COLLECTION,
|
|
embeddings=embeddings,
|
|
url=QDRANT_URL,
|
|
)
|
|
|
|
app = FastAPI(title="RAG Agent")
|
|
|
|
class Document(BaseModel):
|
|
content: str
|
|
|
|
class QueryRequest(BaseModel):
|
|
question: str
|
|
top_k: int = 5
|
|
|
|
@app.post("/add_document")
|
|
def add_document(doc: Document):
|
|
"""
|
|
Добавляет документ в память агента.
|
|
"""
|
|
try:
|
|
vectorstore.add_texts([doc.content])
|
|
console.log(f"[green]Документ добавлен:[/green] {doc.content[:50]}...")
|
|
return {"status": "ok"}
|
|
except Exception as e:
|
|
console.print_exception()
|
|
raise HTTPException(status_code=500, detail=str(e))
|
|
|
|
@app.post("/ask")
|
|
def ask(request: QueryRequest):
|
|
"""
|
|
Отвечает на вопрос, используя RAG.
|
|
"""
|
|
try:
|
|
# Получаем похожие документы
|
|
docs = vectorstore.similarity_search_with_score(request.question, k=request.top_k)
|
|
contexts = [doc.page_content for doc, _ in docs]
|
|
console.log(f"[blue]Найдено контекстов:[/blue] {len(contexts)}")
|
|
|
|
chain = load_qa_chain(llm=llm, chain_type="stuff")
|
|
answer = chain.run(input_documents=[{"content": c} for c in contexts], question=request.question)
|
|
return {"answer": answer}
|
|
except Exception as e:
|
|
console.print_exception()
|
|
raise HTTPException(status_code=500, detail=str(e)) |