Spaces:
Runtime error
Runtime error
File size: 3,130 Bytes
1a2b9e6 55943c5 1a2b9e6 dbb9b6d 1a2b9e6 55943c5 1a2b9e6 55943c5 1a2b9e6 55943c5 1a2b9e6 | 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 | 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
)
|