ganesh-vilje's picture
Deploy to Hugging Face Main
f8f02c0
Raw
History Blame Contribute Delete
2.65 kB
import io
import time as _time
from typing import Any, Dict
import pandas as pd
from .client import api_client
from .parser import parse_model_response
from .schema import TransformedOutput
from .utils import wait_for_completion
def execute_model_step(
current_df: pd.DataFrame,
model_name: str,
compliance_type: str = "firco",
user_id: str = "system",
version: str = "latest",
number_of_reasonings: int = 1
) -> TransformedOutput:
"""
Executes a model integration step for a workflow.
Blocks minimally for 60 seconds (or more if model takes longer)
before retrieving results and transforming them.
"""
# 1. Prepare CSV buffer from dataframe
csv_buffer = io.BytesIO()
current_df.to_csv(csv_buffer, index=False)
csv_bytes = csv_buffer.getvalue()
# Use a unique filename per prediction to avoid S3 key collisions on the remote server
unique_filename = f"workflow_input_{int(_time.time() * 1000)}.csv"
# 2. Call prediction initialization
prediction_id = api_client.run_prediction(
csv_bytes=csv_bytes,
filename=unique_filename,
model_name=model_name,
compliance_type=compliance_type,
source_type="file",
version=version,
number_of_reasonings=number_of_reasonings
)
if not prediction_id:
raise ValueError("Failed to get a valid prediction_id from API.")
# 3 & 4 & 5. Wait, poll, and validate status
def check_status() -> Dict[str, Any]:
return api_client.get_prediction_run(prediction_id)
def is_completed(response: Dict[str, Any]) -> bool:
# The actual workflow status is typically nested under "data"
data_block = response.get("data", response)
status = str(data_block.get("status", "")).lower()
if status in ("failed", "error"):
raise RuntimeError(f"Prediction {prediction_id} failed remotely: {data_block.get('error')}")
return status == "completed"
try:
completed_response = wait_for_completion(
check_status_callable=check_status,
is_completed_callable=is_completed,
poll_interval=5,
max_wait_seconds=600
)
except TimeoutError as e:
raise TimeoutError(f"Prediction {prediction_id} timed out: {e}")
# 6. Parse and transform the final output
transformed_output = parse_model_response(current_df.copy(), completed_response)
# Tag it implicitly with the prediction ID for workflow traceability
transformed_output.prediction_id = prediction_id
return transformed_output