SIENTIAPDE-1094
Refactor activity imports and remove unused base and logger files; update requirements for library versioning
This commit is contained in:
@@ -1,7 +1,7 @@
|
||||
from pytest import mark
|
||||
from unittest.mock import patch, MagicMock, ANY
|
||||
from sientia_do.temporal.activities.postgres import Postgres
|
||||
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
|
||||
@@ -139,9 +139,9 @@ async def test_prepare_activity(_mock_opc_init,
|
||||
|
||||
await activities.prepare_activity(input_data)
|
||||
|
||||
assert activities.notification_handler.base_notification.pipeline_name == input_data[
|
||||
assert activities.notification_handler.base_notification.pipeline == input_data[
|
||||
'workflow_name']
|
||||
assert activities.notification_handler.base_notification.schedule_name == input_data[
|
||||
assert activities.notification_handler.base_notification.trigger == input_data[
|
||||
'schedule_name']
|
||||
assert activities.notification_handler.base_notification.model_name == input_data[
|
||||
'model_name']
|
||||
|
||||
@@ -1,35 +0,0 @@
|
||||
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"
|
||||
@@ -1,159 +0,0 @@
|
||||
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.create_engine")
|
||||
def postgres_activity(_mock_create_engine):
|
||||
return Postgres(
|
||||
host="localhost",
|
||||
port=5432,
|
||||
user="test_user",
|
||||
password="test_password",
|
||||
dbname="test_db",
|
||||
min_connections=1,
|
||||
max_connections=5,
|
||||
logger=MagicMock(),
|
||||
notification_handler=MagicMock()
|
||||
)
|
||||
|
||||
|
||||
@mark.asyncio
|
||||
@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")
|
||||
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
|
||||
@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"]})
|
||||
|
||||
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_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
|
||||
async def test_repeat_last_prediction_success(postgres_activity):
|
||||
query_items = {
|
||||
"schema": "public",
|
||||
"table_name": "predictions",
|
||||
"model": 1
|
||||
}
|
||||
|
||||
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
|
||||
async def test_repeat_last_prediction_error(postgres_activity):
|
||||
query_items = {
|
||||
"schema": "public",
|
||||
"table_name": "predictions",
|
||||
"model": 1
|
||||
}
|
||||
error_msg = "Database error"
|
||||
|
||||
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()
|
||||
@@ -1,14 +1,14 @@
|
||||
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
|
||||
with patch('laborious.utils.repository.model_repository.ModelServing',
|
||||
autospec=True) as mock_model_serving:
|
||||
mock_instance = mock_model_serving.return_value
|
||||
mock_instance.get_transformed_data = MagicMock()
|
||||
|
||||
repo = MLFlowRepository(
|
||||
@@ -19,190 +19,6 @@ def mlflow_repository():
|
||||
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'
|
||||
|
||||
@@ -1,37 +0,0 @@
|
||||
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
|
||||
Reference in New Issue
Block a user