from temporalio import activity, workflow with workflow.unsafe.imports_passed_through(): import traceback from collections.abc import Callable, Mapping from typing import Any from pandas import DataFrame from sientia_do.formatters import create_sample_dict from sientia_do.notifications.handlers import CoreNotificationHandler as NotificationHandler from sientia_do.notifications.models import NotificationLevel from sientia_do.observability.logger import Logger from sientia_do.observability.metrics_controller import MetricsController from sientia_do.observability.sientia_monitoring import SientiaMonitoring from sientia_do.temporal.constants import DATETIME_FORMAT_WITH_TZ, now from laborious import metrics from laborious.utils.filters.conditional_filters import ( filter_empty_data, filter_specific_variables_null_values, ) from laborious.utils.filters.mlflow_filters import api_error_filter, nan_values_filter # Strongly-typed filter function signatures InputFilterFunc = Callable[[DataFrame, dict[str, Any]], bool] ResponseFilterFunc = Callable[[dict[str, Any], dict[str, Any]], bool] ContentFilterFunc = Callable[[DataFrame, dict[str, Any]], bool] # Input filter function mappings input_filter_functions: dict[str, InputFilterFunc] = { 'SPECIFIC_VARIABLES_NULL_VALUES': filter_specific_variables_null_values, 'EMPTY_DATA': filter_empty_data, } # Confidence mappings kept separate from function maps to avoid Union types input_path_confidence: Mapping[str, int] = { 'STOP': -1, 'CONTINUE': 2, 'REPEAT': -1, } # MLFlow response filter function mappings mlflow_response_filter_functions: dict[str, ResponseFilterFunc] = { 'API_ERROR': api_error_filter, } mlflow_response_path_confidence: Mapping[str, int] = { 'STOP': -1, 'CONTINUE': 10, 'REPEAT': -1, } # MLFlow content filter function mappings mlflow_content_filter_functions: dict[str, ContentFilterFunc] = { 'NAN_VALUES': nan_values_filter, 'EMPTY_DATA': filter_empty_data, } mlflow_content_path_confidence: Mapping[str, int] = { 'STOP': -1, 'CONTINUE': 18, 'REPEAT': -1, } class Gates(SientiaMonitoring): """ Data quality gates and filtering activities for the Laborious system. This class implements comprehensive data quality validation and filtering mechanisms that can be applied at different stages of the prediction pipeline. It provides configurable filters with policy-based decision making to ensure data integrity and quality throughout the ML workflow. The class supports multiple filter types and implements a flexible policy system that can be configured for different validation requirements. Each filter returns a path decision (STOP, CONTINUE, REPEAT) along with confidence scores and detailed comments for monitoring and debugging. Attributes: input_filter_functions (dict): Mapping of input filter names to functions mlflow_response_filter_functions (dict): Mapping of MLFlow response filter names to functions mlflow_content_filter_functions (dict): Mapping of MLFlow content filter names to functions """ def __init__( self, logger: Logger, notification_handler: NotificationHandler, metrics_controller: MetricsController, ): """ Initialize data quality gates with logging and notification capabilities. Args: logger: Logger instance for observability and debugging notification_handler: Notification handler for alerts and monitoring Raises: Exception: If BaseActivity initialization fails """ SientiaMonitoring.__init__(self, logger, notification_handler, metrics_controller) def close(self) -> None: """ Close the gates activity and clean up resources. """ SientiaMonitoring.shutdown(self) def __del__(self): self.close() @activity.defn(name='input_gate') async def input_gate(self, input_data: dict[str, Any]) -> tuple[str | None, int, str]: """ Apply input data quality filters and validation. This activity validates input data quality using configurable filters before proceeding with ML operations. It applies multiple filter types and returns a path decision based on the filter results and configured policies. The method implements a comprehensive filtering system that: 1. Applies configured filters to input data 2. Evaluates filter results against policy configurations 3. Determines appropriate path decisions (STOP, CONTINUE, REPEAT) 4. Provides confidence scores and detailed comments 5. Handles errors gracefully with notification integration Args: input_data: Configuration and data for input validation Required keys: - metadata (dict): Workflow execution metadata - filters (dict): Filter configuration and policies - data (dict): Input data to validate - path_priority (list[str]): Priority order for path decisions Returns: tuple: (path_flag, confidence, comment) - path_flag (str | None): Decision path (STOP, CONTINUE, REPEAT, or None) - confidence (int): Confidence score for the decision - comment (str): Detailed explanation of the decision Raises: Exception: If filter execution fails or configuration is invalid """ metadata = input_data['metadata'] self.info('Performing input gate...', metadata) filters = input_data['filters'] data = DataFrame(input_data['data']) path_priority = input_data['path_priority'] filter_output = [] self.debug(f'Input data: {data.head(5).to_string()}', metadata) self.debug(f'Filters: {filters}', metadata) # Apply each configured filter for fil, config in filters.items(): if fil not in input_filter_functions: self.error(f'Filter {fil} not found', metadata) continue try: if input_filter_functions[fil](data, config['config']): self.debug(f'Data not passed the input filter {fil}:{config}', metadata) filter_output.append(config['policy']) except Exception as e: trace = traceback.format_exc() await self.send_notification_async( metadata=metadata, notification_id=f'INTPUT_GATE_ERROR__{fil}', message=f'Error in filter {fil}:{config}: \n {e}', block='input_gate', level=NotificationLevel.ERROR, attachment_content=trace, ) for path_flag in path_priority: if path_flag in filter_output: self.info(f'Input gate result: {path_flag}', metadata) return path_flag, input_path_confidence[path_flag], 'Input data with bad quality' self.info('Nothing was filtered by the input gate', metadata) return None, 0, '' @activity.defn(name='mlflow_response_gate') async def mlflow_response_gate(self, input_data: dict[str, Any]) -> tuple[str | None, int, str]: """ Validate MLFlow API response quality and integrity. This activity validates MLFlow API responses to ensure they meet quality standards before proceeding with further processing. It applies response-specific filters and determines appropriate path decisions based on response quality. The method implements response validation that: 1. Applies MLFlow response-specific filters 2. Evaluates API response quality and integrity 3. Determines path decisions based on response validation results 4. Provides confidence scores and detailed validation comments 5. Handles API errors and response validation failures Args: input_data: Configuration and data for response validation Required keys: - metadata (dict): Workflow execution metadata - filters (dict): Response filter configuration and policies - data (dict): MLFlow API response data to validate - type (str): Type of MLFlow operation (transform, predict) - path_priority (list[str]): Priority order for path decisions Returns: tuple: (path_flag, confidence, comment) - path_flag (str | None): Decision path (STOP, CONTINUE, REPEAT, or None) - confidence (int): Confidence score for the decision - comment (str): Detailed explanation of the decision Raises: Exception: If response validation fails or configuration is invalid """ metadata = input_data['metadata'] self.info('Performing mlflow response gate...', metadata) filters = input_data['filters'] data = input_data['data'] gate_type = input_data['type'] path_priority = input_data['path_priority'] filter_output = [] self.debug(f'Input data: \n {create_sample_dict(data, max_items=5, max_depth=5)}', metadata) self.debug(f'Filters: {filters}', metadata) comments = [] for fil, config in filters.items(): if fil not in mlflow_response_filter_functions: continue try: if mlflow_response_filter_functions[fil](data, config): filter_output.append(config['policy']) comments.append(data['content']['message']) await self.send_notification_async( metadata=metadata, notification_id=f'{gate_type.upper()}_GATE_RESPONSE_FILTER__{fil}', message=data['content']['message'], block='mlflow_gate', level=NotificationLevel.ERROR, attachment_content=data['content']['traceback'], ) except Exception as e: trace = traceback.format_exc() await self.send_notification_async( metadata=metadata, notification_id=f'MLFLOW_GATE_RESPONSE_FILTER__{fil}', message=f'Error in filter {fil}:{config}: \n {e}', block='mlflow_gate', level=NotificationLevel.ERROR, attachment_content=trace, ) for path_flag in path_priority: if path_flag in filter_output: self.info(f'Mlflow response gate result: {path_flag}', metadata) return path_flag, mlflow_response_path_confidence[path_flag], ', '.join(comments) self.info('Nothing was filtered by the mlflow response gate', metadata) return None, 0, '' @activity.defn(name='mlflow_content_gate') async def mlflow_content_gate(self, input_data: dict[str, Any]) -> tuple[str | None, int, str]: """ Validate MLFlow prediction content quality and integrity. This activity validates the content of MLFlow predictions to ensure they meet quality standards before export and persistence. It applies content-specific filters and determines appropriate path decisions based on content quality. The method implements content validation that: 1. Applies MLFlow content-specific filters 2. Evaluates prediction content quality and integrity 3. Determines path decisions based on content validation results 4. Provides confidence scores and detailed validation comments 5. Handles content validation failures and quality issues Args: input_data: Configuration and data for content validation Required keys: - metadata (dict): Workflow execution metadata - filters (dict): Content filter configuration and policies - data (dict): MLFlow prediction content to validate - type (str): Type of MLFlow operation (transform, predict) - path_priority (list[str]): Priority order for path decisions Returns: tuple: (path_flag, confidence, comment) - path_flag (str | None): Decision path (STOP, CONTINUE, REPEAT, or None) - confidence (int): Confidence score for the decision - comment (str): Detailed explanation of the decision Raises: Exception: If content validation fails or configuration is invalid """ metadata = input_data['metadata'] self.info('Performing mlflow content gate...', metadata) filters = input_data['filters'] data = DataFrame(input_data['data']) gate_type = input_data['type'] path_priority = input_data['path_priority'] filter_output = [] self.debug(f'Input data:\n {data.head(5).to_string()}', metadata) self.debug(f'Filters: \n {filters}', metadata) for fil, config in filters.items(): if fil not in mlflow_content_filter_functions: continue try: if mlflow_content_filter_functions[fil](data, config): filter_output.append(config['policy']) await self.send_notification_async( metadata=metadata, notification_id=f'{gate_type.upper()}_GATE_CONTENT_FILTER__{fil}', message=f'Data not passed the content filter {fil}:{config}', block='mlflow_gate', level=NotificationLevel.WARNING, attachment_content=data.to_string(), ) except Exception as e: trace = traceback.format_exc() await self.send_notification_async( metadata=metadata, notification_id=f'MLFLOW_GATE_CONTENT_FILTER__{fil}', message=f'Error in filter {fil}:{config}: \n {e}', block='mlflow_gate', level=NotificationLevel.ERROR, attachment_content=trace, ) for path_flag in path_priority: if path_flag in filter_output: self.info(f'Mlflow content gate result: {path_flag}', metadata) return ( path_flag, mlflow_content_path_confidence[path_flag], 'Transformed data not passed the content filter', ) self.info('Nothing was filtered by the mlflow content gate', metadata) return None, 0, '' def get_prediction_store_policy( self, prediction_store_policy: str, metadata: dict[str, Any] ) -> tuple[str, int]: """ Parse and validate prediction store policy configuration. This method parses prediction store policy strings in the format 'type:value' and validates them against allowed policy types and values. It provides sensible defaults for invalid configurations and logs policy validation failures for operational monitoring. Supported Policy Types: - 'lts': Latest timestamp - sorts data by timestamp descending - 'erl': Earliest timestamp - sorts data by timestamp ascending Args: prediction_store_policy (str): Policy string in format 'type:value' metadata (dict[str, Any]): Context metadata for logging and notifications Returns: tuple[str, int]: (policy_type, policy_value) - policy_type (str): Validated policy type ('lts' or 'erl') - policy_value (int): Number of rows to retain """ policy_elements = prediction_store_policy.split(':') if len(policy_elements) < 2: self.error( f'Invalid prediction store policy: {prediction_store_policy}, using default policy', metadata, ) return 'lts', 1 policy_type = policy_elements[0] policy_value = policy_elements[1] # If the policy_type is not lts or erl, we use the default policy # If the policty_value is not a number or 0, we use the default policy if ( policy_type not in ['lts', 'erl'] or not policy_value.isdigit() or int(policy_value) == 0 ): self.error( f'Invalid prediction store policy: {prediction_store_policy}, using default policy', metadata, ) return 'lts', 1 return policy_type, int(policy_value) @activity.defn(name='format_prediction') async def format_prediction(self, input_data: dict[str, Any]) -> dict[Any, Any]: """ Format prediction data according to configured storage policies. This method formats prediction data for storage and export operations. It applies timestamp-based sorting policies, adds metadata fields, and ensures data consistency before persistence. The method supports multiple storage policies for flexible data retention strategies. Storage Policies: - 'lts:N': Latest timestamp - retains N most recent predictions - 'erl:N': Earliest timestamp - retains N oldest predictions Args: input_data (dict): Input data containing: - data (dict[str, Any]): Raw prediction data to format - timestamp (str): Default timestamp if data lacks timestamp column - model_id (str): Unique identifier for the ML model - prediction_confidence (float): Confidence score for the prediction - prediction_store_policy (str): Storage policy in format 'type:value' Returns: dict: Formatted prediction data ready for storage and export """ metadata = input_data['metadata'] prediction_store_policy = input_data['prediction_store_policy'] self.info('Formatting prediction...', metadata) data = DataFrame(input_data['data']) # Create timestamp column from index and reset index data['timestamp'] = data.index data = data.reset_index(drop=True) self.debug(f'Prediction store policy: {prediction_store_policy}', metadata) self.debug(f'Prediction data: {data.head(5).to_string()}', metadata) policy_type, policy_value = self.get_prediction_store_policy( prediction_store_policy, metadata ) # If data has no timestamp, we use the default timestamp and not sort the data self.info( f'Sorting data by timestamp and applying policy: {policy_type}:{policy_value}', metadata ) # If policy_type is lts, we need to sort the data by timestamp descending and take the first policy_value rows if policy_type == 'lts': self.debug('Sorting data by timestamp descending', metadata) data = data.sort_values(by='timestamp', ascending=False) # If policy_type is erl, we need to sort the data by timestamp ascending and take the first policy_value rows elif policy_type == 'erl': self.debug('Sorting data by timestamp ascending', metadata) data = data.sort_values(by='timestamp', ascending=True) else: self.error(f'Invalid policy type: {policy_type}, using default policy', metadata) raise ValueError(f'Invalid policy type: {policy_type}') data = data.head(int(policy_value)) data['model_id'] = input_data['model_id'] data['prediction_confidence'] = input_data['prediction_confidence'] data['prediction_status'] = 'Good' data['comments'] = '' data = data.sort_values(by='timestamp', ascending=False) data = data.reset_index(drop=True) self.info(f'Prediction formatted: {len(data)} rows', metadata) self.debug(f'Prediction data: {data.head(5).to_string()}', metadata) return data.to_dict() @activity.defn(name='format_default_prediction') async def format_default_prediction(self, input_data: dict[str, Any]) -> dict[Any, Any]: """ Create and format default prediction data for error conditions. This method generates default prediction data when the main prediction pipeline encounters errors or quality issues. It creates a standardized data structure with zero values for predictions and useful metadata for operational monitoring and debugging. The default prediction serves as a fallback mechanism to: 1. Maintain data pipeline continuity during failures 2. Provide operational visibility into prediction quality issues 3. Enable downstream systems to handle error conditions gracefully 4. Support debugging and troubleshooting efforts Args: input_data (dict): Input data containing: - timestamp (str): Timestamp for the default prediction - model_id (str): Unique identifier for the ML model - prediction_confidence (float): Confidence score (typically low for errors) - comment (str): Error description or operational comment Returns: dict: Formatted default prediction data with error indicators """ metadata = input_data['metadata'] self.debug('Formatting default prediction...', metadata) data = DataFrame( { 'prediction': [0], 'response_time': [0], 'timestamp': [input_data['timestamp']], 'model_id': [input_data['model_id']], 'prediction_confidence': [input_data['prediction_confidence']], 'prediction_status': ['Bad'], 'comments': [input_data['comment']], } ) self.info(f'Default prediction formatted: {data.size} rows', metadata) return data.to_dict() @activity.defn(name='format_retrain_report') async def format_retrain_report(self, input_data: dict[str, Any]) -> dict[Any, Any]: """ Format retrain report data according to configured storage policies. """ metadata = input_data['metadata'] self.info('Formatting retrain report...', metadata) experiment_response = input_data['experiment_response'] update_report = input_data['update_report'] model_id = input_data['model_id'] model_name = input_data['model_name'] report = DataFrame( { 'model_id': [model_id], 'model_name': [model_name], 'timestamp': [experiment_response['timestamp']], 'status': [experiment_response['message']], } ) if experiment_response['success']: # Retrain was successfull report['version'] = update_report['version'] report['mlflow_run_id'] = update_report['mlflow_run_id'] report['mlflow_experiment_id'] = update_report['mlflow_experiment_id'] self.debug(f'Retrain report: {report.to_csv()}', metadata) return report.to_dict() @activity.defn(name='get_last_timestamp') async def get_last_timestamp(self, input_data: dict[str, Any]) -> str: """ Extract the most recent timestamp from prediction data. This method analyzes prediction data to find the latest timestamp, enabling incremental processing and data continuity tracking. It handles empty datasets gracefully by returning the current time as a fallback timestamp. The method is essential for: 1. Incremental data processing workflows 2. Data continuity validation 3. Timestamp-based data loading optimization 4. Workflow execution tracking Args: input_data (dict): Input data containing: - data (dict[str, Any]): Prediction data to analyze Returns: str: Formatted timestamp string in UTC with timezone """ metadata = input_data['metadata'] self.info('Getting last timestamp...', metadata) data = DataFrame(input_data['data']) self.debug(f'Input data: {data.head(5).to_string()}', metadata) if data.empty: return now().strftime(DATETIME_FORMAT_WITH_TZ) max_timestamp = max(data['timestamp'].values.tolist()) self.info(f'Last timestamp: {max_timestamp}', metadata) return max_timestamp @activity.defn(name='write_metrics') async def write_metrics(self, input_data: dict[str, Any]): """ Write prediction performance metrics to Prometheus monitoring system. This method records comprehensive metrics for prediction operations, enabling operational monitoring, performance analysis, and alerting. It tracks prediction counts, confidence levels, and response times for each model and pipeline combination. Metrics Recorded: 1. Prediction Count: Incremental counter for successful predictions 2. Confidence Monitor: Current confidence level for predictions 3. Response Time Monitor: Histogram of prediction response times Args: input_data (dict): Input data containing: - metadata (dict[str, Any]): Workflow execution metadata - prediction (dict[str, Any]): Prediction data with metrics Raises: Exception: If metrics writing fails or configuration is invalid """ metadata = input_data['metadata'] prediction = DataFrame(input_data['prediction']) prediction_confidence = prediction['prediction_confidence'].values[0] response_time = prediction['response_time'].values[0] opc_metrics = input_data['opc_metrics'] self.info(f'Writing metrics for model {metadata["model_name"]}', metadata) await self.emit_metric( metric_object=metrics.PREDICTIONS_WRITTEN_COUNT, tags={ 'pod_id': self.pod_id, 'model_name': metadata['model_name'], 'workflow_name': metadata['workflow_name'], }, ) await self.emit_metric( metric_object=metrics.PREDICTION_CONFIDENCE_MONITOR, method='set', tags={ 'pod_id': self.pod_id, 'model_name': metadata['model_name'], 'workflow_name': metadata['workflow_name'], }, value=prediction_confidence, ) await self.emit_metric( metric_object=metrics.PREDICTION_RESPONSE_TIME_MONITOR, method='observe', tags={ 'pod_id': self.pod_id, 'model_name': metadata['model_name'], 'workflow_name': metadata['workflow_name'], }, value=response_time, ) for server_id, tags in opc_metrics.items(): for tag, response_time in tags.items(): if response_time is not None: await self.emit_metric( metric_object=metrics.PREDICTION_OPC_WRITING_RESPONSE_TIME_MONITOR, method='observe', tags={ 'pod_id': self.pod_id, 'model_name': metadata['model_name'], 'workflow_name': metadata['workflow_name'], 'opc_server_id': server_id, 'tag': tag, }, value=response_time, ) await self.emit_metric( metric_object=metrics.PREDICTION_OPC_WRITING_COUNT, tags={ 'pod_id': self.pod_id, 'model_name': metadata['model_name'], 'workflow_name': metadata['workflow_name'], 'opc_server_id': server_id, 'tag': tag, }, ) self.info(f'Metrics written for model {metadata["model_name"]}', metadata)