SIENTIAPDE-1273
Enhance MLFlowRepository and Activities classes with new methods and metrics - Added `check_artifact_exists` method to MLFlowRepository for verifying artifact presence in the MLflow Model Registry. - Implemented `get_prediction_data` method in MLFlowRepository to retrieve prediction data from models. - Updated Activities class to integrate ModelMetrics for improved metrics handling. - Enhanced tests for artifact existence checks and prediction data retrieval, ensuring robust coverage for new functionalities. - Updated various workflows to include `transform_table_name` in input data for better data handling.
This commit is contained in:
641
tests/laborious/activities/test_model_metrics.py
Normal file
641
tests/laborious/activities/test_model_metrics.py
Normal file
@@ -0,0 +1,641 @@
|
||||
from unittest.mock import ANY, AsyncMock, MagicMock, call, patch
|
||||
|
||||
from pandas import DataFrame
|
||||
from pytest import fixture, mark
|
||||
from sientia_do.notifications.models import NotificationLevel
|
||||
|
||||
from laborious.activities.model_metrics import ModelMetrics
|
||||
|
||||
|
||||
@fixture
|
||||
def model_metrics_activity():
|
||||
model_metrics = ModelMetrics(
|
||||
logger=MagicMock(),
|
||||
notification_handler=MagicMock(),
|
||||
metrics_controller=AsyncMock(),
|
||||
)
|
||||
model_metrics.error = MagicMock()
|
||||
model_metrics.debug = MagicMock()
|
||||
model_metrics.info = MagicMock()
|
||||
model_metrics.warning = MagicMock()
|
||||
model_metrics.critical = MagicMock()
|
||||
model_metrics.send_notification = MagicMock()
|
||||
model_metrics.send_notification_async = AsyncMock()
|
||||
model_metrics.emit_metric = AsyncMock()
|
||||
model_metrics.get_core_labels = MagicMock(return_value={'pod_id': 'test_pod', 'model_name': 'test_model', 'workflow_name': 'test_workflow'})
|
||||
model_metrics.observe_lag = AsyncMock()
|
||||
model_metrics.pod_id = 'test_pod'
|
||||
return model_metrics
|
||||
|
||||
|
||||
metadata = {
|
||||
'metadata': {
|
||||
'model_id': 'test_model',
|
||||
'model_name': 'test_model',
|
||||
'workflow_name': 'test_workflow',
|
||||
'schema_name': 'test_schedule',
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
@mark.asyncio
|
||||
async def test_calculate_drift_invalid_chunk_period(model_metrics_activity):
|
||||
# Arrange
|
||||
input_data = {
|
||||
**metadata,
|
||||
'model_name': 'test_model',
|
||||
'model_id': 'test_model_id',
|
||||
'reference_data': None,
|
||||
'target_data': {
|
||||
'timestamp': ['2023-05-26 11:12:27'],
|
||||
'variable': ['feature1'],
|
||||
'value': [1.0],
|
||||
},
|
||||
'target_name': 'target',
|
||||
'drift_metrics': ['ks_test'],
|
||||
'chunk_period': 'invalid',
|
||||
}
|
||||
|
||||
# Act & Assert
|
||||
try:
|
||||
await model_metrics_activity.calculate_drift(input_data)
|
||||
except ValueError as e:
|
||||
assert str(e) == "Invalid chunk period: invalid, must be \"min\" or \"s\""
|
||||
model_metrics_activity.error.assert_called_once_with(
|
||||
'Invalid chunk period: invalid', metadata['metadata']
|
||||
)
|
||||
else:
|
||||
raise AssertionError('Expected ValueError')
|
||||
|
||||
|
||||
@mark.asyncio
|
||||
@patch('laborious.activities.model_metrics.DataFrame')
|
||||
@patch('laborious.activities.model_metrics.to_datetime')
|
||||
async def test_calculate_drift_with_reference_data(
|
||||
mock_to_datetime, mock_dataframe, model_metrics_activity
|
||||
):
|
||||
# Arrange
|
||||
mock_to_datetime.return_value.dt.strftime.return_value = '2023-05-26 11:12:27'
|
||||
mock_to_datetime.return_value.dt.tz_localize.return_value.dt.strftime.return_value = '2023-05-26 11:12:27+00:00'
|
||||
|
||||
mock_drift_df = MagicMock()
|
||||
mock_drift_df.empty = False
|
||||
mock_drift_df.drop.return_value = mock_drift_df
|
||||
mock_drift_df.__getitem__.return_value.isin.return_value = [True]
|
||||
mock_drift_df.__getitem__.return_value = mock_drift_df
|
||||
mock_drift_df.__getitem__.return_value.dt.tz_localize.return_value.dt.strftime.return_value = '2023-05-26 11:12:27+00:00'
|
||||
mock_drift_df.rename.return_value = mock_drift_df
|
||||
mock_drift_df.drop_duplicates.return_value = mock_drift_df
|
||||
mock_drift_df.to_dict.return_value = {'method': ['ks_test'], 'value': [0.5], 'feature': ['feature1'], 'timestamp': ['2023-05-26 11:12:27+00:00'], 'model_id': ['test_model_id'], 'accurate': [True]}
|
||||
|
||||
model_metrics_activity.get_drift_metrics = AsyncMock(return_value=mock_drift_df)
|
||||
|
||||
reference_data = DataFrame({
|
||||
'timestamp': ['2023-05-26 11:12:27'],
|
||||
'target': [1.0],
|
||||
'feature1': [1.0],
|
||||
})
|
||||
|
||||
mock_target_df = MagicMock()
|
||||
mock_target_df.pivot.return_value = mock_target_df
|
||||
mock_target_df.index = ['2023-05-26 11:12:27']
|
||||
mock_target_df.reset_index.return_value = mock_target_df
|
||||
mock_target_df.dropna.return_value = mock_target_df
|
||||
mock_target_df.__getitem__.return_value.apply.return_value = ['2023-05-26 11:12:27']
|
||||
mock_target_df.drop.return_value.columns = ['feature1']
|
||||
mock_dataframe.return_value = mock_target_df
|
||||
|
||||
input_data = {
|
||||
**metadata,
|
||||
'model_name': 'test_model',
|
||||
'model_id': 'test_model_id',
|
||||
'reference_data': reference_data.to_dict(),
|
||||
'target_data': {
|
||||
'timestamp': ['2023-05-26 11:12:27'],
|
||||
'variable': ['feature1'],
|
||||
'value': [1.0],
|
||||
},
|
||||
'target_name': 'target',
|
||||
'drift_metrics': ['ks_test'],
|
||||
'chunk_period': 'min',
|
||||
}
|
||||
|
||||
# Act
|
||||
result = await model_metrics_activity.calculate_drift(input_data)
|
||||
|
||||
# Assert
|
||||
assert isinstance(result, dict)
|
||||
assert result == mock_drift_df.to_dict.return_value
|
||||
model_metrics_activity.info.assert_called()
|
||||
model_metrics_activity.get_drift_metrics.assert_called_once()
|
||||
# Verify transformations were called
|
||||
mock_drift_df.drop.assert_called_once_with(columns=['p_value', 'chunk_start_date', 'chunk_end_date'], inplace=True)
|
||||
mock_drift_df.__getitem__.assert_called()
|
||||
mock_drift_df.rename.assert_called_once_with(columns={'metric': 'method', 'statistic': 'value'}, inplace=True)
|
||||
mock_drift_df.drop_duplicates.assert_called_once_with(subset=['timestamp', 'method', 'feature'], keep='first', inplace=True)
|
||||
mock_drift_df.to_dict.assert_called_once()
|
||||
|
||||
|
||||
@mark.asyncio
|
||||
@patch('laborious.activities.model_metrics.DataFrame')
|
||||
@patch('laborious.activities.model_metrics.to_datetime')
|
||||
async def test_calculate_drift_without_reference_data(
|
||||
mock_to_datetime, mock_dataframe, model_metrics_activity
|
||||
):
|
||||
# Arrange
|
||||
mock_to_datetime.return_value.dt.strftime.return_value = '2023-05-26 11:12:27'
|
||||
mock_to_datetime.return_value.dt.tz_localize.return_value.dt.strftime.return_value = '2023-05-26 11:12:27+00:00'
|
||||
|
||||
mock_drift_df = MagicMock()
|
||||
mock_drift_df.empty = False
|
||||
mock_drift_df.drop.return_value = mock_drift_df
|
||||
mock_drift_df.__getitem__.return_value.isin.return_value = [True]
|
||||
mock_drift_df.__getitem__.return_value = mock_drift_df
|
||||
mock_drift_df.__getitem__.return_value.dt.tz_localize.return_value.dt.strftime.return_value = '2023-05-26 11:12:27+00:00'
|
||||
mock_drift_df.rename.return_value = mock_drift_df
|
||||
mock_drift_df.drop_duplicates.return_value = mock_drift_df
|
||||
mock_drift_df.to_dict.return_value = {'method': ['ks_test'], 'value': [0.5], 'feature': ['feature1'], 'timestamp': ['2023-05-26 11:12:27+00:00'], 'model_id': ['test_model_id'], 'accurate': [False]}
|
||||
|
||||
model_metrics_activity.get_drift_metrics = AsyncMock(return_value=mock_drift_df)
|
||||
|
||||
target_data_dict = {
|
||||
'timestamp': ['2023-05-26 11:12:27', '2023-05-26 11:12:28', '2023-05-26 11:12:29'],
|
||||
'variable': ['feature1', 'feature1', 'feature1'],
|
||||
'value': [1.0, 2.0, 3.0],
|
||||
}
|
||||
|
||||
mock_target_df = MagicMock()
|
||||
mock_target_df.pivot.return_value = mock_target_df
|
||||
mock_target_df.index = ['2023-05-26 11:12:27', '2023-05-26 11:12:28', '2023-05-26 11:12:29']
|
||||
mock_target_df.reset_index.return_value = mock_target_df
|
||||
mock_target_df.dropna.return_value = mock_target_df
|
||||
mock_target_df.sort_values.return_value = mock_target_df
|
||||
mock_target_df.head.return_value = DataFrame({'timestamp': ['2023-05-26 11:12:27'], 'feature1': [1.0]})
|
||||
mock_target_df.__getitem__.return_value.apply.return_value = ['2023-05-26 11:12:27', '2023-05-26 11:12:28', '2023-05-26 11:12:29']
|
||||
mock_target_df.drop.return_value.columns = ['feature1']
|
||||
mock_dataframe.return_value = mock_target_df
|
||||
mock_dataframe.side_effect = lambda x=None: mock_target_df if x is not None else mock_target_df
|
||||
|
||||
input_data = {
|
||||
**metadata,
|
||||
'model_name': 'test_model',
|
||||
'model_id': 'test_model_id',
|
||||
'reference_data': None,
|
||||
'target_data': target_data_dict,
|
||||
'target_name': 'target',
|
||||
'drift_metrics': ['ks_test'],
|
||||
'chunk_period': 's',
|
||||
}
|
||||
|
||||
# Act
|
||||
result = await model_metrics_activity.calculate_drift(input_data)
|
||||
|
||||
# Assert
|
||||
assert isinstance(result, dict)
|
||||
assert result == mock_drift_df.to_dict.return_value
|
||||
model_metrics_activity.warning.assert_called()
|
||||
model_metrics_activity.send_notification_async.assert_called_once_with(
|
||||
metadata=metadata['metadata'],
|
||||
notification_id='MODEL_METRICS_REFERENCE_DATA_WARNING',
|
||||
message='Using 30% first rows of target data as reference data',
|
||||
block='model_metrics',
|
||||
level=NotificationLevel.WARNING,
|
||||
attachment_content=ANY,
|
||||
)
|
||||
# Verify transformations were called
|
||||
mock_drift_df.drop.assert_called_once_with(columns=['p_value', 'chunk_start_date', 'chunk_end_date'], inplace=True)
|
||||
mock_drift_df.__getitem__.assert_called()
|
||||
mock_drift_df.rename.assert_called_once_with(columns={'metric': 'method', 'statistic': 'value'}, inplace=True)
|
||||
mock_drift_df.drop_duplicates.assert_called_once_with(subset=['timestamp', 'method', 'feature'], keep='first', inplace=True)
|
||||
mock_drift_df.to_dict.assert_called_once()
|
||||
|
||||
|
||||
@mark.asyncio
|
||||
@patch('laborious.activities.model_metrics.DataFrame')
|
||||
@patch('laborious.activities.model_metrics.to_datetime')
|
||||
async def test_calculate_drift_empty_drift_df(
|
||||
mock_to_datetime, mock_dataframe, model_metrics_activity
|
||||
):
|
||||
# Arrange
|
||||
mock_to_datetime.return_value.dt.strftime.return_value = '2023-05-26 11:12:27'
|
||||
|
||||
model_metrics_activity.get_drift_metrics = AsyncMock(return_value=DataFrame())
|
||||
|
||||
reference_data = DataFrame({
|
||||
'timestamp': ['2023-05-26 11:12:27'],
|
||||
'target': [1.0],
|
||||
'feature1': [1.0],
|
||||
})
|
||||
|
||||
mock_target_df = MagicMock()
|
||||
mock_target_df.pivot.return_value = mock_target_df
|
||||
mock_target_df.index = ['2023-05-26 11:12:27']
|
||||
mock_target_df.reset_index.return_value = mock_target_df
|
||||
mock_target_df.dropna.return_value = mock_target_df
|
||||
mock_target_df.__getitem__.return_value.apply.return_value = ['2023-05-26 11:12:27']
|
||||
mock_target_df.drop.return_value.columns = ['feature1']
|
||||
mock_dataframe.return_value = mock_target_df
|
||||
|
||||
input_data = {
|
||||
**metadata,
|
||||
'model_name': 'test_model',
|
||||
'model_id': 'test_model_id',
|
||||
'reference_data': reference_data.to_dict(),
|
||||
'target_data': {
|
||||
'timestamp': ['2023-05-26 11:12:27'],
|
||||
'variable': ['feature1'],
|
||||
'value': [1.0],
|
||||
},
|
||||
'target_name': 'target',
|
||||
'drift_metrics': ['ks_test'],
|
||||
'chunk_period': 'min',
|
||||
}
|
||||
|
||||
# Act
|
||||
result = await model_metrics_activity.calculate_drift(input_data)
|
||||
|
||||
# Assert
|
||||
assert result == {}
|
||||
model_metrics_activity.warning.assert_called_with('No drift metrics found', metadata['metadata'])
|
||||
|
||||
|
||||
@mark.asyncio
|
||||
@patch('laborious.activities.model_metrics.DataFrame')
|
||||
@patch('laborious.activities.model_metrics.to_datetime')
|
||||
async def test_calculate_drift_empty_after_timestamp_filter(
|
||||
mock_to_datetime, mock_dataframe, model_metrics_activity
|
||||
):
|
||||
# Arrange
|
||||
mock_to_datetime.return_value.dt.strftime.return_value = '2023-05-26 11:12:27'
|
||||
|
||||
mock_drift_df = MagicMock()
|
||||
mock_drift_df.empty = False
|
||||
mock_drift_df.drop.return_value = mock_drift_df
|
||||
|
||||
# Set up __getitem__ to handle filtering - timestamp access returns series with isin=False
|
||||
# and filtering returns empty DataFrame
|
||||
mock_timestamp_series = MagicMock()
|
||||
mock_timestamp_series.isin.return_value = [False]
|
||||
mock_empty_df = MagicMock()
|
||||
mock_empty_df.empty = True
|
||||
|
||||
def getitem_side_effect(key):
|
||||
if key == 'timestamp':
|
||||
return mock_timestamp_series
|
||||
else:
|
||||
# This is the filtering operation - return empty DataFrame
|
||||
return mock_empty_df
|
||||
|
||||
mock_drift_df.__getitem__.side_effect = getitem_side_effect
|
||||
|
||||
model_metrics_activity.get_drift_metrics = AsyncMock(return_value=mock_drift_df)
|
||||
|
||||
reference_data = DataFrame({
|
||||
'timestamp': ['2023-05-26 11:12:27'],
|
||||
'target': [1.0],
|
||||
'feature1': [1.0],
|
||||
})
|
||||
|
||||
mock_target_df = MagicMock()
|
||||
mock_target_df.pivot.return_value = mock_target_df
|
||||
mock_target_df.index = ['2023-05-26 11:12:27']
|
||||
mock_target_df.reset_index.return_value = mock_target_df
|
||||
mock_target_df.dropna.return_value = mock_target_df
|
||||
mock_target_df.__getitem__.return_value.apply.return_value = ['2023-05-26 11:12:27']
|
||||
mock_target_df.drop.return_value.columns = ['feature1']
|
||||
mock_dataframe.return_value = mock_target_df
|
||||
|
||||
input_data = {
|
||||
**metadata,
|
||||
'model_name': 'test_model',
|
||||
'model_id': 'test_model_id',
|
||||
'reference_data': reference_data.to_dict(),
|
||||
'target_data': {
|
||||
'timestamp': ['2023-05-26 11:12:27'],
|
||||
'variable': ['feature1'],
|
||||
'value': [1.0],
|
||||
},
|
||||
'target_name': 'target',
|
||||
'drift_metrics': ['ks_test'],
|
||||
'chunk_period': 'min',
|
||||
}
|
||||
|
||||
# Act
|
||||
result = await model_metrics_activity.calculate_drift(input_data)
|
||||
|
||||
# Assert
|
||||
assert result == {}
|
||||
model_metrics_activity.warning.assert_called_with(
|
||||
'No drift metrics found after dropping rows where timestamp is not in target data',
|
||||
metadata['metadata']
|
||||
)
|
||||
# Verify transformations were called
|
||||
mock_drift_df.drop.assert_called_once_with(columns=['p_value', 'chunk_start_date', 'chunk_end_date'], inplace=True)
|
||||
mock_drift_df.__getitem__.assert_called()
|
||||
|
||||
|
||||
@mark.asyncio
|
||||
@patch('laborious.activities.model_metrics.DataFrame')
|
||||
@patch('laborious.activities.model_metrics.to_datetime')
|
||||
async def test_calculate_drift_success_min(
|
||||
mock_to_datetime, mock_dataframe, model_metrics_activity
|
||||
):
|
||||
# Arrange
|
||||
mock_to_datetime.return_value.dt.strftime.return_value = '2023-05-26 11:12:27'
|
||||
mock_to_datetime.return_value.dt.tz_localize.return_value.dt.strftime.return_value = '2023-05-26 11:12:27+00:00'
|
||||
|
||||
mock_drift_df = MagicMock()
|
||||
mock_drift_df.empty = False
|
||||
mock_drift_df.drop.return_value = mock_drift_df
|
||||
mock_drift_df.__getitem__.return_value.isin.return_value = [True]
|
||||
mock_drift_df.__getitem__.return_value = mock_drift_df
|
||||
mock_drift_df.__getitem__.return_value.dt.tz_localize.return_value.dt.strftime.return_value = '2023-05-26 11:12:27+00:00'
|
||||
mock_drift_df.rename.return_value = mock_drift_df
|
||||
mock_drift_df.drop_duplicates.return_value = mock_drift_df
|
||||
mock_drift_df.to_dict.return_value = {'method': ['ks_test'], 'value': [0.5], 'feature': ['feature1'], 'timestamp': ['2023-05-26 11:12:27+00:00'], 'model_id': ['test_model_id'], 'accurate': [True]}
|
||||
|
||||
model_metrics_activity.get_drift_metrics = AsyncMock(return_value=mock_drift_df)
|
||||
|
||||
reference_data = DataFrame({
|
||||
'timestamp': ['2023-05-26 11:12:27'],
|
||||
'target': [1.0],
|
||||
'feature1': [1.0],
|
||||
})
|
||||
|
||||
mock_target_df = MagicMock()
|
||||
mock_target_df.pivot.return_value = mock_target_df
|
||||
mock_target_df.index = ['2023-05-26 11:12:27']
|
||||
mock_target_df.reset_index.return_value = mock_target_df
|
||||
mock_target_df.dropna.return_value = mock_target_df
|
||||
mock_target_df.__getitem__.return_value.isin.return_value = [True]
|
||||
mock_target_df.drop.return_value.columns = ['feature1']
|
||||
mock_dataframe.return_value = mock_target_df
|
||||
|
||||
input_data = {
|
||||
**metadata,
|
||||
'model_name': 'test_model',
|
||||
'model_id': 'test_model_id',
|
||||
'reference_data': reference_data.to_dict(),
|
||||
'target_data': {
|
||||
'timestamp': ['2023-05-26 11:12:27'],
|
||||
'variable': ['feature1'],
|
||||
'value': [1.0],
|
||||
},
|
||||
'target_name': 'target',
|
||||
'drift_metrics': ['ks_test'],
|
||||
'chunk_period': 'min',
|
||||
}
|
||||
|
||||
# Act
|
||||
result = await model_metrics_activity.calculate_drift(input_data)
|
||||
|
||||
# Assert
|
||||
assert isinstance(result, dict)
|
||||
assert result == mock_drift_df.to_dict.return_value
|
||||
model_metrics_activity.info.assert_called()
|
||||
model_metrics_activity.get_drift_metrics.assert_called_once()
|
||||
# Verify transformations were called
|
||||
mock_drift_df.drop.assert_called_once_with(columns=['p_value', 'chunk_start_date', 'chunk_end_date'], inplace=True)
|
||||
mock_drift_df.__getitem__.assert_called()
|
||||
mock_drift_df.rename.assert_called_once_with(columns={'metric': 'method', 'statistic': 'value'}, inplace=True)
|
||||
mock_drift_df.drop_duplicates.assert_called_once_with(subset=['timestamp', 'method', 'feature'], keep='first', inplace=True)
|
||||
mock_drift_df.to_dict.assert_called_once()
|
||||
|
||||
|
||||
@mark.asyncio
|
||||
@patch('laborious.activities.model_metrics.DataFrame')
|
||||
@patch('laborious.activities.model_metrics.to_datetime')
|
||||
async def test_calculate_drift_success_s(
|
||||
mock_to_datetime, mock_dataframe, model_metrics_activity
|
||||
):
|
||||
# Arrange
|
||||
mock_to_datetime.return_value.dt.strftime.return_value = '2023-05-26 11:12:27'
|
||||
mock_to_datetime.return_value.dt.tz_localize.return_value.dt.strftime.return_value = '2023-05-26 11:12:27+00:00'
|
||||
|
||||
mock_drift_df = MagicMock()
|
||||
mock_drift_df.empty = False
|
||||
mock_drift_df.drop.return_value = mock_drift_df
|
||||
mock_drift_df.__getitem__.return_value.isin.return_value = [True]
|
||||
mock_drift_df.__getitem__.return_value = mock_drift_df
|
||||
mock_drift_df.__getitem__.return_value.dt.tz_localize.return_value.dt.strftime.return_value = '2023-05-26 11:12:27+00:00'
|
||||
mock_drift_df.rename.return_value = mock_drift_df
|
||||
mock_drift_df.drop_duplicates.return_value = mock_drift_df
|
||||
mock_drift_df.to_dict.return_value = {'method': ['ks_test'], 'value': [0.5], 'feature': ['feature1'], 'timestamp': ['2023-05-26 11:12:27+00:00'], 'model_id': ['test_model_id'], 'accurate': [True]}
|
||||
|
||||
model_metrics_activity.get_drift_metrics = AsyncMock(return_value=mock_drift_df)
|
||||
|
||||
reference_data = DataFrame({
|
||||
'timestamp': ['2023-05-26 11:12:27'],
|
||||
'target': [1.0],
|
||||
'feature1': [1.0],
|
||||
})
|
||||
|
||||
mock_target_df = MagicMock()
|
||||
mock_target_df.pivot.return_value = mock_target_df
|
||||
mock_target_df.index = ['2023-05-26 11:12:27']
|
||||
mock_target_df.reset_index.return_value = mock_target_df
|
||||
mock_target_df.dropna.return_value = mock_target_df
|
||||
mock_target_df.__getitem__.return_value.isin.return_value = [True]
|
||||
mock_target_df.drop.return_value.columns = ['feature1']
|
||||
mock_dataframe.return_value = mock_target_df
|
||||
|
||||
input_data = {
|
||||
**metadata,
|
||||
'model_name': 'test_model',
|
||||
'model_id': 'test_model_id',
|
||||
'reference_data': reference_data.to_dict(),
|
||||
'target_data': {
|
||||
'timestamp': ['2023-05-26 11:12:27'],
|
||||
'variable': ['feature1'],
|
||||
'value': [1.0],
|
||||
},
|
||||
'target_name': 'target',
|
||||
'drift_metrics': ['ks_test'],
|
||||
'chunk_period': 's',
|
||||
}
|
||||
|
||||
# Act
|
||||
result = await model_metrics_activity.calculate_drift(input_data)
|
||||
|
||||
# Assert
|
||||
assert isinstance(result, dict)
|
||||
assert result == mock_drift_df.to_dict.return_value
|
||||
model_metrics_activity.info.assert_called()
|
||||
model_metrics_activity.get_drift_metrics.assert_called_once()
|
||||
# Verify transformations were called
|
||||
mock_drift_df.drop.assert_called_once_with(columns=['p_value', 'chunk_start_date', 'chunk_end_date'], inplace=True)
|
||||
mock_drift_df.__getitem__.assert_called()
|
||||
mock_drift_df.rename.assert_called_once_with(columns={'metric': 'method', 'statistic': 'value'}, inplace=True)
|
||||
mock_drift_df.drop_duplicates.assert_called_once_with(subset=['timestamp', 'method', 'feature'], keep='first', inplace=True)
|
||||
mock_drift_df.to_dict.assert_called_once()
|
||||
|
||||
|
||||
@mark.asyncio
|
||||
@patch('laborious.activities.model_metrics.DataFrame')
|
||||
@patch('laborious.activities.model_metrics.to_datetime')
|
||||
async def test_calculate_drift_get_drift_metrics_error(
|
||||
mock_to_datetime, mock_dataframe, model_metrics_activity
|
||||
):
|
||||
# Arrange
|
||||
mock_to_datetime.return_value.dt.strftime.return_value = '2023-05-26 11:12:27'
|
||||
|
||||
model_metrics_activity.get_drift_metrics = AsyncMock(side_effect=Exception('Get drift metrics error'))
|
||||
|
||||
reference_data = DataFrame({
|
||||
'timestamp': ['2023-05-26 11:12:27'],
|
||||
'target': [1.0],
|
||||
'feature1': [1.0],
|
||||
})
|
||||
|
||||
mock_target_df = MagicMock()
|
||||
mock_target_df.pivot.return_value = mock_target_df
|
||||
mock_target_df.index = ['2023-05-26 11:12:27']
|
||||
mock_target_df.reset_index.return_value = mock_target_df
|
||||
mock_target_df.dropna.return_value = mock_target_df
|
||||
mock_target_df.__getitem__.return_value.apply.return_value = ['2023-05-26 11:12:27']
|
||||
mock_target_df.drop.return_value.columns = ['feature1']
|
||||
mock_dataframe.return_value = mock_target_df
|
||||
|
||||
input_data = {
|
||||
**metadata,
|
||||
'model_name': 'test_model',
|
||||
'model_id': 'test_model_id',
|
||||
'reference_data': reference_data.to_dict(),
|
||||
'target_data': {
|
||||
'timestamp': ['2023-05-26 11:12:27'],
|
||||
'variable': ['feature1'],
|
||||
'value': [1.0],
|
||||
},
|
||||
'target_name': 'target',
|
||||
'drift_metrics': ['ks_test'],
|
||||
'chunk_period': 'min',
|
||||
}
|
||||
|
||||
# Act
|
||||
result = await model_metrics_activity.calculate_drift(input_data)
|
||||
|
||||
# Assert
|
||||
assert result == {}
|
||||
model_metrics_activity.error.assert_called_once_with(
|
||||
'Error getting drift metrics: Get drift metrics error',
|
||||
metadata['metadata']
|
||||
)
|
||||
model_metrics_activity.send_notification_async.assert_called_once_with(
|
||||
metadata=metadata['metadata'],
|
||||
notification_id='MODEL_METRICS_GET_DRIFT_METRICS_ERROR',
|
||||
message='Error getting drift metrics: Get drift metrics error',
|
||||
block='model_metrics',
|
||||
level=NotificationLevel.ERROR,
|
||||
attachment_content=ANY,
|
||||
)
|
||||
|
||||
|
||||
@mark.asyncio
|
||||
@patch('laborious.activities.model_metrics.to_datetime')
|
||||
@patch('laborious.activities.model_metrics.time.time')
|
||||
@patch('laborious.activities.model_metrics.ModelAnalysis')
|
||||
@patch('laborious.activities.model_metrics.metrics')
|
||||
async def test_get_drift_metrics_success(
|
||||
mock_metrics, mock_model_analysis, mock_time, mock_to_datetime, model_metrics_activity
|
||||
):
|
||||
# Arrange
|
||||
mock_time.return_value = 1000.0
|
||||
|
||||
mock_drift_df = DataFrame({
|
||||
'timestamp': ['2023-05-26 11:12:27'],
|
||||
'metric': ['ks_test'],
|
||||
'statistic': [0.5],
|
||||
'feature': ['feature1'],
|
||||
})
|
||||
|
||||
mock_model_analysis.return_value.detect_univariate_drift.return_value = MagicMock()
|
||||
mock_model_analysis.return_value.detect_multivariate_drift.return_value = MagicMock()
|
||||
mock_model_analysis.return_value.get_drift_metrics_dataframe.return_value = mock_drift_df
|
||||
|
||||
reference_data = DataFrame({
|
||||
'timestamp': ['2023-05-26 11:12:27'],
|
||||
'target': [1.0],
|
||||
'feature1': [1.0],
|
||||
})
|
||||
|
||||
target_data = DataFrame({
|
||||
'timestamp': ['2023-05-26 11:12:27'],
|
||||
'target': [1.0],
|
||||
'feature1': [1.0],
|
||||
})
|
||||
|
||||
reference_columns = reference_data.drop(
|
||||
columns=['target', 'timestamp'], errors='ignore'
|
||||
).columns
|
||||
|
||||
# Act
|
||||
result = await model_metrics_activity.get_drift_metrics(
|
||||
reference_data=reference_data,
|
||||
target_data=target_data,
|
||||
target_name='target',
|
||||
reference_columns=reference_columns,
|
||||
drift_metrics=['ks_test'],
|
||||
chunk_period='min',
|
||||
metadata=metadata['metadata'],
|
||||
)
|
||||
|
||||
# Assert
|
||||
assert isinstance(result, DataFrame)
|
||||
model_metrics_activity.debug.assert_called()
|
||||
model_metrics_activity.observe_lag.assert_called()
|
||||
model_metrics_activity.emit_metric.assert_called()
|
||||
|
||||
|
||||
@mark.asyncio
|
||||
@patch('laborious.activities.model_metrics.to_datetime')
|
||||
@patch('laborious.activities.model_metrics.time.time')
|
||||
@patch('laborious.activities.model_metrics.ModelAnalysis')
|
||||
@patch('laborious.activities.model_metrics.metrics')
|
||||
async def test_get_drift_metrics_univariate_error(
|
||||
mock_metrics, mock_model_analysis, mock_time, mock_to_datetime, model_metrics_activity
|
||||
):
|
||||
# Arrange
|
||||
mock_time.return_value = 1000.0
|
||||
|
||||
mock_model_analysis.return_value.detect_univariate_drift.side_effect = Exception('Univariate drift error')
|
||||
|
||||
reference_data = DataFrame({
|
||||
'timestamp': ['2023-05-26 11:12:27'],
|
||||
'target': [1.0],
|
||||
'feature1': [1.0],
|
||||
})
|
||||
|
||||
target_data = DataFrame({
|
||||
'timestamp': ['2023-05-26 11:12:27'],
|
||||
'target': [1.0],
|
||||
'feature1': [1.0],
|
||||
})
|
||||
|
||||
reference_columns = reference_data.drop(
|
||||
columns=['target', 'timestamp'], errors='ignore'
|
||||
).columns
|
||||
|
||||
# Act & Assert
|
||||
try:
|
||||
await model_metrics_activity.get_drift_metrics(
|
||||
reference_data=reference_data,
|
||||
target_data=target_data,
|
||||
target_name='target',
|
||||
reference_columns=reference_columns,
|
||||
drift_metrics=['ks_test'],
|
||||
chunk_period='min',
|
||||
metadata=metadata['metadata'],
|
||||
)
|
||||
except Exception as e:
|
||||
assert str(e) == 'Univariate drift error'
|
||||
model_metrics_activity.error.assert_called_once_with(
|
||||
'Error detecting univariate drift: Univariate drift error',
|
||||
metadata['metadata']
|
||||
)
|
||||
model_metrics_activity.emit_metric.assert_called_with(
|
||||
metric_object=mock_metrics.MODEL_ANALYZE_ERROR_COUNT,
|
||||
tags=ANY
|
||||
)
|
||||
else:
|
||||
raise AssertionError('Expected Exception')
|
||||
|
||||
@@ -190,6 +190,26 @@ def test_get_model_params(mlflow, mlflow_repository):
|
||||
assert output == mlflow.get_run.return_value.data.params
|
||||
|
||||
|
||||
def test_check_artifact_exists_true(mlflow_repository):
|
||||
artifact = MagicMock(path='test_artifact')
|
||||
mlflow_repository.client.list_artifacts.return_value = [artifact]
|
||||
|
||||
result = mlflow_repository.check_artifact_exists('run_id', 'test_artifact', metadata['metadata'])
|
||||
|
||||
assert result is True
|
||||
mlflow_repository.client.list_artifacts.assert_called_once_with('run_id')
|
||||
|
||||
|
||||
def test_check_artifact_exists_false(mlflow_repository):
|
||||
artifact = MagicMock(path='other_artifact')
|
||||
mlflow_repository.client.list_artifacts.return_value = [artifact]
|
||||
|
||||
result = mlflow_repository.check_artifact_exists('run_id', 'test_artifact', metadata['metadata'])
|
||||
|
||||
assert result is False
|
||||
mlflow_repository.client.list_artifacts.assert_called_once_with('run_id')
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@patch('laborious.utils.repository.model_repository.path')
|
||||
@patch('laborious.utils.repository.model_repository.rmtree')
|
||||
@@ -275,6 +295,56 @@ async def test_download_artifacts_error(makedirs, rmtree, path, mlflow_repositor
|
||||
mlflow_repository.observe_lag.assert_not_called()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@patch('laborious.utils.repository.model_repository.mlflow')
|
||||
@patch('laborious.utils.repository.model_repository.pd')
|
||||
@patch('laborious.utils.repository.model_repository.StringIO')
|
||||
async def test_load_artifact_dataframe_success(StringIO, pd, mlflow, mlflow_repository):
|
||||
mlflow_repository.get_model_run_id = MagicMock(return_value='run_id')
|
||||
mlflow_repository.check_artifact_exists = MagicMock(return_value=True)
|
||||
|
||||
mlflow.artifacts.load_text.return_value = 'col1,col2\n1,2\n3,4'
|
||||
|
||||
result = await mlflow_repository.load_artifact_dataframe('model_name', 'artifact_path', metadata['metadata'])
|
||||
|
||||
mlflow_repository.get_model_run_id.assert_called_once_with(model_name='model_name', stage='Production')
|
||||
mlflow_repository.check_artifact_exists.assert_called_once_with('run_id', 'artifact_path', metadata['metadata'])
|
||||
mlflow.artifacts.load_text.assert_called_once_with('runs:/run_id/artifact_path')
|
||||
assert result == pd.read_csv.return_value
|
||||
mlflow_repository.observe_lag.assert_called_once_with(ANY, metrics.MODEL_READ_LAG, ANY)
|
||||
mlflow_repository.emit_metric.assert_called_once_with(
|
||||
metric_object=metrics.MODEL_READ_COUNT, tags=ANY
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_load_artifact_dataframe_not_exists(mlflow_repository):
|
||||
mlflow_repository.get_model_run_id = MagicMock(return_value='run_id')
|
||||
mlflow_repository.check_artifact_exists = MagicMock(return_value=False)
|
||||
|
||||
result = await mlflow_repository.load_artifact_dataframe('model_name', 'artifact_path', metadata['metadata'])
|
||||
|
||||
assert result is None
|
||||
mlflow_repository.get_model_run_id.assert_called_once_with(model_name='model_name', stage='Production')
|
||||
mlflow_repository.check_artifact_exists.assert_called_once_with('run_id', 'artifact_path', metadata['metadata'])
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@patch('laborious.utils.repository.model_repository.mlflow')
|
||||
async def test_load_artifact_dataframe_error(mlflow, mlflow_repository):
|
||||
mlflow_repository.get_model_run_id = MagicMock(return_value='run_id')
|
||||
mlflow_repository.check_artifact_exists = MagicMock(return_value=True)
|
||||
mlflow.artifacts.load_text.side_effect = ValueError('error')
|
||||
|
||||
with pytest.raises(ValueError):
|
||||
await mlflow_repository.load_artifact_dataframe('model_name', 'artifact_path', metadata['metadata'])
|
||||
|
||||
mlflow_repository.emit_metric.assert_called_once_with(
|
||||
metric_object=metrics.MODEL_READ_ERROR_COUNT, tags=ANY
|
||||
)
|
||||
mlflow_repository.observe_lag.assert_not_called()
|
||||
|
||||
|
||||
def test_get_experiment_error(mlflow, mlflow_repository):
|
||||
mlflow.get_experiment_by_name.return_value = None
|
||||
|
||||
@@ -719,6 +789,7 @@ async def test_fit_models_not_df_target_name_none_and_not_in_model(
|
||||
mlflow_repository.detect_and_parse_datetime_index = MagicMock(
|
||||
return_value=MagicMock(drop_duplicates=MagicMock(return_value=MagicMock(columns=[])))
|
||||
)
|
||||
mlflow_repository.get_prediction_data = MagicMock(return_value=DataFrame())
|
||||
|
||||
data = MagicMock()
|
||||
|
||||
@@ -775,9 +846,14 @@ async def test_fit_models_not_df_target_name_none_and_not_in_model(
|
||||
|
||||
prediction_model.fit.assert_called_once_with(pd_merge.return_value)
|
||||
|
||||
mlflow_repository.get_prediction_data.assert_called_once_with(
|
||||
prediction_model, pd_merge.return_value, data_model.fit.return_value.target_variable
|
||||
)
|
||||
|
||||
assert output == {
|
||||
'prediction_model': {'model': prediction_model, 'artifact_path': 'artifact_path'},
|
||||
'data_model': {'model': data_model.fit.return_value, 'artifact_path': 'artifact_path'},
|
||||
'prediction_data': mlflow_repository.get_prediction_data.return_value,
|
||||
}
|
||||
|
||||
|
||||
@@ -801,6 +877,7 @@ async def test_fit_models_df_target_name_not_none_and_in_model(
|
||||
drop_duplicates=MagicMock(return_value=MagicMock(columns=['feat_1']))
|
||||
)
|
||||
)
|
||||
mlflow_repository.get_prediction_data = MagicMock(return_value=DataFrame())
|
||||
|
||||
data = MagicMock()
|
||||
|
||||
@@ -855,9 +932,14 @@ async def test_fit_models_df_target_name_not_none_and_in_model(
|
||||
|
||||
prediction_model.fit.assert_called_once_with(transformed_data)
|
||||
|
||||
mlflow_repository.get_prediction_data.assert_called_once_with(
|
||||
prediction_model, transformed_data, 'feat_1'
|
||||
)
|
||||
|
||||
assert output == {
|
||||
'prediction_model': {'model': prediction_model, 'artifact_path': 'artifact_path'},
|
||||
'data_model': {'model': data_model, 'artifact_path': 'artifact_path'},
|
||||
'prediction_data': mlflow_repository.get_prediction_data.return_value,
|
||||
}
|
||||
|
||||
|
||||
@@ -878,7 +960,8 @@ async def test_log_model_sklearn(mlflow, mlflow_repository):
|
||||
@patch('laborious.utils.repository.model_repository.path')
|
||||
@pytest.mark.asyncio
|
||||
async def test_log_model_pyfunc(path, mlflow, mlflow_repository):
|
||||
model_data = {'model': MagicMock(), 'artifact_path': 'artifact_path'}
|
||||
model_mock = MagicMock()
|
||||
model_data = {'model': model_mock, 'artifact_path': 'artifact_path'}
|
||||
await mlflow_repository.log_model(
|
||||
model_data, 'pyfunc', 'prediction_model', metadata['metadata']
|
||||
)
|
||||
@@ -887,7 +970,7 @@ async def test_log_model_pyfunc(path, mlflow, mlflow_repository):
|
||||
|
||||
path.join.assert_called_once_with('artifact_path', 'code', 'utils')
|
||||
|
||||
model_data['model'].store_model.assert_called_once_with(
|
||||
model_mock.store_model.assert_called_once_with(
|
||||
artifact_path='prediction_model', code_path=[path.join.return_value], to_disk=False
|
||||
)
|
||||
|
||||
@@ -934,9 +1017,11 @@ async def test_create_new_experiment(
|
||||
):
|
||||
model_name = 'model_name'
|
||||
data = MagicMock()
|
||||
prediction_data = MagicMock(spec=DataFrame)
|
||||
retrain_data = {
|
||||
'prediction_model': {'model': MagicMock(), 'artifact_path': 'artifact_path'},
|
||||
'data_model': {'model': MagicMock(), 'artifact_path': 'artifact_path'},
|
||||
'prediction_data': prediction_data,
|
||||
}
|
||||
|
||||
mlflow_repository.get_model_params = MagicMock(
|
||||
@@ -971,7 +1056,11 @@ async def test_create_new_experiment(
|
||||
mlflow_repository.get_experiment.return_value.name
|
||||
)
|
||||
|
||||
data.to_csv.assert_called_once_with('./tmp/artifacts/model_name/retrain_data.csv', index=True)
|
||||
data.to_csv.assert_called_once_with('./tmp/artifacts/model_name/retrain_data.csv', index=False)
|
||||
# Verify prediction_data.to_csv was called with correct arguments
|
||||
prediction_data.to_csv.assert_called_once()
|
||||
assert prediction_data.to_csv.call_args[0][0] == './tmp/artifacts/model_name/evaluation_data.csv'
|
||||
assert prediction_data.to_csv.call_args[1]['index'] is False
|
||||
|
||||
mlflow.start_run.assert_called_once_with(
|
||||
experiment_id=mlflow_repository.get_experiment.return_value.experiment_id,
|
||||
@@ -1000,7 +1089,10 @@ async def test_create_new_experiment(
|
||||
}
|
||||
)
|
||||
|
||||
mlflow.log_artifact.assert_called_once_with('./tmp/artifacts/model_name/retrain_data.csv')
|
||||
mlflow.log_artifact.assert_has_calls([
|
||||
call('./tmp/artifacts/model_name/retrain_data.csv'),
|
||||
call('./tmp/artifacts/model_name/evaluation_data.csv'),
|
||||
])
|
||||
|
||||
force_memory_release.assert_called_once_with(mlflow_repository.logger)
|
||||
|
||||
@@ -1026,9 +1118,11 @@ async def test_create_new_experiment_error(
|
||||
mlflow.start_run.side_effect = ValueError('error')
|
||||
model_name = 'model_name'
|
||||
data = MagicMock()
|
||||
prediction_data = MagicMock(spec=DataFrame)
|
||||
retrain_data = {
|
||||
'prediction_model': {'model': MagicMock(), 'artifact_path': 'artifact_path'},
|
||||
'data_model': {'model': MagicMock(), 'artifact_path': 'artifact_path'},
|
||||
'prediction_data': prediction_data,
|
||||
}
|
||||
|
||||
mlflow_repository.get_model_params = MagicMock(
|
||||
@@ -1371,3 +1465,33 @@ async def test_update_production_model(mlflow_repository):
|
||||
'mlflow_run_id': '0',
|
||||
'mlflow_experiment_id': '0',
|
||||
}
|
||||
|
||||
|
||||
def test_get_prediction_data_dataframe(mlflow_repository):
|
||||
prediction_model = MagicMock()
|
||||
retrain_dataset = DataFrame({'feat_1': [1, 2], 'target': [3, 4]}, index=['idx1', 'idx2'])
|
||||
prediction_model.predict.return_value = DataFrame({'pred': [5, 6]}, index=['idx1', 'idx2'])
|
||||
target_name = 'target'
|
||||
|
||||
result = mlflow_repository.get_prediction_data(prediction_model, retrain_dataset, target_name)
|
||||
|
||||
prediction_model.predict.assert_called_once_with(retrain_dataset)
|
||||
assert 'prediction' in result.columns
|
||||
assert 'target' in result.columns
|
||||
assert 'timestamp' in result.columns
|
||||
assert result.index.tolist() == [0, 1]
|
||||
|
||||
|
||||
def test_get_prediction_data_array(mlflow_repository):
|
||||
prediction_model = MagicMock()
|
||||
retrain_dataset = DataFrame({'feat_1': [1, 2], 'target': [3, 4]}, index=['idx1', 'idx2'])
|
||||
prediction_model.predict.return_value = [5, 6]
|
||||
target_name = 'target'
|
||||
|
||||
result = mlflow_repository.get_prediction_data(prediction_model, retrain_dataset, target_name)
|
||||
|
||||
prediction_model.predict.assert_called_once_with(retrain_dataset)
|
||||
assert 'prediction' in result.columns
|
||||
assert 'target' in result.columns
|
||||
assert 'timestamp' in result.columns
|
||||
assert result.index.tolist() == [0, 1]
|
||||
|
||||
@@ -31,6 +31,7 @@ async def test_run(workflow_mock, prediction_process):
|
||||
'data': {'test': 'data'},
|
||||
'schema': 'test_schema',
|
||||
'table_name': 'test_table',
|
||||
'transform_table_name': 'test_transform_table',
|
||||
'model_id': 1,
|
||||
'input_filters': {'test': 'filter'},
|
||||
'mlflow_transform_filters': {'test': 'filter'},
|
||||
@@ -175,6 +176,7 @@ async def test_run(workflow_mock, prediction_process):
|
||||
'metadata': metadata,
|
||||
'path_flag': 'continue',
|
||||
'data': 'predicted_data',
|
||||
'transformed_data': 'transformed_data',
|
||||
'prediction_confidence': 0.95,
|
||||
'timestamp': '2024-01-01',
|
||||
'model_id': 1,
|
||||
@@ -183,6 +185,7 @@ async def test_run(workflow_mock, prediction_process):
|
||||
'opc_output_config': input_data['opc_output_config'],
|
||||
'schema': input_data['schema'],
|
||||
'table_name': input_data['table_name'],
|
||||
'transform_table_name': input_data['transform_table_name'],
|
||||
'comment': 'Error',
|
||||
'prediction_store_policy': input_data['prediction_store_policy'],
|
||||
},
|
||||
@@ -199,6 +202,7 @@ async def test_run_stop_at_input_gate(workflow_mock, prediction_process):
|
||||
'data': {'test': 'data'},
|
||||
'schema': 'test_schema',
|
||||
'table_name': 'test_table',
|
||||
'transform_table_name': 'test_transform_table',
|
||||
'model_id': 1,
|
||||
'input_filters': {'test': 'filter'},
|
||||
'mlflow_transform_filters': {'test': 'filter'},
|
||||
@@ -257,6 +261,7 @@ async def test_run_stop_at_first_mlflow_response_gate(workflow_mock, prediction_
|
||||
'data': {'test': 'data'},
|
||||
'schema': 'test_schema',
|
||||
'table_name': 'test_table',
|
||||
'transform_table_name': 'test_transform_table',
|
||||
'model_id': 1,
|
||||
'input_filters': {'test': 'filter'},
|
||||
'mlflow_transform_filters': {'test': 'filter'},
|
||||
@@ -352,6 +357,7 @@ async def test_run_stop_at_mlflow_content_gate(workflow_mock, prediction_process
|
||||
'data': {'test': 'data'},
|
||||
'schema': 'test_schema',
|
||||
'table_name': 'test_table',
|
||||
'transform_table_name': 'test_transform_table',
|
||||
'model_id': 1,
|
||||
'input_filters': {'test': 'filter'},
|
||||
'mlflow_transform_filters': {'test': 'filter'},
|
||||
@@ -467,6 +473,7 @@ async def test_run_stop_at_mlflow_last_response_gate(workflow_mock, prediction_p
|
||||
'data': {'test': 'data'},
|
||||
'schema': 'test_schema',
|
||||
'table_name': 'test_table',
|
||||
'transform_table_name': 'test_transform_table',
|
||||
'model_id': 1,
|
||||
'input_filters': {'test': 'filter'},
|
||||
'mlflow_transform_filters': {'test': 'filter'},
|
||||
@@ -626,6 +633,7 @@ async def test_path_flag_handler_stop(workflow_mock, prediction_process):
|
||||
'metadata': metadata,
|
||||
'schema': schema,
|
||||
'table_name': table_name,
|
||||
'transform_table_name': 'test_transform_table',
|
||||
'model_id': model,
|
||||
'last_timestamp': last_timestamp,
|
||||
'model_name': model_name,
|
||||
@@ -664,6 +672,7 @@ async def test_path_flag_handler_repeat(workflow_mock, prediction_process):
|
||||
'metadata': metadata,
|
||||
'schema': schema,
|
||||
'table_name': table_name,
|
||||
'transform_table_name': 'test_transform_table',
|
||||
'model_id': model,
|
||||
'last_timestamp': last_timestamp,
|
||||
'model_name': model_name,
|
||||
@@ -714,6 +723,7 @@ async def test_path_flag_handler_continue(workflow_mock, prediction_process):
|
||||
'metadata': metadata,
|
||||
'schema': schema,
|
||||
'table_name': table_name,
|
||||
'transform_table_name': 'test_transform_table',
|
||||
'model_id': model,
|
||||
'last_timestamp': last_timestamp,
|
||||
'model_name': model_name,
|
||||
@@ -742,6 +752,7 @@ async def test_path_flag_handler_continue(workflow_mock, prediction_process):
|
||||
'model_config': model_config,
|
||||
'schema': schema,
|
||||
'table_name': table_name,
|
||||
'transform_table_name': 'test_transform_table',
|
||||
'comment': 'Prediction Process',
|
||||
'opc_output_config': {'test': 'config'},
|
||||
'prediction_store_policy': prediction_store_policy,
|
||||
@@ -771,6 +782,7 @@ async def test_path_flag_handler_unknown(workflow_mock, prediction_process):
|
||||
**metadata,
|
||||
'schema': schema,
|
||||
'table_name': table_name,
|
||||
'transform_table_name': 'test_transform_table',
|
||||
'model_id': model,
|
||||
'last_timestamp': last_timestamp,
|
||||
'model_name': model_name,
|
||||
|
||||
247
tests/laborious/workflows/test_drift.py
Normal file
247
tests/laborious/workflows/test_drift.py
Normal file
@@ -0,0 +1,247 @@
|
||||
from unittest.mock import ANY, AsyncMock, call, patch
|
||||
|
||||
from pytest import fixture, mark
|
||||
|
||||
from laborious.activities.activities import Activities
|
||||
from laborious.workflows.drift import Drift
|
||||
from sientia_do.temporal.constants import DATETIME_FORMAT_WITH_TZ
|
||||
|
||||
|
||||
@fixture
|
||||
def drift() -> Drift:
|
||||
return Drift()
|
||||
|
||||
|
||||
metadata = {
|
||||
'metadata': {
|
||||
'model_id': 'test_model_id',
|
||||
'model_name': 'test_model',
|
||||
'workflow_name': 'drift',
|
||||
'schedule_name': 'test_schedule',
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
@mark.asyncio
|
||||
@patch('laborious.workflows.drift.workflow', new_callable=AsyncMock)
|
||||
async def test_run(workflow_mock: AsyncMock, drift: Drift):
|
||||
# Arrange
|
||||
input_data = {
|
||||
'schedule_name': 'test_schedule',
|
||||
'model_name': 'test_model',
|
||||
'model_id': 'test_model_id',
|
||||
'schema': 'test_schema',
|
||||
'source_table_name': 'test_source_table',
|
||||
'target_table_name': 'test_target_table',
|
||||
'interval': 60,
|
||||
'target_name': 'test_target',
|
||||
'drift_metrics': ['psi', 'ks'],
|
||||
'chunk_period': 'hour',
|
||||
}
|
||||
|
||||
target_data = {'data': 'test_target_data'}
|
||||
reference_data = {'data': 'test_reference_data'}
|
||||
drift_data = {'drift': 'test_drift_data'}
|
||||
|
||||
workflow_mock.start_local_activity_method.side_effect = [
|
||||
target_data, reference_data
|
||||
]
|
||||
|
||||
workflow_mock.execute_local_activity_method.return_value = drift_data
|
||||
workflow_mock.execute_activity_method = AsyncMock()
|
||||
|
||||
# Act
|
||||
await drift.run(input_data)
|
||||
|
||||
# Assert - Check start_local_activity_method calls
|
||||
expected_gathering_query = f"""
|
||||
SELECT *
|
||||
FROM {input_data['schema']}.{input_data['source_table_name']}
|
||||
WHERE
|
||||
model_id = {input_data['model_id']} AND
|
||||
timestamp > NOW() - INTERVAL '{input_data['interval']} minutes'
|
||||
ORDER BY timestamp ASC
|
||||
"""
|
||||
|
||||
workflow_mock.start_local_activity_method.assert_has_calls(
|
||||
[
|
||||
call(
|
||||
Activities.load_custom_query,
|
||||
{
|
||||
**metadata,
|
||||
'query': expected_gathering_query,
|
||||
'datetime_columns': ['timestamp', 'created_at'],
|
||||
},
|
||||
retry_policy=ANY,
|
||||
start_to_close_timeout=ANY,
|
||||
),
|
||||
call(
|
||||
Activities.get_reference_data,
|
||||
{
|
||||
**metadata,
|
||||
'model_name': input_data['model_name'],
|
||||
},
|
||||
retry_policy=ANY,
|
||||
start_to_close_timeout=ANY,
|
||||
),
|
||||
]
|
||||
)
|
||||
|
||||
# Assert - Check calculate_drift call
|
||||
workflow_mock.execute_local_activity_method.assert_called_once_with(
|
||||
Activities.calculate_drift,
|
||||
{
|
||||
**metadata,
|
||||
'target_data': target_data,
|
||||
'reference_data': reference_data,
|
||||
'model_name': input_data['model_name'],
|
||||
'model_id': input_data['model_id'],
|
||||
'target_name': input_data['target_name'],
|
||||
'drift_metrics': input_data['drift_metrics'],
|
||||
'chunk_period': input_data['chunk_period'],
|
||||
},
|
||||
retry_policy=ANY,
|
||||
start_to_close_timeout=ANY,
|
||||
)
|
||||
|
||||
# Assert - Check export_data_to_postgres call
|
||||
workflow_mock.execute_activity_method.assert_called_once_with(
|
||||
Activities.export_data_to_postgres,
|
||||
{
|
||||
**metadata,
|
||||
'data': drift_data,
|
||||
'schema': input_data['schema'],
|
||||
'table_name': input_data['target_table_name'],
|
||||
'timestamp_conversion': {'column': 'timestamp', 'format': DATETIME_FORMAT_WITH_TZ},
|
||||
},
|
||||
retry_policy=ANY,
|
||||
start_to_close_timeout=ANY,
|
||||
)
|
||||
|
||||
|
||||
@mark.asyncio
|
||||
@patch('laborious.workflows.drift.workflow', new_callable=AsyncMock)
|
||||
async def test_run_empty_target_data(workflow_mock: AsyncMock, drift: Drift):
|
||||
# Arrange
|
||||
input_data = {
|
||||
'schedule_name': 'test_schedule',
|
||||
'model_name': 'test_model',
|
||||
'model_id': 'test_model_id',
|
||||
'schema': 'test_schema',
|
||||
'source_table_name': 'test_source_table',
|
||||
'target_table_name': 'test_target_table',
|
||||
'interval': 60,
|
||||
'target_name': 'test_target',
|
||||
'drift_metrics': ['psi', 'ks'],
|
||||
}
|
||||
|
||||
target_data = None
|
||||
reference_data = {'data': 'test_reference_data'}
|
||||
|
||||
workflow_mock.start_local_activity_method.side_effect = [target_data, reference_data]
|
||||
|
||||
workflow_mock.execute_local_activity_method = AsyncMock()
|
||||
workflow_mock.execute_activity_method = AsyncMock()
|
||||
|
||||
# Act
|
||||
await drift.run(input_data)
|
||||
|
||||
# Assert - Should not call calculate_drift or export
|
||||
workflow_mock.execute_local_activity_method.assert_not_called()
|
||||
workflow_mock.execute_activity_method.assert_not_called()
|
||||
|
||||
|
||||
@mark.asyncio
|
||||
@patch('laborious.workflows.drift.workflow', new_callable=AsyncMock)
|
||||
async def test_run_empty_drift_data(workflow_mock: AsyncMock, drift: Drift):
|
||||
# Arrange
|
||||
input_data = {
|
||||
'schedule_name': 'test_schedule',
|
||||
'model_name': 'test_model',
|
||||
'model_id': 'test_model_id',
|
||||
'schema': 'test_schema',
|
||||
'source_table_name': 'test_source_table',
|
||||
'target_table_name': 'test_target_table',
|
||||
'interval': 60,
|
||||
'target_name': 'test_target',
|
||||
'drift_metrics': ['psi', 'ks'],
|
||||
}
|
||||
|
||||
target_data = {'data': 'test_target_data'}
|
||||
reference_data = {'data': 'test_reference_data'}
|
||||
drift_data = None
|
||||
|
||||
workflow_mock.start_local_activity_method.side_effect = [target_data, reference_data]
|
||||
|
||||
workflow_mock.execute_local_activity_method.return_value = drift_data
|
||||
workflow_mock.execute_activity_method = AsyncMock()
|
||||
|
||||
# Act
|
||||
await drift.run(input_data)
|
||||
|
||||
# Assert - Should call calculate_drift but not export
|
||||
workflow_mock.execute_local_activity_method.assert_called_once_with(
|
||||
Activities.calculate_drift,
|
||||
{
|
||||
**metadata,
|
||||
'target_data': target_data,
|
||||
'reference_data': reference_data,
|
||||
'model_name': input_data['model_name'],
|
||||
'model_id': input_data['model_id'],
|
||||
'target_name': input_data['target_name'],
|
||||
'drift_metrics': input_data['drift_metrics'],
|
||||
'chunk_period': input_data.get('chunk_period', 'min'),
|
||||
},
|
||||
retry_policy=ANY,
|
||||
start_to_close_timeout=ANY,
|
||||
)
|
||||
|
||||
workflow_mock.execute_activity_method.assert_not_called()
|
||||
|
||||
|
||||
@mark.asyncio
|
||||
@patch('laborious.workflows.drift.workflow', new_callable=AsyncMock)
|
||||
async def test_run_default_chunk_period(workflow_mock: AsyncMock, drift: Drift):
|
||||
# Arrange
|
||||
input_data = {
|
||||
'schedule_name': 'test_schedule',
|
||||
'model_name': 'test_model',
|
||||
'model_id': 'test_model_id',
|
||||
'schema': 'test_schema',
|
||||
'source_table_name': 'test_source_table',
|
||||
'target_table_name': 'test_target_table',
|
||||
'interval': 60,
|
||||
'target_name': 'test_target',
|
||||
'drift_metrics': ['psi', 'ks'],
|
||||
# chunk_period not provided, should default to 'min'
|
||||
}
|
||||
|
||||
target_data = {'data': 'test_target_data'}
|
||||
reference_data = {'data': 'test_reference_data'}
|
||||
drift_data = {'drift': 'test_drift_data'}
|
||||
|
||||
workflow_mock.start_local_activity_method.side_effect = [target_data, reference_data]
|
||||
|
||||
workflow_mock.execute_local_activity_method.return_value = drift_data
|
||||
workflow_mock.execute_activity_method = AsyncMock()
|
||||
|
||||
# Act
|
||||
await drift.run(input_data)
|
||||
|
||||
# Assert - Check calculate_drift call with default chunk_period
|
||||
workflow_mock.execute_local_activity_method.assert_called_once_with(
|
||||
Activities.calculate_drift,
|
||||
{
|
||||
**metadata,
|
||||
'target_data': target_data,
|
||||
'reference_data': reference_data,
|
||||
'model_name': input_data['model_name'],
|
||||
'model_id': input_data['model_id'],
|
||||
'target_name': input_data['target_name'],
|
||||
'drift_metrics': input_data['drift_metrics'],
|
||||
'chunk_period': 'min', # Default value
|
||||
},
|
||||
retry_policy=ANY,
|
||||
start_to_close_timeout=ANY,
|
||||
)
|
||||
|
||||
@@ -32,6 +32,7 @@ async def test_run(workflow_mock: AsyncMock, predictions_batch: PredictionsBatch
|
||||
'query': 'SELECT * FROM test',
|
||||
'schema': 'test_schema',
|
||||
'table_name': 'test_table',
|
||||
'transform_table_name': 'test_transform_table',
|
||||
'opc_output_config': 'test_opc_output_config',
|
||||
'datetime_columns': ['timestamp', 'created_at'],
|
||||
'prediction_store_policy': 'erl:1',
|
||||
@@ -59,6 +60,7 @@ async def test_run(workflow_mock: AsyncMock, predictions_batch: PredictionsBatch
|
||||
'data': {'data': 'test_data'},
|
||||
'schema': input_data['schema'],
|
||||
'table_name': input_data['table_name'],
|
||||
'transform_table_name': input_data['transform_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'}}),
|
||||
@@ -71,7 +73,7 @@ async def test_run(workflow_mock: AsyncMock, predictions_batch: PredictionsBatch
|
||||
'model_config': input_data.get('model_config', {}),
|
||||
'path_priority': input_data.get('path_priority', ['STOP', 'CONTINUE', 'REPEAT']),
|
||||
'opc_output_config': input_data.get('opc_output_config', {}),
|
||||
'prediction_store_policy': input_data.get('prediction_store_policy', 'erl:1'),
|
||||
'prediction_store_policy': input_data.get('prediction_store_policy', 'lts:1'),
|
||||
}
|
||||
|
||||
workflow_mock.execute_child_workflow.assert_has_calls(
|
||||
|
||||
Reference in New Issue
Block a user