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
    )