diff --git a/.env b/.env new file mode 100644 index 0000000..af9d68d --- /dev/null +++ b/.env @@ -0,0 +1,4 @@ +# === Simulator Git Repo === +# Use SSH format because the Dockerfile uses SSH to clone +SIMULATOR_GIT_REPO=git@github.com:Aignosi/sientia-dataops-opc_simulator.git +SIMULATOR_GIT_BRANCH=main diff --git a/.github/workflows/quality-gate.yml b/.github/workflows/quality-gate.yml index 84bc294..579b770 100644 --- a/.github/workflows/quality-gate.yml +++ b/.github/workflows/quality-gate.yml @@ -63,7 +63,7 @@ jobs: run: | python -m pip install --upgrade pip pip install -r ${{ steps.prepare-requirements.outputs.PROCESSED_REQUIREMENTS_FILE }} - pip install pytest pytest-cov + pip install pytest pytest-cov pytest-asyncio - name: ⬇️ Setup Node.js 18 uses: actions/setup-node@v4 @@ -105,4 +105,5 @@ jobs: -Dsonar.host.url=$SONAR_HOST_URL \ -Dsonar.token=$SONAR_TOKEN \ -Dsonar.python.version=3.11 \ - -Dsonar.projectVersion=1.0.0 + -Dsonar.projectVersion=1.0.0 \ + -Dsonar.coverage.exclusions=laborious/worker/worker.py diff --git a/.gitignore b/.gitignore index 42893e0..a446ab3 100644 --- a/.gitignore +++ b/.gitignore @@ -12,7 +12,7 @@ docker-compose.override.yml **/deploy/*.yaml scouter/.file_versions/ scouter/pipelines/**/triggers.yaml - +**/postgres_data/** # Ignorar arquivos e diretórios de cache do Python __pycache__/ *.pyc @@ -32,4 +32,8 @@ __pycache__/ *.tmp *.bak *.old -.secret \ No newline at end of file +.secret + +# Ignorar coverage +htmlcov/ +.coverage \ No newline at end of file diff --git a/docker-compose.yml b/docker-compose.yml new file mode 100644 index 0000000..532cee8 --- /dev/null +++ b/docker-compose.yml @@ -0,0 +1,82 @@ +version: '3.8' + +services: + postgres: + image: postgres:15 + container_name: postgres + environment: + POSTGRES_USER: sientia + POSTGRES_PASSWORD: sientia + POSTGRES_DB: sientia + ports: + - "5432:5432" + volumes: + - ./postgres_data:/var/lib/postgresql/data + networks: + - sientia-network + + zookeeper: + image: confluentinc/cp-zookeeper:7.5.1 + container_name: zookeeper + environment: + ZOOKEEPER_CLIENT_PORT: 2181 + ZOOKEEPER_TICK_TIME: 2000 + ports: + - "2181:2181" + networks: + - sientia-network + + kafka: + image: confluentinc/cp-kafka:7.5.1 + container_name: kafka + depends_on: + - zookeeper + ports: + - "9092:9092" + - "29092:29092" + environment: + KAFKA_BROKER_ID: 1 + KAFKA_ZOOKEEPER_CONNECT: zookeeper:2181 + KAFKA_ADVERTISED_LISTENERS: PLAINTEXT://kafka:29092,PLAINTEXT_HOST://localhost:9092 + KAFKA_LISTENER_SECURITY_PROTOCOL_MAP: PLAINTEXT:PLAINTEXT,PLAINTEXT_HOST:PLAINTEXT + KAFKA_INTER_BROKER_LISTENER_NAME: PLAINTEXT + KAFKA_OFFSETS_TOPIC_REPLICATION_FACTOR: 1 + networks: + - sientia-network + + kafka-ui: + image: provectuslabs/kafka-ui:latest + container_name: kafka-ui + ports: + - "8080:8080" + environment: + KAFKA_CLUSTERS_0_NAME: local + KAFKA_CLUSTERS_0_BOOTSTRAPSERVERS: kafka:29092 + networks: + - sientia-network + + simulator: + build: + context: . + dockerfile: simulator/Dockerfile + args: + GIT_REPO: ${SIMULATOR_GIT_REPO} + GIT_BRANCH: ${SIMULATOR_GIT_BRANCH} + container_name: simulator + ports: + - "4840:4840" + depends_on: + - kafka + networks: + - sientia-network + env_file: + - .env + + +networks: + sientia-network: + driver: bridge + +volumes: + postgres_data: + driver: local \ No newline at end of file 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 af7268a..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") - async def prepare_activity(self, schedule_name: str, model_name: str, model_id: str): - await super().prepare_activity(schedule_name, model_name, model_id) + 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 8c6fba8..3adb0e4 100644 --- a/laborious/activities/base.py +++ b/laborious/activities/base.py @@ -1,6 +1,7 @@ +from typing import Any +from logging import Logger from temporalio import activity from sientia_do.notifications.handlers import NotificationHandler -from logging import Logger class BaseActivity: @@ -8,9 +9,18 @@ class BaseActivity: self.logger = logger self.notification_handler = notification_handler - def prepare_activity(self, schedule_name: str, - model_name: str, - model_id: str): - self.notification_handler.base_notification.schedule_name = schedule_name - self.notification_handler.base_notification.model_name = model_name - self.notification_handler.base_notification.model_id = model_id + @activity.defn(name="prepare_activity") + async def prepare_activity(self, input_data: dict[str, Any]): + """ + Prepare the activity for the notification handler. + + Args: + workflow_name (str): The name of the workflow. + schedule_name (str): The name of the schedule. + model_name (str): The name of the model. + model_id (str): The id of the model. + """ + self.notification_handler.base_notification.pipeline_name = input_data['workflow_name'] + self.notification_handler.base_notification.schedule_name = input_data['schedule_name'] + self.notification_handler.base_notification.model_name = input_data['model_name'] + self.notification_handler.base_notification.model_id = input_data['model_id'] diff --git a/laborious/activities/gates.py b/laborious/activities/gates.py index ac07fbf..c4f9dd5 100644 --- a/laborious/activities/gates.py +++ b/laborious/activities/gates.py @@ -5,53 +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 + 'EMPTY_DATA': filter_empty_data, + 'path_confidence': { + 'STOP': -1, + 'CONTINUE': 2, + 'REPEAT': -1 + } } -transform_filter_functions = { - 'response_filter': { - 'API_ERROR': api_error_filter, +mlflow_response_filter_functions = { + 'API_ERROR': api_error_filter, + 'path_confidence': { + 'STOP': -1, + 'CONTINUE': 10, + 'REPEAT': -1 }, - 'content_filter': { - 'NAN_VALUES': nan_values_filter, +} + +mlflow_content_filter_functions = { + 'NAN_VALUES': nan_values_filter, + 'path_confidence': { + '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]: ('stop', -1) if some filter policy is 'stop', ('continue', 2) - if no filter policy is 'stop' and some filter policy is 'continue', - None if no filter is applied. + 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() @@ -63,75 +95,182 @@ class Gates(BaseActivity): attachment_content=trace ) - if 'stop' in filter_output: - return 'stop', -1 - elif 'continue' in filter_output: - return 'continue', 2 + for path_flag in path_priority: + if path_flag in filter_output: + 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_gate") - async def mlflow_gate(self, input_data: dict[str, Any]) -> tuple[str, int]: + @activity.defn(name="mlflow_response_gate") + 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 + being the policy and the second element being the confidence status. + Args: + input_data (dict): The input data. Contains: + filters (dict): The filter configuration to apply. + data (dict[str, Any]): The data to filter. + path_priority (list[str]): The path priority list. + type (str): The type of the gate. + Returns: + 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 + try: + 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'], + block="mlflow_gate", + level=NotificationLevel.WARNING, + attachment_content=data['content']['traceback'] + ) + except Exception as e: + trace = traceback.format_exc() + self.notification_handler.build_and_send_notification( + notification_id=f"MLFLOW_GATE_RESPONSE_FILTER__{fil}", + message=f"Error in filter {fil}:{config}: \n {e}", + block="mlflow_gate", + level=NotificationLevel.ERROR, + attachment_content=trace + ) + + for path_flag in path_priority: + if path_flag in filter_output: + self.logger.debug(f"Mlflow response gate result: {path_flag}") + return path_flag, mlflow_response_filter_functions['path_confidence'][path_flag], \ + ", ".join(comments) + + 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 | None, int, str]: + """ + 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: + filters (dict): The filter configuration to apply. + data (dict[str, Any]): The data to filter. + path_priority (list[str]): The path priority list. + type (str): The type of the gate. + Returns: + 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'] + path_priority = input_data['path_priority'] filter_output = [] - for fil, config in filters.items(): - if transform_filter_functions['response_filter'][fil](data, config): - filter_output.append(config['POLICY']) - self.notification_handler.build_and_send_notification( - notification_id=f"{gate_type.upper()}_GATE_RESPONSE_FILTER__{fil}", - message=data['content']['message'], - block="mlflow_gate", - level=NotificationLevel.WARNING, - attachment_content=data['content']['traceback'] - ) - if 'stop' in filter_output: - return 'stop', -1 - elif 'continue' in filter_output: - return 'continue', 10 - - if gate_type == 'predict': - return None, 0 - - data = DataFrame(data['content']) + self.logger.debug(f"Input data:\n {data}") + self.logger.debug(f"Filters: {filters}") for fil, config in filters.items(): - if transform_filter_functions['content_filter'][fil](data, config): - filter_output.append(config['POLICY']) + if fil not in mlflow_content_filter_functions: + continue + try: + if mlflow_content_filter_functions[fil](data, config): + filter_output.append(config['POLICY']) + self.notification_handler.build_and_send_notification( + notification_id=f"{gate_type.upper()}_GATE_CONTENT_FILTER__{fil}", + message=f"Data not passed the content filter {fil}:{config}", + block="mlflow_gate", + level=NotificationLevel.WARNING, + attachment_content=data.to_string() + ) + except Exception as e: + trace = traceback.format_exc() self.notification_handler.build_and_send_notification( - notification_id=f"{gate_type.upper()}_GATE_CONTENT_FILTER__{fil}", - message=f"Data not passed the content filter {fil}:{config}", + notification_id=f"MLFLOW_GATE_CONTENT_FILTER__{fil}", + message=f"Error in filter {fil}:{config}: \n {e}", block="mlflow_gate", - level=NotificationLevel.WARNING, - attachment_content=data.to_string() + level=NotificationLevel.ERROR, + attachment_content=trace ) - if 'stop' in filter_output: - return 'stop', -1 - elif 'continue' in filter_output: - return 'continue', 18 + for path_flag in path_priority: + if path_flag in filter_output: + 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]) -> str: - data = DataFrame(input_data['data']) + async def format_prediction(self, input_data: dict[str, Any]) -> dict[str, Any]: + """ + Formats the prediction data. + Args: + input_data (dict): The input data. Contains: + data (dict[str, Any]): The data to format. + timestamp (str): The timestamp of the data. + model_id (str): The id of the model. + prediction_confidence (float): The confidence of the prediction. + 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() @activity.defn(name="format_default_prediction") - async def format_default_prediction(self, input_data: dict[str, Any]) -> str: + async def format_default_prediction(self, input_data: dict[str, Any]) -> dict[str, Any]: + """ + Creates and formats the default prediction data, with zero value in prediction, + and usefull information in the other fields. + + Args: + input_data (dict): The input data. Contains: + timestamp (str): The timestamp of the data. + model_id (str): The id of the model. + prediction_confidence (float): The confidence of the prediction. + comment (str): The comment of the prediction. + Returns: + dict: The formatted data. + """ + + self.logger.debug("Formatting default prediction...") + return DataFrame({ 'prediction': [0], 'response_time': [0], @@ -139,5 +278,20 @@ 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") + async def get_last_timestamp(self, input_data: dict[str, Any]) -> str: + """ + Gets the last timestamp of the data. + Args: + input_data (dict): The input data. Contains: + data (dict[str, Any]): The data to get the last timestamp from. + Returns: + 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 af3408a..ed96192 100644 --- a/laborious/activities/mlflow.py +++ b/laborious/activities/mlflow.py @@ -5,28 +5,37 @@ 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.model_repository import ModelMonitoringRepository - 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 self.mlflow_password = mlflow_password - self.model_monitoring_repository = ModelMonitoringRepository( + self.model_monitoring_repository = MLFlowRepository( f"{mlflow_host}:{mlflow_port}", mlflow_username, mlflow_password ) @activity.defn(name="request_transform") - async def request_transform(self, input_data: dict[str, Any]) -> tuple[dict[str, Any], str]: + async def request_transform(self, input_data: dict[str, Any]) -> dict[str, Any]: + """ + Access MLFlow model to get the transformed data. + Args: + 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 in minutes. + Returns: + dict[str, Any]: The transformed data. + """ self.logger.info('Transforming data...') data = DataFrame(input_data['data']) model_name = input_data['model_name'] @@ -44,12 +53,22 @@ class MLFlow(BaseActivity): response_data = self.model_monitoring_repository.transform( model_name, data, model_retention) - timestamp = max(data['timestamp'].values.tolist()) + self.logger.debug(response_data) - return response_data, timestamp + return response_data @activity.defn(name="request_predict") - async def request_predict(self, input_data: dict[str, Any]) -> tuple[dict[str, Any], str]: + async def request_predict(self, input_data: dict[str, Any]) -> dict[str, Any]: + """ + Access MLFlow model to get the predicted data. + Args: + input_data (dict): The input data. Contains: + data (dict[str, Any]): The data to predict. + model_name (str): The name of the model. + model_retention (int): The retention of the model. + Returns: + dict[str, Any]: The predicted data. + """ self.logger.info('Predicting data...') data = DataFrame(input_data['data']) model_name = input_data['model_name'] @@ -62,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 eae202f..e6a1dcc 100644 --- a/laborious/activities/opc.py +++ b/laborious/activities/opc.py @@ -1,76 +1,101 @@ -import traceback -from pandas import DataFrame 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]): + """ + Write prediction and confidence data to OPC servers. The two writing + operations are optional and independent of each other. + + 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_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) - 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 - 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 - ) - self.logger.error(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' + ) diff --git a/laborious/activities/postgres.py b/laborious/activities/postgres.py index 7ef2309..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,25 +25,26 @@ 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() @activity.defn(name="load_custom_query") - async def load_custom_query(self, query: str) -> dict[str, dict]: + async def load_custom_query(self, query: str) -> dict[str, Any]: """ Loads data from a custom query. @@ -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]): @@ -135,31 +144,38 @@ class Postgres(BaseActivity): Exports data to a postgres table. Args: - input_data (dict[str, Any]): The data to export. + input_data (dict[str, Any]): The data to export. Contains: + schema (str): The schema of the table. + table_name (str): The name of the table. + 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/base_filter.py b/laborious/utils/filters/base_filter.py deleted file mode 100644 index f5d9948..0000000 --- a/laborious/utils/filters/base_filter.py +++ /dev/null @@ -1,13 +0,0 @@ -from pandas import DataFrame - -class Filter: - def __init__(self, id: str): - self.id = id - self.warnings = [] - - @staticmethod - def method(df: DataFrame) -> DataFrame: - raise NotImplementedError - - def warning(self, message: str): - self.warnings.append(f'[{self.id}] - {message}') \ No newline at end of file diff --git a/laborious/utils/filters/conditional_filters.py b/laborious/utils/filters/conditional_filters.py index d819a3d..57cd3fd 100644 --- a/laborious/utils/filters/conditional_filters.py +++ b/laborious/utils/filters/conditional_filters.py @@ -1,14 +1,11 @@ -from typing import List - -from laborious.utils.filters.base_filter import Filter 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/filters/mlflow_filters.py b/laborious/utils/filters/mlflow_filters.py index 3f5db3b..9936018 100644 --- a/laborious/utils/filters/mlflow_filters.py +++ b/laborious/utils/filters/mlflow_filters.py @@ -1,6 +1,5 @@ import numpy as np from pandas import DataFrame -from laborious.utils.filters.base_filter import Filter def api_error_filter(response: dict, _config: dict): @@ -15,7 +14,7 @@ def api_error_filter(response: dict, _config: dict): def nan_values_filter(predictions: DataFrame, _config: dict): data = predictions.replace({None: np.nan}).drop( - columns=['timestamp'], errors='ignore') + columns=['timestamp'], errors='ignore').infer_objects(copy=False) if data.isna().all().all(): return True 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/model_repository.py b/laborious/utils/repository/model_repository.py index 5efe83b..3896411 100644 --- a/laborious/utils/repository/model_repository.py +++ b/laborious/utils/repository/model_repository.py @@ -16,7 +16,7 @@ from sientia.ModelServing import ModelServing from pathlib import Path -class ModelMonitoringRepository(): +class MLFlowRepository(): def __init__(self, host, username, password): self.model_serving = ModelServing(tracking_uri=host, diff --git a/laborious/utils/repository/opc_repository.py b/laborious/utils/repository/opc_repository.py index 99fbc08..e674354 100644 --- a/laborious/utils/repository/opc_repository.py +++ b/laborious/utils/repository/opc_repository.py @@ -2,27 +2,40 @@ from pathlib import Path from asyncua.sync import Client from asyncua.crypto.security_policies import SecurityPolicyBasic256 from asyncua.ua import DataValue, Variant, VariantType -from datetime import datetime -from time import sleep from logging import Logger -from statistics import mean, median -from typing import Callable - +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 @@ -30,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): @@ -73,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. @@ -94,19 +102,106 @@ 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): + if self.client is None: + return self.client.disconnect() self.client = None self.logger.info('Disconnected from OPC server') def __del__(self): - self.disconnect() + try: + self.disconnect() + except Exception as e: + self.logger.error(f"Error in destructor: {e}") - 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 new file mode 100644 index 0000000..bdf168d --- /dev/null +++ b/laborious/worker/worker.py @@ -0,0 +1,96 @@ +from temporalio import workflow, client +from temporalio.worker import Worker + +with workflow.unsafe.imports_passed_through(): + import os + 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 = get_logger(__name__) + + logger.info('Starting Worker...') + + logger.info('Starting Notification Handler...') + + notification_handler = NotificationHandler( + servers=os.getenv('KAFKA_SERVERS', 'http://localhost:9092'), + logger=logger, + project_name=os.getenv('PROJECT_NAME', 'laborious'), + pipeline_name='-', + trigger_name='-', + model_name='-', + model='-' + ) + + logger.info('Starting Activities...') + + activities = Activities( + postgres_config=build_postgres_config(), + mlflow_config=build_mlflow_config(), + opc_config=build_opc_config(), + logger=logger, + notification_handler=notification_handler + ) + + 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-queue', + workflows=[PredictionsBatch, PredictionProcess, + FormatAndExportPrediction], + activities=[ + # Base + activities.prepare_activity, + # MLFlow + activities.request_predict, + activities.request_transform, + # Gates + activities.input_gate, + activities.mlflow_response_gate, + activities.mlflow_content_gate, + activities.format_prediction, + activities.format_default_prediction, + activities.get_last_timestamp, + # OPC + activities.write_opc_data, + # Postgres + activities.load_custom_query, + activities.repeat_last_prediction, + activities.export_data_to_postgres + ] + ) + ] + + handlers = [] + for w in workers: + handlers.append(w.run()) + + logger.info('Workers started successfully') + + await asyncio.gather(*handlers) + +if __name__ == '__main__': + asyncio.run(main()) diff --git a/laborious/workflows/predictions_batch.py b/laborious/workflows/predictions_batch.py index c3c54b0..425f59f 100644 --- a/laborious/workflows/predictions_batch.py +++ b/laborious/workflows/predictions_batch.py @@ -1,85 +1,89 @@ from temporalio import workflow with workflow.unsafe.imports_passed_through(): - from laborious.activities.postgres import Postgres - from laborious.activities.mlflow import MLFlow - from laborious.activities.gates import Gates - from laborious.activities.opc import OPC + 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( - Postgres.prepare_activity, + 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'] - } + '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( - Postgres.load_custom_query, - input_data['query'] + data = await workflow.execute_local_activity_method( + Activities.load_custom_query, + input_data['query'], + retry_policy=retry_policy, + start_to_close_timeout=timedelta(seconds=60) ) - path_flag, confidence = await workflow.execute_activity_method( - Gates.input_gate, - { - 'filters': input_data['filters'], - 'data': data - } - ) - - if path_flag == 'stop': - return - - if path_flag == 'continue': - # repeat last prediction - await workflow.execute_activity_method( - Postgres.repeat_last_prediction, - { - 'schema': input_data['schema'], - 'table_name': input_data['table_name'], - 'model': input_data['model'] + # 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' } - ) - return - - response_data, last_timestamp = await workflow.execute_activity_method( - MLFlow.transform_data, - { - 'data': data, - 'model_name': input_data['model_name'], - 'model_retention': input_data['model_retention'] - } - ) - - path_flag, confidence = await workflow.execute_activity_method( - Gates.mlflow_gate, - { - 'filters': input_data['filters'], - 'data': response_data, - 'type': 'transform' - } - ) - - if path_flag == 'stop': - return - - if path_flag is None: - # procced with prediction - response_data = await workflow.execute_activity_method( - MLFlow.request_predict, - { - 'data': response_data, - 'model_name': input_data['model_name'], - 'model_retention': input_data['model_retention'] + }), + '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', {}) + } - path_flag + await workflow.execute_child_workflow( + '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 8591f8a..67cad49 100644 --- a/laborious/workflows/sub_workflows/format_and_export_prediction.py +++ b/laborious/workflows/sub_workflows/format_and_export_prediction.py @@ -3,38 +3,71 @@ 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") class FormatAndExportPrediction(): @workflow.run 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 + 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(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. + """ path_flag = input_data['path_flag'] data = input_data['data'] - confidence = input_data['confidence'] + 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': confidence, - } + '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': confidence, + 'prediction_confidence': prediction_confidence, 'comment': input_data['comment'] - } + }, + retry_policy=retry_policy, + start_to_close_timeout=timedelta(seconds=60) ) # write to postgres @@ -44,16 +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 5b40362..de84328 100644 --- a/laborious/workflows/sub_workflows/prediction_process.py +++ b/laborious/workflows/sub_workflows/prediction_process.py @@ -1,59 +1,228 @@ from temporalio import workflow with workflow.unsafe.imports_passed_through(): - from laborious.activities.postgres import Postgres - from laborious.activities.mlflow import MLFlow - from laborious.activities.gates import Gates - from laborious.activities.opc import OPC + 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]): - data = input_data['data'] + """ + This workflow runs a prediction process based on the input data. - path_flag, _confidence = await workflow.execute_activity_method( - Gates.input_gate, + 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'] + model_id = input_data['model_id'] + model_name = input_data['model_name'] + model_retention = input_data['model_retention'] + + last_timestamp = await workflow.execute_local_activity_method( + Activities.get_last_timestamp, { - 'filters': input_data['filters'], 'data': data - } + }, + retry_policy=retry_policy, + start_to_close_timeout=timedelta(minutes=1), ) - if path_flag == 'stop': + path_flag, confidence, comment = await workflow.execute_local_activity_method( + Activities.input_gate, + { + '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, input_data, confidence, last_timestamp, comment + ): return - if path_flag == 'continue': - # repeat last prediction - await workflow.execute_activity_method( - Postgres.repeat_last_prediction, - { - 'schema': input_data['schema'], - 'table_name': input_data['table_name'], - 'model': input_data['model'] - } - ) - return - - response_data, last_timestamp = await workflow.execute_activity_method( - MLFlow.transform_data, + response_data = await workflow.execute_local_activity_method( + Activities.request_transform, { 'data': data, - 'model_name': input_data['model_name'], - 'model_retention': input_data['model_retention'] - } + '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( - Gates.mlflow_gate, + path_flag, confidence, comment = await workflow.execute_local_activity_method( + Activities.mlflow_response_gate, { - 'filters': input_data['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, input_data, confidence, last_timestamp, comment + ): + return + + path_flag, confidence, comment = await workflow.execute_local_activity_method( + Activities.mlflow_content_gate, + { + 'filters': input_data['mlflow_transform_filters'], + 'data': response_data, + '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, input_data, confidence, last_timestamp, comment + ): + return + + 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, comment = await workflow.execute_local_activity_method( + Activities.mlflow_response_gate, + { + 'filters': input_data['mlflow_predict_filters'], + 'data': response_data, + '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, input_data, confidence, last_timestamp, comment + ): + return + + await workflow.execute_child_workflow( + 'format_and_export_prediction', + { + 'path_flag': path_flag, + 'data': response_data['content'], + 'prediction_confidence': confidence, + 'timestamp': response_data['timestamp'], + 'model_id': model_id, + 'model_name': model_name, + 'model_retention': model_retention, + 'opc_output_config': input_data['opc_output_config'] } ) - if path_flag == 'stop': - return + async def path_flag_handler(self, data: dict[str, Any], path_flag: 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. + 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_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 (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. + """ + + 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': + # repeat last prediction + await workflow.execute_activity_method( + Activities.repeat_last_prediction, + { + 'schema': schema, + 'table_name': table_name, + 'model_id': model_id + }, + retry_policy=retry_policy, + start_to_close_timeout=timedelta(minutes=1), + ) + return True + + elif path_flag == 'CONTINUE': + # call write 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_id, + 'model_name': model_name, + 'model_retention': model_retention, + 'schema': schema, + 'table_name': table_name, + 'comment': comment, + 'opc_output_config': input_data['opc_output_config'] + } + ) + return True + + return False diff --git a/requirements.txt b/requirements.txt index a03f80a..90e049b 100644 --- a/requirements.txt +++ b/requirements.txt @@ -1,24 +1,7 @@ temporalio psycopg2-binary +sqlalchemy asyncua -mlflow==2.10.1 -scikit-learn==1.4.2 -scipy==1.13.0 -shap==0.46.0 -catboost==1.2.5 -hyperopt==0.2.7 -kaleido==0.2.1 -xgboost==2.0.2 -cloudpickle==3.0.0 -pathspec -kafka-python -sshtunnel +redis git+ssh://git@github.com/Aignosi/sientia-dataops-library.git git+ssh://git@github.com/Aignosi/sientia-mlops-library.git -async-timeout -kubernetes -boto3 -scikit-optimize==0.10.2 -pandas==2.2.2 -PyYAML==6.0.1 -numpy==1.26.4 \ No newline at end of file diff --git a/simulator/Dockerfile b/simulator/Dockerfile new file mode 100644 index 0000000..d467676 --- /dev/null +++ b/simulator/Dockerfile @@ -0,0 +1,30 @@ +# syntax=docker/dockerfile:1.4 + +FROM python:3.11-slim + +# Enable use of SSH agent/socket +# This line enables SSH during build +# (don't forget the syntax header above) +RUN apt-get update && apt-get install -y git openssh-client && rm -rf /var/lib/apt/lists/* + +# Use build-time SSH mount for Git clone +# The SSH key will NOT remain in the image +# IMPORTANT: this block requires BuildKit +# and the --ssh flag during docker build + +# SSH config to skip host key check (safe in CI/local dev) +RUN mkdir -p /root/.ssh && echo "StrictHostKeyChecking no" > /root/.ssh/config + +WORKDIR /app + +# Clone using SSH +ARG GIT_REPO +ARG GIT_BRANCH=main + +# Mount SSH key just for this RUN +RUN --mount=type=ssh git clone --branch ${GIT_BRANCH} ${GIT_REPO} . + +# Install requirements if exists +RUN if [ -f requirements.txt ]; then pip install --no-cache-dir -r requirements.txt; fi + +CMD ["python", "server.py"] diff --git a/simulator/redis-feeder.py b/simulator/redis-feeder.py new file mode 100644 index 0000000..3f7e36a --- /dev/null +++ b/simulator/redis-feeder.py @@ -0,0 +1,55 @@ +import redis +import json +import os + +# Redis connection settings +redis_host = "localhost" +redis_port = 6379 + +# Connect to Redis +r = redis.Redis(host=redis_host, port=redis_port, + decode_responses=True, username='default', password='bdnZOpcyiL') + +# Define the key pattern to target +pattern = "slot:opc_tags:*" + +# Step 1: Find and delete matching keys +print("🔍 Searching for keys matching:", pattern) +for key in r.scan_iter(match=pattern): + r.delete(key) + print(f"❌ Deleted: {key}") + +# Step 2: Insert new data +# Example new OPC tag data +new_data = { + "slot:opc_tags:1": { + "server1": { + "name": "server1", + "url": "opc.tcp://sientia-opc-simulator-service.sientia-opc.svc.cluster.local:4840", + "server_uri": "http://opcua-server.simulator", + "tags": { + 'ns=2;i=2': { + 'tag_name': 'Counter', + 'frequency': 1000, + 'topics': ['opcua', 'counter'], + }, + 'ns=2;i=3': { + 'tag_name': 'Rollout', + 'frequency': 1000, + "topics": ['opcua', 'rollout'], + }, + 'ns=2;i=4': { + 'tag_name': 'Square', + 'frequency': 1000, + "topics": ['opcua'], + }, + } + } + } +} + +for key, val in new_data.items(): + r.set(key, json.dumps(val)) + print(f"✅ Set: {key} -> {val}") + +print("🚀 OPC tag keys replaced successfully.") diff --git a/tests/laborious/activities/test_activities.py b/tests/laborious/activities/test_activities.py new file mode 100644 index 0000000..98be011 --- /dev/null +++ b/tests/laborious/activities/test_activities.py @@ -0,0 +1,149 @@ +from pytest import mark +from unittest.mock import patch, MagicMock, ANY +from laborious.activities.activities import Activities +from laborious.activities.postgres import Postgres +from laborious.activities.mlflow import MLFlow +from laborious.activities.gates import Gates +from laborious.activities.opc import OPC + + +@patch('laborious.activities.activities.Postgres.__init__') +@patch('laborious.activities.activities.MLFlow.__init__') +@patch('laborious.activities.activities.OPC.__init__') +@patch('laborious.activities.activities.Gates.__init__') +def test___init__(mock_gates_init, mock_opc_init, mock_mlflow_init, mock_postgres_init): + + postgres_config = { + 'host': 'localhost', + 'port': 5432, + 'user': 'postgres', + 'password': 'postgres', + 'dbname': 'postgres', + 'min_connections': 1, + 'max_connections': 10 + } + + mlflow_config = { + 'host': 'localhost', + 'port': 5000, + 'username': 'mlflow', + 'password': 'mlflow' + } + + opc_config = { + 'bootstrap_servers': 'localhost:9092', + 'polling_time': 1000, + 'group_id': 'test-group' + } + + logger = MagicMock() + notification_handler = MagicMock() + + activities = Activities( + postgres_config=postgres_config, + mlflow_config=mlflow_config, + opc_config=opc_config, + logger=logger, + notification_handler=notification_handler + ) + + assert isinstance(activities, Activities) + assert isinstance(activities, Postgres) + assert isinstance(activities, MLFlow) + assert isinstance(activities, OPC) + assert isinstance(activities, Gates) + + mock_postgres_init.assert_called_once_with( + ANY, + host=postgres_config['host'], + port=postgres_config['port'], + user=postgres_config['user'], + password=postgres_config['password'], + dbname=postgres_config['dbname'], + min_connections=postgres_config['min_connections'], + max_connections=postgres_config['max_connections'], + logger=logger, + notification_handler=notification_handler + ) + + mock_mlflow_init.assert_called_once_with( + ANY, + mlflow_host=mlflow_config['host'], + mlflow_port=mlflow_config['port'], + mlflow_username=mlflow_config['username'], + mlflow_password=mlflow_config['password'], + logger=logger, + notification_handler=notification_handler + ) + + mock_opc_init.assert_called_once_with( + ANY, + opc_servers=opc_config, + logger=logger, + notification_handler=notification_handler + ) + + mock_gates_init.assert_called_once_with( + ANY, + logger=logger, + notification_handler=notification_handler + ) + + +@mark.asyncio +@patch('laborious.activities.activities.Postgres.__init__') +@patch('laborious.activities.activities.MLFlow.__init__') +@patch('laborious.activities.activities.OPC.__init__') +async def test_prepare_activity(_mock_opc_init, + _mock_mlflow_init, _mock_postgres_init): + postgres_config = { + 'host': 'localhost', + 'port': 5432, + 'user': 'postgres', + 'password': 'postgres', + 'dbname': 'postgres', + 'min_connections': 1, + 'max_connections': 10 + } + + mlflow_config = { + 'host': 'localhost', + 'port': 5000, + 'username': 'mlflow', + 'password': 'mlflow' + } + + opc_config = { + 'bootstrap_servers': 'localhost:9092', + 'polling_time': 1000, + 'group_id': 'test-group' + } + + logger = MagicMock() + notification_handler = MagicMock() + + activities = Activities( + postgres_config=postgres_config, + mlflow_config=mlflow_config, + opc_config=opc_config, + logger=logger, + notification_handler=notification_handler + ) + + input_data = { + 'workflow_name': 'test-workflow-name', + 'schedule_name': 'test-schedule-name', + 'model_name': 'test-model-name', + 'model_id': 'test-model-id' + } + + await activities.prepare_activity(input_data) + + assert activities.notification_handler.base_notification.pipeline_name == input_data[ + 'workflow_name'] + assert activities.notification_handler.base_notification.schedule_name == input_data[ + 'schedule_name'] + assert activities.notification_handler.base_notification.model_name == input_data[ + 'model_name'] + assert activities.notification_handler.base_notification.model_id == input_data[ + 'model_id'] diff --git a/tests/laborious/activities/test_base.py b/tests/laborious/activities/test_base.py new file mode 100644 index 0000000..6978acb --- /dev/null +++ b/tests/laborious/activities/test_base.py @@ -0,0 +1,35 @@ +from unittest.mock import MagicMock +from laborious.activities.base import BaseActivity +from pytest import fixture, mark +from sientia_do.notifications.models import Notification + + +@fixture +def base_activity(): + return BaseActivity( + logger=MagicMock(), + notification_handler=MagicMock(), + ) + + +@mark.asyncio +async def test_prepare_activity(base_activity): + base_activity.notification_handler.base_notification = Notification( + project="project", + pipeline="pipeline", + trigger="-", + model_name="-", + 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 8edca83..6b61c81 100644 --- a/tests/laborious/activities/test_gates.py +++ b/tests/laborious/activities/test_gates.py @@ -1,308 +1,369 @@ -from unittest.mock import ANY, MagicMock, patch -from pandas import DataFrame +from unittest.mock import MagicMock, ANY, patch from pytest import fixture, mark - -from laborious.activities.gates import Gates from sientia_do.notifications.models import NotificationLevel +from laborious.activities.gates import Gates @fixture -def gates(): +def gates_activity(): return Gates( logger=MagicMock(), - notification_handler=MagicMock() + notification_handler=MagicMock(), ) @mark.asyncio -@patch('laborious.activities.gates.filter_functions') -async def test_input_gate_specific_variables_null_values_with_stop_policy_only( - filter_functions_mock, - gates -): - specific_variables_null_values_mock = MagicMock(return_value=True) - empty_data_mock = MagicMock(return_value=False) - - def functions_side_effect(x): - if x == 'SPECIFIC_VARIABLES_NULL_VALUES': - return specific_variables_null_values_mock - return empty_data_mock - - filter_functions_mock.__getitem__.side_effect = functions_side_effect - +async def test_input_gate_invalid_filter(gates_activity): + # Arrange input_data = { 'filters': { - 'SPECIFIC_VARIABLES_NULL_VALUES': { - 'POLICY': 'stop', - 'VARIABLES': ['variable2'] - } + 'INVALID_FILTER': {'POLICY': 'STOP'} }, - 'data': { - 'variable': ['variable1', 'variable2'], - 'value': [1, 2] - } + 'data': {'value': [1, 2, 3]}, + 'path_priority': ['STOP', 'CONTINUE', 'REPEAT'] } - result = await gates.input_gate(input_data) - assert result == ('stop', -1) + # Act + result = await gates_activity.input_gate(input_data) - input_args = specific_variables_null_values_mock.call_args - assert input_args[0][0].equals(DataFrame( - {'variable': ['variable1', 'variable2'], 'value': [1, 2]})) - assert input_args[0][1] == input_data['filters']['SPECIFIC_VARIABLES_NULL_VALUES'] - - empty_data_mock.assert_not_called() + # Assert + assert result == (None, 0, "") + gates_activity.logger.error.assert_called_once_with( + "Filter INVALID_FILTER not found" + ) @mark.asyncio -@patch('laborious.activities.gates.filter_functions') -async def test_input_gate_specific_variables_null_values_with_continue_policy_only( - filter_functions_mock, - gates -): - specific_variables_null_values_mock = MagicMock(return_value=True) - empty_data_mock = MagicMock(return_value=False) - - def functions_side_effect(x): - if x == 'SPECIFIC_VARIABLES_NULL_VALUES': - return specific_variables_null_values_mock - return empty_data_mock - - filter_functions_mock.__getitem__.side_effect = functions_side_effect - +@patch('laborious.activities.gates.input_filter_functions') +async def test_input_gate_filter_exception(mock_input_filter_functions, gates_activity): + # Arrange + mock_input_filter_functions.__contains__.return_value = True + mock_input_filter_functions.__getitem__.return_value = MagicMock( + side_effect=Exception("Test error")) input_data = { 'filters': { - 'SPECIFIC_VARIABLES_NULL_VALUES': { - 'POLICY': 'continue', - 'VARIABLES': ['variable2'] - } + 'EMPTY_DATA': {'POLICY': 'STOP'} }, - 'data': { - 'variable': ['variable1', 'variable2'], - 'value': [1, 2] - } + 'data': {'value': []}, + 'path_priority': ['STOP', 'CONTINUE', 'REPEAT'] } - result = await gates.input_gate(input_data) - assert result == ('continue', 2) + # Act + result = await gates_activity.input_gate(input_data) - input_args = specific_variables_null_values_mock.call_args - assert input_args[0][0].equals(DataFrame( - {'variable': ['variable1', 'variable2'], 'value': [1, 2]})) - assert input_args[0][1] == input_data['filters']['SPECIFIC_VARIABLES_NULL_VALUES'] - - empty_data_mock.assert_not_called() - - -@mark.asyncio -@patch('laborious.activities.gates.filter_functions') -async def test_input_gate_specific_variables_null_values_no_filtered( - filter_functions_mock, - gates -): - specific_variables_null_values_mock = MagicMock(return_value=False) - empty_data_mock = MagicMock(return_value=False) - - def functions_side_effect(x): - if x == 'SPECIFIC_VARIABLES_NULL_VALUES': - return specific_variables_null_values_mock - return empty_data_mock - - filter_functions_mock.__getitem__.side_effect = functions_side_effect - - input_data = { - 'filters': { - 'SPECIFIC_VARIABLES_NULL_VALUES': { - 'POLICY': 'stop', - 'VARIABLES': ['variable2'] - } - }, - 'data': { - 'variable': ['variable1', 'variable2'], - 'value': [1, 2] - } - } - - result = await gates.input_gate(input_data) - 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( - {'variable': ['variable1', 'variable2'], 'value': [1, 2]})) - assert specific_variables_null_values_input_args[0][ - 1] == input_data['filters']['SPECIFIC_VARIABLES_NULL_VALUES'] - - empty_data_mock.assert_not_called() - - -@mark.asyncio -@patch('laborious.activities.gates.filter_functions') -async def test_input_gate_one_stop_policy( - filter_functions_mock, - gates -): - specific_variables_null_values_mock = MagicMock(return_value=True) - empty_data_mock = MagicMock(return_value=True) - - def functions_side_effect(x): - if x == 'SPECIFIC_VARIABLES_NULL_VALUES': - return specific_variables_null_values_mock - return empty_data_mock - - filter_functions_mock.__getitem__.side_effect = functions_side_effect - - input_data = { - 'filters': { - 'SPECIFIC_VARIABLES_NULL_VALUES': { - 'POLICY': 'stop', - 'VARIABLES': ['variable2'] - }, - 'EMPTY_DATA': { - 'POLICY': 'continue', - } - }, - 'data': { - 'variable': ['variable1', 'variable2'], - 'value': [1, 2] - } - } - - result = await gates.input_gate(input_data) - assert result == ('stop', -1) - - specific_variables_null_values_input_args = specific_variables_null_values_mock.call_args - assert specific_variables_null_values_input_args[0][0].equals(DataFrame( - {'variable': ['variable1', 'variable2'], 'value': [1, 2]})) - assert specific_variables_null_values_input_args[0][ - 1] == input_data['filters']['SPECIFIC_VARIABLES_NULL_VALUES'] - - empty_data_input_args = empty_data_mock.call_args - assert empty_data_input_args[0][0].equals(DataFrame( - {'variable': ['variable1', 'variable2'], 'value': [1, 2]})) - assert empty_data_input_args[0][1] == input_data['filters']['EMPTY_DATA'] - - -@mark.asyncio -@patch('laborious.activities.gates.filter_functions') -async def test_input_gate_one_continue_policy( - filter_functions_mock, - gates -): - specific_variables_null_values_mock = MagicMock(return_value=False) - empty_data_mock = MagicMock(return_value=True) - - def functions_side_effect(x): - if x == 'SPECIFIC_VARIABLES_NULL_VALUES': - return specific_variables_null_values_mock - return empty_data_mock - - filter_functions_mock.__getitem__.side_effect = functions_side_effect - - input_data = { - 'filters': { - 'SPECIFIC_VARIABLES_NULL_VALUES': { - 'POLICY': 'stop', - 'VARIABLES': ['variable2'] - }, - 'EMPTY_DATA': { - 'POLICY': 'continue', - } - }, - 'data': { - 'variable': ['variable1', 'variable2'], - 'value': [1, 2] - } - } - - result = await gates.input_gate(input_data) - assert result == ('continue', 2) - - specific_variables_null_values_input_args = specific_variables_null_values_mock.call_args - assert specific_variables_null_values_input_args[0][0].equals(DataFrame( - {'variable': ['variable1', 'variable2'], 'value': [1, 2]})) - assert specific_variables_null_values_input_args[0][ - 1] == input_data['filters']['SPECIFIC_VARIABLES_NULL_VALUES'] - - empty_data_input_args = empty_data_mock.call_args - assert empty_data_input_args[0][0].equals(DataFrame( - {'variable': ['variable1', 'variable2'], 'value': [1, 2]})) - assert empty_data_input_args[0][1] == input_data['filters']['EMPTY_DATA'] - - -@mark.asyncio -@patch('laborious.activities.gates.filter_functions') -async def test_input_gate_no_filtered( - filter_functions_mock, - gates -): - specific_variables_null_values_mock = MagicMock(return_value=False) - empty_data_mock = MagicMock(return_value=False) - - def functions_side_effect(x): - if x == 'SPECIFIC_VARIABLES_NULL_VALUES': - return specific_variables_null_values_mock - return empty_data_mock - - filter_functions_mock.__getitem__.side_effect = functions_side_effect - - input_data = { - 'filters': { - 'SPECIFIC_VARIABLES_NULL_VALUES': { - 'POLICY': 'stop', - 'VARIABLES': ['variable2'] - }, - 'EMPTY_DATA': { - 'POLICY': 'continue', - } - }, - 'data': { - 'variable': ['variable1', 'variable2'], - 'value': [1, 2] - } - } - - result = await gates.input_gate(input_data) - 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( - {'variable': ['variable1', 'variable2'], 'value': [1, 2]})) - - empty_data_input_args = empty_data_mock.call_args - assert empty_data_input_args[0][0].equals(DataFrame( - {'variable': ['variable1', 'variable2'], 'value': [1, 2]})) - assert empty_data_input_args[0][1] == input_data['filters']['EMPTY_DATA'] - - -@mark.asyncio -@patch('laborious.activities.gates.filter_functions') -async def test_input_gate_error( - filter_functions_mock, - gates -): - filter_functions_mock.__getitem__.side_effect = KeyError('test') - - input_data = { - 'filters': { - 'SPECIFIC_VARIABLES_NULL_VALUES': { - 'POLICY': 'stop', - 'VARIABLES': ['variable2'] - } - }, - 'data': { - 'variable': ['variable1', 'variable2'], - 'value': [1, 2] - } - } - - result = await gates.input_gate(input_data) - assert result == (None, 0) - - gates.notification_handler.build_and_send_notification.assert_called_once_with( - notification_id='INTPUT_GATE_ERROR__SPECIFIC_VARIABLES_NULL_VALUES', - message="Error in filter SPECIFIC_VARIABLES_NULL_VALUES:{'POLICY': 'stop', 'VARIABLES': ['variable2']}: \n 'test'", - block='input_gate', + # Assert + assert result == (None, 0, "") + gates_activity.notification_handler.build_and_send_notification.assert_called_once_with( + notification_id="INTPUT_GATE_ERROR__EMPTY_DATA", + message="Error in filter EMPTY_DATA:{'POLICY': 'STOP'}: \n Test error", + block="input_gate", level=NotificationLevel.ERROR, attachment_content=ANY ) + + +@mark.asyncio +async def test_input_gate_no_filters(gates_activity): + # Arrange + input_data = { + 'filters': {}, + 'data': {'value': [1, 2, 3]}, + 'path_priority': ['CONTINUE', 'STOP', 'REPEAT'] + } + + # Act + result = await gates_activity.input_gate(input_data) + + # Assert + assert result == (None, 0, "") + gates_activity.logger.debug.assert_called() + + +@mark.asyncio +async def test_input_gate_with_filter(gates_activity): + # Arrange + input_data = { + 'filters': { + 'EMPTY_DATA': {'POLICY': 'STOP'} + }, + 'data': {'value': []}, + 'path_priority': ['STOP', 'CONTINUE', 'REPEAT'] + } + + # Act + result = await gates_activity.input_gate(input_data) + + # Assert + assert result == ('STOP', -1, "Input data with bad quality") + gates_activity.logger.debug.assert_called() + + +@mark.asyncio +async def test_mlflow_response_gate_invalid_filter(gates_activity): + # Arrange + input_data = { + 'filters': { + 'INVALID_FILTER': {'POLICY': 'STOP'} + }, + 'data': {'content': {'message': 'success'}}, + 'type': 'test', + 'path_priority': ['STOP', 'CONTINUE', 'REPEAT'] + } + + # Act + result = await gates_activity.mlflow_response_gate(input_data) + + # Assert + assert result == (None, 0, "") + + +@mark.asyncio +@patch('laborious.activities.gates.mlflow_response_filter_functions') +async def test_mlflow_response_gate_filter_exception(mock_mlflow_response_filter_functions, + gates_activity): + # Arrange + mock_mlflow_response_filter_functions.__contains__.return_value = True + mock_mlflow_response_filter_functions.__getitem__.return_value = MagicMock( + side_effect=Exception("Test error")) + input_data = { + 'filters': { + 'INVALID_FILTER': {'POLICY': 'STOP'} + }, + 'data': {'content': {'message': 'success'}}, + 'type': 'test', + 'path_priority': ['STOP', 'CONTINUE', 'REPEAT'] + } + + # Act + result = await gates_activity.mlflow_response_gate(input_data) + + # Assert + assert result == (None, 0, "") + gates_activity.notification_handler.build_and_send_notification.assert_called_once_with( + notification_id="MLFLOW_GATE_RESPONSE_FILTER__INVALID_FILTER", + message="Error in filter INVALID_FILTER:{'POLICY': 'STOP'}: \n Test error", + block="mlflow_gate", + level=NotificationLevel.ERROR, + attachment_content=ANY + ) + + +@mark.asyncio +async def test_mlflow_response_gate_no_filters(gates_activity): + # Arrange + input_data = { + 'filters': {}, + 'data': {'content': {'message': 'success'}}, + 'type': 'test', + 'path_priority': ['CONTINUE', 'STOP', 'REPEAT'] + } + + # Act + result = await gates_activity.mlflow_response_gate(input_data) + + # Assert + assert result == (None, 0, "") + gates_activity.logger.debug.assert_called() + + +@mark.asyncio +async def test_mlflow_response_gate_with_filter(gates_activity): + # Arrange + input_data = { + 'filters': { + 'API_ERROR': {'POLICY': 'STOP'} + }, + 'data': { + 'success': False, + 'content': { + 'message': 'API error occurred', + 'traceback': 'error trace' + } + }, + 'type': 'test', + 'path_priority': ['STOP', 'CONTINUE', 'REPEAT'] + } + + # Act + result = await gates_activity.mlflow_response_gate(input_data) + + # Assert + assert result == ('STOP', -1, "API error occurred") + gates_activity.logger.debug.assert_called() + gates_activity.notification_handler.build_and_send_notification.assert_called() + + +@mark.asyncio +async def test_mlflow_content_gate_invalid_filter(gates_activity): + # Arrange + input_data = { + 'filters': { + 'INVALID_FILTER': {'POLICY': 'STOP'} + }, + 'data': {'value': [1, 2, 3]}, + 'type': 'test', + 'path_priority': ['STOP', 'CONTINUE', 'REPEAT'] + } + + # Act + result = await gates_activity.mlflow_content_gate(input_data) + + # Assert + assert result == (None, 0, "") + + +@mark.asyncio +@patch('laborious.activities.gates.mlflow_content_filter_functions') +async def test_mlflow_content_gate_filter_exception(mock_mlflow_content_filter_functions, + gates_activity): + # Arrange + mock_mlflow_content_filter_functions.__contains__.return_value = True + mock_mlflow_content_filter_functions.__getitem__.return_value = MagicMock( + side_effect=Exception("Test error")) + input_data = { + 'filters': { + 'API_ERROR': {'POLICY': 'STOP'} + }, + 'data': { + 'success': False, + 'content': { + 'message': 'API error occurred', + 'traceback': 'error trace' + } + }, + 'type': 'test', + 'path_priority': ['STOP', 'CONTINUE', 'REPEAT'] + } + + # Act + result = await gates_activity.mlflow_content_gate(input_data) + + # Assert + assert result == (None, 0, "") + gates_activity.logger.debug.assert_called() + gates_activity.notification_handler.build_and_send_notification.assert_called_once_with( + notification_id="MLFLOW_GATE_CONTENT_FILTER__API_ERROR", + message="Error in filter API_ERROR:{'POLICY': 'STOP'}: \n Test error", + block="mlflow_gate", + level=NotificationLevel.ERROR, + attachment_content=ANY + ) + + +@mark.asyncio +async def test_mlflow_content_gate_no_filters(gates_activity): + # Arrange + input_data = { + 'filters': {}, + 'data': {'value': [1, 2, 3]}, + 'type': 'test', + 'path_priority': ['CONTINUE', 'STOP', 'REPEAT'] + } + + # Act + result = await gates_activity.mlflow_content_gate(input_data) + + # Assert + assert result == (None, 0, "") + gates_activity.logger.debug.assert_called() + + +@mark.asyncio +async def test_mlflow_content_gate_with_filter(gates_activity): + # Arrange + input_data = { + 'filters': { + 'NAN_VALUES': {'POLICY': 'STOP'} + }, + 'data': {'value': [None, None, None]}, + 'type': 'test', + 'path_priority': ['STOP', 'CONTINUE', 'REPEAT'] + } + + # Act + result = await gates_activity.mlflow_content_gate(input_data) + + # Assert + assert result == ( + 'STOP', -1, "Transformed data not passed the content filter") + gates_activity.logger.debug.assert_called() + gates_activity.notification_handler.build_and_send_notification.assert_called() + + +@mark.asyncio +async def test_format_prediction(gates_activity): + # Arrange + input_data = { + 'data': {'prediction': [1], 'response_time': [0.1]}, + 'timestamp': '2023-05-26 11:12:27', + 'model_id': 'test_model', + 'prediction_confidence': 0.9 + } + + # Act + result = await gates_activity.format_prediction(input_data) + + # Assert + assert result['prediction'] == {0: 1} + assert result['response_time'] == {0: ANY} + assert result['timestamp'] == {0: '2023-05-26 11:12:27'} + assert result['model_id'] == {0: 'test_model'} + assert result['prediction_confidence'] == {0: 0.9} + assert result['prediction_status'] == {0: 'Good'} + assert result['comments'] == {0: ""} + gates_activity.logger.debug.assert_called() + + +@mark.asyncio +async def test_format_default_prediction(gates_activity): + # Arrange + input_data = { + 'timestamp': '2023-05-26 11:12:27', + 'model_id': 'test_model', + 'prediction_confidence': 0.1, + 'comment': 'Test comment' + } + + # Act + result = await gates_activity.format_default_prediction(input_data) + + # Assert + assert result['prediction'] == {0: 0} + assert result['response_time'] == {0: 0} + assert result['timestamp'] == {0: '2023-05-26 11:12:27'} + assert result['model_id'] == {0: 'test_model'} + assert result['prediction_confidence'] == {0: 0.1} + assert result['prediction_status'] == {0: 'Bad'} + assert result['comments'] == {0: 'Test comment'} + gates_activity.logger.debug.assert_called() + + +@mark.asyncio +async def test_get_last_timestamp_with_data(gates_activity): + # Arrange + input_data = { + 'data': { + 'timestamp': ['2023-05-26 11:12:27', '2023-05-26 11:12:28'] + } + } + + # Act + result = await gates_activity.get_last_timestamp(input_data) + + # Assert + assert result == '2023-05-26 11:12:28' + + +@mark.asyncio +async def test_get_last_timestamp_no_data(gates_activity): + # Arrange + input_data = { + 'data': {} + } + + # Act + result = await gates_activity.get_last_timestamp(input_data) + + # Assert + assert isinstance(result, str) # Should be a timestamp string + assert len(result) > 0 diff --git a/tests/laborious/activities/test_mlflow.py b/tests/laborious/activities/test_mlflow.py new file mode 100644 index 0000000..5834cb4 --- /dev/null +++ b/tests/laborious/activities/test_mlflow.py @@ -0,0 +1,120 @@ +from unittest.mock import MagicMock, patch + +import numpy as np +from pytest import fixture, mark +from laborious.activities.mlflow import MLFlow + + +@patch("laborious.activities.mlflow.MLFlowRepository") +def test___init__(mock_mlflow_repository): + mlflow = MLFlow( + mlflow_host="http://localhost", + mlflow_port=5000, + mlflow_username="admin", + mlflow_password="admin", + logger=MagicMock(), + notification_handler=MagicMock() + ) + + assert mlflow.mlflow_host == "http://localhost" + assert mlflow.mlflow_port == 5000 + assert mlflow.mlflow_username == "admin" + assert mlflow.mlflow_password == "admin" + + mock_mlflow_repository.assert_called_once_with( + "http://localhost:5000", "admin", "admin" + ) + + +@fixture +@patch("laborious.activities.mlflow.MLFlowRepository") +def mlflow(mock_mlflow_repository): + return MLFlow( + mlflow_host="http://localhost:5000", + mlflow_port=5000, + mlflow_username="admin", + mlflow_password="admin", + logger=MagicMock(), + notification_handler=MagicMock() + ) + + +@mark.asyncio +@patch("laborious.activities.mlflow.DataFrame") +@patch("laborious.activities.mlflow.max") +async def test_request_transform(mock_max, mock_dataframe, mlflow): + mock_max.return_value = '2024-01-02' + # Mock input data + input_data = { + 'data': [ + {'timestamp': '2024-01-01', 'variable': 'var1', 'value': 1.0}, + {'timestamp': '2024-01-01', 'variable': 'var2', 'value': 2.0}, + {'timestamp': '2024-01-02', 'variable': 'var1', 'value': 3.0}, + {'timestamp': '2024-01-02', 'variable': 'var2', 'value': 4.0} + ], + 'model_name': 'test_model', + 'model_retention': 30 + } + + # Mock the transform response + expected_response = {'prediction': [0.5, 0.6]} + mlflow.model_monitoring_repository.transform.return_value = expected_response + + # Call the method + response_data = await mlflow.request_transform(input_data) + + # Verify the data was correctly transformed + mock_dataframe.assert_called_once_with(input_data['data']) + mock_dataframe.return_value.pivot.assert_called_once_with( + index='timestamp', columns='variable', values='value' + ) + mock_dataframe = mock_dataframe.return_value.pivot.return_value + mock_dataframe.fillna.assert_called_once_with(np.nan, inplace=True) + mock_dataframe.reset_index.assert_called_once() + mock_dataframe.columns.name = None + + # Verify the response + assert response_data == expected_response + + # Verify the repository was called with correct arguments + mlflow.model_monitoring_repository.transform.assert_called_once_with( + 'test_model', mock_dataframe, 30 + ) + + +@mark.asyncio +@patch("laborious.activities.mlflow.DataFrame") +@patch("laborious.activities.mlflow.max") +async def test_request_predict(mock_max, mock_dataframe, mlflow): + mock_max.return_value = '2024-01-02' + # Mock input data + input_data = { + 'data': [ + {'timestamp': '2024-01-01', 'variable': 'var1', 'value': 1.0}, + {'timestamp': '2024-01-01', 'variable': 'var2', 'value': 2.0}, + {'timestamp': '2024-01-02', 'variable': 'var1', 'value': 3.0}, + {'timestamp': '2024-01-02', 'variable': 'var2', 'value': 4.0} + ], + 'model_name': 'test_model', + 'model_retention': 30 + } + + # Mock the predict response + expected_response = {'prediction': [0.5, 0.6]} + mlflow.model_monitoring_repository.predict.return_value = expected_response + + # Call the method + response_data = await mlflow.request_predict(input_data) + + mock_dataframe.assert_called_once_with(input_data['data']) + mock_dataframe.return_value.replace.assert_called_once_with( + np.nan, None, inplace=True + ) + + # Verify the response + assert response_data == expected_response + + # Verify the repository was called with correct arguments + mlflow.model_monitoring_repository.predict.assert_called_once_with( + 'test_model', mock_dataframe.return_value, 30 + ) diff --git a/tests/laborious/activities/test_opc.py b/tests/laborious/activities/test_opc.py new file mode 100644 index 0000000..c4e26ee --- /dev/null +++ b/tests/laborious/activities/test_opc.py @@ -0,0 +1,191 @@ +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 + + +@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( + opc_servers=servers, + logger=mock_logger, + notification_handler=mock_notification_handler + ) + + 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_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, + ) + ]) + + server1.connect.assert_called_once() + server2.connect.assert_called_once() + + +@fixture +@patch("laborious.activities.opc.OpcRepository") +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( + opc_servers=servers, + logger=MagicMock(), + notification_handler=MagicMock() + ) + + +WRITE_DATA_CASES = [ + ('tag1', 'int', 50), + ('tag2', 'float', 50.5), + ('tag3', 'bool', True), + ('tag4', 'string', 'test'), +] + + +@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) + + +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", + level=NotificationLevel.ERROR, + attachment_content=ANY + ) + opc.logger.error.assert_called_once() + + +@mark.asyncio +async def test_write_opc_data_success(opc): + # Arrange + input_data = { + 'data': { + 'prediction': [0.75], + 'prediction_confidence': [0.95] + }, + 'opc_output_config': { + 'server1': { + 'prediction_tags': { + 'tag1': {'data_type': 'float'} + }, + 'confidence_tags': { + 'tag2': {'data_type': 'float'} + } + } + } + } + + # Act + opc.write_data = MagicMock() + await opc.write_opc_data(input_data) + + # Assert + 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 +async def test_write_opc_data_empty_config(opc): + # Arrange + input_data = { + 'data': { + 'prediction': [0.75], + 'prediction_confidence': [0.95] + }, + 'opc_servers': ['server1'], + 'opc_output_config': { + 'prediction_tags': {}, + 'confidence_tags': {} + } + } + + # Act + await opc.write_opc_data(input_data) + + # Assert + opc.opc_repository['server1'].write_data.assert_not_called() diff --git a/tests/laborious/activities/test_postgres.py b/tests/laborious/activities/test_postgres.py index 98e74e5..e4a4545 100644 --- a/tests/laborious/activities/test_postgres.py +++ b/tests/laborious/activities/test_postgres.py @@ -1,132 +1,159 @@ -from unittest.mock import ANY, MagicMock, patch -from pandas import DataFrame -from pytest import fixture -from pytest import mark -from sientia_do.notifications.models import NotificationLevel - +from unittest.mock import MagicMock, patch +from pytest import fixture, mark +import pandas as pd from laborious.activities.postgres import Postgres @fixture -@patch("laborious.activities.postgres.ThreadedConnectionPool") -def postgres_client(mock_pool): +@patch("laborious.activities.postgres.create_engine") +def postgres_activity(_mock_create_engine): return Postgres( host="localhost", port=5432, - user="postgres", - password="postgres", - dbname="postgres", + user="test_user", + password="test_password", + dbname="test_db", min_connections=1, - max_connections=10, + max_connections=5, logger=MagicMock(), - notification_handler=MagicMock(), + notification_handler=MagicMock() ) @mark.asyncio -@patch("laborious.activities.postgres.read_sql_query", - return_value=DataFrame([{"a": 1, "b": 2}])) -async def test_load_custom_query_success(mock_read_sql_query, postgres_client): - query = "SELECT * FROM test" - result = await postgres_client.load_custom_query(query) - assert result is not None - assert len(result) > 0 - assert result == {'a': {0: 1}, 'b': {0: 2}} - postgres_client.notification_handler.build_and_send_notification.assert_not_called() +@patch("laborious.activities.postgres.read_sql_query") +async def test_load_custom_query_none_data(mock_read_sql_query, postgres_activity): + query = "SELECT * FROM test_table LIMIT 1" + mock_read_sql_query.return_value = None + + result = await postgres_activity.load_custom_query(query) + + assert isinstance(result, dict) + assert len(result) == 0 @mark.asyncio -@patch("laborious.activities.postgres.read_sql_query", - side_effect=Exception("Error fetching data from query")) -async def test_load_custom_query_error(mock_read_sql_query, postgres_client): - query = "SELECT * FROM test" - result = await postgres_client.load_custom_query(query) - assert result == {} - postgres_client.notification_handler.build_and_send_notification.assert_called_once_with( - notification_id="ERROR_LOADING_CUSTOM_QUERY", - message="Error fetching data from query: Error fetching data from query", - block="load_custom_query", - level=NotificationLevel.ERROR, - attachment_content=ANY - ) +@patch("laborious.activities.postgres.read_sql_query") +async def test_load_custom_query_date_converted(mock_read_sql_query, postgres_activity): + query = "SELECT * FROM test_table LIMIT 1" + mock_data = pd.DataFrame({"column1": [1], "column2": ["test"]}) + mock_data['date'] = pd.to_datetime('2022-01-01') + + mock_read_sql_query.return_value = mock_data + + result = await postgres_activity.load_custom_query(query) + + assert isinstance(result, dict) + assert len(result) == 3 + assert "column1" in result + assert "column2" in result + assert "date" in result + assert result['date'] == {0: '2022-01-01 00:00:00'} @mark.asyncio -async def test_repeat_last_prediction_success(postgres_client): - query_items = {"schema": "test", "table_name": "test", "model": 1} - await postgres_client.repeat_last_prediction(query_items) - postgres_client.notification_handler.build_and_send_notification.assert_not_called() - postgres_client.pool.getconn.assert_called_once() - postgres_client.pool.putconn.assert_called_once() +@patch("laborious.activities.postgres.read_sql_query") +async def test_load_custom_query_success(mock_read_sql_query, postgres_activity): + query = "SELECT * FROM test_table LIMIT 1" + mock_data = pd.DataFrame({"column1": [1], "column2": ["test"]}) - postgres_client.pool.getconn.return_value.cursor.assert_called_once() - postgres_client.pool.getconn.return_value.cursor.return_value.execute.assert_called_once_with( - f""" - INSERT INTO \"{query_items['schema']}\".{query_items['table_name']} (model_id, prediction, timestamp, response_time, prediction_status, prediction_confidence, created_at) - SELECT model_id, prediction, timestamp, response_time, prediction_status, prediction_confidence, NOW() - FROM \"{query_items['schema']}\".{query_items['table_name']} - WHERE model_id = {query_items['model']} - ORDER BY timestamp DESC - LIMIT 1; - """ - ) - postgres_client.pool.getconn.return_value.commit.assert_called_once() - postgres_client.pool.getconn.return_value.cursor.return_value.close.assert_called_once() + mock_read_sql_query.return_value = mock_data + + result = await postgres_activity.load_custom_query(query) + + assert isinstance(result, dict) + assert len(result) == 2 + assert "column1" in result + assert "column2" in result + postgres_activity.logger.info.assert_called() @mark.asyncio -async def test_repeat_last_prediction_error(postgres_client): - postgres_client.pool.getconn.return_value.cursor.return_value.execute.side_effect = Exception( - "Error repeating last prediction") - query_items = {"schema": "test", "table_name": "test", "model": 1} - await postgres_client.repeat_last_prediction(query_items) - postgres_client.notification_handler.build_and_send_notification.assert_called_once_with( - notification_id="ERROR_REPEATING_LAST_PREDICTION", - message="Error repeating last prediction: Error repeating last prediction", - block="repeat_last_prediction", - level=NotificationLevel.ERROR, - attachment_content=ANY - ) - postgres_client.pool.getconn.assert_called_once() - postgres_client.pool.putconn.assert_called_once() +async def test_load_custom_query_error(postgres_activity): + query = "SELECT * FROM non_existent_table" + error_msg = "Table not found" + + with patch("laborious.activities.postgres.read_sql_query", side_effect=ValueError(error_msg)): + result = await postgres_activity.load_custom_query(query) + + assert isinstance(result, dict) + assert len(result) == 0 + postgres_activity.notification_handler.build_and_send_notification.assert_called_once() + postgres_activity.logger.error.assert_called() @mark.asyncio -@patch("laborious.activities.postgres.DataFrame") -async def test_export_data_to_postgres_success(mock_dataframe, postgres_client): - data = {"schema": "test", "table_name": "test", - "data": {"a": [1, 2, 3], "b": [4, 5, 6]}} - await postgres_client.export_data_to_postgres(data) - postgres_client.notification_handler.build_and_send_notification.assert_not_called() - postgres_client.pool.getconn.assert_called_once() - postgres_client.pool.putconn.assert_called_once() +async def test_repeat_last_prediction_success(postgres_activity): + query_items = { + "schema": "public", + "table_name": "predictions", + "model": 1 + } - mock_dataframe.assert_called_once_with(data["data"]) - mock_dataframe.return_value.to_sql.assert_called_once_with( - data["table_name"], - postgres_client.pool.getconn.return_value, - schema=data["schema"], - if_exists="append", - index=False - ) - postgres_client.pool.getconn.return_value.commit.assert_called_once() + with patch("sqlalchemy.orm.session.Session.execute") as mock_execute: + await postgres_activity.repeat_last_prediction(query_items) + + mock_execute.assert_called_once() + postgres_activity.logger.info.assert_called() @mark.asyncio -@patch("laborious.activities.postgres.DataFrame", return_value=MagicMock( - to_sql=MagicMock(side_effect=Exception("Error exporting data to postgres")) -)) -async def test_export_data_to_postgres_error(mock_dataframe, postgres_client): - data = {"schema": "test", "table_name": "test", - "data": {"a": [1, 2, 3], "b": [4, 5, 6]}} - await postgres_client.export_data_to_postgres(data) - postgres_client.notification_handler.build_and_send_notification.assert_called_once_with( - notification_id="ERROR_EXPORTING_DATA_TO_POSTGRES", - message="Error exporting data to postgres: Error exporting data to postgres", - block="export_data_to_postgres", - level=NotificationLevel.ERROR, - attachment_content=ANY - ) +async def test_repeat_last_prediction_error(postgres_activity): + query_items = { + "schema": "public", + "table_name": "predictions", + "model": 1 + } + error_msg = "Database error" - postgres_client.pool.getconn.assert_called_once() - postgres_client.pool.putconn.assert_called_once() + with patch("sqlalchemy.orm.session.Session.execute", side_effect=ValueError(error_msg)): + await postgres_activity.repeat_last_prediction(query_items) + + postgres_activity.notification_handler.build_and_send_notification.assert_called_once() + postgres_activity.logger.error.assert_called() + + +@mark.asyncio +async def test_export_data_to_postgres_success(postgres_activity): + input_data = { + "schema": "public", + "table_name": "test_table", + "data": pd.DataFrame({"column1": [1, 2], "column2": ["a", "b"]}) + } + + with patch("laborious.activities.postgres.DataFrame.to_sql") as mock_to_sql: + await postgres_activity.export_data_to_postgres(input_data) + + mock_to_sql.assert_called_once() + postgres_activity.logger.debug.assert_called() + + +@mark.asyncio +async def test_export_data_to_postgres_error(postgres_activity): + input_data = { + "schema": "public", + "table_name": "test_table", + "data": pd.DataFrame({"column1": [1, 2], "column2": ["a", "b"]}) + } + error_msg = "Export failed" + + with patch("laborious.activities.postgres.DataFrame.to_sql", side_effect=ValueError(error_msg)): + await postgres_activity.export_data_to_postgres(input_data) + + postgres_activity.notification_handler.build_and_send_notification.assert_called_once() + postgres_activity.logger.error.assert_called() + + +@mark.asyncio +async def test_close(postgres_activity): + postgres_activity.close() + + postgres_activity.engine.dispose.assert_called_once() + + +@mark.asyncio +async def test_del(postgres_activity): + postgres_activity.close = MagicMock() + postgres_activity.__del__() + + postgres_activity.close.assert_called_once() diff --git a/tests/laborious/utils/filters/test_conditional_filters.py b/tests/laborious/utils/filters/test_conditional_filters.py index 8fec8be..edcbcd6 100644 --- a/tests/laborious/utils/filters/test_conditional_filters.py +++ b/tests/laborious/utils/filters/test_conditional_filters.py @@ -1,26 +1,30 @@ -""" from pandas import DataFrame +from pandas import DataFrame -from laborious.utils.filters.conditional_filters import filter_specific_variables_null_values, filter_empty_data +from laborious.utils.filters.conditional_filters import ( + filter_specific_variables_null_values, + filter_empty_data +) def test_filter_specific_variables_null_values(): assert filter_specific_variables_null_values( DataFrame( {'variable': ['variable1', 'variable2'], 'value': [1, 2]}), - 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]}), - 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(): assert filter_empty_data( - DataFrame({'variable': ['variable1', 'variable2'], 'value': [1, 2]})) == False """ + DataFrame({'variable': ['variable1', 'variable2'], 'value': [1, 2]}), + {}) is False diff --git a/tests/laborious/utils/filters/test_mlflow_filters.py b/tests/laborious/utils/filters/test_mlflow_filters.py new file mode 100644 index 0000000..f9c61e9 --- /dev/null +++ b/tests/laborious/utils/filters/test_mlflow_filters.py @@ -0,0 +1,22 @@ +from pandas import DataFrame +from laborious.utils.filters.mlflow_filters import api_error_filter, nan_values_filter + + +def test_api_error_filter_invalid_response(): + assert api_error_filter(None, {}) == True # NOSONAR + + +def test_api_error_filter_valid_response_fail(): + assert api_error_filter({'success': False}, {}) == True + + +def test_api_error_filter_valid_response_success(): + assert api_error_filter({'success': True}, {}) == False + + +def test_nan_values_filter_all_nan_values(): + assert nan_values_filter(DataFrame({'variable': [None, None]}), {}) == True + + +def test_nan_values_filter_no_nan_values(): + assert nan_values_filter(DataFrame({'variable': [1, 2]}), {}) == False diff --git a/tests/laborious/utils/repository/test_model_repository.py b/tests/laborious/utils/repository/test_model_repository.py new file mode 100644 index 0000000..abf4ab8 --- /dev/null +++ b/tests/laborious/utils/repository/test_model_repository.py @@ -0,0 +1,278 @@ +from unittest.mock import ANY, MagicMock, patch +import numpy as np +from pandas import DataFrame +import pytest +from laborious.utils.repository.model_repository import MLFlowRepository + + +@pytest.fixture +def mlflow_repository(): + with patch('laborious.utils.repository.model_repository.ModelServing', autospec=True) as MockModelServing: + mock_instance = MockModelServing.return_value + mock_instance.get_transformed_data = MagicMock() + + repo = MLFlowRepository( + host='http://localhost:5000', + username='admin', + password='admin' + ) + return repo + + +def test_get_current_data_df(mlflow_repository): + current_data = { + 'prediction': [1, 3], + 'target': [1, 1], + } + mlflow_repository.model_serving.get_transformed_data.return_value = { + 'var1': [1, 2], + 'var2': [2, np.nan], + } + expected = DataFrame({ + 'var1': [1], + 'var2': [2], + 'prediction': [1], + 'target': [1], + }) + output = mlflow_repository.get_current_data_df(current_data, + 'model', 'target') + + mlflow_repository.model_serving.get_transformed_data.assert_called_once_with( + 'model', current_data, by='model') + + diff = output.compare(expected) + assert diff.empty + + +def test_get_artifact(mlflow_repository): + mlflow_repository.get_artifact( + 'destination', 'search_by', 'run_id', 'model', 'artifact' + ) + mlflow_repository.model_serving.get_artifact.assert_called_once_with( + destination='destination', + search_by='search_by', + run_id='run_id', + model_name='model', + artifact_name='artifact' + ) + + +def test_calculate_model_metrics(mlflow_repository): + mlflow_repository.model_serving.get_model_metrics.return_value = 'data' + real_data = 'real_data' + predictions = 'predictions' + flag = 'flag' + output = mlflow_repository.calculate_model_metrics( + real_data, predictions, flag + ) + mlflow_repository.model_serving.get_model_metrics.assert_called_once_with( + reference_data=None, + real_data=real_data, + predictions=predictions, + type_flag=flag + ) + assert output == 'data' + + +@patch('laborious.utils.repository.model_repository.mlflow') +def test_get_experiment_by_run_id(mlflow, mlflow_repository): + mlflow.get_run.return_value = MagicMock( + info=MagicMock( + experiment_id='0', + ) + ) + mlflow.get_experiment.return_value = MagicMock() + mlflow.get_experiment.return_value.name = 'test' + + output = mlflow_repository.get_experiment_by_run_id('0') + assert output == 'test' + mlflow.get_run.assert_called_once_with('0') + mlflow.get_experiment.assert_called_once_with('0') + + +@patch('laborious.utils.repository.model_repository.mlflow') +def test_get_next_run_name(mlflow, mlflow_repository): + mlflow.search_runs.return_value = [1, 2, 3] + output = mlflow_repository.get_next_run_name('run') + assert output == 'run-4' + mlflow.search_runs.assert_called_once_with( + experiment_names=['run'], + order_by=['start_time desc'], + ) + + +@patch('laborious.utils.repository.model_repository.mlflow') +def test_get_experiment_success(mlflow, mlflow_repository): + mlflow.get_experiment_by_name.return_value = MagicMock( + experiment_id='0') + + output = mlflow_repository.get_experiment('test') + + assert output == 0 + + +@patch('laborious.utils.repository.model_repository.mlflow') +def test_get_experiment_error(mlflow, mlflow_repository): + mlflow.get_experiment_by_name.return_value = None + + try: + mlflow_repository.get_experiment('test') + except ValueError as e: + assert str(e) == 'Experiment test not found' + else: + assert False + + +@patch('laborious.utils.repository.model_repository.mlflow') +def test_get_experiment_last_run(mlflow, mlflow_repository): + mlflow.search_runs.return_value = DataFrame({ + 'params.retrain': ['True', 'False', 'True', 'False'], + 'end_time': ['2021-01-01', '2021-01-02', '2021-01-03', '2021-01-04'], + 'run_id': ['0', '1', '2', '3'], + }) + + output = mlflow_repository.get_experiment_last_run(0) + + mlflow.search_runs.assert_called_once_with( + experiment_ids=[0], + filter_string="", + output_format="pandas", + ) + + assert output == '2' + + +@patch('laborious.utils.repository.model_repository.mlflow') +def test_update_production_model_by_run_id(mlflow, mlflow_repository): + client_mock = MagicMock() + mlflow.tracking.MlflowClient.return_value = client_mock + + client_mock.get_registered_model.return_value = MagicMock( + latest_versions=[ + MagicMock(version='1'), + MagicMock(version='2'), + MagicMock(version='3'), + ] + ) + output = mlflow_repository.update_production_model_by_run_id('0', 'test') + + mlflow.register_model.assert_called_once_with( + "runs:/0/prediction_model", + 'test', + ) + + mlflow.tracking.MlflowClient.assert_called_once() + client_mock.get_registered_model.assert_called_once_with('test') + client_mock.transition_model_version_stage.assert_called_once_with( + name='test', + version='3', + stage='Production', + archive_existing_versions=True, + ) + + assert output == { + 'model_name': 'test', + 'version': '3', + 'mlflow_run_id': '0', + } + + +def test_update_production_model(mlflow_repository): + connector = mlflow_repository + + with patch.object(connector, 'get_experiment', + return_value='0') as get_experiment: + with patch.object(connector, 'get_experiment_last_run', + return_value='2') as get_experiment_last_run: + with patch.object(connector, 'update_production_model_by_run_id', + return_value={'model_name': 'test', 'version': '3', + 'mlflow_run_id': '0'}) as update_production_model_by_run_id: + + output = connector.update_production_model('0', 'test') + + get_experiment.assert_called_once_with('0') + get_experiment_last_run.assert_called_once_with('0') + update_production_model_by_run_id.assert_called_once_with( + '2', 'test') + + assert output == { + 'model_name': 'test', + 'version': '3', + 'mlflow_run_id': '0', + 'mlflow_experiment_id': '0', + } + + +def test_transform_success(mlflow_repository): + data = 'data' + model_name = 'model' + + output = mlflow_repository.transform(model_name, data, 1) + + mlflow_repository.model_serving.get_cached_transform.assert_called_once_with( + model_name, data, 1) + + assert output == { + 'success': True, + 'content': mlflow_repository.model_serving.get_cached_transform.return_value.to_dict.return_value + } + + +def test_transform_error(mlflow_repository): + data = 'data' + model_name = 'model' + + mlflow_repository.model_serving.get_cached_transform.side_effect = Exception( + 'error') + + output = mlflow_repository.transform(model_name, data, 1) + + mlflow_repository.model_serving.get_cached_transform.assert_called_once_with( + model_name, data, 1) + + assert output == { + 'success': False, + 'content': { + 'message': 'error', + 'traceback': ANY + } + } + + +def test_predict_success(mlflow_repository): + data = 'data' + model_name = 'model' + mlflow_repository.model_serving.get_cached_predict.return_value = np.array( + [2, 3] + ) + + output = mlflow_repository.predict(model_name, data, 1) + + mlflow_repository.model_serving.get_cached_predict.assert_called_once_with( + model_name, data, 1) + + assert output['success'] == True + assert output['content'] == {'prediction': { + 0: 2, 1: 3}, 'response_time': ANY} + + +def test_predict_error(mlflow_repository): + data = 'data' + model_name = 'model' + + mlflow_repository.model_serving.get_cached_predict = MagicMock( + side_effect=Exception('error') + ) + + output = mlflow_repository.predict(model_name, data, 1) + + mlflow_repository.model_serving.get_cached_predict.assert_called_once_with( + model_name, data, 1) + + assert output == { + 'success': False, + 'content': { + 'message': 'error', + 'traceback': ANY + } + } 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..ae9dd89 --- /dev/null +++ b/tests/laborious/utils/repository/test_opc_repository.py @@ -0,0 +1,259 @@ +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_validate_connection_failed(opc_repository): + opc_repository.client = MagicMock() + opc_repository.error_count = 0 + + output = opc_repository.validate_connection() + assert output is True + + +def test_write_data_validate_connection_do_nothing(opc_repository): + opc_repository.validate_connection = MagicMock(return_value=True) + 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_called_once_with("ns=2;s=TestNode") + + +def test_write_data_validate_connection_failed(opc_repository): + opc_repository.validate_connection = MagicMock(return_value=False) + opc_repository.client = MagicMock() + opc_repository.error_count = 0 + 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/utils/test_connectors_config.py b/tests/laborious/utils/test_connectors_config.py new file mode 100644 index 0000000..b137bc2 --- /dev/null +++ b/tests/laborious/utils/test_connectors_config.py @@ -0,0 +1,133 @@ +from os import environ +from laborious.utils.connectors_config import (build_mlflow_config, + build_opc_config, + build_postgres_config) + + +def test_build_mlflow_config_with_env_vars(): + # Arrange + environ['MLFLOW_HOST'] = 'http://test-host' + environ['MLFLOW_PORT'] = '8080' + environ['MLFLOW_USERNAME'] = 'test-user' + environ['MLFLOW_PASSWORD'] = 'test-pass' + + # Act + config = build_mlflow_config() + + # Assert + assert config['host'] == 'http://test-host' + assert config['port'] == 8080 + assert config['username'] == 'test-user' + assert config['password'] == 'test-pass' + + +def test_build_mlflow_config_with_defaults(): + # Arrange + # Clear any existing env vars + environ.pop('MLFLOW_HOST', None) + environ.pop('MLFLOW_PORT', None) + environ.pop('MLFLOW_USERNAME', None) + environ.pop('MLFLOW_PASSWORD', None) + + # Act + config = build_mlflow_config() + + # Assert + assert config['host'] == 'http://localhost' + assert config['port'] == 5080 + assert config['username'] == 'aignosi' + assert config['password'] == 'aignosi' + + +def test_build_opc_config_with_env_vars(): + # Arrange + environ['OPC_CONFIG'] = '{"opc": {"name": "test-opc", "url": "opc.tcp://test:4840"}}' + + # Act + config = build_opc_config() + + # Assert + assert config['opc']['name'] == 'test-opc' + assert config['opc']['url'] == 'opc.tcp://test:4840' + + +def test_build_opc_config_with_individual_env_vars(): + # Arrange + environ.pop('OPC_CONFIG', None) + environ['OPC_NAME'] = 'test-name' + environ['OPC_URL'] = 'opc.tcp://test:4840' + environ['OPC_SERVER_URI'] = 'opc.tcp://test:4840' + environ['OPC_RECONNECTION_INTERVAL'] = '300' + + # Act + config = build_opc_config() + + # Assert + assert config['opc']['name'] == 'test-name' + assert config['opc']['url'] == 'opc.tcp://test:4840' + assert config['opc']['server_uri'] == 'opc.tcp://test:4840' + assert config['opc']['reconnection_interval'] == 300 + + +def test_build_opc_config_with_defaults(): + # Arrange + environ.pop('OPC_CONFIG', None) + environ.pop('OPC_NAME', None) + environ.pop('OPC_URL', None) + environ.pop('OPC_SERVER_URI', None) + environ.pop('OPC_RECONNECTION_INTERVAL', None) + + # Act + config = build_opc_config() + + # Assert + assert config['opc']['name'] == 'opc' + assert config['opc']['url'] == 'opc.tcp://localhost:4840' + assert config['opc']['server_uri'] == 'opc.tcp://localhost:4840' + assert config['opc']['reconnection_interval'] == 120 + + +def test_build_postgres_config_with_env_vars(): + # Arrange + environ['POSTGRES_HOST'] = 'test-host' + environ['POSTGRES_PORT'] = '5433' + environ['POSTGRES_USER'] = 'test-user' + environ['POSTGRES_PASSWORD'] = 'test-pass' + environ['POSTGRES_DBNAME'] = 'test-db' + environ['POSTGRES_MIN_CONNECTIONS'] = '10' + environ['POSTGRES_MAX_CONNECTIONS'] = '30' + + # Act + config = build_postgres_config() + + # Assert + assert config['host'] == 'test-host' + assert config['port'] == 5433 + assert config['user'] == 'test-user' + assert config['password'] == 'test-pass' + assert config['dbname'] == 'test-db' + assert config['min_connections'] == 10 + assert config['max_connections'] == 30 + + +def test_build_postgres_config_with_defaults(): + # Arrange + environ.pop('POSTGRES_HOST', None) + environ.pop('POSTGRES_PORT', None) + environ.pop('POSTGRES_USER', None) + environ.pop('POSTGRES_PASSWORD', None) + environ.pop('POSTGRES_DBNAME', None) + environ.pop('POSTGRES_MIN_CONNECTIONS', None) + environ.pop('POSTGRES_MAX_CONNECTIONS', None) + + # Act + config = build_postgres_config() + + # Assert + assert config['host'] == 'localhost' + assert config['port'] == 5432 + assert config['user'] == 'sientia' + assert config['password'] == 'sientia' + assert config['dbname'] == 'sientia' + assert config['min_connections'] == 5 + assert config['max_connections'] == 20 diff --git a/tests/laborious/utils/test_logger.py b/tests/laborious/utils/test_logger.py new file mode 100644 index 0000000..cb68cb4 --- /dev/null +++ b/tests/laborious/utils/test_logger.py @@ -0,0 +1,37 @@ +import os +from unittest.mock import patch +import logging +import pytest +from laborious.utils.logger import get_logger + + +@pytest.fixture +def mock_env_vars(): + with patch.dict(os.environ, {}, clear=True): + yield + + +@pytest.mark.usefixtures("mock_env_vars") +@patch('laborious.utils.logger.logging.Formatter') +@patch('laborious.utils.logger.logging.StreamHandler') +def test_get_logger_defaults(mock_stream_handler, mock_formatter): + """Test logger creation with default settings""" + # Mock the StreamHandler and Formatter + + logger = get_logger('test_logger') + + # Verify logger settings + assert logger.name == 'test_logger' + assert logger.level == logging.INFO + + # Verify handler configuration + mock_stream_handler.return_value.setLevel.assert_called_once_with('INFO') + mock_stream_handler.return_value.setFormatter.assert_called_once() + + # Verify formatter configuration + mock_formatter.assert_called_once_with( + '%(asctime)s - %(name)s - %(levelname)s - %(message)s' + ) + + # Verify handler was added to logger + assert len(logger.handlers) == 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 new file mode 100644 index 0000000..f3d5024 --- /dev/null +++ b/tests/laborious/workflows/subworkflows/test_format_and_export_prediction.py @@ -0,0 +1,127 @@ +from unittest.mock import call, patch, AsyncMock, ANY +from pytest import mark, fixture + +from laborious.activities.activities import Activities +from laborious.workflows.sub_workflows.format_and_export_prediction import FormatAndExportPrediction + + +@fixture +def format_and_export_prediction(): + return FormatAndExportPrediction() + + +@mark.asyncio +@patch("laborious.workflows.sub_workflows.format_and_export_prediction.workflow", new_callable=AsyncMock) +async def test_run_none_path_flag(workflow_mock, format_and_export_prediction): + + input_data = { + "path_flag": None, + "data": {"test": "data"}, + "timestamp": "2021-01-01", + "model_id": 1, + "prediction_confidence": 0, + "schema": "test_schema", + "table_name": "test_table", + "opc_servers": ["test_server"], + "opc_output_config": {"test": "config"} + } + + await format_and_export_prediction.run(input_data) + + workflow_mock.execute_local_activity_method.assert_has_calls([ + call( + Activities.format_prediction, + { + 'data': input_data['data'], + '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( + Activities.export_data_to_postgres, + { + 'schema': input_data['schema'], + 'table_name': input_data['table_name'], + '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( + Activities.write_opc_data, + { + 'opc_output_config': input_data['opc_output_config'], + '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 == 2 + assert workflow_mock.execute_local_activity_method.call_count == 1 + + +@mark.asyncio +@patch("laborious.workflows.sub_workflows.format_and_export_prediction.workflow", new_callable=AsyncMock) +async def test_run_default_path_flag(workflow_mock, format_and_export_prediction): + + input_data = { + "path_flag": "default", + "data": {"test": "data"}, + "timestamp": "2021-01-01", + "model_id": 1, + "prediction_confidence": 0, + "schema": "test_schema", + "table_name": "test_table", + "opc_servers": ["test_server"], + "opc_output_config": {"test": "config"}, + "comment": "test_comment" + } + + await format_and_export_prediction.run(input_data) + + workflow_mock.execute_local_activity_method.assert_has_calls([ + call( + Activities.format_default_prediction, + { + 'timestamp': input_data['timestamp'], + '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([ + call( + Activities.export_data_to_postgres, + { + 'schema': input_data['schema'], + 'table_name': input_data['table_name'], + '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( + Activities.write_opc_data, + { + 'opc_output_config': input_data['opc_output_config'], + '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 == 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 new file mode 100644 index 0000000..4aebc6f --- /dev/null +++ b/tests/laborious/workflows/subworkflows/test_prediction_process.py @@ -0,0 +1,512 @@ +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 + + +@fixture +def prediction_process(): + return PredictionProcess() + + +@mark.asyncio +@patch("laborious.workflows.sub_workflows.prediction_process.workflow", new_callable=AsyncMock) +async def test_run(workflow_mock, prediction_process): + prediction_process.path_flag_handler = AsyncMock(return_value=False) + # Arrange + input_data = { + 'data': {'test': 'data'}, + 'schema': 'test_schema', + 'table_name': 'test_table', + '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', + 'path_priority': ['continue', 'repeat', 'stop'], + 'opc_output_config': {'test': 'config'}, + } + + # Mock the activity responses + workflow_mock.execute_local_activity_method.side_effect = [ + '2024-01-01', # get_last_timestamp + ('continue', 0.95, "Input data with bad quality"), # input_gate + {'content': 'transformed_data', 'timestamp': '2024-01-01'}, # transform_data + # 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 + # mlflow_response_gate (predict) + ('continue', 0.95, "Error"), + ] + + # Act + await prediction_process.run(input_data) + + # Assert + 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['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'] + }, 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['mlflow_transform_filters'], + 'data': {'content': 'transformed_data', 'timestamp': '2024-01-01'}, + '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['mlflow_transform_filters'], + 'data': {'content': 'transformed_data', 'timestamp': '2024-01-01'}, + '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'] + }, 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['mlflow_predict_filters'], + 'data': {'content': 'predicted_data', 'timestamp': '2024-01-01'}, + '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', + { + 'path_flag': 'continue', + 'data': 'predicted_data', + 'prediction_confidence': 0.95, + 'timestamp': '2024-01-01', + 'model_id': 1, + 'model_name': 'test_model_name', + 'model_retention': '30', + 'opc_output_config': input_data['opc_output_config'] + } + ) + + +@mark.asyncio +@patch("laborious.workflows.sub_workflows.prediction_process.workflow", new_callable=AsyncMock) +async def test_run_stop_at_input_gate(workflow_mock, prediction_process): + prediction_process.path_flag_handler = AsyncMock(return_value=True) + # Arrange + input_data = { + 'data': {'test': 'data'}, + 'schema': 'test_schema', + 'table_name': 'test_table', + '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', + 'path_priority': ['continue', 'repeat', 'stop'], + 'opc_output_config': {'test': 'config'} + } + + # Mock the activity responses + workflow_mock.execute_local_activity_method.side_effect = [ + '2024-01-01', # get_last_timestamp + ('stop', 0.95, "Input data with bad quality"), # input_gate + ] + + # Act + await prediction_process.run(input_data) + + # Assert + 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['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() + + +@mark.asyncio +@patch("laborious.workflows.sub_workflows.prediction_process.workflow", new_callable=AsyncMock) +async def test_run_stop_at_first_mlflow_response_gate(workflow_mock, prediction_process): + prediction_process.path_flag_handler = AsyncMock(side_effect=[False, True]) + # Arrange + input_data = { + 'data': {'test': 'data'}, + 'schema': 'test_schema', + 'table_name': 'test_table', + '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', + 'path_priority': ['continue', 'repeat', 'stop'], + 'opc_output_config': {'test': 'config'} + } + + # Mock the activity responses + workflow_mock.execute_local_activity_method.side_effect = [ + '2024-01-01', # get_last_timestamp + ('repeat', 0.95, "Input data with bad quality"), # input_gate + {'content': 'transformed_data', 'timestamp': '2024-01-01'}, # transform_data + ('continue', 0.95, "Error"), # mlflow_response_gate (transform) + ] + + # Act + await prediction_process.run(input_data) + + # Assert + 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['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']}, + 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['mlflow_transform_filters'], + 'data': {'content': 'transformed_data', 'timestamp': '2024-01-01'}, + 'type': 'transform', + 'path_priority': input_data['path_priority'] + }, retry_policy=ANY, start_to_close_timeout=ANY) + ]) + workflow_mock.execute_child_workflow.assert_not_called() + + +@mark.asyncio +@patch("laborious.workflows.sub_workflows.prediction_process.workflow", new_callable=AsyncMock) +async def test_run_stop_at_mlflow_content_gate(workflow_mock, prediction_process): + prediction_process.path_flag_handler = AsyncMock( + side_effect=[False, False, True]) + # Arrange + input_data = { + 'data': {'test': 'data'}, + 'schema': 'test_schema', + 'table_name': 'test_table', + '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', + 'path_priority': ['continue', 'repeat', 'stop'], + 'opc_output_config': {'test': 'config'} + } + + # Mock the activity responses + workflow_mock.execute_local_activity_method.side_effect = [ + '2024-01-01', # get_last_timestamp + ('continue', 0.95, "Input data with bad quality"), # input_gate + {'content': 'transformed_data', 'timestamp': '2024-01-01'}, # transform_data + # mlflow_response_gate (transform) + ('continue', 0.95, "Error"), + # 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_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['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'] + }, 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['mlflow_transform_filters'], + 'data': {'content': 'transformed_data', 'timestamp': '2024-01-01'}, + '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['mlflow_transform_filters'], + 'data': {'content': 'transformed_data', 'timestamp': '2024-01-01'}, + 'type': 'transform', + 'path_priority': input_data['path_priority'] + }, retry_policy=ANY, start_to_close_timeout=ANY)]) + workflow_mock.execute_child_workflow.assert_not_called() + + +@mark.asyncio +@patch("laborious.workflows.sub_workflows.prediction_process.workflow", new_callable=AsyncMock) +async def test_run_stop_at_mlflow_last_response_gate(workflow_mock, prediction_process): + prediction_process.path_flag_handler = AsyncMock( + side_effect=[False, False, False, True]) + # Arrange + input_data = { + 'data': {'test': 'data'}, + 'schema': 'test_schema', + 'table_name': 'test_table', + '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', + 'path_priority': ['continue', 'repeat', 'stop'], + 'opc_output_config': {'test': 'config'} + } + + # Mock the activity responses + workflow_mock.execute_local_activity_method.side_effect = [ + '2024-01-01', # get_last_timestamp + ('continue', 0.95, "Input data with bad quality"), # input_gate + {'content': 'transformed_data', 'timestamp': '2024-01-01'}, # transform_data + # 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, "Error"), # mlflow_response_gate (predict) + ] + + # Act + await prediction_process.run(input_data) + + # Assert + 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['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'] + }, 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['mlflow_transform_filters'], + 'data': {'content': 'transformed_data', 'timestamp': '2024-01-01'}, + '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['mlflow_transform_filters'], + 'data': {'content': 'transformed_data', 'timestamp': '2024-01-01'}, + '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'] + }, 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['mlflow_predict_filters'], + 'data': {'content': 'predicted_data', 'timestamp': '2024-01-01'}, + 'type': 'predict', + 'path_priority': input_data['path_priority'] + }, retry_policy=ANY, start_to_close_timeout=ANY)]) + workflow_mock.execute_child_workflow.assert_not_called() + + +@mark.asyncio +@patch("laborious.workflows.sub_workflows.prediction_process.workflow", new_callable=AsyncMock) +async def test_path_flag_handler_stop(workflow_mock, prediction_process): + # Arrange + data = {'test': 'data'} + path_flag = 'STOP' + confidence = 0.95 + schema = 'test_schema' + table_name = 'test_table' + model = 'test_model' + last_timestamp = '2024-01-01' + model_name = 'test_model_name' + model_retention = '30' + + # Act + result = await prediction_process.path_flag_handler( + data, path_flag, { + 'schema': schema, + 'table_name': table_name, + 'model_id': model, + 'last_timestamp': last_timestamp, + 'model_name': model_name, + 'model_retention': model_retention + }, confidence, last_timestamp, "" + ) + + # Assert + assert result is True + workflow_mock.execute_local_activity_method.assert_not_called() + workflow_mock.execute_child_workflow.assert_not_called() + + +@mark.asyncio +@patch("laborious.workflows.sub_workflows.prediction_process.workflow", new_callable=AsyncMock) +async def test_path_flag_handler_repeat(workflow_mock, prediction_process): + # Arrange + data = {'test': 'data'} + path_flag = 'repeat' + confidence = 0.95 + schema = 'test_schema' + table_name = 'test_table' + model = 'test_model' + last_timestamp = '2024-01-01' + model_name = 'test_model_name' + model_retention = '30' + + # Act + result = await prediction_process.path_flag_handler( + data, path_flag, { + 'schema': schema, + 'table_name': table_name, + 'model_id': model, + 'last_timestamp': last_timestamp, + 'model_name': model_name, + 'model_retention': model_retention + }, confidence, last_timestamp, "" + ) + + # Assert + assert result is True + workflow_mock.execute_activity_method.assert_called_once_with( + Activities.repeat_last_prediction, + { + 'schema': schema, + 'table_name': table_name, + 'model_id': model + }, + retry_policy=ANY, + start_to_close_timeout=ANY + ) + workflow_mock.execute_child_workflow.assert_not_called() + + +@mark.asyncio +@patch("laborious.workflows.sub_workflows.prediction_process.workflow", new_callable=AsyncMock) +async def test_path_flag_handler_continue(workflow_mock, prediction_process): + # Arrange + data = {'test': 'data'} + path_flag = 'CONTINUE' + confidence = 0.95 + schema = 'test_schema' + table_name = 'test_table' + model = 'test_model' + last_timestamp = '2024-01-01' + model_name = 'test_model_name' + model_retention = '30' + + # Act + result = await prediction_process.path_flag_handler( + data, path_flag, { + 'schema': schema, + 'table_name': table_name, + 'model_id': model, + 'last_timestamp': last_timestamp, + 'model_name': model_name, + 'model_retention': model_retention, + 'opc_output_config': {'test': 'config'} + }, confidence, last_timestamp, 'Prediction Process' + ) + + # Assert + assert result is True + workflow_mock.execute_activity_method.assert_not_called() + workflow_mock.execute_child_workflow.assert_called_once_with( + 'format_and_export_prediction', + { + 'path_flag': path_flag, + 'data': data, + 'prediction_confidence': confidence, + 'timestamp': last_timestamp, + 'model_id': model, + 'model_name': model_name, + 'model_retention': model_retention, + 'schema': schema, + 'table_name': table_name, + 'comment': 'Prediction Process', + 'opc_output_config': {'test': 'config'} + } + ) + + +@mark.asyncio +@patch("laborious.workflows.sub_workflows.prediction_process.workflow", new_callable=AsyncMock) +async def test_path_flag_handler_unknown(workflow_mock, prediction_process): + # Arrange + data = {'test': 'data'} + path_flag = 'unknown' + confidence = 0.95 + schema = 'test_schema' + table_name = 'test_table' + model = 'test_model' + last_timestamp = '2024-01-01' + model_name = 'test_model_name' + model_retention = '30' + + # Act + result = await prediction_process.path_flag_handler( + data, path_flag, { + 'schema': schema, + 'table_name': table_name, + 'model_id': model, + 'last_timestamp': last_timestamp, + 'model_name': model_name, + 'model_retention': model_retention, + 'opc_output_config': {'test': 'config'} + }, confidence, last_timestamp, "" + ) + + # Assert + assert result is False + workflow_mock.execute_activity_method.assert_not_called() + workflow_mock.execute_child_workflow.assert_not_called() diff --git a/tests/laborious/workflows/test_predictions_batch.py b/tests/laborious/workflows/test_predictions_batch.py new file mode 100644 index 0000000..0ca45e1 --- /dev/null +++ b/tests/laborious/workflows/test_predictions_batch.py @@ -0,0 +1,81 @@ +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', { + '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', {}) + } + + workflow_mock.execute_child_workflow.assert_has_calls([ + call( + 'prediction_process', prediction_input) + ])