AnimeRAGSystem / src /api /main.py
Pushkar02-n's picture
Complete MVP with awesome UI
dbb9b6d
Raw
History Blame
3.13 kB
from fastapi import FastAPI, HTTPException, status
from fastapi.middleware.cors import CORSMiddleware
from pydantic import BaseModel, Field
import uvicorn
from src.api.schemas import RecommendationRequest, RecommendationResponse
import time
from config import settings
import traceback
from src.retrieval.rag_pipeline import AnimeRAGPipeline
app = FastAPI(title="Anime Recommendation API",
description="RAG-powered anime recommendation system",
version="1.0.0")
app.add_middleware(
CORSMiddleware,
allow_origins=["*"],
allow_credentials=True,
allow_methods=["*"],
allow_headers=["*"]
)
pipeline = None
def get_pipeline():
"""Lazy initialization of pipeline"""
global pipeline
if pipeline is None:
pipeline = AnimeRAGPipeline(retriever_k=10)
return pipeline
@app.get("/")
async def root():
"""Healthcheck Endpoint"""
return {
"status": "online",
"message": "Anime recommendation API",
"version": "1.0.0"
}
@app.post("/recommend", response_model=RecommendationResponse)
async def get_recommendations(request: RecommendationRequest):
"""
Get anime recommendation based on user query
Example request:
```json
{
"query": "Anime similar to Death Note but lighter",
"n_results": 5,
"min_score": 7.5
}
```
"""
try:
rag_pipeline = get_pipeline()
rag_pipeline.recommendation_n = request.n_results
filters = {}
if request.min_score:
filters["min_score"] = request.min_score
if request.genre_filter:
filters["genre_filter"] = request.genre_filter
start_time = time.time()
result = rag_pipeline.recommend(
user_query=request.query,
filters=filters if filters else None
)
end_time = time.time()
print(f"Retrieved anime : \n{result["retrieved_count"]}")
print(f"Result Recommendations: \n{result["recommendations"][:20]}")
return RecommendationResponse(
query=result["query"],
recommendations=result["recommendations"],
retrieved_count=result["retrieved_count"],
metadata={
"model": settings.model_name,
"retriever_k": rag_pipeline.retriever_k,
"Time taken for LLM + vector search": str(end_time - start_time)
}
)
except Exception as e:
traceback.print_exc()
raise HTTPException(status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
detail=f"Error processing request: {str(e)}")
@app.get("/stats")
async def get_stats():
"""Get system statistics"""
rag_pipeline = get_pipeline()
return {
"total_anime": rag_pipeline.retriever.collection.count(),
"embedding_model": "all-MiniLM-L6-v2",
"llm_model": settings.model_name,
"retrieval_k": rag_pipeline.retriever_k
}
if __name__ == "__main__":
uvicorn.run(
"src.api.main:app",
host="0.0.0.0",
port=8000,
reload=True
)