Spaces:
Runtime error
Runtime error
| 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 | |