File size: 4,615 Bytes
20f1ed0
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
import os
from fastapi import FastAPI, UploadFile, File, HTTPException
from fastapi.responses import StreamingResponse
from pydantic import BaseModel
from typing import List
import uvicorn
import shutil
import tempfile
import json

from embeddings.embedder import get_embeddings
from vectorstore.faiss_store import FAISSVectorStore
from reranker.cross_encoder import CrossEncoderReranker
from rag.retriever import AdvancedRetriever
from rag.memory import ConversationMemory
from rag.chain import RAGChain

app = FastAPI(title="MedRAG AI Backend", description="Medical Report Q&A Assistant API")

# ==========================================
# Shared RAG Core Initialization
# ==========================================
embeddings_model = get_embeddings()
vectorstore = FAISSVectorStore(embeddings_model)
vectorstore.load_index()

reranker = CrossEncoderReranker()
memory = ConversationMemory()
retriever = AdvancedRetriever(vectorstore, reranker)

# The same exact RAGChain used in Gradio, exposing identical logic
rag_chain = RAGChain(retriever, memory)
# ==========================================

class AskRequest(BaseModel):
    query: str
    stream: bool = True

@app.get("/health")
def health_check():
    """Health check endpoint."""
    return {"status": "ok", "index_loaded": vectorstore.vectorstore is not None}

@app.post("/upload")
async def upload_documents(files: List[UploadFile] = File(...)):
    """Uploads PDFs, processes them, and indexes them into FAISS."""
    try:
        temp_dir = tempfile.mkdtemp()
        filepaths = []
        
        # Save uploaded files temporarily
        for file in files:
            file_path = os.path.join(temp_dir, file.filename)
            with open(file_path, "wb") as buffer:
                shutil.copyfileobj(file.file, buffer)
            filepaths.append(file_path)
            
        # Delegate document processing/indexing to the shared RAGChain
        result = rag_chain.index_documents(filepaths)
        
        # Cleanup
        shutil.rmtree(temp_dir)
        
        return result
        
    except Exception as e:
        raise HTTPException(status_code=500, detail=str(e))

@app.post("/ask")
async def ask_question(request: AskRequest):
    """Answers a question based on uploaded documents."""
    if vectorstore.vectorstore is None:
        raise HTTPException(status_code=400, detail="No documents indexed. Please upload documents first.")
        
    if not request.stream:
        response_data = rag_chain.ask(request.query, stream=False)
        context = [{"source": doc.metadata.get('filename'), "page": doc.metadata.get('page'), "content": doc.page_content, "score": float(score)} for doc, score in response_data["context_docs"]]
        
        return {
            "answer": response_data["answer"],
            "context": context,
            "standalone_query": response_data["standalone_query"],
            "time_taken": response_data["time_taken"]
        }

    # Streaming setup
    response_data = rag_chain.ask(request.query, stream=True)
    streamer = response_data["streamer"]
    
    context = [{"source": doc.metadata.get('filename'), "page": doc.metadata.get('page'), "content": doc.page_content, "score": float(score)} for doc, score in response_data["context_docs"]]
    
    def event_stream():
        # First chunk contains metadata
        metadata = {
            "type": "metadata",
            "context": context,
            "standalone_query": response_data["standalone_query"]
        }
        yield f"data: {json.dumps(metadata)}\n\n"
        
        # Subsequent chunks contain the streamed text
        full_answer = ""
        for text in streamer:
            full_answer += text
            yield f"data: {json.dumps({'type': 'chunk', 'text': text})}\n\n"
            
        # Update memory after stream completes
        rag_chain.memory.add_user_message(request.query)
        rag_chain.memory.add_assistant_message(full_answer.strip())
        
        yield "data: [DONE]\n\n"

    return StreamingResponse(event_stream(), media_type="text/event-stream")

@app.post("/summarize")
def summarize():
    """Generates a summary of all uploaded documents using Map-Reduce."""
    try:
        summary = rag_chain.summarize()
        return {"summary": summary}
    except Exception as e:
         raise HTTPException(status_code=400, detail=str(e))

@app.post("/clear_history")
def clear_history():
    rag_chain.memory.clear()
    return {"message": "Chat history cleared."}

if __name__ == "__main__":
    uvicorn.run("fastapi_app:app", host="0.0.0.0", port=8000, reload=True)