Merge pull request #1 from Aignosi/SIENTIAPDE-994-implementar-os-workflows-mapeados-utilizando-as-workers-e-activities-apropriadas

Sientiapde 994 implementar os workflows mapeados utilizando as workers e activities apropriadas
This commit is contained in:
vitor-aignosi
2025-05-26 14:32:12 -03:00
committed by GitHub
42 changed files with 3671 additions and 794 deletions

4
.env Normal file
View File

@@ -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

View File

@@ -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

8
.gitignore vendored
View File

@@ -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
.secret
# Ignorar coverage
htmlcov/
.coverage

82
docker-compose.yml Normal file
View File

@@ -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

34
input_sample.json Normal file
View File

@@ -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": {}
}

View File

@@ -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)

View File

@@ -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']

View File

@@ -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())

View File

@@ -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

View File

@@ -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'
)

View File

@@ -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()

View File

@@ -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'))
}
}

View File

@@ -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}')

View File

@@ -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

View File

@@ -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

View File

@@ -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}.")

22
laborious/utils/logger.py Normal file
View File

@@ -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

View File

@@ -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
)

View File

@@ -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,

View File

@@ -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

View File

@@ -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())

View File

@@ -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)

View File

@@ -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

View File

@@ -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

View File

@@ -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

30
simulator/Dockerfile Normal file
View File

@@ -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"]

55
simulator/redis-feeder.py Normal file
View File

@@ -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.")

View File

@@ -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']

View File

@@ -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"

View File

@@ -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

View File

@@ -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
)

View File

@@ -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()

View File

@@ -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()

View File

@@ -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

View File

@@ -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

View File

@@ -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
}
}

View File

@@ -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 - <class 'float'> 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

View File

@@ -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

View File

@@ -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

View File

@@ -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

View File

@@ -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()

View File

@@ -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)
])