afulara/PythonicRAG-FastAPI-React
0
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) 