Spaces:
Sleeping
Sleeping
Upload main.py
Browse files
main.py
CHANGED
|
@@ -478,12 +478,24 @@ async def modelfit(
|
|
| 478 |
|
| 479 |
@app.post("/modelforecast/")
|
| 480 |
async def modelforecast(
|
| 481 |
-
|
| 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")
|