SIENTIAPDE-1646

Update dependencies and refactor MLFlow activities

- Replaced direct GitHub dependencies in `requirements.txt` with specific versioned packages for `sientia_do` and `sientia_model`.
- Refactored imports in `activities.py` to streamline the code structure.
- Enhanced the `MLFlow` class in `mlflow.py` by introducing a method to resolve model aliases, improving flexibility in model lookups.
- Simplified shutdown logic in `worker.py` for better readability.
- Added new tests for MLFlow activities and improved existing test coverage for data handling and model retraining processes.
This commit is contained in:
vitor-aignosi
2026-05-06 15:31:25 -03:00
parent 1ce8b9d3a7
commit 424be007ef
12 changed files with 625 additions and 29 deletions

View File

@@ -10,14 +10,13 @@ with workflow.unsafe.imports_passed_through():
from sientia_model.model_repository.mlflow_repository import SientiaMLflowRepository from sientia_model.model_repository.mlflow_repository import SientiaMLflowRepository
from sientia_model.model_repository.plugin_store import PluginStore from sientia_model.model_repository.plugin_store import PluginStore
from laborious.utils.connectors_config import build_mlflow_config
from laborious.activities.api import API from laborious.activities.api import API
from laborious.activities.gates import Gates from laborious.activities.gates import Gates
from laborious.activities.mlflow import MLFlow from laborious.activities.mlflow import MLFlow
from laborious.activities.model_metrics import ModelMetrics from laborious.activities.model_metrics import ModelMetrics
from laborious.activities.opc import OPC from laborious.activities.opc import OPC
from laborious.activities.storage import Storage from laborious.activities.storage import Storage
from laborious.utils.connectors_config import build_mlflow_config
class Activities(Storage, MLFlow, Gates, OPC, ModelMetrics, API): class Activities(Storage, MLFlow, Gates, OPC, ModelMetrics, API):

View File

