Team Ai
Apppublic

afulara/PythonicRAG-FastAPI-React

sourceHugging Faceapache-2.0updated 1y agoView on Hugging Face
0likes
main.py229 linesDownload Raw Back to backend
1from fastapi import FastAPI, UploadFile, File, HTTPException, WebSocket2from fastapi.middleware.cors import CORSMiddleware3from fastapi.responses import StreamingResponse4from fastapi.staticfiles import StaticFiles5from pydantic import BaseModel6from typing import List, Optional, Dict, AsyncGenerator7import os8from dotenv import load_dotenv9from aimakerspace.vectordatabase import VectorDatabase10from aimakerspace.openai_utils.embedding import EmbeddingModel11from aimakerspace.text_utils import CharacterTextSplitter, PDFLoader12from aimakerspace.openai_utils.prompts import (13    UserRolePrompt,14    SystemRolePrompt,15    AssistantRolePrompt,16)17from aimakerspace.openai_utils.chatmodel import ChatOpenAI18import asyncio19import tempfile20import shutil21import json22from uuid import uuid423 24# Load environment variables25load_dotenv()26 27app = FastAPI()28 29# Mount static files30app.mount("/", StaticFiles(directory="static", html=True), name="static")31 32# Configure CORS33app.add_middleware(34    CORSMiddleware,35    allow_origins=["http://localhost:3000"],36    allow_credentials=True,37    allow_methods=["*"],38    allow_headers=["*"],39)40 41# Initialize components42text_splitter = CharacterTextSplitter(chunk_size=1000, chunk_overlap=200)43chat_openai = ChatOpenAI()44 45# Define prompts46system_template = """\47You are a helpful assistant that provides concise, direct answers based on the provided context. 48If the answer cannot be found in the context, simply say "I don't know" or "The information is not available in the provided context."49Keep your answers brief and to the point."""50system_role_prompt = SystemRolePrompt(system_template)51 52user_prompt_template = """\53Context:54{context}55 56Question:57{question}58 59Answer the question concisely based on the context above."""60user_role_prompt = UserRolePrompt(user_prompt_template)61 62# Session management63sessions: Dict[str, Dict] = {}64 65class Query(BaseModel):66    text: str67    k: int = 468 69class DocumentResponse(BaseModel):70    text: str71    type: str  # 'answer' or 'context'72    score: Optional[float] = None73 74class RetrievalAugmentedQAPipeline:75    def __init__(self, llm: ChatOpenAI, vector_db_retriever: VectorDatabase) -> None:76        self.llm = llm77        self.vector_db_retriever = vector_db_retriever78 79    async def arun_pipeline(self, user_query: str, k: int = 4) -> AsyncGenerator[str, None]:80        # Get top k most relevant chunks81        context_list = self.vector_db_retriever.search_by_text(user_query, k=k)82        83        # Format context84        context_prompt = ""85        for context in context_list:86            context_prompt += context[0] + "\n"87 88        # Format prompts89        formatted_system_prompt = system_role_prompt.create_message()90        formatted_user_prompt = user_role_prompt.create_message(91            question=user_query, 92            context=context_prompt93        )94 95        # Stream only the LLM response96        async for chunk in self.llm.astream([formatted_system_prompt, formatted_user_prompt]):97            yield json.dumps({98                "type": "token",99                "text": chunk100            })101 102        # Send context information once at the end103        yield json.dumps({104            "type": "context",105            "context": [{"text": text, "score": score} for text, score in context_list]106        })107 108def process_file(file_path: str, file_name: str):109    if file_name.lower().endswith('.pdf'):110        loader = PDFLoader(file_path)111    else:112        raise HTTPException(status_code=400, detail="Only PDF files are supported")113        114    documents = loader.load_documents()115    texts = text_splitter.split_texts(documents)116    return texts117 118@app.post("/upload")119async def upload_document(file: UploadFile = File(...)):120    if not file.filename.lower().endswith('.pdf'):121        raise HTTPException(status_code=400, detail="Only PDF files are supported")122    123    try:124        # Read the file content directly into memory125        content = await file.read()126        127        # Create a temporary file in a directory we know exists128        temp_dir = "/tmp"  # Using /tmp which is writable in most environments129        os.makedirs(temp_dir, exist_ok=True)130        131        temp_path = os.path.join(temp_dir, f"upload_{file.filename}")132        133        # Write the content to the temporary file134        with open(temp_path, 'wb') as temp_file:135            temp_file.write(content)136        137        try:138            # Process the file139            texts = process_file(temp_path, file.filename)140            141            # Create a new session142            session_id = str(uuid4())143            vector_db = VectorDatabase()144            await vector_db.abuild_from_list(texts)145            146            # Store session data147            sessions[session_id] = {148                "vector_db": vector_db,149                "texts": texts150            }151            152            return {153                "session_id": session_id,154                "message": f"Document processed successfully. Added {len(texts)} chunks to the database."155            }156            157        finally:158            # Clean up the temporary file159            try:160                if os.path.exists(temp_path):161                    os.unlink(temp_path)162            except Exception as e:163                print(f"Warning: Could not delete temporary file: {e}")164                165    except Exception as e:166        raise HTTPException(status_code=500, detail=f"Error processing file: {str(e)}")167 168@app.post("/query/{session_id}")169async def query_documents(session_id: str, query: Query):170    if session_id not in sessions:171        raise HTTPException(status_code=404, detail="Session not found")172    173    try:174        session = sessions[session_id]175        vector_db = session["vector_db"]176        177        # Initialize RAG pipeline178        rag_pipeline = RetrievalAugmentedQAPipeline(179            llm=chat_openai,180            vector_db_retriever=vector_db181        )182        183        # Create streaming response184        async def generate():185            async for chunk in rag_pipeline.arun_pipeline(query.text, query.k):186                yield f"data: {chunk}\n\n"187        188        return StreamingResponse(189            generate(),190            media_type="text/event-stream"191        )192    except Exception as e:193        raise HTTPException(status_code=500, detail=str(e))194 195@app.websocket("/ws/{session_id}")196async def websocket_endpoint(websocket: WebSocket, session_id: str):197    await websocket.accept()198    199    if session_id not in sessions:200        await websocket.close(code=1008, reason="Session not found")201        return202    203    try:204        session = sessions[session_id]205        vector_db = session["vector_db"]206        207        while True:208            data = await websocket.receive_text()209            query = json.loads(data)210            211            # Initialize RAG pipeline212            rag_pipeline = RetrievalAugmentedQAPipeline(213                llm=chat_openai,214                vector_db_retriever=vector_db215            )216            217            # Stream response218            async for chunk in rag_pipeline.arun_pipeline(query["text"], query.get("k", 4)):219                await websocket.send_text(json.dumps({220                    "type": "token" if isinstance(chunk, str) else "context",221                    "text": chunk if isinstance(chunk, str) else chunk222                }))223            224    except Exception as e:225        await websocket.close(code=1011, reason=str(e))226 227if __name__ == "__main__":228    import uvicorn229    uvicorn.run(app, host="0.0.0.0", port=9000)