diff --git a/.gitignore b/.gitignore index 478b169..ed7b8b9 100644 --- a/.gitignore +++ b/.gitignore @@ -210,4 +210,4 @@ __marimo__/ raw/ processed/ -.DS_store \ No newline at end of file +.DS_store diff --git a/backend/api.py b/backend/api.py index 414484e..36eb72a 100644 --- a/backend/api.py +++ b/backend/api.py @@ -1,15 +1,24 @@ -import fastapi +from fastapi import FastAPI, HTTPException, status import boto3 import os import datetime -import time import random +import pickle +import joblib +import json +import numpy as np from typing import List from pydantic import BaseModel -from dotenv import load_dotenv +import io -# Load environment variables from .env file -load_dotenv() +class MyFavorites(BaseModel): + items: List[str] + userid: str + +class PredictionResponse(BaseModel): +# user_id: int +# req: MyFavorites + recs: List[str] # titles of the recommended books ''' To-Do: @@ -19,6 +28,20 @@ ''' ## Helper Functions +def load_supporting_tables_from_s3(table_name: str): + """ + Download and load supporting tables from S3 into memory without persisting it to disk. + """ + s3 = boto3.client("s3") + s3_bucket = 'readcrumbs' + + # Download model object as bytes into memory + response = s3.get_object(Bucket=s3_bucket, Key=table_name) + table_bytes = response['Body'].read().decode('utf-8') + table = json.loads(table_bytes) + + return table + def load_model_from_s3(model_name: str): """ Download and load an ML model file from S3 into memory without persisting it to disk. @@ -40,39 +63,27 @@ def load_model_from_s3(model_name: str): Returns: The loaded model object. """ - s3_bucket = os.environ.get("S3_MODEL_BUCKET") - if not s3_bucket: - raise ValueError("S3_MODEL_BUCKET environment variable not set.") - - region = os.environ.get("AWS_REGION", "us-east-1") - - # Get AWS credentials from environment variables - aws_access_key_id = os.environ.get("AWS_ACCESS_KEY_ID") - aws_secret_access_key = os.environ.get("AWS_SECRET_ACCESS_KEY") - aws_session_token = os.environ.get("AWS_SESSION_TOKEN") + s3_bucket = 'readcrumbs' - # Create boto3 client with explicit credentials if available - if aws_access_key_id and aws_secret_access_key: - s3 = boto3.client( - "s3", - aws_access_key_id=aws_access_key_id, - aws_secret_access_key=aws_secret_access_key, - aws_session_token=aws_session_token, - region_name=region - ) - else: - # Fall back to default credential chain (IAM roles, ~/.aws/credentials, etc.) - s3 = boto3.client("s3", region_name=region) - # Download model object as bytes into memory - response = s3.get_object(Bucket=s3_bucket, Key=model_name) - model_bytes = response['Body'].read() - import pickle - model = pickle.loads(model_bytes) + s3_client = boto3.client('s3') + response = s3_client.get_object(Bucket=s3_bucket, Key=model_name) +# model_data = response['Body'].read() + buffer = io.BytesIO(response['Body'].read()) + model = joblib.load(buffer) +# model_file = io.BytesIO(model_data) +# model = pickle.load(io.BytesIO(model_data)) +# model = joblib.load(loaded_model) +# model = pickle.loads(loaded_model) return model -def predict_using_model(model, data): - pass +def predict_using_model(data: MyFavorites, n_recs: int = 10): + my_favs_ids = [title_to_index[f] for f in data.items] + fav_vectors = [model.item_factors[i] for i in my_favs_ids] + #Average the vectors + avg_vec = np.average(np.stack(fav_vectors), axis=0) + recommendations = np.argsort(np.dot(avg_vec, model.item_factors.T))[:n_recs] + return [index_to_title[i] for i in recommendations] def serialize_for_dynamodb(data): """ @@ -97,7 +108,8 @@ def get_dynamodb_table(): """ table_name = os.environ.get("DDB_TABLE") if not table_name: - raise ValueError("DDB_TABLE environment variable not set.") +# raise ValueError("DDB_TABLE environment variable not set.") + table_name = "readcrumbs-logs" region = os.environ.get("AWS_REGION", "us-east-1") @@ -177,23 +189,18 @@ def save_to_ddb(data): # Map user_id to pred-id (the table's primary key field name) # Keep user_id in the data as well for reference - serialized_data['pred-id'] = serialized_data['user_id'] - - # DynamoDB put_item will create a new item if pred-id doesn't exist, - # or update/replace the existing item if pred-id already exists +# serialized_data['pred-id'] = serialized_data['user_id'] + response = table.put_item(Item=serialized_data) return response ## API -class MyFavorites(BaseModel): - items: List[str] - -class PredictionResponse(BaseModel): - user_id: int - req: MyFavorites - prediction: str -app = fastapi.FastAPI() +app = FastAPI() + +model = load_model_from_s3("models/als_model-small-v1.pkl") +index_to_title = load_supporting_tables_from_s3("data/v1/index_to_title.json") +title_to_index = load_supporting_tables_from_s3("data/v1/title_to_index.json") @app.get("/health") def health_check(): @@ -209,13 +216,33 @@ def get_random(): """ random_item = get_random_item_from_ddb() if random_item is None: - raise fastapi.HTTPException(status_code=404, detail="No items found in table") + raise HTTPException(status_code=404, detail="No items found in table") return random_item @app.post("/predict") -def predict(request: PredictionResponse): - # Automatically add current timestamp - data = request.model_dump() - data['timestamp'] = datetime.datetime.now(datetime.timezone.utc) - save_to_ddb(data) - return {"status": "ok"} \ No newline at end of file +def predict(request: MyFavorites) -> PredictionResponse: + # Create a dictionary from the request and other metrics. + if len(request.items) < 1 or type(request.items) != List: + raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail="Must enter at least one favorite book.") + + try: + recs = predict_using_model(request.items) + logs = { + "items": request.items, + "user_id": request.userid, + "timestamp": datetime.datetime.now(datetime.timezone.utc), + "prediction": recs + } + + save_to_ddb(logs) + + return { + "recs": recs + } + except Exception as e: + raise HTTPException(status_code=status.HTTP_500_INTERNAL_SERVER_ERROR, detail=str(e)) +# data = request.model_dump() +# data['timestamp'] = datetime.datetime.now(datetime.timezone.utc) +# data['prediction'] = predict_using_model(model, data) +# save_to_ddb(data) +# return {"status": "ok"} \ No newline at end of file