@@ -12,7 +12,6 @@ with workflow.unsafe.imports_passed_through():
import numpy as np import numpy as np
import pandas as pd import pandas as pd
from pandas import DataFrame, to_datetime from pandas import DataFrame, to_datetime
from sklearn.model_selection import train_test_split
from sientia_do.notifications.handlers import CoreNotificationHandler as NotificationHandler from sientia_do.notifications.handlers import CoreNotificationHandler as NotificationHandler
from sientia_do.notifications.models import NotificationLevel from sientia_do.notifications.models import NotificationLevel
from sientia_do.observability.logger import Logger from sientia_do.observability.logger import Logger
@@ -27,6 +26,7 @@ with workflow.unsafe.imports_passed_through():
from sientia_do.utils.formatters import create_sample_dict from sientia_do.utils.formatters import create_sample_dict
from sientia_model.model_repository.mlflow_repository import SientiaMLflowRepository from sientia_model.model_repository.mlflow_repository import SientiaMLflowRepository
from sientia_model.model_repository.plugin_store import PluginStore from sientia_model.model_repository.plugin_store import PluginStore
from sklearn.model_selection import train_test_split
from laborious.utils.dataframe_debug import build_dataframe_debug_message from laborious.utils.dataframe_debug import build_dataframe_debug_message
from laborious.utils.models.minio_dataframe_payload import MinioDataFramePayload from laborious.utils.models.minio_dataframe_payload import MinioDataFramePayload
@@ -53,6 +53,7 @@ class MLFlow(MinioManager):
""" """
_MAX_DEBUG_DATAFRAME_ROWS = 100 _MAX_DEBUG_DATAFRAME_ROWS = 100
_DEFAULT_MODEL_ALIAS = 'production'
def __init__( def __init__(
self, self,
@@ -191,6 +192,21 @@ class MLFlow(MinioManager):
latest = max(versions, key=lambda v: int(v.version)) latest = max(versions, key=lambda v: int(v.version))
return str(latest.version) return str(latest.version)
def _resolve_model_alias(self, model_config: dict[str, Any] | None = None) -> str:
"""
Resolve which MLflow alias should be used for model lookup/promotion.
Args:
- model_config: Optional model configuration that may include ``alias``.
Return:
str: Alias name trimmed and normalized; defaults to ``production``.
"""
if not model_config:
return self._DEFAULT_MODEL_ALIAS
alias = str(model_config.get('alias', self._DEFAULT_MODEL_ALIAS)).strip()
return alias or self._DEFAULT_MODEL_ALIAS
@activity.defn(name='request_transform') @activity.defn(name='request_transform')
async def request_transform(self, input_data: dict[str, Any]) -> MinioDataFramePayload: async def request_transform(self, input_data: dict[str, Any]) -> MinioDataFramePayload:
""" """
@@ -217,6 +233,7 @@ class MLFlow(MinioManager):
model_name = input_data['model_name'] model_name = input_data['model_name']
model_config = input_data.get('model_config', {}) model_config = input_data.get('model_config', {})
model_alias = self._resolve_model_alias(model_config)
self._debug_dataframe('Raw input data:', data, metadata) self._debug_dataframe('Raw input data:', data, metadata)
@@ -238,7 +255,7 @@ class MLFlow(MinioManager):
try: try:
wrapper = self.mlflow_repository.get_cached_model( wrapper = self.mlflow_repository.get_cached_model(
model_name=model_name, model_name=model_name,
alias='production', alias=model_alias,
retention_minutes=model_config.get('retention_minutes', 0), retention_minutes=model_config.get('retention_minutes', 0),
metadata=metadata, metadata=metadata,
) )
@@ -317,6 +334,7 @@ class MLFlow(MinioManager):
model_name = input_data['model_name'] model_name = input_data['model_name']
model_config = input_data.get('model_config', {}) model_config = input_data.get('model_config', {})
model_alias = self._resolve_model_alias(model_config)
self._debug_dataframe('Input data for prediction:', data, metadata) self._debug_dataframe('Input data for prediction:', data, metadata)
@@ -332,7 +350,7 @@ class MLFlow(MinioManager):
try: try:
wrapper = self.mlflow_repository.get_cached_model( wrapper = self.mlflow_repository.get_cached_model(
model_name=model_name, model_name=model_name,
alias='production', alias=model_alias,
retention_minutes=model_config.get('retention_minutes', 0), retention_minutes=model_config.get('retention_minutes', 0),
metadata=metadata, metadata=metadata,
) )
@@ -490,15 +508,16 @@ class MLFlow(MinioManager):
} }
try: try:
model_alias = self._resolve_model_alias(model_config)
mv_src = self.mlflow_repository._client.get_model_version_by_alias( mv_src = self.mlflow_repository._client.get_model_version_by_alias(
name=model_name, name=model_name,
alias='production', alias=model_alias,
) )
source_run_id = mv_src.run_id source_run_id = mv_src.run_id
wrapper = self.mlflow_repository.get_cached_model( wrapper = self.mlflow_repository.get_cached_model(
model_name=model_name, model_name=model_name,
alias='production', alias=model_alias,
retention_minutes=0, retention_minutes=0,
metadata=metadata, metadata=metadata,
) )
@@ -595,10 +614,11 @@ class MLFlow(MinioManager):
version = self._resolve_model_version_for_run(run_id) version = self._resolve_model_version_for_run(run_id)
promote_alias = self._resolve_model_alias(input_data.get('model_config'))
self.mlflow_repository.promote_to_alias( self.mlflow_repository.promote_to_alias(
model_name=model_name, model_name=model_name,
version=version, version=version,
alias='production', alias=promote_alias,
metadata=metadata, metadata=metadata,
) )
@@ -643,9 +663,10 @@ class MLFlow(MinioManager):
model_name = input_data['model_name'] model_name = input_data['model_name']
try: try:
model_alias = self._resolve_model_alias(input_data.get('model_config'))
mv = self.mlflow_repository._client.get_model_version_by_alias( mv = self.mlflow_repository._client.get_model_version_by_alias(
name=model_name, name=model_name,
alias='production', alias=model_alias,
) )
run_id = mv.run_id run_id = mv.run_id

View File

@@ -264,9 +264,7 @@ async def main():
logger.custom_error(f'An unhandled exception occurred: {e}', metadata) logger.custom_error(f'An unhandled exception occurred: {e}', metadata)
exit_code = 1 exit_code = 1
finally: finally:
if notification_handler:
notification_handler.shutdown() notification_handler.shutdown()
if activities:
await activities.shutdown() await activities.shutdown()
metrics.APP_UP.labels(pod_id=POD_ID).set(0) # Mark app as DOWN metrics.APP_UP.labels(pod_id=POD_ID).set(0) # Mark app as DOWN
sys.exit(exit_code) sys.exit(exit_code)

19
requirements-local.txt Normal file
View File

@@ -0,0 +1,19 @@
temporalio
psycopg2-binary
sqlalchemy
asyncua
redis
git+ssh://git@github.com/Aignosi/sientia-dataops-library.git@1.12.0
#git+ssh://git@github.com/Aignosi/sientia-model-library.git@0.8.2
/home/grezewave/Documents/projects/sientia/sientia-model-library
prometheus-client
botocore
boto3
s3fs
pyarrow
kaleido
hyperopt
shap
pycurl
scipy<1.14.0
scikit-learn==1.5.2

View File

