Vadashuk commited on
Commit
d1f513a
·
verified ·
1 Parent(s): 633e14e

Upload main.py

Browse files
Files changed (1) hide show
  1. main.py +13 -1
main.py CHANGED
@@ -478,12 +478,24 @@ async def modelfit(
478
 
479
  @app.post("/modelforecast/")
480
  async def modelforecast(
481
- df_forecast: List[Dict] = Body(...),
482
  ):
483
  global stored_model, stored_result_join
484
  if stored_model is None or stored_result_join is None:
485
  raise HTTPException(status_code=400, detail="Model not fitted. Call /modelfit/ first.")
486
  try:
 
 
 
 
 
 
 
 
 
 
 
 
487
  df_forecast1 = pd.DataFrame(df_forecast)
488
  forecast_df = stored_model.forecast(extern_self=stored_result_join, df_forecast=df_forecast1)
489
  return forecast_df.to_dict(orient="records")
 
478
 
479
  @app.post("/modelforecast/")
480
  async def modelforecast(
481
+ payload: Any = Body(...),
482
  ):
483
  global stored_model, stored_result_join
484
  if stored_model is None or stored_result_join is None:
485
  raise HTTPException(status_code=400, detail="Model not fitted. Call /modelfit/ first.")
486
  try:
487
+ # Accept both payload styles:
488
+ # 1) raw list: [...]
489
+ # 2) wrapped dict: {"df_forecast": [...]}
490
+ if isinstance(payload, dict) and "df_forecast" in payload:
491
+ df_forecast = payload.get("df_forecast")
492
+ else:
493
+ df_forecast = payload
494
+ if not isinstance(df_forecast, list):
495
+ raise HTTPException(
496
+ status_code=422,
497
+ detail="Payload must be a list of rows or {'df_forecast': [rows]}",
498
+ )
499
  df_forecast1 = pd.DataFrame(df_forecast)
500
  forecast_df = stored_model.forecast(extern_self=stored_result_join, df_forecast=df_forecast1)
501
  return forecast_df.to_dict(orient="records")