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:
@@ -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):
|
||||||
|
|||||||
@@ -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
|
||||||
|
|
||||||
|
|||||||
@@ -264,10 +264,8 @@ 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()
|
await activities.shutdown()
|
||||||
if activities:
|
|
||||||
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
19
requirements-local.txt
Normal 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
|
||||||
@@ -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
|
||||||
|
|||||||
@@ -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,
|
||||||
|
)
|
||||||
|
|||||||
@@ -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()
|
||||||
|
|||||||
@@ -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 = {
|
||||||
|
|||||||
@@ -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})
|
||||||
|
|||||||
12
tests/laborious/utils/test_dataframe_debug.py
Normal file
12
tests/laborious/utils/test_dataframe_debug.py
Normal 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
|
||||||
243
tests/laborious/worker/test_worker.py
Normal file
243
tests/laborious/worker/test_worker.py
Normal 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
|
||||||
@@ -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,
|
||||||
|
)
|
||||||
|
|||||||
Reference in New Issue
Block a user