@@ -3,8 +3,8 @@ psycopg2-binary
sqlalchemy sqlalchemy
asyncua asyncua
redis redis
git+ssh://git@github.com/Aignosi/sientia-dataops-library.git@1.12.0 sientia_do==1.12.0
git+ssh://git@github.com/Aignosi/sientia-model-library.git@0.8.1 sientia_model==0.8.2
prometheus-client prometheus-client
botocore botocore
boto3 boto3

View File

@@ -229,3 +229,74 @@ async def test_shutdown(
mock_gates_init.close.assert_called_once() mock_gates_init.close.assert_called_once()
mock_model_metrics_init.close.assert_called_once() mock_model_metrics_init.close.assert_called_once()
mock_api_init.close.assert_called_once() mock_api_init.close.assert_called_once()
@patch('laborious.activities.activities.SientiaMLflowRepository')
@patch('laborious.activities.activities.build_mlflow_config')
@patch('laborious.activities.activities.Storage.__init__')
@patch('laborious.activities.activities.MLFlow.__init__')
@patch('laborious.activities.activities.OPC.__init__')
@patch('laborious.activities.activities.Gates.__init__')
@patch('laborious.activities.activities.ModelMetrics.__init__')
@patch('laborious.activities.activities.API.__init__')
@patch('laborious.activities.activities.MinioRepository')
@patch('laborious.activities.activities.MetricsController')
def test___init___builds_mlflow_repository_when_not_provided(
mock_metrics_controller,
mock_minio_repository,
_mock_api_init,
_mock_model_metrics_init,
_mock_gates_init,
_mock_opc_init,
_mock_mlflow_init,
_mock_storage_init,
mock_build_mlflow_config,
mock_mlflow_repository_cls,
):
postgres_config = {
'host': 'localhost',
'port': 5432,
'user': 'postgres',
'password': 'postgres',
'dbname': 'postgres',
'min_connections': 1,
'max_connections': 10,
}
minio_config = {
'endpoint_url': 'localhost:9000',
'access_key': 'minio',
'secret_key': 'minio123',
'default_bucket': 'test',
'retention_hours': 24,
'secure': False,
}
opc_config = {'bootstrap_servers': 'localhost:9092', 'polling_time': 1000, 'group_id': 'test'}
pi_web_api_config = {'base_url': 'https://pi', 'auth_type': 'bearer', 'auth_token': 'token'}
logger = MagicMock()
notification_handler = MagicMock()
plugin_store = MagicMock()
mock_build_mlflow_config.return_value = {
'url': 'http://mlflow:80',
'username': 'u',
'password': 'p',
}
Activities(
postgres_config=postgres_config,
plugin_store=plugin_store,
minio_config=minio_config,
opc_config=opc_config,
pi_web_api_config=pi_web_api_config,
logger=logger,
notification_handler=notification_handler,
)
mock_build_mlflow_config.assert_called_once()
mock_mlflow_repository_cls.assert_called_once_with(
host='http://mlflow:80',
username='u',
password='p',
logger=logger,
notification_handler=notification_handler,
metrics_controller=mock_metrics_controller.return_value,
)

View File

@@ -328,6 +328,27 @@ async def test_write_pi_web_api_data_empty_tags(mock_dataframe, api, base_input_
} }
@mark.asyncio
@patch('laborious.activities.api.DataFrame')
async def test_write_pi_web_api_data_updates_confidence_and_comments(
mock_dataframe, api, base_input_data
):
mock_dataframe.return_value = _create_mock_dataframe()
api.pi_web_api_client.write_value.side_effect = [
[{'WebId': 'web_id_1', 'Errors': []}],
[{'WebId': 'web_id_2', 'Errors': []}],
]
with patch.object(
api,
'process_pi_web_api_response',
new=AsyncMock(side_effect=[(0.33, 'PI warning'), (0, '')]),
) as process_mock:
result = await api.write_pi_web_api_data(base_input_data)
assert process_mock.await_count == 2
assert result is not None
@mark.asyncio @mark.asyncio
async def test_close(api): async def test_close(api):
api.close() api.close()

View File

