import os import time import json import joblib import logging import traceback from typing import List, Any import numpy as np import pandas as pd from flask import Flask, request, jsonify MODEL_PATH = os.environ.get( "MODEL_PATH", "/content/superkart_best_tuned_RandomForest.joblib" ) ALLOWED_EXTENSIONS = {"csv"} EXPECTED_COLUMNS: List[str] = None ID_COLUMNS = ["Product_Id", "Store_Id"] logging.basicConfig(level=logging.INFO, format="%(asctime)s [%(levelname)s] %(message)s") log = logging.getLogger("superkart_api") if not os.path.exists(MODEL_PATH): raise FileNotFoundError( f"MODEL_PATH not found: {MODEL_PATH}\n" "Make sure your Drive is mounted and the path is correct." ) t0 = time.time() pipe = joblib.load(MODEL_PATH) load_secs = time.time() - t0 log.info(f"Loaded pipeline from {MODEL_PATH} in {load_secs:.2f}s") assert hasattr(pipe, "predict"), "Loaded object does not have .predict(...)" assert "pre" in pipe.named_steps and "model" in pipe.named_steps, \ "Pipeline must contain 'pre' and 'model' steps." try: raw_feature_names = getattr(pipe.named_steps["pre"], "feature_names_in_", None) except Exception: raw_feature_names = None app = Flask(__name__) def _allowed_file(filename: str) -> bool: return "." in filename and filename.rsplit(".", 1)[1].lower() in ALLOWED_EXTENSIONS def _align_columns(df: pd.DataFrame) -> pd.DataFrame: global raw_feature_names if raw_feature_names is not None: for col in raw_feature_names: if col not in df.columns: df[col] = np.nan df = df[list(raw_feature_names)] if EXPECTED_COLUMNS: for col in EXPECTED_COLUMNS: if col not in df.columns: df[col] = np.nan df = df[EXPECTED_COLUMNS] return df def _to_dataframe_from_json(payload: Any) -> pd.DataFrame: if isinstance(payload, dict): df = pd.DataFrame([payload]) elif isinstance(payload, list): if not payload or not all(isinstance(x, dict) for x in payload): raise ValueError("Payload must be a dict or list of dicts.") df = pd.DataFrame(payload) else: raise ValueError("JSON payload must be object or list of objects.") return _align_columns(df) def _predict_df(df: pd.DataFrame) -> np.ndarray: preds = pipe.predict(df) return np.asarray(preds).reshape(-1).astype(float) @app.get("/") def index(): return jsonify({ "service": "SuperKart Forecast API", "status": "ok", "endpoints": ["/health", "/model-info", "/predict", "/predict-csv"] }), 200 @app.get("/health") def health(): return jsonify({ "status": "ok", "model_path": MODEL_PATH, "loaded_in_seconds": round(load_secs, 3), "training_features": list(raw_feature_names) if raw_feature_names is not None else None }) @app.get("/model-info") def model_info(): mdl = pipe.named_steps["model"] try: params = mdl.get_params() except Exception: params = str(mdl) return jsonify({ "type": mdl.__class__.__name__, "params": params }) @app.post("/predict") def predict(): try: payload = request.get_json(silent=True) if payload is None: return jsonify({"error": "Invalid or empty JSON body."}), 400 df = _to_dataframe_from_json(payload) preds = _predict_df(df) echo_ids = {idc: df[idc].tolist() for idc in ID_COLUMNS if idc in df.columns} return jsonify({"n": int(len(preds)), **({"ids": echo_ids} if echo_ids else {}), "predictions": preds.tolist()}) except Exception as e: log.error("Predict error: %s\n%s", e, traceback.format_exc()) return jsonify({"error": str(e)}), 500 @app.post("/predict-csv") def predict_csv(): try: if "file" not in request.files: return jsonify({"error": "No file part named 'file'."}), 400 file = request.files["file"] if file.filename == "": return jsonify({"error": "Empty filename."}), 400 if not _allowed_file(file.filename): return jsonify({"error": "Only .csv files are allowed."}), 400 df = pd.read_csv(file) df = _align_columns(df) preds = _predict_df(df) ids = {idc: df[idc].astype(str).tolist() for idc in ID_COLUMNS if idc in df.columns} return jsonify({"n": int(len(preds)), **({"ids": ids} if ids else {}), "predictions": preds.tolist()}) except Exception as e: log.error("Predict-CSV error: %s\n%s", e, traceback.format_exc()) return jsonify({"error": str(e)}), 500 if __name__ == "__main__": port = int(os.environ.get("PORT", "8000")) app.run(host="0.0.0.0", port=port, debug=False)