SIENTIAPDE-1222
Refactor datetime index handling in MLFlow and MLFlowRepository - Moved the detect_and_parse_datetime_index method from MLFlow to MLFlowRepository for better organization and reusability. - Updated the method to include enhanced logging and error handling for invalid datetime formats. - Adjusted the transform method in MLFlowRepository to utilize the new datetime index parsing logic. - Added unit tests for both valid and invalid datetime index cases to ensure robustness.
This commit is contained in:
@@ -1,8 +1,9 @@
|
||||
from datetime import datetime
|
||||
from unittest.mock import ANY, MagicMock, patch
|
||||
|
||||
import numpy as np
|
||||
from pandas import DataFrame
|
||||
from pytest import fixture, mark
|
||||
from pandas import DataFrame, Timestamp
|
||||
from pytest import fixture, mark, raises
|
||||
from laborious.activities.mlflow import MLFlow
|
||||
from sientia_do.notifications.models import NotificationLevel
|
||||
|
||||
@@ -58,7 +59,7 @@ metadata = {
|
||||
@mark.asyncio
|
||||
@patch("laborious.activities.mlflow.DataFrame")
|
||||
@patch("laborious.activities.mlflow.max")
|
||||
async def test_request_transform(mock_max, mock_dataframe, mlflow):
|
||||
async def test_request_transform_success(mock_max, mock_dataframe, mlflow):
|
||||
mock_max.return_value = '2024-01-02'
|
||||
# Mock input data
|
||||
input_data = {
|
||||
|
||||
@@ -2,7 +2,8 @@ from unittest.mock import ANY, MagicMock, call, patch
|
||||
import numpy as np
|
||||
from pandas import DataFrame
|
||||
import pytest
|
||||
from laborious.utils.repository import model_repository
|
||||
from datetime import datetime, timezone
|
||||
from pandas import Timestamp
|
||||
from laborious.utils.repository.model_repository import MLFlowRepository
|
||||
|
||||
|
||||
@@ -22,18 +23,113 @@ def mlflow_repository():
|
||||
return repo
|
||||
|
||||
|
||||
metadata = {
|
||||
"metadata": {
|
||||
"model_id": "test_model",
|
||||
"model_name": "test_model",
|
||||
"workflow_name": "test_workflow",
|
||||
"schema_name": "test_schedule",
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
invalid_cases = [
|
||||
(
|
||||
{
|
||||
'value': {
|
||||
'2024-01-01 12:00:00': 1,
|
||||
2024: 2
|
||||
}
|
||||
}
|
||||
),
|
||||
(
|
||||
{
|
||||
'value': {
|
||||
'2024-01-01': 1,
|
||||
'2024-01-02': 2
|
||||
}
|
||||
}
|
||||
),
|
||||
(
|
||||
{
|
||||
'value': {
|
||||
1: 1,
|
||||
2: 2
|
||||
}
|
||||
}
|
||||
)
|
||||
]
|
||||
|
||||
|
||||
@pytest.mark.parametrize("data", invalid_cases)
|
||||
def test_detect_and_parse_datetime_index_error_cases(mlflow_repository, data):
|
||||
input_data = DataFrame(
|
||||
data
|
||||
)
|
||||
|
||||
with pytest.raises(ValueError) as e:
|
||||
mlflow_repository.detect_and_parse_datetime_index(
|
||||
input_data, metadata['metadata'])
|
||||
|
||||
assert str(e) == "Index must be all timestamp like column. Valid formats are: pandas Timestamp, datetime, string in format %Y-%m-%d %H:%M:%S"
|
||||
|
||||
|
||||
valid_cases = [
|
||||
(
|
||||
{
|
||||
'value': {
|
||||
'2024-01-01 12:00:00+0000': 1,
|
||||
'2024-01-02 12:00:00+0000': 2
|
||||
}
|
||||
}, ['2024-01-01 12:00:00+0000', '2024-01-02 12:00:00+0000']
|
||||
),
|
||||
(
|
||||
{
|
||||
'value': {
|
||||
datetime(2025, 1, 1, 12, 0, 0, tzinfo=timezone.utc): 1,
|
||||
datetime(2025, 1, 2, 12, 0, 0, tzinfo=timezone.utc): 2
|
||||
}
|
||||
}, ['2025-01-01 12:00:00+0000', '2025-01-02 12:00:00+0000']
|
||||
),
|
||||
(
|
||||
{
|
||||
'value': {
|
||||
Timestamp(2026, 1, 1, 12, 0, 0, tzinfo=timezone.utc): 1,
|
||||
Timestamp(2026, 1, 2, 12, 0, 0, tzinfo=timezone.utc): 2
|
||||
}
|
||||
}, ['2026-01-01 12:00:00+0000', '2026-01-02 12:00:00+0000']
|
||||
),
|
||||
]
|
||||
|
||||
|
||||
@pytest.mark.parametrize("data,expected", valid_cases)
|
||||
def test_detect_and_parse_datetime_index_valid_format(mlflow_repository, data, expected):
|
||||
input_data = DataFrame(data)
|
||||
|
||||
response = mlflow_repository.detect_and_parse_datetime_index(
|
||||
input_data, metadata['metadata'])
|
||||
|
||||
assert response.index.tolist() == expected
|
||||
|
||||
|
||||
def test_transform_success(mlflow_repository):
|
||||
data = 'data'
|
||||
model_name = 'model'
|
||||
|
||||
output = mlflow_repository.transform(model_name, data, 1)
|
||||
mlflow_repository.detect_and_parse_datetime_index = MagicMock()
|
||||
|
||||
output = mlflow_repository.transform(
|
||||
model_name, data, {}, metadata['metadata'])
|
||||
|
||||
mlflow_repository.model_serving.get_cached_transform.assert_called_once_with(
|
||||
model_name, data, 1)
|
||||
model_name, data, 0, 'sklearn', False, 'model', 'predict')
|
||||
|
||||
mlflow_repository.detect_and_parse_datetime_index.assert_called_once_with(
|
||||
mlflow_repository.model_serving.get_cached_transform.return_value, metadata['metadata'])
|
||||
|
||||
assert output == {
|
||||
'success': True,
|
||||
'content': mlflow_repository.model_serving.get_cached_transform.return_value.to_dict.return_value
|
||||
'content': mlflow_repository.detect_and_parse_datetime_index.return_value.to_dict.return_value
|
||||
}
|
||||
|
||||
|
||||
@@ -44,10 +140,11 @@ def test_transform_error(mlflow_repository):
|
||||
mlflow_repository.model_serving.get_cached_transform.side_effect = Exception(
|
||||
'error')
|
||||
|
||||
output = mlflow_repository.transform(model_name, data, 1)
|
||||
output = mlflow_repository.transform(
|
||||
model_name, data, {}, metadata['metadata'])
|
||||
|
||||
mlflow_repository.model_serving.get_cached_transform.assert_called_once_with(
|
||||
model_name, data, 1)
|
||||
model_name, data, 0, 'sklearn', False, 'model', 'predict')
|
||||
|
||||
assert output == {
|
||||
'success': False,
|
||||
|
||||
Reference in New Issue
Block a user