diff --git a/input_sample.json b/input_sample.json new file mode 100644 index 0000000..9fffe37 --- /dev/null +++ b/input_sample.json @@ -0,0 +1,34 @@ +{ + "schedule_name": "scouter-opcua-pipeline", + "model_name": "Demo Model", + "model_id": 1, + "query": "SELECT * FROM sientia_data.laborious_data order by \"timestamp\" desc limit 30;", + "schema": "sientia_data", + "table_name": "predictions", + "retention_time": 3600, + "model_retention": 120, + "path_priority": ["STOP", "CONTINUE", "REPEAT"], + "input_filters": { + "SPECIFIC_VARIABLES_NULL_VALUES": { + "POLICY": "STOP", + "VARIABLES": ["Counter"] + }, + "EMPTY_DATA": { + "POLICY": "STOP" + } + }, + "mlflow_transform_filters": { + "API_ERROR": { + "POLICY": "CONTINUE" + }, + "NAN_VALUES": { + "POLICY": "CONTINUE" + } + }, + "mlflow_predict_filters": { + "API_ERROR": { + "POLICY": "CONTINUE" + } + }, + "opc_output_config": {} +} \ No newline at end of file diff --git a/laborious/activities/activities.py b/laborious/activities/activities.py index c7d1821..5ba2cee 100644 --- a/laborious/activities/activities.py +++ b/laborious/activities/activities.py @@ -40,15 +40,10 @@ class Activities(Postgres, MLFlow, Gates, OPC): notification_handler=notification_handler) OPC.__init__(self, - name=opc_config['name'], - url=opc_config['url'], - server_uri=opc_config['server_uri'], - cert_path=opc_config['cert_path'], - private_key_path=opc_config['private_key_path'], - server_cert_path=opc_config['server_cert_path'], + opc_servers=opc_config, logger=logger, notification_handler=notification_handler) @activity.defn(name="prepare_activity") - def prepare_activity(self, input_data: dict[str, Any]): - super().prepare_activity(input_data) + async def prepare_activity(self, input_data: dict[str, Any]): + await super().prepare_activity(input_data) diff --git a/laborious/activities/base.py b/laborious/activities/base.py index 2742457..3adb0e4 100644 --- a/laborious/activities/base.py +++ b/laborious/activities/base.py @@ -1,3 +1,4 @@ +from typing import Any from logging import Logger from temporalio import activity from sientia_do.notifications.handlers import NotificationHandler @@ -9,7 +10,7 @@ class BaseActivity: self.notification_handler = notification_handler @activity.defn(name="prepare_activity") - def prepare_activity(self, input_data: dict[str, Any]): + async def prepare_activity(self, input_data: dict[str, Any]): """ Prepare the activity for the notification handler. diff --git a/laborious/activities/gates.py b/laborious/activities/gates.py index 4d9eb26..0b021b5 100644 --- a/laborious/activities/gates.py +++ b/laborious/activities/gates.py @@ -5,67 +5,85 @@ with workflow.unsafe.imports_passed_through(): import traceback from logging import Logger from sientia_do.notifications.handlers import NotificationHandler - from laborious.activities.base import BaseActivity - from typing import Any - from laborious.utils.filters.conditional_filters import filter_empty_data, filter_specific_variables_null_values - from pandas import DataFrame from sientia_do.notifications.models import NotificationLevel + from laborious.activities.base import BaseActivity from laborious.utils.filters.mlflow_filters import nan_values_filter, api_error_filter + from typing import Any + from laborious.utils.filters.conditional_filters import ( + filter_empty_data, + filter_specific_variables_null_values + ) + from pandas import DataFrame + from datetime import datetime input_filter_functions = { 'SPECIFIC_VARIABLES_NULL_VALUES': filter_specific_variables_null_values, 'EMPTY_DATA': filter_empty_data, 'path_confidence': { - 'stop': -1, - 'continue': 2, - 'repeat': -1 + 'STOP': -1, + 'CONTINUE': 2, + 'REPEAT': -1 } } mlflow_response_filter_functions = { 'API_ERROR': api_error_filter, 'path_confidence': { - 'stop': -1, - 'continue': 10, - 'repeat': -1 + 'STOP': -1, + 'CONTINUE': 10, + 'REPEAT': -1 }, } mlflow_content_filter_functions = { 'NAN_VALUES': nan_values_filter, 'path_confidence': { - 'stop': -1, - 'continue': 18, - 'repeat': -1 + 'STOP': -1, + 'CONTINUE': 18, + 'REPEAT': -1 } } class Gates(BaseActivity): def __init__(self, logger: Logger, notification_handler: NotificationHandler): - super().__init__(logger, notification_handler) + BaseActivity.__init__(self, logger, notification_handler) @activity.defn(name="input_gate") - async def input_gate(self, input_data: dict[str, Any]) -> tuple[str, int]: + async def input_gate(self, input_data: dict[str, Any]) -> tuple[str | None, int, str]: """ Filters the data based on the filters. The return value is a tuple with the first element being the policy and the second element being the confidence status. Args: input_data (dict): The input data. Contains: filters (dict): The filters to apply. + The key is the filter name and the value is the filter configuration. data (dict[str, Any]): The data to filter. path_priority (list[str]): The path priority. Returns: - tuple[str, int]: (policy, confidence) based in priority list and filter configuration and functions. + tuple[str | None, int, str]: (policy, confidence) based in priority + list and filter configuration and functions. """ + + self.logger.debug("Performing input gate...") + filters = input_data['filters'] data = DataFrame(input_data['data']) path_priority = input_data['path_priority'] filter_output = [] + + self.logger.debug(f"Input data:\n {data.to_string()}") + self.logger.debug(f"Filters: {filters}") + for fil, config in filters.items(): + if fil not in input_filter_functions: + self.logger.error(f"Filter {fil} not found") + continue try: if input_filter_functions[fil](data, config): + self.logger.debug( + f"Data not passed the input filter {fil}:{config}") filter_output.append(config['POLICY']) except Exception as e: trace = traceback.format_exc() @@ -79,14 +97,18 @@ class Gates(BaseActivity): for path_flag in path_priority: if path_flag in filter_output: - return path_flag, input_filter_functions['path_confidence'][path_flag] + self.logger.debug(f"Input gate result: {path_flag}") + return path_flag, input_filter_functions['path_confidence'][path_flag], \ + "Input data with bad quality" - return None, 0 + self.logger.debug("Nothing was filtered by the input gate") + return None, 0, "" @activity.defn(name="mlflow_response_gate") - async def mlflow_response_gate(self, input_data: dict[str, Any]) -> tuple[str, int]: + async def mlflow_response_gate(self, input_data: dict[str, Any]) -> tuple[str | None, int, str]: """ - Filters the data based on the mlflow response filters. The return value is a tuple with the first element + Filters the data based on the mlflow response filters. + The return value is a tuple with the first element being the policy and the second element being the confidence status. Args: input_data (dict): The input data. Contains: @@ -95,17 +117,29 @@ class Gates(BaseActivity): path_priority (list[str]): The path priority list. type (str): The type of the gate. Returns: - tuple[str, int]: (policy, confidence) based in priority list and filter configuration and functions. + tuple[str | None, int, str]: (policy, confidence) based in priority list + and filter configuration and functions. """ + + self.logger.debug("Performing mlflow response gate...") + filters = input_data['filters'] data = input_data['data'] gate_type = input_data['type'] path_priority = input_data['path_priority'] filter_output = [] + + self.logger.debug(f"Input data:\n {data}") + self.logger.debug(f"Filters: {filters}") + + comments = [] for fil, config in filters.items(): + if fil not in mlflow_response_filter_functions: + continue if mlflow_response_filter_functions[fil](data, config): filter_output.append(config['POLICY']) + comments.append(data['content']['message']) self.notification_handler.build_and_send_notification( notification_id=f"{gate_type.upper()}_GATE_RESPONSE_FILTER__{fil}", message=data['content']['message'], @@ -116,14 +150,18 @@ class Gates(BaseActivity): for path_flag in path_priority: if path_flag in filter_output: - return path_flag, mlflow_response_filter_functions['path_confidence'][path_flag] + self.logger.debug(f"Mlflow response gate result: {path_flag}") + return path_flag, mlflow_response_filter_functions['path_confidence'][path_flag], \ + ", ".join(comments) - return None, 0 + self.logger.debug("Nothing was filtered by the mlflow response gate") + return None, 0, "" @activity.defn(name="mlflow_content_gate") - async def mlflow_content_gate(self, input_data: dict[str, Any]) -> tuple[str, int]: + async def mlflow_content_gate(self, input_data: dict[str, Any]) -> tuple[str | None, int, str]: """ - Filters the data based on the mlflow content filters. The return value is a tuple with the first element + Filters the data based on the mlflow content filters. + The return value is a tuple with the first element being the policy and the second element being the confidence status. Args: input_data (dict): The input data. Contains: @@ -132,9 +170,12 @@ class Gates(BaseActivity): path_priority (list[str]): The path priority list. type (str): The type of the gate. Returns: - tuple[str, int]: (policy, confidence) based in priority list and filter configuration and functions. + tuple[str | None, int, str]: (policy, confidence) based in priority + list and filter configuration and functions. """ + self.logger.debug("Performing mlflow content gate...") + filters = input_data['filters'] data = DataFrame(input_data['data']) gate_type = input_data['type'] @@ -142,7 +183,12 @@ class Gates(BaseActivity): filter_output = [] + self.logger.debug(f"Input data:\n {data}") + self.logger.debug(f"Filters: {filters}") + for fil, config in filters.items(): + if fil not in mlflow_content_filter_functions: + continue if mlflow_content_filter_functions[fil](data, config): filter_output.append(config['POLICY']) self.notification_handler.build_and_send_notification( @@ -155,9 +201,12 @@ class Gates(BaseActivity): for path_flag in path_priority: if path_flag in filter_output: - return path_flag, mlflow_content_filter_functions['path_confidence'][path_flag] + self.logger.debug(f"Mlflow content gate result: {path_flag}") + return path_flag, mlflow_content_filter_functions['path_confidence'][path_flag], \ + "Transformed data not passed the content filter" - return None, 0 + self.logger.debug("Nothing was filtered by the mlflow content gate") + return None, 0, "" @activity.defn(name="format_prediction") async def format_prediction(self, input_data: dict[str, Any]) -> dict[str, Any]: @@ -172,12 +221,14 @@ class Gates(BaseActivity): Returns: dict: The formatted data. """ + self.logger.debug("Formatting prediction...") + data = DataFrame(input_data['data']) data['timestamp'] = input_data['timestamp'] data['model_id'] = input_data['model_id'] data['prediction_confidence'] = input_data['prediction_confidence'] data['prediction_status'] = 'Good' - data['comment'] = "" + data['comments'] = "" data.sort_values(by='timestamp', inplace=True) return data.to_dict() @@ -198,6 +249,8 @@ class Gates(BaseActivity): dict: The formatted data. """ + self.logger.debug("Formatting default prediction...") + return DataFrame({ 'prediction': [0], 'response_time': [0], @@ -205,7 +258,7 @@ class Gates(BaseActivity): 'model_id': [input_data['model_id']], 'prediction_confidence': [input_data['prediction_confidence']], 'prediction_status': ['Bad'], - 'comment': [input_data['comment']] + 'comments': [input_data['comment']] }).to_dict() @activity.defn(name="get_last_timestamp") @@ -219,4 +272,6 @@ class Gates(BaseActivity): str: The last timestamp of the data. """ data = DataFrame(input_data['data']) + if data.empty: + return datetime.now().strftime('%Y-%m-%d %H:%M:%S') return max(data['timestamp'].values.tolist()) diff --git a/laborious/activities/mlflow.py b/laborious/activities/mlflow.py index d1b8d4c..ed96192 100644 --- a/laborious/activities/mlflow.py +++ b/laborious/activities/mlflow.py @@ -5,17 +5,16 @@ from temporalio import activity, workflow with workflow.unsafe.imports_passed_through(): from laborious.activities.base import BaseActivity + from laborious.utils.repository.model_repository import MLFlowRepository from typing import Any from logging import Logger from sientia_do.notifications.handlers import NotificationHandler - from laborious.utils.repository.model_repository import MLFlowRepository - from sientia_do.notifications.models import NotificationLevel class MLFlow(BaseActivity): def __init__(self, mlflow_host: str, mlflow_port: int, mlflow_username: str, mlflow_password: str, logger: Logger, notification_handler: NotificationHandler): - super().__init__(logger, notification_handler) + BaseActivity.__init__(self, logger, notification_handler) self.mlflow_host = mlflow_host self.mlflow_port = mlflow_port self.mlflow_username = mlflow_username @@ -33,7 +32,7 @@ class MLFlow(BaseActivity): input_data (dict): The input data. Contains: data (dict[str, Any]): The data to transform. model_name (str): The name of the model. - model_retention (int): The retention of the model. + model_retention (int): The retention of the model in minutes. Returns: dict[str, Any]: The transformed data. """ @@ -54,6 +53,8 @@ class MLFlow(BaseActivity): response_data = self.model_monitoring_repository.transform( model_name, data, model_retention) + self.logger.debug(response_data) + return response_data @activity.defn(name="request_predict") @@ -80,4 +81,6 @@ class MLFlow(BaseActivity): response_data = self.model_monitoring_repository.predict( model_name, data, model_retention) + self.logger.debug(response_data) + return response_data diff --git a/laborious/activities/opc.py b/laborious/activities/opc.py index b6d2c38..e6a1dcc 100644 --- a/laborious/activities/opc.py +++ b/laborious/activities/opc.py @@ -4,40 +4,55 @@ from temporalio import activity, workflow with workflow.unsafe.imports_passed_through(): from logging import Logger from sientia_do.notifications.handlers import NotificationHandler + from sientia_do.notifications.models import NotificationLevel from laborious.activities.base import BaseActivity from laborious.utils.repository.opc_repository import OpcRepository from typing import Any - from sientia_do.notifications.models import NotificationLevel import traceback from pandas import DataFrame class OPC(BaseActivity): - def __init__(self, - name: str, url: str, server_uri: str, - cert_path: str, private_key_path: str, server_cert_path: str, + def __init__(self, opc_servers: dict[str, dict[str, Any]], logger: Logger, notification_handler: NotificationHandler): self.logger = logger self.notification_handler = notification_handler - self.name = name - self.url = url - self.server_uri = server_uri - self.cert_path = cert_path - self.private_key_path = private_key_path - self.server_cert_path = server_cert_path + self.opc_servers = opc_servers - self.opc_repository = OpcRepository( - name=self.name, - url=self.url, - logger=self.logger, - server_uri=self.server_uri, - cert_path=self.cert_path, - private_key_path=self.private_key_path, - server_cert_path=self.server_cert_path - ) + self.opc_repository = {} + for name, server in opc_servers.items(): + self.opc_repository[name] = OpcRepository( + name=name, + url=server['url'], + logger=self.logger, + server_uri=server['server_uri'], + cert_path=server['cert_path'], + private_key_path=server['private_key_path'], + server_cert_path=server['server_cert_path'], + notification_handler=self.notification_handler, + reconnection_interval=server['reconnection_interval'], + ) + self.opc_repository[name].connect() - self.opc_repository.connect() + BaseActivity.__init__(self, logger, notification_handler) + + def write_data(self, server: str, tag: str, data: Any, + data_type: str, tag_type: str): + try: + self.opc_repository[server].write_data( + tag, data, data_type) + self.logger.debug(f"Wrote {tag_type} to {tag}") + except Exception as e: + trace = traceback.format_exc() + self.notification_handler.build_and_send_notification( + notification_id=f"WRITE_OPC_{tag_type.upper()}_ERROR", + message=f"Error writing data to OPC server: {e}", + block="write_opc_data", + level=NotificationLevel.ERROR, + attachment_content=trace + ) + self.logger.error(trace) @activity.defn(name='write_opc_data') async def write_opc_data(self, input_data: dict[str, Any]): @@ -47,46 +62,40 @@ class OPC(BaseActivity): Args: input_data (dict[str, Any]): The input data. Contains the following keys: - data (dict[str, Any]): The dataframe that contains the data to write to the OPC servers. - opc_servers (list[str]): The OPC servers to write to. - opc_output_config (dict[str, Any]): The OPC writing configuration. Contains: + - data (dict[str, Any]): The dataframe that contains the data to write + to the OPC servers. + - opc_output_config (dict[str, Any]): The OPC writing configuration. + The keys are the OPC server names and the values contain: prediction_tags (dict[str, Any]): The tags to write to the OPC servers. confidence_tags (dict[str, Any]): The tags to write to the OPC servers. Returns: """ + self.logger.debug("Writing data to OPC servers...") data = DataFrame(input_data['data']) - _opc_servers = input_data['opc_servers'] opc_output_config = input_data['opc_output_config'] + self.logger.debug(data) - if 'prediction_tags' in opc_output_config: - for tag, config in opc_output_config['prediction_tags'].items(): - try: - self.opc_repository.write_data( - tag, data.head(1)['prediction'].values[0], config['data_type']) - except Exception as e: - trace = traceback.format_exc() - self.notification_handler.build_and_send_notification( - notification_id="WRITE_OPC_PREDICTION_ERROR", - message=f"Error writing data to OPC server: {e}", - block="write_opc_data", - level=NotificationLevel.ERROR, - attachment_content=trace - ) - self.logger.error(trace) + for server, config in opc_output_config.items(): + if self.opc_repository.get(server) is None: + self.logger.error(f"OPC server {server} not found") + continue - if 'confidence_tags' in opc_output_config: - for tag, config in opc_output_config['confidence_tags'].items(): - try: - self.opc_repository.write_data( - tag, data.head(1)['prediction_confidence'].values[0], config['data_type']) - except Exception as e: - trace = traceback.format_exc() - self.notification_handler.build_and_send_notification( - notification_id="WRITE_OPC_CONFIDENCE_ERROR", - message=f"Error writing data to OPC server: {e}", - block="write_opc_data", - level=NotificationLevel.ERROR, - attachment_content=trace + if 'prediction_tags' in config: + for tag, tag_config in config['prediction_tags'].items(): + self.write_data( + server=server, + tag=tag, + data=data.head(1)['prediction'].values[0], + data_type=tag_config['data_type'], + tag_type='prediction' + ) + if 'confidence_tags' in config: + for tag, tag_config in config['confidence_tags'].items(): + self.write_data( + server=server, + tag=tag, + data=data.head(1)['prediction_confidence'].values[0], + data_type=tag_config['data_type'], + tag_type='confidence' ) - self.logger.error(trace) diff --git a/laborious/activities/postgres.py b/laborious/activities/postgres.py index 7d8c7f6..9fc0ff7 100644 --- a/laborious/activities/postgres.py +++ b/laborious/activities/postgres.py @@ -3,6 +3,9 @@ from temporalio import workflow, activity from laborious.activities.base import BaseActivity with workflow.unsafe.imports_passed_through(): + from sqlalchemy import create_engine + from sqlalchemy.orm import sessionmaker + from sqlalchemy.pool import QueuePool from psycopg2.pool import ThreadedConnectionPool from pandas import read_sql_query, DataFrame from logging import Logger @@ -22,19 +25,20 @@ class Postgres(BaseActivity): self.password = password self.dbname = dbname - self.pool = ThreadedConnectionPool( - minconn=min_connections, - maxconn=max_connections, - host=self.host, - port=self.port, - user=self.user, - password=self.password, - dbname=self.dbname) + # Create SQLAlchemy engine with connection pooling + self.engine = create_engine( + f'postgresql://{user}:{password}@{host}:{port}/{dbname}', + poolclass=QueuePool, + pool_size=min_connections, + max_overflow=max_connections - min_connections, + pool_pre_ping=True + ) + self.session_factory = sessionmaker(bind=self.engine) - super().__init__(logger, notification_handler) + BaseActivity.__init__(self, logger, notification_handler) def close(self): - self.pool.closeall() + self.engine.dispose() def __del__(self): self.close() @@ -52,28 +56,36 @@ class Postgres(BaseActivity): """ self.logger.info(f"Fetching data from query: {query}") - conn = self.pool.getconn() - try: - data = read_sql_query(query, conn) + data = None + with self.session_factory() as session: + try: + data = read_sql_query(query, self.engine) - except Exception as e: - trace = traceback.format_exc() - self.notification_handler.build_and_send_notification( - notification_id="ERROR_LOADING_CUSTOM_QUERY", - message=f"Error fetching data from query: {e}", - block="load_custom_query", - level=NotificationLevel.ERROR, - attachment_content=trace - ) + except Exception as e: + trace = traceback.format_exc() + self.notification_handler.build_and_send_notification( + notification_id="ERROR_LOADING_CUSTOM_QUERY", + message=f"Error fetching data from query: {e}", + block="load_custom_query", + level=NotificationLevel.ERROR, + attachment_content=trace + ) - self.logger.error(trace) + self.logger.error(trace) + return {} + finally: + session.close() + + if data is None: return {} - finally: - self.pool.putconn(conn) + + # Converts any datetime datatype columns to string + for col in data.select_dtypes(include=['datetime64']).columns: + data[col] = data[col].dt.strftime('%Y-%m-%d %H:%M:%S') self.logger.info(f"Fetched {len(data)} rows") - self.logger.debug(f"Data: {data.to_string()}") + self.logger.debug(f"Data: \n{data.to_string()}") return data.to_dict() @@ -86,7 +98,7 @@ class Postgres(BaseActivity): query_items (dict[str, str]): The query items. Contains: schema (str): The schema of the table. table_name (str): The name of the table. - model (str): The model to repeat the prediction for. + model (int): The model to repeat the prediction for. Returns: None @@ -106,28 +118,25 @@ class Postgres(BaseActivity): self.logger.info(f"Repeating last prediction for model {model}") self.logger.debug(f"Query: {repeat_query}") - conn = self.pool.getconn() + with self.session_factory() as session: + try: + session.execute(repeat_query) + session.commit() - try: - cursor = conn.cursor() - cursor.execute(repeat_query) - conn.commit() - cursor.close() + except Exception as e: + trace = traceback.format_exc() + self.notification_handler.build_and_send_notification( + notification_id="ERROR_REPEATING_LAST_PREDICTION", + message=f"Error repeating last prediction: {e}", + block="repeat_last_prediction", + level=NotificationLevel.ERROR, + attachment_content=trace + ) - except Exception as e: - trace = traceback.format_exc() - self.notification_handler.build_and_send_notification( - notification_id="ERROR_REPEATING_LAST_PREDICTION", - message=f"Error repeating last prediction: {e}", - block="repeat_last_prediction", - level=NotificationLevel.ERROR, - attachment_content=trace - ) + self.logger.error(trace) - self.logger.error(trace) - - finally: - self.pool.putconn(conn) + finally: + session.close() @activity.defn(name="export_data_to_postgres") async def export_data_to_postgres(self, input_data: dict[str, Any]): @@ -141,28 +150,32 @@ class Postgres(BaseActivity): data (DataFrame): The data to export. """ + self.logger.debug( + f"Exporting data to postgres: {input_data['data']}") + schema = input_data["schema"] table_name = input_data["table_name"] data = DataFrame(input_data["data"]) - conn = self.pool.getconn() + with self.session_factory() as session: + try: + data.to_sql(table_name, self.engine, schema=schema, + if_exists="append", index=False) + session.commit() - try: - data.to_sql(table_name, conn, schema=schema, - if_exists="append", index=False) - conn.commit() + except Exception as e: + trace = traceback.format_exc() + self.notification_handler.build_and_send_notification( + notification_id="ERROR_EXPORTING_DATA_TO_POSTGRES", + message=f"Error exporting data to postgres: {e}", + block="export_data_to_postgres", + level=NotificationLevel.ERROR, + attachment_content=trace + ) - except Exception as e: - trace = traceback.format_exc() - self.notification_handler.build_and_send_notification( - notification_id="ERROR_EXPORTING_DATA_TO_POSTGRES", - message=f"Error exporting data to postgres: {e}", - block="export_data_to_postgres", - level=NotificationLevel.ERROR, - attachment_content=trace - ) + self.logger.error(trace) - self.logger.error(trace) - - finally: - self.pool.putconn(conn) + else: + self.logger.debug("Data exported to postgres") + finally: + session.close() diff --git a/laborious/utils/connectors_config.py b/laborious/utils/connectors_config.py new file mode 100644 index 0000000..80ed24d --- /dev/null +++ b/laborious/utils/connectors_config.py @@ -0,0 +1,42 @@ +from os import getenv +import json + + +def build_postgres_config(): + return { + 'host': getenv('POSTGRES_HOST', 'localhost'), + 'port': int(getenv('POSTGRES_PORT', '5432')), + 'user': getenv('POSTGRES_USER', 'sientia'), + 'password': getenv('POSTGRES_PASSWORD', 'sientia'), + 'dbname': getenv('POSTGRES_DBNAME', 'sientia'), + 'min_connections': int(getenv('POSTGRES_MIN_CONNECTIONS', '5')), + 'max_connections': int(getenv('POSTGRES_MAX_CONNECTIONS', '20')) + } + + +def build_mlflow_config(): + return { + 'host': getenv('MLFLOW_HOST', 'http://localhost'), + 'port': int(getenv('MLFLOW_PORT', '5080')), + 'username': getenv('MLFLOW_USERNAME', 'aignosi'), + 'password': getenv('MLFLOW_PASSWORD', 'aignosi') + } + + +def build_opc_config(): + opc_raw = getenv('OPC_CONFIG', None) + + if opc_raw: + return json.loads(opc_raw) + + return { + 'opc': { + 'name': getenv('OPC_NAME', 'opc'), + 'url': getenv('OPC_URL', 'opc.tcp://localhost:4840'), + 'server_uri': getenv('OPC_SERVER_URI', 'opc.tcp://localhost:4840'), + 'cert_path': getenv('OPC_CERT_PATH', None), + 'private_key_path': getenv('OPC_PRIVATE_KEY_PATH', None), + 'server_cert_path': getenv('OPC_SERVER_CERT_PATH', None), + 'reconnection_interval': int(getenv('OPC_RECONNECTION_INTERVAL', '120')) + } + } diff --git a/laborious/utils/filters/conditional_filters.py b/laborious/utils/filters/conditional_filters.py index 6638982..57cd3fd 100644 --- a/laborious/utils/filters/conditional_filters.py +++ b/laborious/utils/filters/conditional_filters.py @@ -1,13 +1,11 @@ -from typing import List - from pandas import DataFrame def filter_specific_variables_null_values(data: DataFrame, config: dict) -> bool: """ - Returns True if the data is empty, False otherwise. + Returns True if the specific columns have null values, False otherwise. """ - return data[ + return not data[ data['variable'].isin(config['VARIABLES']) & data['value'].isna()].empty diff --git a/laborious/utils/git_clone.py b/laborious/utils/git_clone.py deleted file mode 100644 index 8cfc362..0000000 --- a/laborious/utils/git_clone.py +++ /dev/null @@ -1,30 +0,0 @@ -import os -from git import Repo -from urllib.parse import quote - -# Lê variáveis de ambiente -GIT_TOKEN = os.getenv("GIT_TOKEN") -GIT_EMAIL = os.getenv("GIT_EMAIL") -REPO_URL = os.getenv("REPO_URL") # ex: "github.com/usuario/repositorio.git" -CLONE_DIR = os.getenv("CLONE_DIR", "./repo_clonado") - -if not GIT_TOKEN or not GIT_EMAIL or not REPO_URL: - raise EnvironmentError("As variáveis GIT_TOKEN, GIT_EMAIL e REPO_URL devem estar definidas.") - -# Escapa o token (caso contenha caracteres especiais) -safe_token = quote(GIT_TOKEN) - -# Monta URL com autenticação via token -repo_url_with_auth = f"https://{safe_token}@{REPO_URL}" - -# Clona o repositório -print(f"Clonando repositório em {CLONE_DIR}...") -Repo.clone_from(repo_url_with_auth, CLONE_DIR) -print("Repositório clonado com sucesso.") - -# Opcional: configura o e-mail globalmente no Git (ou dentro do repo) -repo = Repo(CLONE_DIR) -with repo.config_writer() as git_config: - git_config.set_value("user", "email", GIT_EMAIL) - -print(f"E-mail configurado como {GIT_EMAIL}.") diff --git a/laborious/utils/logger.py b/laborious/utils/logger.py new file mode 100644 index 0000000..42a9cfd --- /dev/null +++ b/laborious/utils/logger.py @@ -0,0 +1,22 @@ +from os import getenv +import logging +import sys + + +def get_logger(name: str): + log_level = getenv('LOG_LEVEL', 'INFO').upper() + + logger = logging.getLogger(name) + logger.setLevel(log_level) + stream_handler = logging.StreamHandler(sys.stdout) + stream_handler.setLevel(log_level) + + stream_handler.setFormatter( + logging.Formatter( + '%(asctime)s - %(name)s - %(levelname)s - %(message)s' + ) + ) + + logger.addHandler(stream_handler) + + return logger diff --git a/laborious/utils/policies.py b/laborious/utils/policies.py new file mode 100644 index 0000000..8c7449a --- /dev/null +++ b/laborious/utils/policies.py @@ -0,0 +1,9 @@ +from datetime import timedelta +from temporalio.common import RetryPolicy + +retry_policy = RetryPolicy( + initial_interval=timedelta(seconds=1), + backoff_coefficient=2.0, + maximum_interval=timedelta(minutes=1), + maximum_attempts=1 +) diff --git a/laborious/utils/repository/opc_repository.py b/laborious/utils/repository/opc_repository.py index 7d84f85..722bddf 100644 --- a/laborious/utils/repository/opc_repository.py +++ b/laborious/utils/repository/opc_repository.py @@ -3,20 +3,39 @@ from asyncua.sync import Client from asyncua.crypto.security_policies import SecurityPolicyBasic256 from asyncua.ua import DataValue, Variant, VariantType from logging import Logger +from datetime import datetime +from sientia_do.notifications.handlers import NotificationHandler +from sientia_do.notifications.models import NotificationLevel +import traceback data_type_map = { - 'float': VariantType.Float, - 'double': VariantType.Double, - 'int': VariantType.Int32, - 'bool': VariantType.Boolean, - 'str': VariantType.String, - 'datetime': VariantType.DateTime, + 'float': { + 'converter': float, + 'opc_type': VariantType.Float, + }, + 'double': { + 'converter': float, + 'opc_type': VariantType.Double, + }, + 'int': { + 'converter': int, + 'opc_type': VariantType.Int32, + }, + 'bool': { + 'converter': bool, + 'opc_type': VariantType.Boolean, + }, + 'str': { + 'converter': str, + 'opc_type': VariantType.String, + } } class OpcRepository(): - def __init__(self, name: str, url: str, logger: Logger, server_uri: str, - cert_path: str = None, private_key_path: str = None, server_cert_path: str = None): + def __init__(self, name: str, url: str, logger: Logger, notification_handler: NotificationHandler, + reconnection_interval: int = 60, server_uri: str = None, cert_path: str = None, + private_key_path: str = None, server_cert_path: str = None): self.url = url self.name = name self.server_uri = server_uri @@ -24,8 +43,10 @@ class OpcRepository(): self.private_key_path = private_key_path self.server_cert_path = server_cert_path self.logger = logger - self.non_receive_count = 0 - + self.error_count = 0 + self.reconnection_interval = reconnection_interval + self.last_reconnection_time = None + self.notification_handler = notification_handler self.client = None def set_security(self): @@ -67,13 +88,6 @@ class OpcRepository(): self.client.secure_channel_timeout = 10000000 self.client.session_timeout = 10000000 - def connect(self): - self.client = Client(self.url) - if self.security: - self.set_security() - self.logger.info('Starting connection...') - self.client.connect() - def connect(self): """ Establishes a connection to the OPC server. @@ -88,7 +102,24 @@ class OpcRepository(): if self.cert_path: self.set_security() self.logger.info('Starting connection...') - self.client.connect() + return self.try_connect() + + def try_connect(self): + try: + self.last_reconnection_time = datetime.now() + self.client.connect() + return True + except Exception as e: + trace = traceback.format_exc() + self.notification_handler.build_and_send_notification( + notification_id=f"OPC_CONNECTION_ERROR_{self.name}", + message=f"Failed to connect to OPC server: {e}", + block="opc_repository", + level=NotificationLevel.ERROR, + attachment_content=trace + ) + self.logger.error(trace) + return False def disconnect(self): self.client.disconnect() @@ -98,9 +129,74 @@ class OpcRepository(): def __del__(self): self.disconnect() - def write_data(self, node, value, data_type, logger): - node = self.client.get_node(node) - data = float(value) - logger.info(f'Writing {data} - {type(data)} to {node}') - ua_data = DataValue(Variant(data, data_type_map[data_type])) - node.write_value(ua_data) + def validate_connection(self): + if self.client is None: + return self.connect() + + if self.error_count > 5: + self.logger.warning( + f"OPC server {self.name} will be disconnected due to multiple errors") + try: + self.disconnect() + except Exception as e: + trace = traceback.format_exc() + self.logger.error(f"Failed to disconnect from OPC server: {e}") + self.logger.error(trace) + self.logger.info( + f"Attempting to reconnect to OPC server {self.name}...") + return self.connect() + + if hasattr(self.client, 'aio_obj') and self.client.aio_obj.uaclient.protocol is None or \ + (hasattr(self.client.aio_obj.uaclient, 'protocol') and + self.client.aio_obj.uaclient.protocol.state == "closed"): + + self.logger.error( + f"OPC server {self.name} is not connected") + if (datetime.now() - self.last_reconnection_time).total_seconds( + ) > self.reconnection_interval: + self.logger.error( + f"Trying to reconnect to OPC server {self.name}...") + return self.try_connect() + + return False + + return True + + def write_data(self, node, value, data_type): + if not self.validate_connection(): + return + try: + node = self.client.get_node(node) + except Exception as e: + trace = traceback.format_exc() + self.notification_handler.build_and_send_notification( + notification_id=f"OPC_WRITE_GET_NODE_ERROR_{self.name}", + message=f"Failed to get node from OPC server: {e}", + block="opc_repository", + level=NotificationLevel.ERROR, + attachment_content=trace + ) + self.logger.error(trace) + self.error_count += 1 + return + + data = data_type_map[data_type]['converter'](value) + self.logger.info(f'Writing {data} - {type(data)} to {node}') + ua_data = DataValue( + Variant(data, data_type_map[data_type]['opc_type'])) + + try: + node.write_value(ua_data) + except Exception as e: + trace = traceback.format_exc() + self.notification_handler.build_and_send_notification( + notification_id=f"OPC_WRITE_DATA_ERROR_{self.name}", + message=f"Failed to write data to OPC server: {e}", + block="opc_repository", + level=NotificationLevel.ERROR, + attachment_content=trace + ) + self.logger.error(trace) + self.error_count += 1 + return + self.error_count = 0 diff --git a/laborious/worker/worker.py b/laborious/worker/worker.py index 35fe543..bdf168d 100644 --- a/laborious/worker/worker.py +++ b/laborious/worker/worker.py @@ -2,30 +2,32 @@ from temporalio import workflow, client from temporalio.worker import Worker with workflow.unsafe.imports_passed_through(): - from laborious.workflows.predictions_batch import PredictionsBatch - from laborious.activities.activities import Activities import os - import logging + import asyncio + from laborious.workflows.predictions_batch import PredictionsBatch + from laborious.workflows.sub_workflows.prediction_process import PredictionProcess + from laborious.workflows.sub_workflows.format_and_export_prediction import \ + FormatAndExportPrediction + from laborious.activities.activities import Activities + from laborious.utils.logger import get_logger + from laborious.utils.connectors_config import ( + build_postgres_config, + build_mlflow_config, + build_opc_config + ) from sientia_do.notifications.handlers import NotificationHandler async def main(): host = os.getenv('TEMPORAL_HOST', 'localhost:7233') - logger = logging.getLogger(__name__) - stream_handler = logging.StreamHandler() - stream_handler.setLevel( - os.getenv('LOG_LEVEL', 'INFO').upper() - ) - stream_handler.setFormatter( - logging.Formatter( - '%(asctime)s - %(name)s - %(levelname)s - %(message)s' - ) - ) + logger = get_logger(__name__) - logger.addHandler(stream_handler) + logger.info('Starting Worker...') + + logger.info('Starting Notification Handler...') notification_handler = NotificationHandler( - servers=os.getenv('NOTIFICATION_SERVERS', 'http://localhost:29092'), + servers=os.getenv('KAFKA_SERVERS', 'http://localhost:9092'), logger=logger, project_name=os.getenv('PROJECT_NAME', 'laborious'), pipeline_name='-', @@ -34,46 +36,31 @@ async def main(): model='-' ) - postgres_config = { - 'host': os.getenv('POSTGRES_HOST', 'localhost'), - 'port': int(os.getenv('POSTGRES_PORT', '5432')), - 'user': os.getenv('POSTGRES_USER', 'sientia'), - 'password': os.getenv('POSTGRES_PASSWORD', 'sientia'), - 'dbname': os.getenv('POSTGRES_DBNAME', 'sientia'), - 'min_connections': int(os.getenv('POSTGRES_MIN_CONNECTIONS', '5')), - 'max_connections': int(os.getenv('POSTGRES_MAX_CONNECTIONS', '20')) - } - - mlflow_config = { - 'host': os.getenv('MLFLOW_HOST', 'localhost'), - 'port': int(os.getenv('MLFLOW_PORT', '5000')), - 'username': os.getenv('MLFLOW_USERNAME', 'aignosi'), - 'password': os.getenv('MLFLOW_PASSWORD', 'aignosi') - } - - opc_config = { - 'name': os.getenv('OPC_NAME', 'opc'), - 'url': os.getenv('OPC_URL', 'opc.tcp://localhost:4840'), - 'server_uri': os.getenv('OPC_SERVER_URI', 'opc.tcp://localhost:4840'), - 'cert_path': os.getenv('OPC_CERT_PATH', None), - 'private_key_path': os.getenv('OPC_PRIVATE_KEY_PATH', None), - 'server_cert_path': os.getenv('OPC_SERVER_CERT_PATH', None) - } + logger.info('Starting Activities...') activities = Activities( - postgres_config=postgres_config, - mlflow_config=mlflow_config, - opc_config=opc_config, + postgres_config=build_postgres_config(), + mlflow_config=build_mlflow_config(), + opc_config=build_opc_config(), logger=logger, notification_handler=notification_handler ) - temporal_client = await client.Client.connect(target_host=host) + logger.info('Starting Temporal Client...') + + temporal_client = await client.Client.connect( + target_host=host, + namespace=os.getenv('TEMPORAL_NAMESPACE', 'default') + ) + + logger.info('Starting Workers...') + workers = [ Worker( temporal_client, - task_queue='predictions', - workflows=[PredictionsBatch], + task_queue='predictions-queue', + workflows=[PredictionsBatch, PredictionProcess, + FormatAndExportPrediction], activities=[ # Base activities.prepare_activity, @@ -97,9 +84,13 @@ async def main(): ) ] + handlers = [] for w in workers: - await w.run() + handlers.append(w.run()) + + logger.info('Workers started successfully') + + await asyncio.gather(*handlers) if __name__ == '__main__': - import asyncio asyncio.run(main()) diff --git a/laborious/workflows/predictions_batch.py b/laborious/workflows/predictions_batch.py index 55eff80..425f59f 100644 --- a/laborious/workflows/predictions_batch.py +++ b/laborious/workflows/predictions_batch.py @@ -3,29 +3,87 @@ from temporalio import workflow with workflow.unsafe.imports_passed_through(): from laborious.activities.activities import Activities from typing import Any + from laborious.utils.policies import retry_policy + from datetime import timedelta @workflow.defn(name="predictions_batch") class PredictionsBatch(): @workflow.run async def run(self, input_data: dict[str, Any]): + """ + This workflow runs a batch of predictions based on the input data. - await workflow.execute_activity_method( + The workflow executes in two main steps: + 1. Prepares the activity with schedule and model information + 2. Loads data using a custom query and executes the prediction process + + Args: + input_data (dict[str, Any]): The input data for the workflow. + Contains the following keys: + schedule_name (str): The name of the schedule. + model_name (str): The name of the model. + model_id (int): The id of the model. + query (str): The SQL query to be executed to load data. + schema (dict, optional): The schema definition for the data. + table_name (str, optional): The name of the table to process. + input_filters (dict, optional): Filters to be applied during prediction. + mlflow_transform_filters (dict, optional): Filters to be applied during prediction. + mlflow_predict_filters (dict, optional): Filters to be applied during prediction. + model_retention (int, optional): The model retention period in minutes. + path_priority (list[str]): The path priority. + Returns: + None + + Raises: + Exception: If any of the required parameters are missing or if the workflow fails. + """ + + await workflow.execute_local_activity_method( Activities.prepare_activity, { 'schedule_name': input_data['schedule_name'], 'model_name': input_data['model_name'], 'model_id': input_data['model_id'], 'workflow_name': 'predictions_batch' - } + }, + retry_policy=retry_policy, + start_to_close_timeout=timedelta(seconds=60) ) - data = await workflow.execute_activity_method( + data = await workflow.execute_local_activity_method( Activities.load_custom_query, - input_data['query'] + input_data['query'], + retry_policy=retry_policy, + start_to_close_timeout=timedelta(seconds=60) ) - input_data['data'] = data + # Prepare input for prediction_process workflow + prediction_input = { + 'data': data, + 'schema': input_data['schema'], + 'table_name': input_data['table_name'], + 'model_id': input_data['model_id'], + 'model_name': input_data['model_name'], + 'input_filters': input_data.get('input_filters', { + 'EMPTY_DATA': { + 'POLICY': 'STOP' + } + }), + 'mlflow_transform_filters': input_data.get('mlflow_transform_filters', { + 'API_ERROR': { + 'POLICY': 'STOP' + } + }), + 'mlflow_predict_filters': input_data.get('mlflow_predict_filters', { + 'API_ERROR': { + 'POLICY': 'STOP' + } + }), + 'model_retention': input_data.get('model_retention', 60), + 'path_priority': input_data.get('path_priority', ['STOP', 'CONTINUE', 'REPEAT']), + 'opc_output_config': input_data.get('opc_output_config', {}) + } await workflow.execute_child_workflow( - 'prediction_process', input_data) + 'prediction_process', prediction_input) diff --git a/laborious/workflows/sub_workflows/format_and_export_prediction.py b/laborious/workflows/sub_workflows/format_and_export_prediction.py index 7d6d2ec..67cad49 100644 --- a/laborious/workflows/sub_workflows/format_and_export_prediction.py +++ b/laborious/workflows/sub_workflows/format_and_export_prediction.py @@ -3,6 +3,8 @@ from temporalio import workflow with workflow.unsafe.imports_passed_through(): from laborious.activities.activities import Activities from typing import Any + from datetime import timedelta + from laborious.utils.policies import retry_policy @workflow.defn(name="format_and_export_prediction") @@ -11,23 +13,25 @@ class FormatAndExportPrediction(): async def run(self, input_data: dict[str, Any]): """ This workflow formats and exports predictions based on path_flag: - - If path_flag is None: formats prediction using input data, timestamp, model_id and confidence - - If path_flag exists: creates default prediction with timestamp, model_id, confidence and comment + - If path_flag is None: formats prediction + using input data, timestamp, model_id and confidence + - If path_flag exists: creates default prediction + with timestamp, model_id, confidence and comment Finally exports formatted prediction to postgres table Args: - input_data(dict[str, Any]): The input data for the workflow. Contains the following keys: - - path_flag(str): The path flag to determine the type of prediction to format - - data(dict[str, Any]): The data to format - - prediction_confidence(float): The prediction confidence to be registered - - timestamp(str): The timestamp of the prediction, synchronized with the data - - model_id(str): The model id of the prediction - - model_name(str): The model name of the prediction - - model_retention(str): The model retention of the prediction - - comment(str): The comment to be registered - - schema(str): The schema of the prediction - - table_name(str): The table name of the prediction - - opc_servers(list[str]): The opc servers of the prediction - - opc_output_config(dict[str, Any]): The opc output config of the prediction + input_data(dict[str, Any]): The input data for the workflow. + Contains the following keys: + path_flag(str): The path flag to determine the type of prediction to format + data(dict[str, Any]): The data to format + prediction_confidence(float): The prediction confidence to be registered + timestamp(str): The timestamp of the prediction, synchronized with the data + model_id(int): The model id of the prediction + model_name(str): The model name of the prediction + model_retention(str): The model retention of the prediction + comment(str): The comment to be registered + schema(str): The schema of the prediction + table_name(str): The table name of the prediction + opc_output_config(dict[str, Any]): The opc output config of the prediction Returns: bool: True if the workflow was successful, False otherwise. @@ -36,28 +40,34 @@ class FormatAndExportPrediction(): data = input_data['data'] prediction_confidence = input_data['prediction_confidence'] + print(f"Input data: {input_data}") + if path_flag is None: # proceed with formatting and exporting - prediction = await workflow.execute_activity_method( + prediction = await workflow.execute_local_activity_method( Activities.format_prediction, { 'data': data, 'timestamp': input_data['timestamp'], 'model_id': input_data['model_id'], 'prediction_confidence': prediction_confidence, - } + }, + retry_policy=retry_policy, + start_to_close_timeout=timedelta(seconds=60) ) else: # create default prediction - prediction = await workflow.execute_activity_method( + prediction = await workflow.execute_local_activity_method( Activities.format_default_prediction, { 'timestamp': input_data['timestamp'], 'model_id': input_data['model_id'], 'prediction_confidence': prediction_confidence, 'comment': input_data['comment'] - } + }, + retry_policy=retry_policy, + start_to_close_timeout=timedelta(seconds=60) ) # write to postgres @@ -67,17 +77,20 @@ class FormatAndExportPrediction(): 'schema': input_data['schema'], 'table_name': input_data['table_name'], 'data': prediction - } + }, + retry_policy=retry_policy, + start_to_close_timeout=timedelta(seconds=60) ) # write to opc opc_holder = workflow.execute_activity_method( Activities.write_opc_data, { - 'opc_servers': input_data['opc_servers'], 'opc_output_config': input_data['opc_output_config'], 'data': prediction - } + }, + retry_policy=retry_policy, + start_to_close_timeout=timedelta(seconds=60) ) await postgres_holder diff --git a/laborious/workflows/sub_workflows/prediction_process.py b/laborious/workflows/sub_workflows/prediction_process.py index 05b3350..f6d2809 100644 --- a/laborious/workflows/sub_workflows/prediction_process.py +++ b/laborious/workflows/sub_workflows/prediction_process.py @@ -3,101 +3,144 @@ from temporalio import workflow with workflow.unsafe.imports_passed_through(): from laborious.activities.activities import Activities from typing import Any + from laborious.utils.policies import retry_policy + from datetime import timedelta @workflow.defn(name="prediction_process") class PredictionProcess(): @workflow.run async def run(self, input_data: dict[str, Any]): + """ + This workflow runs a prediction process based on the input data. + + The workflow executes in two main steps: + 1. Prepares the activity with schedule and model information + 2. Loads data using a custom query and executes the prediction process + + Args: + input_data (dict[str, Any]): The input data for the workflow. + Contains the following keys: + data (dict[str, Any]): The data to be used for the prediction. + schema (str): The schema of the table. + table_name (str): The name of the table. + model_id (int): The id of the model. + input_filters (dict, optional): Filters to be applied during prediction. + mlflow_transform_filters (dict, optional): Filters to be applied during prediction. + mlflow_predict_filters (dict, optional): Filters to be applied during prediction. + model_name (str): The name of the model. + model_retention (int, optional): The model retention period in minutes. + path_priority (list[str]): The path priority. + opc_output_config (dict[str, Any]): The opc output config of the prediction. + Returns: + None + + Raises: + Exception: If any of the required parameters are missing or if the workflow fails. + """ + data = input_data['data'] schema = input_data['schema'] table_name = input_data['table_name'] - model = input_data['model'] - filters = input_data['filters'] + model_id = input_data['model_id'] model_name = input_data['model_name'] model_retention = input_data['model_retention'] - last_timestamp = await workflow.execute_activity_method( + last_timestamp = await workflow.execute_local_activity_method( Activities.get_last_timestamp, { 'data': data - } + }, + retry_policy=retry_policy, + start_to_close_timeout=timedelta(minutes=1), ) - path_flag, confidence = await workflow.execute_activity_method( + path_flag, confidence, comment = await workflow.execute_local_activity_method( Activities.input_gate, { - 'filters': input_data['filters'], - 'data': data - } + 'filters': input_data['input_filters'], + 'data': data, + 'path_priority': input_data['path_priority'] + }, + retry_policy=retry_policy, + start_to_close_timeout=timedelta(minutes=1), ) if await self.path_flag_handler( - data, path_flag, confidence, schema, table_name, - model, last_timestamp, model_name, model_retention + data, path_flag, input_data, confidence, last_timestamp, comment ): return - response_data = await workflow.execute_activity_method( + response_data = await workflow.execute_local_activity_method( Activities.request_transform, { 'data': data, 'model_name': model_name, 'model_retention': model_retention - } + }, + retry_policy=retry_policy, + start_to_close_timeout=timedelta(minutes=1), ) - path_flag, confidence = await workflow.execute_activity_method( + path_flag, confidence, comment = await workflow.execute_local_activity_method( Activities.mlflow_response_gate, { - 'filters': filters, + 'filters': input_data['mlflow_transform_filters'], 'data': response_data, - 'type': 'transform' - } + 'type': 'transform', + 'path_priority': input_data['path_priority'] + }, + retry_policy=retry_policy, + start_to_close_timeout=timedelta(minutes=1), ) if await self.path_flag_handler( - data, path_flag, confidence, schema, table_name, - model, last_timestamp, model_name, model_retention + data, path_flag, input_data, confidence, last_timestamp, comment ): return - path_flag, confidence = await workflow.execute_activity_method( + path_flag, confidence, comment = await workflow.execute_local_activity_method( Activities.mlflow_content_gate, { - 'filters': filters, + 'filters': input_data['mlflow_transform_filters'], 'data': response_data, - 'type': 'transform' - } + 'type': 'transform', + 'path_priority': input_data['path_priority'] + }, + retry_policy=retry_policy, + start_to_close_timeout=timedelta(minutes=1), ) if await self.path_flag_handler( - data, path_flag, confidence, schema, table_name, - model, last_timestamp, model_name, model_retention + data, path_flag, input_data, confidence, last_timestamp, comment ): return - response_data = await workflow.execute_activity_method( + response_data = await workflow.execute_local_activity_method( Activities.request_predict, { 'data': response_data, 'model_name': model_name, 'model_retention': model_retention - } + }, + retry_policy=retry_policy, + start_to_close_timeout=timedelta(minutes=1), ) - path_flag, confidence = await workflow.execute_activity_method( + path_flag, confidence, comment = await workflow.execute_local_activity_method( Activities.mlflow_response_gate, { - 'filters': filters, + 'filters': input_data['mlflow_predict_filters'], 'data': response_data, - 'type': 'predict' - } + 'type': 'predict', + 'path_priority': input_data['path_priority'] + }, + retry_policy=retry_policy, + start_to_close_timeout=timedelta(minutes=1), ) if await self.path_flag_handler( - data, path_flag, confidence, schema, table_name, - model, last_timestamp, model_name, model_retention + data, path_flag, input_data, confidence, last_timestamp, comment ): return @@ -108,60 +151,78 @@ class PredictionProcess(): 'data': response_data['content'], 'prediction_confidence': confidence, 'timestamp': response_data['timestamp'], - 'model_id': model, + 'model_id': model_id, 'model_name': model_name, - 'model_retention': model_retention + 'model_retention': model_retention, + 'opc_output_config': input_data['opc_output_config'] } ) async def path_flag_handler(self, data: dict[str, Any], path_flag: str, - confidence: int, schema: str, table_name: str, - model: str, last_timestamp: str, model_name: str, - model_retention: str): + input_data: dict[str, Any], confidence: int, + last_timestamp: str, comment: str): """ This function handles the path flag and the confidence of the prediction. - It returns True if the prediction should be stopped. If path_flag is 'repeat', it repeats the last prediction. - If path_flag is 'continue', it calls the write workflow. If path_flag is 'stop', it stops the prediction process. + It returns True if the prediction should be stopped. If path_flag is 'repeat', + it repeats the last prediction. + If path_flag is 'continue', it calls the write workflow. If path_flag is 'stop', + it stops the prediction process. Args: data (dict[str, Any]): The data to be used for the prediction. path_flag (str): The path flag to determine the type of prediction to format confidence (int): The confidence of the prediction schema (str): The schema of the prediction table_name (str): The table name of the prediction - model (str): The model id of the prediction + model_id (int): The model id of the prediction last_timestamp (str): The timestamp of the last prediction model_name (str): The model name of the prediction - model_retention (str): The model retention of the prediction + model_retention (int): The model retention of the prediction + comment (str): The comment of the prediction Returns: bool: True if the prediction should be stopped, False otherwise. """ - if path_flag == 'stop': + + schema = input_data['schema'] + table_name = input_data['table_name'] + model_id = input_data['model_id'] + model_name = input_data['model_name'] + model_retention = input_data['model_retention'] + + path_flag = path_flag.upper() if path_flag else None + + if path_flag == 'STOP': return True - elif path_flag == 'repeat': + elif path_flag == 'REPEAT': # repeat last prediction await workflow.execute_activity_method( Activities.repeat_last_prediction, { 'schema': schema, 'table_name': table_name, - 'model': model - } + 'model_id': model_id + }, + retry_policy=retry_policy, + start_to_close_timeout=timedelta(minutes=1), ) return True - elif path_flag == 'continue': + elif path_flag == 'CONTINUE': # call write workflow - workflow.execute_child_workflow( + await workflow.execute_child_workflow( 'format_and_export_prediction', { 'path_flag': path_flag, 'data': data, 'prediction_confidence': confidence, 'timestamp': last_timestamp, - 'model_id': model, + 'model_id': model_id, 'model_name': model_name, - 'model_retention': model_retention + 'model_retention': model_retention, + 'schema': schema, + 'table_name': table_name, + 'comment': comment, + 'opc_output_config': input_data['opc_output_config'] } ) return True diff --git a/requirements.txt b/requirements.txt index 1ccf40c..90e049b 100644 --- a/requirements.txt +++ b/requirements.txt @@ -1,5 +1,7 @@ temporalio psycopg2-binary +sqlalchemy asyncua +redis git+ssh://git@github.com/Aignosi/sientia-dataops-library.git git+ssh://git@github.com/Aignosi/sientia-mlops-library.git diff --git a/tests/laborious/activities/test_base.py b/tests/laborious/activities/test_base.py index 9b704ae..6978acb 100644 --- a/tests/laborious/activities/test_base.py +++ b/tests/laborious/activities/test_base.py @@ -1,6 +1,6 @@ from unittest.mock import MagicMock from laborious.activities.base import BaseActivity -from pytest import fixture +from pytest import fixture, mark from sientia_do.notifications.models import Notification @@ -12,7 +12,8 @@ def base_activity(): ) -def test_prepare_activity(base_activity): +@mark.asyncio +async def test_prepare_activity(base_activity): base_activity.notification_handler.base_notification = Notification( project="project", pipeline="pipeline", @@ -21,12 +22,14 @@ def test_prepare_activity(base_activity): model_id="-", ) - base_activity.prepare_activity( - schedule_name="test_schedule", - model_name="test_model", - model_id="test_model_id", - ) + await base_activity.prepare_activity({ + 'workflow_name': 'test_workflow', + 'schedule_name': 'test_schedule', + 'model_name': 'test_model', + 'model_id': 'test_model_id' + }) assert base_activity.notification_handler.base_notification.schedule_name == "test_schedule" assert base_activity.notification_handler.base_notification.model_name == "test_model" assert base_activity.notification_handler.base_notification.model_id == "test_model_id" + assert base_activity.notification_handler.base_notification.pipeline_name == "test_workflow" diff --git a/tests/laborious/activities/test_gates.py b/tests/laborious/activities/test_gates.py index 87c8254..03c1977 100644 --- a/tests/laborious/activities/test_gates.py +++ b/tests/laborious/activities/test_gates.py @@ -51,7 +51,7 @@ async def test_input_gate_specific_variables_null_values_with_stop_policy_only( } result = await gates.input_gate(input_data) - assert result == ('stop', -1) + assert result == ('stop', -1, 'Input data with bad quality') input_args = specific_variables_null_values_mock.call_args assert input_args[0][0].equals(DataFrame( @@ -98,7 +98,7 @@ async def test_input_gate_specific_variables_null_values_with_continue_policy_on } result = await gates.input_gate(input_data) - assert result == ('continue', 2) + assert result == ('continue', 2, 'Input data with bad quality') input_args = specific_variables_null_values_mock.call_args assert input_args[0][0].equals(DataFrame( @@ -145,7 +145,7 @@ async def test_input_gate_specific_variables_null_values_no_filtered( } result = await gates.input_gate(input_data) - assert result == (None, 0) + assert result == (None, 0, '') specific_variables_null_values_input_args = specific_variables_null_values_mock.call_args @@ -197,7 +197,7 @@ async def test_input_gate_one_stop_policy( } result = await gates.input_gate(input_data) - assert result == ('stop', -1) + assert result == ('stop', -1, 'Input data with bad quality') specific_variables_null_values_input_args = specific_variables_null_values_mock.call_args assert specific_variables_null_values_input_args[0][0].equals(DataFrame( @@ -251,7 +251,7 @@ async def test_input_gate_one_continue_policy( } result = await gates.input_gate(input_data) - assert result == ('continue', 2) + assert result == ('continue', 2, 'Input data with bad quality') specific_variables_null_values_input_args = specific_variables_null_values_mock.call_args assert specific_variables_null_values_input_args[0][0].equals(DataFrame( @@ -305,7 +305,7 @@ async def test_input_gate_no_filtered( } result = await gates.input_gate(input_data) - assert result == (None, 0) + assert result == (None, 0, '') specific_variables_null_values_input_args = specific_variables_null_values_mock.call_args assert specific_variables_null_values_input_args[0][0].equals(DataFrame( @@ -340,7 +340,7 @@ async def test_input_gate_error( } result = await gates.input_gate(input_data) - assert result == (None, 0) + assert result == (None, 0, '') gates.notification_handler.build_and_send_notification.assert_called_once_with( notification_id='INTPUT_GATE_ERROR__SPECIFIC_VARIABLES_NULL_VALUES', @@ -389,7 +389,7 @@ async def test_mlflow_response_gate_no_filtered( } result = await gates.mlflow_response_gate(input_data) - assert result == (None, 0) + assert result == (None, 0, '') api_error_filter_mock.assert_called_once_with( input_data['data'], @@ -433,7 +433,7 @@ async def test_mlflow_response_gate_filtered( } result = await gates.mlflow_response_gate(input_data) - assert result == ('continue', 255) + assert result == ('continue', 255, "Error") api_error_filter_mock.assert_called_once_with( input_data['data'], @@ -480,7 +480,7 @@ async def test_mlflow_content_gate_no_filtered( } result = await gates.mlflow_content_gate(input_data) - assert result == (None, 0) + assert result == (None, 0, '') nan_values_filter_mock_args = nan_values_filter_mock.call_args assert nan_values_filter_mock_args[0][0].equals(DataFrame( @@ -504,7 +504,8 @@ async def test_mlflow_content_gate_filtered( if x == 'path_confidence': return transform_filter_path_confidence - mlflow_content_filter_functions_mock.__getitem__.side_effect = transform_filter_functions_side_effect + mlflow_content_filter_functions_mock.__getitem__.side_effect = \ + transform_filter_functions_side_effect input_data = { 'filters': { @@ -521,7 +522,8 @@ async def test_mlflow_content_gate_filtered( } result = await gates.mlflow_content_gate(input_data) - assert result == ('repeat', -1) + assert result == ( + 'repeat', -1, "Transformed data not passed the content filter") nan_values_filter_mock_args = nan_values_filter_mock.call_args assert nan_values_filter_mock_args[0][0].equals(DataFrame( diff --git a/tests/laborious/activities/test_mlflow.py b/tests/laborious/activities/test_mlflow.py index e470564..5834cb4 100644 --- a/tests/laborious/activities/test_mlflow.py +++ b/tests/laborious/activities/test_mlflow.py @@ -61,7 +61,7 @@ async def test_request_transform(mock_max, mock_dataframe, mlflow): mlflow.model_monitoring_repository.transform.return_value = expected_response # Call the method - response_data, timestamp = await mlflow.request_transform(input_data) + response_data = await mlflow.request_transform(input_data) # Verify the data was correctly transformed mock_dataframe.assert_called_once_with(input_data['data']) @@ -75,7 +75,6 @@ async def test_request_transform(mock_max, mock_dataframe, mlflow): # Verify the response assert response_data == expected_response - assert timestamp == '2024-01-02' # Verify the repository was called with correct arguments mlflow.model_monitoring_repository.transform.assert_called_once_with( diff --git a/tests/laborious/activities/test_opc.py b/tests/laborious/activities/test_opc.py index 2cde0af..c4e26ee 100644 --- a/tests/laborious/activities/test_opc.py +++ b/tests/laborious/activities/test_opc.py @@ -1,111 +1,120 @@ -from unittest.mock import patch, MagicMock - +from unittest.mock import patch, MagicMock, ANY, call from pytest import fixture, mark +from laborious.activities.opc import NotificationLevel + from laborious.activities.opc import OPC -from sientia_do.notifications.models import NotificationLevel -from unittest.mock import ANY @patch("laborious.activities.opc.OpcRepository") def test___init__(mock_opc_repository): + mock_logger = MagicMock() + server1 = MagicMock() + server2 = MagicMock() + mock_opc_repository.side_effect = [server1, server2] + mock_notification_handler = MagicMock() + servers = { + 'server1': { + 'url': 'http://localhost:8080', + 'server_uri': 'opc.tcp://localhost:4840', + 'cert_path': '', + 'private_key_path': '', + 'server_cert_path': '', + 'reconnection_interval': 60, + }, + 'server2': { + 'url': 'http://localhost:8080', + 'server_uri': 'opc.tcp://localhost:4840', + 'cert_path': '', + 'private_key_path': '', + 'server_cert_path': '', + 'reconnection_interval': 60, + } + } opc = OPC( - name="test", - url="http://localhost:8080", - server_uri="opc.tcp://localhost:4840", - cert_path="", - private_key_path="", - server_cert_path="", - logger=MagicMock(), - notification_handler=MagicMock() + opc_servers=servers, + logger=mock_logger, + notification_handler=mock_notification_handler ) - assert opc.name == "test" - assert opc.url == "http://localhost:8080" - assert opc.server_uri == "opc.tcp://localhost:4840" - assert opc.cert_path == "" - assert opc.private_key_path == "" - assert opc.server_cert_path == "" - assert opc.opc_repository == mock_opc_repository.return_value + assert opc.opc_servers == servers + assert opc.logger == mock_logger + assert opc.notification_handler == mock_notification_handler + assert opc.opc_repository['server1'] == server1 + assert opc.opc_repository['server2'] == server2 - mock_opc_repository.assert_called_once_with( - name="test", - url="http://localhost:8080", - server_uri="opc.tcp://localhost:4840", - cert_path="", - private_key_path="", - server_cert_path="", - logger=opc.logger, - ) + mock_opc_repository.assert_has_calls([ + call( + name="server1", + url="http://localhost:8080", + logger=mock_logger, + server_uri="opc.tcp://localhost:4840", + cert_path="", + private_key_path="", + server_cert_path="", + notification_handler=mock_notification_handler, + reconnection_interval=60, + ), + ]) + mock_opc_repository.assert_has_calls([ + call( + name="server2", + url="http://localhost:8080", + logger=mock_logger, + server_uri="opc.tcp://localhost:4840", + cert_path="", + private_key_path="", + server_cert_path="", + notification_handler=mock_notification_handler, + reconnection_interval=60, + ) + ]) - opc.opc_repository.connect.assert_called_once() + server1.connect.assert_called_once() + server2.connect.assert_called_once() @fixture @patch("laborious.activities.opc.OpcRepository") -def opc(mock_opc_repository): +def opc(_mock_opc_repository): + servers = { + 'server1': { + 'url': 'http://localhost:8080', + 'server_uri': 'opc.tcp://localhost:4840', + 'cert_path': '', + 'private_key_path': '', + 'server_cert_path': '', + 'reconnection_interval': 60, + } + } return OPC( - name="test", - url="http://localhost:8080", - server_uri="opc.tcp://localhost:4840", - cert_path="", - private_key_path="", - server_cert_path="", + opc_servers=servers, logger=MagicMock(), notification_handler=MagicMock() ) -@mark.asyncio -async def test_write_opc_data_success(opc): - # Arrange - input_data = { - 'data': { - 'prediction': [0.75], - 'prediction_confidence': [0.95] - }, - 'opc_servers': ['server1'], - 'opc_output_config': { - 'prediction_tags': { - 'tag1': {'data_type': 'float'} - }, - 'confidence_tags': { - 'tag2': {'data_type': 'float'} - } - } - } - - # Act - await opc.write_opc_data(input_data) - - # Assert - opc.opc_repository.write_data.assert_any_call('tag1', 0.75, 'float') - opc.opc_repository.write_data.assert_any_call('tag2', 0.95, 'float') - assert opc.opc_repository.write_data.call_count == 2 +WRITE_DATA_CASES = [ + ('tag1', 'int', 50), + ('tag2', 'float', 50.5), + ('tag3', 'bool', True), + ('tag4', 'string', 'test'), +] -@mark.asyncio -async def test_write_opc_data_prediction_error(opc): - # Arrange - input_data = { - 'data': { - 'prediction': [0.75], - 'prediction_confidence': [0.95] - }, - 'opc_servers': ['server1'], - 'opc_output_config': { - 'prediction_tags': { - 'tag1': {'data_type': 'float'} - } - } - } +@mark.parametrize('tag,data_type,data', WRITE_DATA_CASES) +def test_write_data_success(opc, tag, data_type, data): + opc.write_data(server='server1', tag=tag, data=data, + data_type=data_type, tag_type='prediction') + opc.opc_repository['server1'].write_data.assert_called_once_with( + tag, data, data_type) - opc.opc_repository.write_data.side_effect = Exception("Test error") - # Act - await opc.write_opc_data(input_data) - - # Assert - opc.notification_handler.build_and_send_notification.assert_called_with( +def test_write_data_exception(opc): + opc.opc_repository['server1'].write_data.side_effect = Exception( + "Test error") + opc.write_data(server='server1', tag='tag1', data=50, + data_type='int', tag_type='prediction') + opc.notification_handler.build_and_send_notification.assert_called_once_with( notification_id="WRITE_OPC_PREDICTION_ERROR", message="Error writing data to OPC server: Test error", block="write_opc_data", @@ -116,44 +125,48 @@ async def test_write_opc_data_prediction_error(opc): @mark.asyncio -async def test_write_opc_data_confidence_error(opc): +async def test_write_opc_data_success(opc): # Arrange input_data = { 'data': { 'prediction': [0.75], 'prediction_confidence': [0.95] }, - 'opc_servers': ['server1'], 'opc_output_config': { - 'prediction_tags': { - 'tag1': {'data_type': 'float'} - }, - 'confidence_tags': { - 'tag2': {'data_type': 'float'} + 'server1': { + 'prediction_tags': { + 'tag1': {'data_type': 'float'} + }, + 'confidence_tags': { + 'tag2': {'data_type': 'float'} + } } } } - # Make first call succeed but second fail - def side_effect(*args, **kwargs): - if args[0] == 'tag2': - raise ValueError("Test error") - return None - - opc.opc_repository.write_data.side_effect = side_effect - # Act + opc.write_data = MagicMock() await opc.write_opc_data(input_data) # Assert - opc.notification_handler.build_and_send_notification.assert_called_with( - notification_id="WRITE_OPC_CONFIDENCE_ERROR", - message="Error writing data to OPC server: Test error", - block="write_opc_data", - level=NotificationLevel.ERROR, - attachment_content=ANY - ) - opc.logger.error.assert_called_once() + opc.write_data.assert_has_calls([ + call( + server='server1', + tag='tag1', + data=0.75, + data_type='float', + tag_type='prediction' + )]) + opc.write_data.assert_has_calls([ + call( + server='server1', + tag='tag2', + data=0.95, + data_type='float', + tag_type='confidence' + ) + ]) + assert opc.write_data.call_count == 2 @mark.asyncio @@ -175,4 +188,4 @@ async def test_write_opc_data_empty_config(opc): await opc.write_opc_data(input_data) # Assert - opc.opc_repository.write_data.assert_not_called() + opc.opc_repository['server1'].write_data.assert_not_called() diff --git a/tests/laborious/utils/filters/repository/test_opc_repository.py b/tests/laborious/utils/filters/repository/test_opc_repository.py deleted file mode 100644 index 070f001..0000000 --- a/tests/laborious/utils/filters/repository/test_opc_repository.py +++ /dev/null @@ -1,106 +0,0 @@ -from unittest.mock import Mock, patch, MagicMock -from pathlib import Path -from asyncua.sync import Client -from asyncua.crypto.security_policies import SecurityPolicyBasic256 -from asyncua.ua import DataValue, Variant, VariantType -from pytest import fixture -from laborious.utils.repository.opc_repository import OpcRepository - - -@fixture -def mock_logger(): - return Mock() - - -@fixture -def opc_repository(mock_logger): - return OpcRepository( - name="test_repo", - url="opc.tcp://localhost:4840", - logger=mock_logger, - server_uri="urn:test:server", - cert_path="/path/to/cert.pem", - private_key_path="/path/to/key.pem", - server_cert_path="/path/to/server_cert.pem" - ) - - -@fixture -def mock_client(): - with patch('laborious.utils.repository.opc_repository.Client') as mock: - client_instance = MagicMock() - mock.return_value = client_instance - yield client_instance - - -def test_init(opc_repository): - assert opc_repository.name == "test_repo" - assert opc_repository.url == "opc.tcp://localhost:4840" - assert opc_repository.server_uri == "urn:test:server" - assert opc_repository.cert_path == "/path/to/cert.pem" - assert opc_repository.private_key_path == "/path/to/key.pem" - assert opc_repository.server_cert_path == "/path/to/server_cert.pem" - assert opc_repository.non_receive_count == 0 - assert opc_repository.client is None - - -def test_set_security(opc_repository, mock_client): - opc_repository.client = mock_client - opc_repository.set_security() - - mock_client.application_uri = "urn:test:server" - mock_client.set_security.assert_called_once_with( - SecurityPolicyBasic256, - certificate="/path/to/cert.pem", - private_key="/path/to/key.pem", - server_certificate="/path/to/server_cert.pem" - ) - assert mock_client.secure_channel_timeout == 10000000 - assert mock_client.session_timeout == 10000000 - - -def test_set_security_missing_certificates(opc_repository): - opc_repository.cert_path = None - opc_repository.private_key_path = None - - try: - opc_repository.set_security() - except ValueError as e: - assert str( - e) == "Certificate and private key paths must be provided for secure connection." - - -def test_connect_with_security(opc_repository, mock_client): - opc_repository.connect() - - mock_client.connect.assert_called_once() - assert opc_repository.client == mock_client - - -def test_connect_without_security(opc_repository, mock_client): - opc_repository.cert_path = None - opc_repository.connect() - - mock_client.connect.assert_called_once() - assert opc_repository.client == mock_client - - -def test_disconnect(opc_repository, mock_client): - opc_repository.client = mock_client - opc_repository.disconnect() - - mock_client.disconnect.assert_called_once() - assert opc_repository.client is None - - -def test_write_data(opc_repository, mock_client, mock_logger): - opc_repository.client = mock_client - mock_node = MagicMock() - mock_client.get_node.return_value = mock_node - - opc_repository.write_data("ns=2;s=TestNode", 42.0, "float", mock_logger) - - mock_client.get_node.assert_called_once_with("ns=2;s=TestNode") - mock_node.write_value.assert_called_once() - mock_logger.info.assert_called_once_with( - "Writing 42.0 - to " + str(mock_node)) diff --git a/tests/laborious/utils/filters/test_conditional_filters.py b/tests/laborious/utils/filters/test_conditional_filters.py index acd073a..16ddd45 100644 --- a/tests/laborious/utils/filters/test_conditional_filters.py +++ b/tests/laborious/utils/filters/test_conditional_filters.py @@ -7,18 +7,18 @@ def test_filter_specific_variables_null_values(): assert filter_specific_variables_null_values( DataFrame( {'variable': ['variable1', 'variable2'], 'value': [1, 2]}), - config={'VARIABLES': ['variable2']}) == True + config={'VARIABLES': ['variable2']}) is False def test_filter_specific_variables_null_values_with_null_values(): assert filter_specific_variables_null_values( DataFrame( {'variable': ['variable1', 'variable2'], 'value': [1, None]}), - config={'VARIABLES': ['variable2']}) == False + config={'VARIABLES': ['variable2']}) is True def test_filter_empty_data(): - assert filter_empty_data(DataFrame(), {}) == True + assert filter_empty_data(DataFrame(), {}) is True def test_filter_empty_data_with_data(): diff --git a/tests/laborious/utils/filters/repository/test_model_repository.py b/tests/laborious/utils/repository/test_model_repository.py similarity index 100% rename from tests/laborious/utils/filters/repository/test_model_repository.py rename to tests/laborious/utils/repository/test_model_repository.py diff --git a/tests/laborious/utils/repository/test_opc_repository.py b/tests/laborious/utils/repository/test_opc_repository.py new file mode 100644 index 0000000..3dc7f11 --- /dev/null +++ b/tests/laborious/utils/repository/test_opc_repository.py @@ -0,0 +1,242 @@ +from unittest.mock import Mock, patch, MagicMock, ANY, call +from asyncua.crypto.security_policies import SecurityPolicyBasic256 +from pytest import fixture +from laborious.utils.repository.opc_repository import OpcRepository +from sientia_do.notifications.models import NotificationLevel +from datetime import datetime + + +@fixture +def mock_logger(): + return Mock() + + +@fixture +def opc_repository(mock_logger): + return OpcRepository( + name="test_repo", + url="opc.tcp://localhost:4840", + logger=mock_logger, + notification_handler=Mock(), + reconnection_interval=60, + server_uri="urn:test:server", + cert_path="/path/to/cert.pem", + private_key_path="/path/to/key.pem", + server_cert_path="/path/to/server_cert.pem" + ) + + +@fixture +def mock_client(): + with patch('laborious.utils.repository.opc_repository.Client') as mock: + client_instance = MagicMock() + mock.return_value = client_instance + yield client_instance + + +def test_init(opc_repository): + assert opc_repository.name == "test_repo" + assert opc_repository.url == "opc.tcp://localhost:4840" + assert opc_repository.server_uri == "urn:test:server" + assert opc_repository.cert_path == "/path/to/cert.pem" + assert opc_repository.private_key_path == "/path/to/key.pem" + assert opc_repository.server_cert_path == "/path/to/server_cert.pem" + assert opc_repository.reconnection_interval == 60 + assert opc_repository.client is None + assert opc_repository.last_reconnection_time is None + assert opc_repository.error_count == 0 + + +def test_set_security(opc_repository, mock_client): + opc_repository.client = mock_client + opc_repository.set_security() + + mock_client.application_uri = "urn:test:server" + mock_client.set_security.assert_called_once_with( + SecurityPolicyBasic256, + certificate="/path/to/cert.pem", + private_key="/path/to/key.pem", + server_certificate="/path/to/server_cert.pem" + ) + assert mock_client.secure_channel_timeout == 10000000 + assert mock_client.session_timeout == 10000000 + + +def test_set_security_missing_certificates(opc_repository): + opc_repository.cert_path = None + opc_repository.private_key_path = None + + try: + opc_repository.set_security() + except ValueError as e: + assert str( + e) == "Certificate and private key paths must be provided for secure connection." + + +def test_connect_with_security(opc_repository, mock_client): + opc_repository.try_connect = MagicMock() + opc_repository.connect() + + opc_repository.try_connect.assert_called_once() + assert opc_repository.client == mock_client + + +def test_connect_without_security(opc_repository, mock_client): + opc_repository.cert_path = None + opc_repository.try_connect = MagicMock() + opc_repository.set_security = MagicMock() + opc_repository.connect() + + opc_repository.try_connect.assert_called_once() + opc_repository.set_security.assert_not_called() + assert opc_repository.client == mock_client + + +def test_try_connect_sucess(opc_repository): + opc_repository.last_reconnection_time = None + opc_repository.client = MagicMock() + opc_repository.try_connect() + opc_repository.client.connect.assert_called_once() + assert opc_repository.last_reconnection_time is not None + + +def test_try_connect_fail(opc_repository): + opc_repository.last_reconnection_time = None + opc_repository.client = MagicMock() + opc_repository.client.connect.side_effect = Exception("Test error") + + opc_repository.try_connect() + + opc_repository.client.connect.assert_called_once() + assert opc_repository.last_reconnection_time is not None + opc_repository.notification_handler.build_and_send_notification.assert_called_once_with( + notification_id=f"OPC_CONNECTION_ERROR_{opc_repository.name}", + message="Failed to connect to OPC server: Test error", + block="opc_repository", + level=NotificationLevel.ERROR, + attachment_content=ANY + ) + + +def test_disconnect(opc_repository, mock_client): + opc_repository.client = mock_client + opc_repository.disconnect() + + mock_client.disconnect.assert_called_once() + assert opc_repository.client is None + + +def test_validate_connection_none_client(opc_repository): + opc_repository.client = None + opc_repository.connect = MagicMock() + response = opc_repository.validate_connection() + assert response + opc_repository.connect.assert_called_once() + + +def test_validate_connection_error_count_disconnect_error(opc_repository): + opc_repository.error_count = 6 + opc_repository.client = MagicMock() + opc_repository.disconnect = MagicMock(side_effect=Exception("Test error")) + opc_repository.connect = MagicMock() + + response = opc_repository.validate_connection() + assert response == opc_repository.connect.return_value + opc_repository.disconnect.assert_called_once() + opc_repository.connect.assert_called_once() + opc_repository.logger.error.assert_has_calls( + [ + call("Failed to disconnect from OPC server: Test error"), + ] + ) + + +@patch('laborious.utils.repository.opc_repository.hasattr', return_value=True) +@patch('laborious.utils.repository.opc_repository.datetime', + MagicMock(now=MagicMock(return_value=datetime(2025, 1, 1, 0, 0, 0)))) +def test_validate_connection_lost_not_time_to_reconect(_mock_datetime, opc_repository): + opc_repository.error_count = 0 + opc_repository.client = MagicMock() + opc_repository.client.aio_obj.uaclient.protocol = None + opc_repository.last_reconnection_time = datetime(2025, 1, 1, 0, 0, 0) + opc_repository.try_connect = MagicMock() + + response = opc_repository.validate_connection() + opc_repository.try_connect.assert_not_called() + assert response is False + + +@patch('laborious.utils.repository.opc_repository.hasattr', return_value=True) +@patch('laborious.utils.repository.opc_repository.datetime', + MagicMock(now=MagicMock(return_value=datetime(2025, 1, 1, 1, 0, 0)))) +def test_validate_connection_lost_time_to_reconect(_mock_datetime, opc_repository): + opc_repository.error_count = 0 + opc_repository.client = MagicMock() + opc_repository.client.aio_obj.uaclient.protocol = None + opc_repository.last_reconnection_time = datetime(2025, 1, 1, 0, 0, 0) + opc_repository.try_connect = MagicMock() + + response = opc_repository.validate_connection() + opc_repository.try_connect.assert_called_once() + assert response == opc_repository.try_connect.return_value + + +def test_write_data_validate_connection_failed(opc_repository): + opc_repository.validate_connection = MagicMock(return_value=False) + opc_repository.client = MagicMock() + opc_repository.write_data("ns=2;s=TestNode", 42.0, "float") + opc_repository.validate_connection.assert_called_once() + opc_repository.client.get_node.assert_not_called() + + +def test_write_data_get_node_failed(opc_repository): + opc_repository.validate_connection = MagicMock(return_value=True) + opc_repository.client = MagicMock() + opc_repository.error_count = 0 + opc_repository.client.get_node.side_effect = Exception("Test error") + opc_repository.write_data("ns=2;s=TestNode", 42.0, "float") + opc_repository.validate_connection.assert_called_once() + opc_repository.client.get_node.assert_called_once_with("ns=2;s=TestNode") + opc_repository.notification_handler.build_and_send_notification.assert_called_once_with( + notification_id=f"OPC_WRITE_GET_NODE_ERROR_{opc_repository.name}", + message="Failed to get node from OPC server: Test error", + block="opc_repository", + level=NotificationLevel.ERROR, + attachment_content=ANY + ) + assert opc_repository.error_count == 1 + + +def test_write_data(opc_repository, mock_client): + opc_repository.validate_connection = MagicMock(return_value=True) + opc_repository.client = mock_client + mock_node = MagicMock() + mock_client.get_node.return_value = mock_node + + opc_repository.write_data("ns=2;s=TestNode", 42.0, "float") + + mock_client.get_node.assert_called_once_with("ns=2;s=TestNode") + mock_node.write_value.assert_called_once() + opc_repository.logger.info.assert_called_once_with( + "Writing 42.0 - to " + str(mock_node)) + + +def test_write_data_write_value_failed(opc_repository, mock_client): + opc_repository.validate_connection = MagicMock(return_value=True) + opc_repository.client = mock_client + mock_node = MagicMock() + opc_repository.error_count = 0 + mock_client.get_node.return_value = mock_node + mock_node.write_value.side_effect = Exception("Test error") + opc_repository.write_data("ns=2;s=TestNode", 42.0, "float") + opc_repository.validate_connection.assert_called_once() + mock_client.get_node.assert_called_once_with("ns=2;s=TestNode") + mock_node.write_value.assert_called_once() + opc_repository.notification_handler.build_and_send_notification.assert_called_once_with( + notification_id=f"OPC_WRITE_DATA_ERROR_{opc_repository.name}", + message="Failed to write data to OPC server: Test error", + block="opc_repository", + level=NotificationLevel.ERROR, + attachment_content=ANY + ) + assert opc_repository.error_count == 1 diff --git a/tests/laborious/workflows/subworkflows/test_format_and_export_prediction.py b/tests/laborious/workflows/subworkflows/test_format_and_export_prediction.py index 31e5db1..a762283 100644 --- a/tests/laborious/workflows/subworkflows/test_format_and_export_prediction.py +++ b/tests/laborious/workflows/subworkflows/test_format_and_export_prediction.py @@ -1,4 +1,4 @@ -from unittest.mock import call, patch, AsyncMock +from unittest.mock import call, patch, AsyncMock, ANY from pytest import mark, fixture from laborious.activities.activities import Activities @@ -28,7 +28,7 @@ async def test_run_none_path_flag(workflow_mock, format_and_export_prediction): await format_and_export_prediction.run(input_data) - workflow_mock.execute_activity_method.assert_has_calls([ + workflow_mock.execute_local_activity_method.assert_has_calls([ call( Activities.format_prediction, { @@ -36,7 +36,9 @@ async def test_run_none_path_flag(workflow_mock, format_and_export_prediction): 'timestamp': input_data['timestamp'], 'model_id': input_data['model_id'], 'prediction_confidence': input_data['prediction_confidence'] - } + }, + retry_policy=ANY, + start_to_close_timeout=ANY )]) workflow_mock.execute_activity_method.assert_has_calls([ call( @@ -44,8 +46,10 @@ async def test_run_none_path_flag(workflow_mock, format_and_export_prediction): { 'schema': input_data['schema'], 'table_name': input_data['table_name'], - 'data': workflow_mock.execute_activity_method.return_value - } + 'data': workflow_mock.execute_local_activity_method.return_value + }, + retry_policy=ANY, + start_to_close_timeout=ANY )]) workflow_mock.execute_activity_method.assert_has_calls([ call( @@ -53,12 +57,15 @@ async def test_run_none_path_flag(workflow_mock, format_and_export_prediction): { 'opc_servers': input_data['opc_servers'], 'opc_output_config': input_data['opc_output_config'], - 'data': workflow_mock.execute_activity_method.return_value - } + 'data': workflow_mock.execute_local_activity_method.return_value + }, + retry_policy=ANY, + start_to_close_timeout=ANY ) ]) - assert workflow_mock.execute_activity_method.call_count == 3 + assert workflow_mock.execute_activity_method.call_count == 2 + assert workflow_mock.execute_local_activity_method.call_count == 1 @mark.asyncio @@ -80,7 +87,7 @@ async def test_run_default_path_flag(workflow_mock, format_and_export_prediction await format_and_export_prediction.run(input_data) - workflow_mock.execute_activity_method.assert_has_calls([ + workflow_mock.execute_local_activity_method.assert_has_calls([ call( Activities.format_default_prediction, { @@ -88,7 +95,9 @@ async def test_run_default_path_flag(workflow_mock, format_and_export_prediction 'model_id': input_data['model_id'], 'prediction_confidence': input_data['prediction_confidence'], 'comment': input_data['comment'] - } + }, + retry_policy=ANY, + start_to_close_timeout=ANY ) ]) workflow_mock.execute_activity_method.assert_has_calls([ @@ -97,8 +106,10 @@ async def test_run_default_path_flag(workflow_mock, format_and_export_prediction { 'schema': input_data['schema'], 'table_name': input_data['table_name'], - 'data': workflow_mock.execute_activity_method.return_value - } + 'data': workflow_mock.execute_local_activity_method.return_value + }, + retry_policy=ANY, + start_to_close_timeout=ANY ) ]) workflow_mock.execute_activity_method.assert_has_calls([ @@ -107,9 +118,12 @@ async def test_run_default_path_flag(workflow_mock, format_and_export_prediction { 'opc_servers': input_data['opc_servers'], 'opc_output_config': input_data['opc_output_config'], - 'data': workflow_mock.execute_activity_method.return_value - } + 'data': workflow_mock.execute_local_activity_method.return_value + }, + retry_policy=ANY, + start_to_close_timeout=ANY ) ]) - assert workflow_mock.execute_activity_method.call_count == 3 + assert workflow_mock.execute_activity_method.call_count == 2 + assert workflow_mock.execute_local_activity_method.call_count == 1 diff --git a/tests/laborious/workflows/subworkflows/test_prediction_process.py b/tests/laborious/workflows/subworkflows/test_prediction_process.py index 67c81e1..853c992 100644 --- a/tests/laborious/workflows/subworkflows/test_prediction_process.py +++ b/tests/laborious/workflows/subworkflows/test_prediction_process.py @@ -1,4 +1,4 @@ -from unittest.mock import AsyncMock, patch, call +from unittest.mock import AsyncMock, patch, call, ANY from pytest import fixture, mark from laborious.activities.activities import Activities from laborious.workflows.sub_workflows.prediction_process import PredictionProcess @@ -18,65 +18,78 @@ async def test_run(workflow_mock, prediction_process): 'data': {'test': 'data'}, 'schema': 'test_schema', 'table_name': 'test_table', - 'model': 'test_model', - 'filters': {'test': 'filter'}, + 'model_id': 1, + 'input_filters': {'test': 'filter'}, + 'mlflow_transform_filters': {'test': 'filter'}, + 'mlflow_predict_filters': {'test': 'filter'}, 'model_name': 'test_model_name', - 'model_retention': '30' + 'model_retention': '30', + 'path_priority': ['continue', 'repeat', 'stop'], + 'opc_output_config': {'test': 'config'}, } # Mock the activity responses - workflow_mock.execute_activity_method.side_effect = [ + workflow_mock.execute_local_activity_method.side_effect = [ '2024-01-01', # get_last_timestamp - ('continue', 0.95), # input_gate + ('continue', 0.95, "Input data with bad quality"), # input_gate {'content': 'transformed_data', 'timestamp': '2024-01-01'}, # transform_data - ('continue', 0.95), # mlflow_response_gate (transform) - ('continue', 0.95), # mlflow_content_gate (transform) + # mlflow_response_gate (transform) + ('continue', 0.95, "Error"), + # mlflow_content_gate (transform) + ('continue', 0.95, "Transformed data not passed the content filter"), {'content': 'predicted_data', 'timestamp': '2024-01-01'}, # request_predict - ('continue', 0.95), # mlflow_response_gate (predict) + # mlflow_response_gate (predict) + ('continue', 0.95, "Error"), ] # Act await prediction_process.run(input_data) # Assert - assert workflow_mock.execute_activity_method.call_count == 7 - workflow_mock.execute_activity_method.assert_has_calls([ - call(Activities.get_last_timestamp, {'data': input_data['data']})]) - workflow_mock.execute_activity_method.assert_has_calls([ + assert workflow_mock.execute_local_activity_method.call_count == 7 + + workflow_mock.execute_local_activity_method.assert_has_calls([ + call(Activities.get_last_timestamp, {'data': input_data['data']}, + retry_policy=ANY, start_to_close_timeout=ANY)]) + workflow_mock.execute_local_activity_method.assert_has_calls([ call(Activities.input_gate, { - 'filters': input_data['filters'], - 'data': input_data['data'] - })]) - workflow_mock.execute_activity_method.assert_has_calls([ + 'filters': input_data['input_filters'], + 'data': input_data['data'], + 'path_priority': input_data['path_priority'] + }, retry_policy=ANY, start_to_close_timeout=ANY)]) + workflow_mock.execute_local_activity_method.assert_has_calls([ call(Activities.request_transform, { 'data': input_data['data'], 'model_name': input_data['model_name'], 'model_retention': input_data['model_retention'] - })]) - workflow_mock.execute_activity_method.assert_has_calls([ + }, retry_policy=ANY, start_to_close_timeout=ANY)]) + workflow_mock.execute_local_activity_method.assert_has_calls([ call(Activities.mlflow_response_gate, { - 'filters': input_data['filters'], + 'filters': input_data['mlflow_transform_filters'], 'data': {'content': 'transformed_data', 'timestamp': '2024-01-01'}, - 'type': 'transform' - })]) - workflow_mock.execute_activity_method.assert_has_calls([ + 'type': 'transform', + 'path_priority': input_data['path_priority'] + }, retry_policy=ANY, start_to_close_timeout=ANY)]) + workflow_mock.execute_local_activity_method.assert_has_calls([ call(Activities.mlflow_content_gate, { - 'filters': input_data['filters'], + 'filters': input_data['mlflow_transform_filters'], 'data': {'content': 'transformed_data', 'timestamp': '2024-01-01'}, - 'type': 'transform' - })]) - workflow_mock.execute_activity_method.assert_has_calls([ + 'type': 'transform', + 'path_priority': input_data['path_priority'] + }, retry_policy=ANY, start_to_close_timeout=ANY)]) + workflow_mock.execute_local_activity_method.assert_has_calls([ call(Activities.request_predict, { 'data': {'content': 'transformed_data', 'timestamp': '2024-01-01'}, 'model_name': input_data['model_name'], 'model_retention': input_data['model_retention'] - })]) - workflow_mock.execute_activity_method.assert_has_calls([ + }, retry_policy=ANY, start_to_close_timeout=ANY)]) + workflow_mock.execute_local_activity_method.assert_has_calls([ call(Activities.mlflow_response_gate, { - 'filters': input_data['filters'], + 'filters': input_data['mlflow_predict_filters'], 'data': {'content': 'predicted_data', 'timestamp': '2024-01-01'}, - 'type': 'predict' - })]) + 'type': 'predict', + 'path_priority': input_data['path_priority'] + }, retry_policy=ANY, start_to_close_timeout=ANY)]) workflow_mock.execute_child_workflow.assert_called_once_with( 'format_and_export_prediction', @@ -85,9 +98,10 @@ async def test_run(workflow_mock, prediction_process): 'data': 'predicted_data', 'prediction_confidence': 0.95, 'timestamp': '2024-01-01', - 'model_id': 'test_model', + 'model_id': 1, 'model_name': 'test_model_name', - 'model_retention': '30' + 'model_retention': '30', + 'opc_output_config': input_data['opc_output_config'] } ) @@ -101,27 +115,34 @@ async def test_run_stop_at_input_gate(workflow_mock, prediction_process): 'data': {'test': 'data'}, 'schema': 'test_schema', 'table_name': 'test_table', - 'model': 'test_model', - 'filters': {'test': 'filter'}, + 'model_id': 1, + 'input_filters': {'test': 'filter'}, + 'mlflow_transform_filters': {'test': 'filter'}, + 'mlflow_predict_filters': {'test': 'filter'}, 'model_name': 'test_model_name', - 'model_retention': '30' + 'model_retention': '30', + 'path_priority': ['continue', 'repeat', 'stop'], + 'opc_output_config': {'test': 'config'} } # Mock the activity responses - workflow_mock.execute_activity_method.side_effect = [ + workflow_mock.execute_local_activity_method.side_effect = [ '2024-01-01', # get_last_timestamp - ('stop', 0.95), # input_gate + ('stop', 0.95, "Input data with bad quality"), # input_gate ] # Act await prediction_process.run(input_data) # Assert - assert workflow_mock.execute_activity_method.call_count == 2 - workflow_mock.execute_activity_method.assert_has_calls([ - call(Activities.get_last_timestamp, {'data': input_data['data']}), + assert workflow_mock.execute_local_activity_method.call_count == 2 + workflow_mock.execute_local_activity_method.assert_has_calls([ + call(Activities.get_last_timestamp, { + 'data': input_data['data']}, retry_policy=ANY, start_to_close_timeout=ANY), call(Activities.input_gate, { - 'filters': input_data['filters'], 'data': input_data['data']}) + 'filters': input_data['input_filters'], + 'data': input_data['data'], + 'path_priority': input_data['path_priority']}, retry_policy=ANY, start_to_close_timeout=ANY) ]) workflow_mock.execute_child_workflow.assert_not_called() @@ -135,42 +156,53 @@ async def test_run_stop_at_first_mlflow_response_gate(workflow_mock, prediction_ 'data': {'test': 'data'}, 'schema': 'test_schema', 'table_name': 'test_table', - 'model': 'test_model', - 'filters': {'test': 'filter'}, + 'model_id': 1, + 'input_filters': {'test': 'filter'}, + 'mlflow_transform_filters': {'test': 'filter'}, + 'mlflow_predict_filters': {'test': 'filter'}, 'model_name': 'test_model_name', - 'model_retention': '30' + 'model_retention': '30', + 'path_priority': ['continue', 'repeat', 'stop'], + 'opc_output_config': {'test': 'config'} } # Mock the activity responses - workflow_mock.execute_activity_method.side_effect = [ + workflow_mock.execute_local_activity_method.side_effect = [ '2024-01-01', # get_last_timestamp - ('repeat', 0.95), # input_gate + ('repeat', 0.95, "Input data with bad quality"), # input_gate {'content': 'transformed_data', 'timestamp': '2024-01-01'}, # transform_data - ('continue', 0.95), # mlflow_response_gate (transform) + ('continue', 0.95, "Error"), # mlflow_response_gate (transform) ] # Act await prediction_process.run(input_data) # Assert - assert workflow_mock.execute_activity_method.call_count == 4 - workflow_mock.execute_activity_method.assert_has_calls([ - call(Activities.get_last_timestamp, {'data': input_data['data']})]) - workflow_mock.execute_activity_method.assert_has_calls([ + assert workflow_mock.execute_local_activity_method.call_count == 4 + workflow_mock.execute_local_activity_method.assert_has_calls([ + call(Activities.get_last_timestamp, {'data': input_data['data']}, + retry_policy=ANY, start_to_close_timeout=ANY)]) + workflow_mock.execute_local_activity_method.assert_has_calls([ call(Activities.input_gate, { - 'filters': input_data['filters'], 'data': input_data['data']})]) - workflow_mock.execute_activity_method.assert_has_calls([ + 'filters': input_data['input_filters'], + 'data': input_data['data'], + 'path_priority': input_data['path_priority']}, + retry_policy=ANY, start_to_close_timeout=ANY)]) + workflow_mock.execute_local_activity_method.assert_has_calls([ call(Activities.request_transform, { 'data': input_data['data'], 'model_name': input_data['model_name'], - 'model_retention': input_data['model_retention'] - })]) - workflow_mock.execute_activity_method.assert_has_calls([ + 'model_retention': input_data['model_retention']}, + retry_policy=ANY, start_to_close_timeout=ANY) + ]) + workflow_mock.execute_local_activity_method.assert_has_calls([ call(Activities.mlflow_response_gate, { - 'filters': input_data['filters'], + 'filters': input_data['mlflow_transform_filters'], 'data': {'content': 'transformed_data', 'timestamp': '2024-01-01'}, - 'type': 'transform' - })]) + 'type': 'transform', + 'path_priority': input_data['path_priority'] + }, retry_policy=ANY, start_to_close_timeout=ANY) + ]) workflow_mock.execute_child_workflow.assert_not_called() @@ -184,49 +216,61 @@ async def test_run_stop_at_mlflow_content_gate(workflow_mock, prediction_process 'data': {'test': 'data'}, 'schema': 'test_schema', 'table_name': 'test_table', - 'model': 'test_model', - 'filters': {'test': 'filter'}, + 'model_id': 1, + 'input_filters': {'test': 'filter'}, + 'mlflow_transform_filters': {'test': 'filter'}, + 'mlflow_predict_filters': {'test': 'filter'}, 'model_name': 'test_model_name', - 'model_retention': '30' + 'model_retention': '30', + 'path_priority': ['continue', 'repeat', 'stop'], + 'opc_output_config': {'test': 'config'} } # Mock the activity responses - workflow_mock.execute_activity_method.side_effect = [ + workflow_mock.execute_local_activity_method.side_effect = [ '2024-01-01', # get_last_timestamp - ('continue', 0.95), # input_gate + ('continue', 0.95, "Input data with bad quality"), # input_gate {'content': 'transformed_data', 'timestamp': '2024-01-01'}, # transform_data - ('continue', 0.95), # mlflow_response_gate (transform) - ('continue', 0.95), # mlflow_content_gate (transform) + ('continue', 0.95, "Error"), # mlflow_response_gate (transform) + # mlflow_content_gate (transform) + ('continue', 0.95, "Transformed data not passed the content filter"), ] # Act await prediction_process.run(input_data) # Assert - assert workflow_mock.execute_activity_method.call_count == 5 - workflow_mock.execute_activity_method.assert_has_calls([ - call(Activities.get_last_timestamp, {'data': input_data['data']})]) - workflow_mock.execute_activity_method.assert_has_calls([ + assert workflow_mock.execute_local_activity_method.call_count == 5 + + workflow_mock.execute_local_activity_method.assert_has_calls([ + call(Activities.get_last_timestamp, {'data': input_data['data']}, + retry_policy=ANY, start_to_close_timeout=ANY)]) + workflow_mock.execute_local_activity_method.assert_has_calls([ call(Activities.input_gate, { - 'filters': input_data['filters'], 'data': input_data['data']})]) - workflow_mock.execute_activity_method.assert_has_calls([ + 'filters': input_data['input_filters'], + 'data': input_data['data'], + 'path_priority': input_data['path_priority']}, + retry_policy=ANY, start_to_close_timeout=ANY)]) + workflow_mock.execute_local_activity_method.assert_has_calls([ call(Activities.request_transform, { 'data': input_data['data'], 'model_name': input_data['model_name'], 'model_retention': input_data['model_retention'] - })]) - workflow_mock.execute_activity_method.assert_has_calls([ + }, retry_policy=ANY, start_to_close_timeout=ANY)]) + workflow_mock.execute_local_activity_method.assert_has_calls([ call(Activities.mlflow_response_gate, { - 'filters': input_data['filters'], + 'filters': input_data['mlflow_transform_filters'], 'data': {'content': 'transformed_data', 'timestamp': '2024-01-01'}, - 'type': 'transform' - })]) - workflow_mock.execute_activity_method.assert_has_calls([ + 'type': 'transform', + 'path_priority': input_data['path_priority'] + }, retry_policy=ANY, start_to_close_timeout=ANY)]) + workflow_mock.execute_local_activity_method.assert_has_calls([ call(Activities.mlflow_content_gate, { - 'filters': input_data['filters'], + 'filters': input_data['mlflow_transform_filters'], 'data': {'content': 'transformed_data', 'timestamp': '2024-01-01'}, - 'type': 'transform' - })]) + 'type': 'transform', + 'path_priority': input_data['path_priority'] + }, retry_policy=ANY, start_to_close_timeout=ANY)]) workflow_mock.execute_child_workflow.assert_not_called() @@ -240,63 +284,75 @@ async def test_run_stop_at_mlflow_last_response_gate(workflow_mock, prediction_p 'data': {'test': 'data'}, 'schema': 'test_schema', 'table_name': 'test_table', - 'model': 'test_model', - 'filters': {'test': 'filter'}, + 'model_id': 1, + 'input_filters': {'test': 'filter'}, + 'mlflow_transform_filters': {'test': 'filter'}, + 'mlflow_predict_filters': {'test': 'filter'}, 'model_name': 'test_model_name', - 'model_retention': '30' + 'model_retention': '30', + 'path_priority': ['continue', 'repeat', 'stop'], + 'opc_output_config': {'test': 'config'} } # Mock the activity responses - workflow_mock.execute_activity_method.side_effect = [ + workflow_mock.execute_local_activity_method.side_effect = [ '2024-01-01', # get_last_timestamp - ('continue', 0.95), # input_gate + ('continue', 0.95, "Input data with bad quality"), # input_gate {'content': 'transformed_data', 'timestamp': '2024-01-01'}, # transform_data - ('continue', 0.95), # mlflow_response_gate (transform) - ('continue', 0.95), # mlflow_content_gate (transform) + ('continue', 0.95, "Error"), # mlflow_response_gate (transform) + # mlflow_content_gate (transform) + ('continue', 0.95, "Transformed data not passed the content filter"), {'content': 'predicted_data', 'timestamp': '2024-01-01'}, # request_predict - ('continue', 0.95), # mlflow_response_gate (predict) + ('continue', 0.95, "Error"), # mlflow_response_gate (predict) ] # Act await prediction_process.run(input_data) # Assert - assert workflow_mock.execute_activity_method.call_count == 7 - workflow_mock.execute_activity_method.assert_has_calls([ - call(Activities.get_last_timestamp, {'data': input_data['data']})]) - workflow_mock.execute_activity_method.assert_has_calls([ + assert workflow_mock.execute_local_activity_method.call_count == 7 + workflow_mock.execute_local_activity_method.assert_has_calls([ + call(Activities.get_last_timestamp, {'data': input_data['data']}, + retry_policy=ANY, start_to_close_timeout=ANY)]) + workflow_mock.execute_local_activity_method.assert_has_calls([ call(Activities.input_gate, { - 'filters': input_data['filters'], 'data': input_data['data']})]) - workflow_mock.execute_activity_method.assert_has_calls([ + 'filters': input_data['input_filters'], + 'data': input_data['data'], + 'path_priority': input_data['path_priority']}, + retry_policy=ANY, start_to_close_timeout=ANY)]) + workflow_mock.execute_local_activity_method.assert_has_calls([ call(Activities.request_transform, { 'data': input_data['data'], 'model_name': input_data['model_name'], 'model_retention': input_data['model_retention'] - })]) - workflow_mock.execute_activity_method.assert_has_calls([ + }, retry_policy=ANY, start_to_close_timeout=ANY)]) + workflow_mock.execute_local_activity_method.assert_has_calls([ call(Activities.mlflow_response_gate, { - 'filters': input_data['filters'], + 'filters': input_data['mlflow_transform_filters'], 'data': {'content': 'transformed_data', 'timestamp': '2024-01-01'}, - 'type': 'transform' - })]) - workflow_mock.execute_activity_method.assert_has_calls([ + 'type': 'transform', + 'path_priority': input_data['path_priority'] + }, retry_policy=ANY, start_to_close_timeout=ANY)]) + workflow_mock.execute_local_activity_method.assert_has_calls([ call(Activities.mlflow_content_gate, { - 'filters': input_data['filters'], + 'filters': input_data['mlflow_transform_filters'], 'data': {'content': 'transformed_data', 'timestamp': '2024-01-01'}, - 'type': 'transform' - })]) - workflow_mock.execute_activity_method.assert_has_calls([ + 'type': 'transform', + 'path_priority': input_data['path_priority'] + }, retry_policy=ANY, start_to_close_timeout=ANY)]) + workflow_mock.execute_local_activity_method.assert_has_calls([ call(Activities.request_predict, { 'data': {'content': 'transformed_data', 'timestamp': '2024-01-01'}, 'model_name': input_data['model_name'], 'model_retention': input_data['model_retention'] - })]) - workflow_mock.execute_activity_method.assert_has_calls([ + }, retry_policy=ANY, start_to_close_timeout=ANY)]) + workflow_mock.execute_local_activity_method.assert_has_calls([ call(Activities.mlflow_response_gate, { - 'filters': input_data['filters'], + 'filters': input_data['mlflow_predict_filters'], 'data': {'content': 'predicted_data', 'timestamp': '2024-01-01'}, - 'type': 'predict' - })]) + 'type': 'predict', + 'path_priority': input_data['path_priority'] + }, retry_policy=ANY, start_to_close_timeout=ANY)]) workflow_mock.execute_child_workflow.assert_not_called() @@ -305,7 +361,7 @@ async def test_run_stop_at_mlflow_last_response_gate(workflow_mock, prediction_p async def test_path_flag_handler_stop(workflow_mock, prediction_process): # Arrange data = {'test': 'data'} - path_flag = 'stop' + path_flag = 'STOP' confidence = 0.95 schema = 'test_schema' table_name = 'test_table' @@ -317,12 +373,12 @@ async def test_path_flag_handler_stop(workflow_mock, prediction_process): # Act result = await prediction_process.path_flag_handler( data, path_flag, confidence, schema, table_name, - model, last_timestamp, model_name, model_retention + model, last_timestamp, model_name, model_retention, "" ) # Assert assert result is True - workflow_mock.execute_activity_method.assert_not_called() + workflow_mock.execute_local_activity_method.assert_not_called() workflow_mock.execute_child_workflow.assert_not_called() @@ -343,7 +399,7 @@ async def test_path_flag_handler_repeat(workflow_mock, prediction_process): # Act result = await prediction_process.path_flag_handler( data, path_flag, confidence, schema, table_name, - model, last_timestamp, model_name, model_retention + model, last_timestamp, model_name, model_retention, "" ) # Assert @@ -353,8 +409,10 @@ async def test_path_flag_handler_repeat(workflow_mock, prediction_process): { 'schema': schema, 'table_name': table_name, - 'model': model - } + 'model_id': model + }, + retry_policy=ANY, + start_to_close_timeout=ANY ) workflow_mock.execute_child_workflow.assert_not_called() @@ -364,7 +422,7 @@ async def test_path_flag_handler_repeat(workflow_mock, prediction_process): async def test_path_flag_handler_continue(workflow_mock, prediction_process): # Arrange data = {'test': 'data'} - path_flag = 'continue' + path_flag = 'CONTINUE' confidence = 0.95 schema = 'test_schema' table_name = 'test_table' @@ -376,7 +434,7 @@ async def test_path_flag_handler_continue(workflow_mock, prediction_process): # Act result = await prediction_process.path_flag_handler( data, path_flag, confidence, schema, table_name, - model, last_timestamp, model_name, model_retention + model, last_timestamp, model_name, model_retention, 'Prediction Process' ) # Assert @@ -391,7 +449,10 @@ async def test_path_flag_handler_continue(workflow_mock, prediction_process): 'timestamp': last_timestamp, 'model_id': model, 'model_name': model_name, - 'model_retention': model_retention + 'model_retention': model_retention, + 'schema': schema, + 'table_name': table_name, + 'comment': 'Prediction Process' } ) @@ -413,7 +474,7 @@ async def test_path_flag_handler_unknown(workflow_mock, prediction_process): # Act result = await prediction_process.path_flag_handler( data, path_flag, confidence, schema, table_name, - model, last_timestamp, model_name, model_retention + model, last_timestamp, model_name, model_retention, "" ) # Assert diff --git a/tests/laborious/workflows/subworkflows/test_predictions_batch.py b/tests/laborious/workflows/subworkflows/test_predictions_batch.py deleted file mode 100644 index 065ca31..0000000 --- a/tests/laborious/workflows/subworkflows/test_predictions_batch.py +++ /dev/null @@ -1,48 +0,0 @@ -from unittest.mock import AsyncMock, call, patch -from pytest import fixture, mark -from laborious.activities.activities import Activities -from laborious.workflows.predictions_batch import PredictionsBatch - - -@fixture -def predictions_batch() -> PredictionsBatch: - return PredictionsBatch() - - -@mark.asyncio -@patch('laborious.workflows.predictions_batch.workflow', new_callable=AsyncMock) -async def test_run(workflow_mock: AsyncMock, predictions_batch: PredictionsBatch): - workflow_mock.execute_activity_method.return_value = { - 'data': 'test_data' - } - input_data = { - 'schedule_name': 'test_schedule', - 'model_name': 'test_model', - 'model_id': 'test_model_id', - 'query': 'SELECT * FROM test' - } - - await predictions_batch.run(input_data) - - workflow_mock.execute_activity_method.assert_has_calls([ - call( - Activities.prepare_activity, - { - 'schedule_name': input_data['schedule_name'], - 'model_name': input_data['model_name'], - 'model_id': input_data['model_id'] - } - ) - ]) - - workflow_mock.execute_activity_method.assert_has_calls([ - call( - Activities.load_custom_query, - input_data['query'] - ) - ]) - - workflow_mock.execute_child_workflow.assert_has_calls([ - call( - 'prediction_process', input_data) - ]) diff --git a/tests/laborious/workflows/test_predictions_batch.py b/tests/laborious/workflows/test_predictions_batch.py new file mode 100644 index 0000000..01cbc22 --- /dev/null +++ b/tests/laborious/workflows/test_predictions_batch.py @@ -0,0 +1,68 @@ +from unittest.mock import AsyncMock, call, patch, ANY +from pytest import fixture, mark +from laborious.activities.activities import Activities +from laborious.workflows.predictions_batch import PredictionsBatch + + +@fixture +def predictions_batch() -> PredictionsBatch: + return PredictionsBatch() + + +@mark.asyncio +@patch('laborious.workflows.predictions_batch.workflow', new_callable=AsyncMock) +async def test_run(workflow_mock: AsyncMock, predictions_batch: PredictionsBatch): + workflow_mock.execute_local_activity_method.return_value = { + 'data': 'test_data' + } + input_data = { + 'schedule_name': 'test_schedule', + 'model_name': 'test_model', + 'model_id': 'test_model_id', + 'query': 'SELECT * FROM test', + 'schema': 'test_schema', + 'table_name': 'test_table', + 'opc_output_config': 'test_opc_output_config' + } + + await predictions_batch.run(input_data) + + workflow_mock.execute_local_activity_method.assert_has_calls([ + call( + Activities.prepare_activity, + { + 'schedule_name': input_data['schedule_name'], + 'model_name': input_data['model_name'], + 'model_id': input_data['model_id'], + 'workflow_name': 'predictions_batch' + }, + retry_policy=ANY, + start_to_close_timeout=ANY + ) + ]) + + workflow_mock.execute_local_activity_method.assert_has_calls([ + call( + Activities.load_custom_query, + input_data['query'], + retry_policy=ANY, + start_to_close_timeout=ANY + ) + ]) + prediction_input = { + 'data': {'data': 'test_data'}, + 'schema': input_data['schema'], + 'table_name': input_data['table_name'], + 'model_id': input_data['model_id'], + 'model_name': input_data['model_name'], + 'input_filters': input_data.get('input_filters', {}), + 'mlflow_transform_filters': input_data.get('mlflow_transform_filters', {}), + 'mlflow_predict_filters': input_data.get('mlflow_predict_filters', {}), + 'model_retention': input_data.get('model_retention', 60), + 'path_priority': input_data.get('path_priority', ['STOP', 'CONTINUE', 'REPEAT']) + } + + workflow_mock.execute_child_workflow.assert_has_calls([ + call( + 'prediction_process', prediction_input) + ])