singhankur01's picture
Update app.py
5eba371 verified
Raw
History Blame
6.66 kB
import os
import json
import re
import asyncio
import time
import numpy as np
from cachetools import TTLCache
from contextlib import asynccontextmanager
from dotenv import load_dotenv
from operator import itemgetter
from fastapi import FastAPI, Depends, HTTPException, Header
from fastapi.responses import JSONResponse
from utils.DocsLoader import load_and_chunk
from utils.Schemas import RunRequest, RunResponse
from langchain_community.vectorstores import FAISS
from langchain_google_genai import ChatGoogleGenerativeAI
from langchain_huggingface import HuggingFaceEmbeddings
from langchain_community.retrievers import BM25Retriever
from langchain.retrievers import EnsembleRetriever
from sklearn.metrics.pairwise import cosine_similarity
from langchain.prompts import PromptTemplate
# Load environment variables
load_dotenv()
# Cache setup
document_cache = TTLCache(maxsize=5, ttl=300)
cache_lock = asyncio.Lock()
# --- 1. Lifespan Event Handler ---
ml_models = {}
@asynccontextmanager
async def lifespan(app: FastAPI):
print("🚀 Initializing models and prompt template...")
try:
GOOGLE_API_KEY = os.getenv("gemini_api_key")
if not GOOGLE_API_KEY:
raise RuntimeError("CRITICAL: Missing GOOGLE_API_KEY in environment secrets!")
# Optimized embedding model with batching
ml_models["embedder"] = HuggingFaceEmbeddings(
model_name="BAAI/bge-base-en-v1.5",
encode_kwargs={
'batch_size': 64,
'show_progress_bar': False
}
)
# Faster LLM with constrained output
ml_models["llm"] = ChatGoogleGenerativeAI(
model="gemini-1.5-pro",
api_key=GOOGLE_API_KEY,
temperature=0.1,
max_output_tokens=300
)
# Improved prompt template
ml_models["prompt_template"] = PromptTemplate.from_template("""
**Role**: Insurance Policy Expert
**Context**:
{context}
**Query**: {full_query}
**Instructions**:
1. If query contains age/gender/procedure/location/duration:
- Output ONLY JSON: {{"decision":"approved/rejected","amount":"₹X","justification":"Clause reference"}}
2. Else: Provide concise answer
3. NEVER mention document sources
4. If unsure, respond: "Insufficient information"
**Response**:
""")
print("✅ Models and prompt loaded successfully!")
except Exception as e:
print(f"❌ Lifespan error: {str(e)}")
raise e
yield
print("🧹 Cleaning up.")
ml_models.clear()
# --- 2. FastAPI App Instance ---
app = FastAPI(title="HackRX RAG Server", lifespan=lifespan)
# --- 3. API Key Verification ---
TEAM_API_KEY = os.getenv("TEAM_API_KEY")
def verify_api_key(authorization: str = Header(...)):
if not authorization.startswith("Bearer "):
raise HTTPException(status_code=401, detail="Invalid Authorization header format")
token = authorization.split("Bearer ")[1]
if token != TEAM_API_KEY:
raise HTTPException(status_code=403, detail="Invalid or missing API key")
# --- 4. Parsing Helper ---
def parse_llm_response(content: str) -> str:
try:
# Clean JSON response
content_cleaned = re.sub(r"^```json|```$", "", content.strip(), flags=re.IGNORECASE).strip()
data = json.loads(content_cleaned)
if "decision" in data:
return (
f"Decision: {data.get('decision', 'N/A').upper()}\n"
f"Amount: {data.get('amount', 'Not specified')}\n"
f"Justification: {data.get('justification', 'No justification provided')}"
)
elif "response" in data:
return data["response"]
return "Response format error"
except json.JSONDecodeError:
return content.strip()
except Exception as e:
return f"Response processing error: {str(e)}"
# --- 5. Retrieval Optimization ---
async def get_relevant_docs(question: str, vectorstore: FAISS, keyword_retriever: BM25Retriever):
# Parallel retrieval
dense_docs, sparse_docs = await asyncio.gather(
vectorstore.asimilarity_search(question, k=6),
asyncio.get_event_loop().run_in_executor(
None,
keyword_retriever.get_relevant_documents,
question
)
)
# Combine and deduplicate
all_docs = dense_docs + sparse_docs
unique_docs = {doc.page_content: doc for doc in all_docs}.values()
# Fast reranking
query_embedding = ml_models["embedder"].embed_query(question)
doc_texts = [doc.page_content for doc in unique_docs]
doc_embeddings = ml_models["embedder"].embed_documents(doc_texts)
similarities = cosine_similarity([query_embedding], doc_embeddings)[0]
sorted_indices = np.argsort(similarities)[::-1][:5] # Top 5
return [list(unique_docs)[i] for i in sorted_indices]
# --- 6. Main API Endpoint ---
@app.post("/api/v1/hackrx/run", response_model=RunResponse, dependencies=[Depends(verify_api_key)])
async def run_hackrx(req: RunRequest):
start_time = time.time()
# Cache document processing
async with cache_lock:
if req.documents in document_cache:
print("♻️ Using cached document")
chunks = document_cache[req.documents]
else:
chunks = load_and_chunk(str(req.documents))
document_cache[req.documents] = chunks
if not chunks:
return JSONResponse({"error": "No documents processed"}, status_code=400)
# Create vector store and retrievers
vectorstore = await FAISS.afrom_documents(chunks, ml_models["embedder"])
keyword_retriever = BM25Retriever.from_documents(chunks)
keyword_retriever.k = 6
# Process questions in parallel
async def process_question(q: str):
relevant_docs = await get_relevant_docs(q, vectorstore, keyword_retriever)
context = "\n".join([d.page_content for d in relevant_docs])
# Generate response
prompt = ml_models["prompt_template"].format_prompt(
full_query=q,
context=context
)
result = await ml_models["llm"].ainvoke(prompt)
return parse_llm_response(result.content)
answers = await asyncio.gather(*(process_question(q) for q in req.questions))
# Performance logging
proc_time = time.time() - start_time
print(f"⏱️ Processed {len(req.questions)} questions in {proc_time:.2f}s")
return JSONResponse({"answers": answers}, status_code=200)
@app.get("/", include_in_schema=False)
def root():
return {"message": "HackRX API operational. Use /api/v1/hackrx/run"}