@@ -1,3 +1,4 @@
from datetime import datetime
from unittest.mock import ANY, AsyncMock, MagicMock, patch from unittest.mock import ANY, AsyncMock, MagicMock, patch
import numpy as np import numpy as np
@@ -108,6 +109,47 @@ metadata = {
} }
def test_detect_and_parse_datetime_index_empty(mlflow):
df = pd.DataFrame()
out = mlflow._detect_and_parse_datetime_index(df, metadata['metadata'])
assert out.empty
def test_detect_and_parse_datetime_index_mixed_types_error(mlflow):
idx = pd.Index([pd.Timestamp('2020-01-01', tz='UTC'), 'x'])
df = pd.DataFrame({'a': [1, 2]}, index=idx)
with raises(ValueError):
mlflow._detect_and_parse_datetime_index(df, metadata['metadata'])
def test_detect_and_parse_datetime_index_invalid_string_error(mlflow):
idx = pd.Index(['bad-format'])
df = pd.DataFrame({'a': [1]}, index=idx)
with raises(ValueError):
mlflow._detect_and_parse_datetime_index(df, metadata['metadata'])
def test_detect_and_parse_datetime_index_unsupported_type_error(mlflow):
idx = pd.Index([pd.Period('2020-01', freq='M')])
df = pd.DataFrame({'a': [1]}, index=idx)
with raises(ValueError):
mlflow._detect_and_parse_datetime_index(df, metadata['metadata'])
def test_detect_and_parse_datetime_index_datetime_success(mlflow):
idx = pd.Index([datetime(2020, 1, 1, 0, 0, 0)])
df = pd.DataFrame({'a': [1]}, index=idx)
out = mlflow._detect_and_parse_datetime_index(df, metadata['metadata'])
assert out.index[0].endswith('+0000')
def test_detect_and_parse_datetime_index_timestamp_with_tz_success(mlflow):
idx = pd.DatetimeIndex([pd.Timestamp('2020-01-01 00:00:00', tz='UTC')])
df = pd.DataFrame({'a': [1]}, index=idx)
out = mlflow._detect_and_parse_datetime_index(df, metadata['metadata'])
assert out.index[0].endswith('+0000')
@mark.asyncio @mark.asyncio
@patch( @patch(
'laborious.activities.mlflow.MinioDataFramePayload.from_dataframe', 'laborious.activities.mlflow.MinioDataFramePayload.from_dataframe',
@@ -160,6 +202,41 @@ async def test_request_transform_success(mock_from_dataframe, mlflow):
assert response_data == mock_from_dataframe.return_value assert response_data == mock_from_dataframe.return_value
@mark.asyncio
@patch(
'laborious.activities.mlflow.MinioDataFramePayload.from_dataframe',
new_callable=AsyncMock,
)
async def test_request_transform_success_without_transform_meta(mock_from_dataframe, mlflow):
ts = pd.Timestamp('2020-01-01', tz='UTC')
raw = pd.DataFrame(
{
'variable': ['v1'],
'timestamp': [ts],
'value': [1.0],
'created_at': [ts],
}
)
out_idx = pd.Index([ts.strftime(DATETIME_FORMAT_WITH_TZ)], name=None)
out_df = pd.DataFrame({'v1': [1.0]}, index=out_idx)
wrapper = MagicMock()
wrapper.transform.return_value = (out_df, {})
mlflow.mlflow_repository.get_cached_model.return_value = wrapper
payload = AsyncMock()
payload.retrieve = AsyncMock(return_value=raw)
input_data = {
**metadata,
'data': payload,
'model_name': 'test_model',
'model_config': {},
}
await mlflow.request_transform(input_data)
mock_from_dataframe.assert_called_once()
@mark.asyncio @mark.asyncio
@patch( @patch(
'laborious.activities.mlflow.MinioDataFramePayload.from_dataframe', 'laborious.activities.mlflow.MinioDataFramePayload.from_dataframe',
@@ -233,6 +310,39 @@ async def test_request_predict(mock_to_datetime, mock_from_dataframe, mlflow):
assert response_data == mock_from_dataframe.return_value assert response_data == mock_from_dataframe.return_value
@mark.asyncio
@patch(
'laborious.activities.mlflow.MinioDataFramePayload.from_dataframe',
new_callable=AsyncMock,
)
@patch('laborious.activities.mlflow.to_datetime')
async def test_request_predict_success_dataframe_and_meta(
mock_to_datetime, mock_from_dataframe, mlflow
):
wrapper = MagicMock()
pred_df = pd.DataFrame({'raw': [0.3]})
wrapper.predict.return_value = (pred_df, {'m': 1})
mlflow.mlflow_repository.get_cached_model.return_value = wrapper
data_mock = MagicMock()
data_mock.index = pd.DatetimeIndex([pd.Timestamp('2020-01-01', tz='UTC')])
payload = AsyncMock()
payload.retrieve = AsyncMock(return_value=data_mock)
input_data = {
**metadata,
'data': payload,
'model_name': 'test_model',
'model_config': {},
}
await mlflow.request_predict(input_data)
assert list(pred_df.columns) == ['prediction', 'response_time']
mlflow.info.assert_any_call("Wrapper predict metadata: {'m': 1}", metadata['metadata'])
mock_from_dataframe.assert_called_once()
@mark.asyncio @mark.asyncio
@patch( @patch(
'laborious.activities.mlflow.MinioDataFramePayload.from_dataframe', 'laborious.activities.mlflow.MinioDataFramePayload.from_dataframe',
@@ -275,7 +385,7 @@ async def test_request_predict_failure(mock_to_datetime, mock_from_dataframe, ml
async def test_retrain_model_success_data_success_retrain( async def test_retrain_model_success_data_success_retrain(
mock_to_datetime, mock_rmtree, mock_mkdtemp, mock_log_artifact, mlflow mock_to_datetime, mock_rmtree, mock_mkdtemp, mock_log_artifact, mlflow
): ):
mock_mkdtemp.return_value = '/tmp/x' mock_mkdtemp.return_value = 'tmp'
mv_alias = MagicMock() mv_alias = MagicMock()
mv_alias.run_id = 'source-run' mv_alias.run_id = 'source-run'
@@ -326,7 +436,7 @@ async def test_retrain_model_success_data_success_retrain(
async def test_retrain_model_success_with_payload_data( async def test_retrain_model_success_with_payload_data(
mock_to_datetime, mock_rmtree, mock_mkdtemp, mock_log_artifact, mlflow mock_to_datetime, mock_rmtree, mock_mkdtemp, mock_log_artifact, mlflow
): ):
mock_mkdtemp.return_value = '/tmp/x' mock_mkdtemp.return_value = 'tmp'
mv_alias = MagicMock(run_id='src') mv_alias = MagicMock(run_id='src')
mlflow.mlflow_repository._client.get_model_version_by_alias.return_value = mv_alias mlflow.mlflow_repository._client.get_model_version_by_alias.return_value = mv_alias
wrapper = MagicMock() wrapper = MagicMock()
@@ -362,6 +472,48 @@ async def test_retrain_model_success_with_payload_data(
assert response['success'] is True assert response['success'] is True
@mark.asyncio
@patch('laborious.activities.mlflow.mlflow.log_artifact')
@patch('laborious.activities.mlflow.tempfile.mkdtemp')
@patch('laborious.activities.mlflow.rmtree')
@patch('laborious.activities.mlflow.to_datetime')
async def test_retrain_model_success_full_retrain_branch(
mock_to_datetime, mock_rmtree, mock_mkdtemp, mock_log_artifact, mlflow
):
mock_mkdtemp.return_value = 'tmp'
mv_alias = MagicMock(run_id='src')
mlflow.mlflow_repository._client.get_model_version_by_alias.return_value = mv_alias
wrapper = MagicMock()
mlflow.mlflow_repository.get_cached_model.return_value = wrapper
mock_cm = MagicMock()
mock_cm.__enter__.return_value = MagicMock(run_id='r', experiment_id='e')
mock_cm.__exit__.return_value = False
mlflow.mlflow_repository.start_run.return_value = mock_cm
ts = pd.Timestamp('2020-01-01', tz='UTC')
raw_data = pd.DataFrame(
{
'variable': ['target', 'f1', 'target', 'f1'],
'timestamp': [ts, ts, ts + pd.Timedelta(days=1), ts + pd.Timedelta(days=1)],
'value': [1.0, 2.0, 3.0, 4.0],
}
)
payload = AsyncMock()
payload.retrieve = AsyncMock(return_value=raw_data)
response = await mlflow.retrain_model(
{
**metadata,
'data': payload,
'model_name': 'test_model',
'model_config': {'target': 'target', 'full_retrain': True, 'validation_fraction': 0.5},
}
)
wrapper.train.assert_called_once()
assert response['success'] is True
@mark.asyncio @mark.asyncio
@patch('laborious.activities.mlflow.to_datetime') @patch('laborious.activities.mlflow.to_datetime')
async def test_retrain_model_success_data_fail_retrain(mock_to_datetime, mlflow): async def test_retrain_model_success_data_fail_retrain(mock_to_datetime, mlflow):
@@ -393,7 +545,6 @@ async def test_retrain_model_success_data_fail_retrain(mock_to_datetime, mlflow)
) )
assert response['success'] is False assert response['success'] is False
mlflow.send_notification_async.assert_called_once()
assert 'retrain failed' in response['message'] assert 'retrain failed' in response['message']
@@ -415,19 +566,17 @@ async def test_retrain_model_data_error(mlflow):
@mark.asyncio @mark.asyncio
async def test_retrain_model_missing_target(mlflow): async def test_retrain_model_missing_target(mlflow):
raw_data = MagicMock(columns=['variable', 'timestamp', 'value']) ts = pd.Timestamp('2020-01-01', tz='UTC')
raw_data.__getitem__.return_value.max.return_value = 'ts' raw_data = pd.DataFrame(
{
'variable': ['f1', 'f2'],
'timestamp': [ts, ts],
'value': [1.0, 2.0],
}
)
payload = AsyncMock() payload = AsyncMock()
payload.retrieve = AsyncMock(return_value=raw_data) payload.retrieve = AsyncMock(return_value=raw_data)
pivoted = MagicMock()
raw_data.drop_duplicates.return_value = raw_data
raw_data.pivot.return_value = pivoted
pivoted.fillna = MagicMock()
pivoted.columns.name = None
pivoted.index = MagicMock()
pivoted.__setitem__ = MagicMock()
response = await mlflow.retrain_model( response = await mlflow.retrain_model(
{ {
**metadata, **metadata,
@@ -561,6 +710,24 @@ async def test_get_reference_data_not_found(mlflow):
assert result is None assert result is None
@mark.asyncio
async def test_get_reference_data_missing_csv_file_returns_none(mlflow):
input_data = {
**metadata,
'model_name': 'test_model',
}
mv = MagicMock(run_id='run1')
mlflow.mlflow_repository._client.get_model_version_by_alias.return_value = mv
with patch('laborious.activities.mlflow.tempfile.mkdtemp', return_value='tmp'):
with patch('laborious.activities.mlflow.rmtree'):
with patch('laborious.activities.mlflow.Path') as mp:
mp.return_value.rglob.return_value = []
result = await mlflow.get_reference_data(input_data)
assert result is None
@mark.asyncio @mark.asyncio
async def test_get_reference_data_exception(mlflow): async def test_get_reference_data_exception(mlflow):
input_data = { input_data = {

View File

@@ -190,6 +190,21 @@ async def test_from_dataframe_inline():
assert result.last_timestamp == '2024-01-01' assert result.last_timestamp == '2024-01-01'
@pytest.mark.asyncio
@patch('laborious.utils.models.minio_dataframe_payload.OFFLOAD_THRESHOLD_BYTES', 10**9)
async def test_from_dataframe_inline_uses_provided_last_timestamp():
minio = AsyncMock()
df = _mock_dataframe({'timestamp': ['2024-01-01'], 'value': [42]})
result = await MinioDataFramePayload.from_dataframe(
dataframe=df,
minio_repo=minio,
model_name='m',
operation='initial',
last_timestamp='2024-01-02',
)
assert result.last_timestamp == '2024-01-02'
@pytest.mark.asyncio @pytest.mark.asyncio
@patch('laborious.utils.models.minio_dataframe_payload.now') @patch('laborious.utils.models.minio_dataframe_payload.now')
@patch('laborious.utils.models.minio_dataframe_payload.OFFLOAD_THRESHOLD_BYTES', 0) @patch('laborious.utils.models.minio_dataframe_payload.OFFLOAD_THRESHOLD_BYTES', 0)
@@ -264,3 +279,9 @@ def test_from_dict_passthrough_existing_instance():
original = MinioDataFramePayload(last_timestamp='2024-01-01', data={'a': 1}, bucket='b') original = MinioDataFramePayload(last_timestamp='2024-01-01', data={'a': 1}, bucket='b')
result = MinioDataFramePayload.from_dict(original) result = MinioDataFramePayload.from_dict(original)
assert result is original assert result is original
def test_debug_with_logger_calls_custom_debug():
logger = MagicMock()
MinioDataFramePayload._debug(logger, 'msg', {'a': 1})
logger.custom_debug.assert_called_once_with('msg', {'a': 1})

View File

@@ -0,0 +1,12 @@
from pandas import DataFrame
from laborious.utils.dataframe_debug import build_dataframe_debug_message
def test_build_dataframe_debug_message_skips_large_dataframe():
df = DataFrame({'a': [1, 2, 3]})
msg = build_dataframe_debug_message('payload', df, max_rows=1)
assert 'skipped because dataframe has 3 rows' in msg
assert '(max: 1)' in msg

View File

@@ -0,0 +1,243 @@
from unittest.mock import AsyncMock, MagicMock, patch
from pytest import mark, raises
from laborious.worker import worker
def _build_fake_activities():
inst = MagicMock()
inst.init_opc = AsyncMock()
inst.shutdown = AsyncMock()
inst.load_query_with_minio_offload = MagicMock()
inst.retrain_model = MagicMock()
inst.update_production_model = MagicMock()
inst.format_retrain_report = MagicMock()
inst.export_data_to_postgres = MagicMock()
inst.load_custom_query = MagicMock()
inst.calculate_simple_metrics = MagicMock()
inst.get_reference_data = MagicMock()
inst.calculate_drift = MagicMock()
inst.request_predict = MagicMock()
inst.request_transform = MagicMock()
inst.input_gate = MagicMock()
inst.mlflow_response_gate = MagicMock()
inst.mlflow_content_gate = MagicMock()
inst.format_transformed_data = MagicMock()
inst.format_prediction = MagicMock()
inst.format_default_prediction = MagicMock()
inst.write_opc_data = MagicMock()
inst.cleanup_minio_objects_expired = MagicMock()
inst.repeat_last_prediction = MagicMock()
inst.write_metrics = MagicMock()
inst.write_pi_web_api_data = MagicMock()
return inst
def _build_fake_worker(async_result=None, async_error: Exception | None = None):
w = MagicMock()
async def _run():
if async_error is not None:
raise async_error
return async_result
w.run = MagicMock(side_effect=_run)
return w
@patch('laborious.worker.worker.start_http_server')
def test_start_prometheus_server_success(mock_start_http):
with patch.object(worker.metrics.APP_UP, 'labels') as labels:
gauge = MagicMock()
labels.return_value = gauge
with patch('laborious.worker.worker.os.getenv', return_value='9090'):
worker.start_prometheus_server()
mock_start_http.assert_called_once_with(9090)
gauge.set.assert_called_once_with(1)
@patch('laborious.worker.worker.start_http_server', side_effect=RuntimeError('nope'))
def test_start_prometheus_server_error_exits(_mock_start_http):
with patch('laborious.worker.worker.os._exit', side_effect=SystemExit(1)) as m_exit:
with raises(SystemExit):
worker.start_prometheus_server()
m_exit.assert_called_once_with(1)
@mark.asyncio
async def test_main_missing_runtime_exits_fast(monkeypatch):
monkeypatch.setenv('RUNTIME', '')
with (
patch('laborious.worker.worker.start_prometheus_server'),
patch('laborious.worker.worker.get_logger') as m_logger,
patch('laborious.worker.worker.NotificationHandler'),
patch('laborious.worker.worker.MetricsController'),
patch.object(worker.metrics.APP_UP, 'labels') as labels,
patch('laborious.worker.worker.sys.exit', side_effect=SystemExit(1)),
):
labels.return_value = MagicMock()
with raises(SystemExit):
await worker.main()
assert m_logger.return_value.custom_critical.called
@mark.asyncio
async def test_main_plugin_install_failure(monkeypatch):
monkeypatch.setenv('RUNTIME', 'single')
fake_activities = _build_fake_activities()
fake_plugin = MagicMock()
fake_plugin.install_runtime = AsyncMock(side_effect=RuntimeError('install failed'))
with (
patch('laborious.worker.worker.start_prometheus_server'),
patch('laborious.worker.worker.get_logger'),
patch(
'laborious.worker.worker.build_mongodb_config',
return_value={'connection_string': 'cs', 'database_name': 'db'},
),
patch('laborious.worker.worker.NotificationHandler'),
patch('laborious.worker.worker.MetricsController'),
patch(
'laborious.worker.worker.build_plugin_store_config',
return_value={
'base_url': '',
'owner': '',
'repo': '',
'username': None,
'password': None,
'branch': None,
'cache_ttl_seconds': None,
'pypi_index_url': '',
'pypi_username': None,
'pypi_password': None,
},
),
patch('laborious.worker.worker.PluginStore', return_value=fake_plugin),
patch('laborious.worker.worker.Activities', return_value=fake_activities),
patch.object(worker.metrics.APP_UP, 'labels') as labels,
patch('laborious.worker.worker.sys.exit', side_effect=SystemExit(1)),
):
labels.return_value = MagicMock()
with raises(SystemExit):
await worker.main()
@mark.asyncio
async def test_main_success_exit_zero(monkeypatch):
monkeypatch.setenv('RUNTIME', 'single')
fake_activities = _build_fake_activities()
fake_plugin = MagicMock()
fake_plugin.install_runtime = AsyncMock(return_value=None)
fake_workers = [_build_fake_worker() for _ in range(4)]
with (
patch('laborious.worker.worker.start_prometheus_server'),
patch('laborious.worker.worker.get_logger'),
patch(
'laborious.worker.worker.build_mongodb_config',
return_value={'connection_string': 'cs', 'database_name': 'db'},
),
patch('laborious.worker.worker.NotificationHandler') as m_notif_cls,
patch('laborious.worker.worker.MetricsController'),
patch(
'laborious.worker.worker.build_plugin_store_config',
return_value={
'base_url': '',
'owner': '',
'repo': '',
'username': None,
'password': None,
'branch': None,
'cache_ttl_seconds': None,
'pypi_index_url': '',
'pypi_username': None,
'pypi_password': None,
},
),
patch('laborious.worker.worker.PluginStore', return_value=fake_plugin),
patch('laborious.worker.worker.Activities', return_value=fake_activities),
patch('laborious.worker.worker.build_postgres_config', return_value={}),
patch('laborious.worker.worker.build_minio_config', return_value={}),
patch('laborious.worker.worker.build_opc_config', return_value={}),
patch('laborious.worker.worker.build_api_config', return_value={}),
patch('laborious.worker.worker.PrometheusConfig', return_value=MagicMock()),
patch('laborious.worker.worker.TelemetryConfig', return_value=MagicMock()),
patch('laborious.worker.worker.Runtime', return_value=MagicMock()),
patch(
'laborious.worker.worker.client.Client.connect', new=AsyncMock(return_value=MagicMock())
),
patch('laborious.worker.worker.prepare_worker', side_effect=fake_workers) as m_prepare,
patch.object(worker.metrics.APP_UP, 'labels') as labels,
patch('laborious.worker.worker.sys.exit', side_effect=SystemExit(0)),
):
labels.return_value = MagicMock()
with raises(SystemExit):
await worker.main()
notif = m_notif_cls.return_value
notif.shutdown.assert_called_once()
fake_activities.shutdown.assert_awaited_once()
assert m_prepare.call_count == 4
prepare_calls = m_prepare.call_args_list
assert prepare_calls[0].kwargs['runtime'] == 'single'
assert prepare_calls[3].kwargs['runtime'] == 'single'
@mark.asyncio
async def test_main_worker_gather_error_exits_one(monkeypatch):
monkeypatch.setenv('RUNTIME', 'single')
fake_activities = _build_fake_activities()
fake_plugin = MagicMock()
fake_plugin.install_runtime = AsyncMock(return_value=None)
fake_workers = [
_build_fake_worker(async_error=RuntimeError('boom')),
_build_fake_worker(),
_build_fake_worker(),
_build_fake_worker(),
]
with (
patch('laborious.worker.worker.start_prometheus_server'),
patch('laborious.worker.worker.get_logger') as m_logger,
patch(
'laborious.worker.worker.build_mongodb_config',
return_value={'connection_string': 'cs', 'database_name': 'db'},
),
patch('laborious.worker.worker.NotificationHandler'),
patch('laborious.worker.worker.MetricsController'),
patch(
'laborious.worker.worker.build_plugin_store_config',
return_value={
'base_url': '',
'owner': '',
'repo': '',
'username': None,
'password': None,
'branch': None,
'cache_ttl_seconds': None,
'pypi_index_url': '',
'pypi_username': None,
'pypi_password': None,
},
),
patch('laborious.worker.worker.PluginStore', return_value=fake_plugin),
patch('laborious.worker.worker.Activities', return_value=fake_activities),
patch('laborious.worker.worker.build_postgres_config', return_value={}),
patch('laborious.worker.worker.build_minio_config', return_value={}),
patch('laborious.worker.worker.build_opc_config', return_value={}),
patch('laborious.worker.worker.build_api_config', return_value={}),
patch('laborious.worker.worker.PrometheusConfig', return_value=MagicMock()),
patch('laborious.worker.worker.TelemetryConfig', return_value=MagicMock()),
patch('laborious.worker.worker.Runtime', return_value=MagicMock()),
patch(
'laborious.worker.worker.client.Client.connect', new=AsyncMock(return_value=MagicMock())
),
patch('laborious.worker.worker.prepare_worker', side_effect=fake_workers),
patch.object(worker.metrics.APP_UP, 'labels') as labels,
patch('laborious.worker.worker.sys.exit', side_effect=SystemExit(1)),
):
labels.return_value = MagicMock()
with raises(SystemExit):
await worker.main()
assert m_logger.return_value.custom_error.called

View File

@@ -1,6 +1,6 @@
from unittest.mock import ANY, AsyncMock, MagicMock, call, patch from unittest.mock import ANY, AsyncMock, MagicMock, call, patch
from pytest import fixture, mark from pytest import fixture, mark, raises
from laborious.activities.activities import Activities from laborious.activities.activities import Activities
from laborious.workflows.sub_workflows.prediction_process import PredictionProcess from laborious.workflows.sub_workflows.prediction_process import PredictionProcess
@@ -840,3 +840,27 @@ async def test_run_with_cleanup_prefixes(workflow_mock, prediction_process):
retry_policy=ANY, retry_policy=ANY,
start_to_close_timeout=ANY, start_to_close_timeout=ANY,
) )
@mark.asyncio
@patch('laborious.workflows.sub_workflows.prediction_process.workflow', new_callable=AsyncMock)
async def test_run_always_cleans_up_on_pipeline_exception(workflow_mock, prediction_process):
input_data = {
'metadata': metadata,
'data': {'last_timestamp': '2024-01-01'},
'model_id': 1,
'model_name': 'm',
'model_config': {},
'save_transform': False,
}
prediction_process._run_prediction_pipeline = AsyncMock(side_effect=RuntimeError('boom'))
with raises(RuntimeError):
await prediction_process.run(input_data)
workflow_mock.execute_activity_method.assert_called_once_with(
Activities.cleanup_minio_objects_expired,
{**metadata, 'data': input_data['data']},
retry_policy=ANY,
start_to_close_timeout=ANY,
)