diff --git a/laborious/activities/mlflow.py b/laborious/activities/mlflow.py index aae09ee..896f05a 100644 --- a/laborious/activities/mlflow.py +++ b/laborious/activities/mlflow.py @@ -63,44 +63,6 @@ class MLFlow(BaseActivity): f"{mlflow_host}:{mlflow_port}", mlflow_username, mlflow_password, logger.base_logger ) - def detect_and_parse_datetime_index(self, data: DataFrame, metadata: dict) -> DataFrame: - """ - Detect and parse datetime index from data. index must be a timestamp like column. - This function must detect the timestamp type (pandas Timestamp or datetime) and convert it to DATETIME_FORMAT_WITH_TZ. - If the index is a string, must be in format DATETIME_FORMAT_WITH_TZ. - If another type or format, must raise an error. - """ - index = data.index - - # Get type of first element of index - index_type = type(index[0]) - - self.info(f"Index type: {index_type}", metadata) - - message = f"Index must be all timestamp like column. Valid formats are: pandas Timestamp, datetime, string in format {DATETIME_FORMAT_WITH_TZ}" - - # Check if all in index are of the same type - if not all(isinstance(i, index_type) for i in index): - raise ValueError( - f"{message}") - - # Check type and converts to DATETIME_FORMAT_WITH_TZ - if index_type == str: - # Validate format of string and return error if not valid - try: - to_datetime(data.index) - except ValueError: - raise ValueError( - f"{message}") - - elif index_type == datetime or index_type == Timestamp: - data.index = data.index.strftime(DATETIME_FORMAT_WITH_TZ) - else: - raise ValueError( - f"{message}") - - return data - @activity.defn(name="request_transform") async def request_transform(self, input_data: dict[str, Any]) -> dict[str, Any]: """ @@ -164,35 +126,6 @@ class MLFlow(BaseActivity): self.debug( f"Transform raw response data: \n {create_sample_dict(response_data, max_items=5, max_depth=5)}", metadata) - if response_data['success']: - - response_dataframe = DataFrame(response_data['content']) - if len(response_dataframe) == 0: - return response_data - try: - response_dataframe = self.detect_and_parse_datetime_index( - response_dataframe, metadata) - response_dataframe['timestamp'] = to_datetime( - response_dataframe.index, format=DATETIME_FORMAT_WITH_TZ) - response_dataframe['timestamp'] = response_dataframe['timestamp'].dt.strftime( - DATETIME_FORMAT) - except ValueError as e: - trace = traceback.format_exc() - self.send_notification( - metadata=metadata, - notification_id='TRANSFORM_DATA_INDEX_ERROR', - message=f'Error parsing trasnformed data index: {e}', - block='transform', - level=NotificationLevel.ERROR, - attachment_content=trace - ) - self.error(trace, metadata=metadata) - raise e - - response_dataframe.to_csv('response_data.csv') - - response_data['content'] = response_dataframe.to_dict() - self.debug( f"Transform response data: \n {create_sample_dict(response_data, max_items=5, max_depth=5)}", metadata) diff --git a/laborious/utils/repository/model_repository.py b/laborious/utils/repository/model_repository.py index 0145554..3107d1c 100644 --- a/laborious/utils/repository/model_repository.py +++ b/laborious/utils/repository/model_repository.py @@ -16,17 +16,57 @@ import pandas as pd import mlflow from os import makedirs, path, remove from sientia.ModelServing import ModelServing +from sientia_do.temporal.constants import DATETIME_FORMAT_WITH_TZ +from sientia_do.observability.logger import Logger class MLFlowRepository(): - def __init__(self, host, username, password, logger): + def __init__(self, host, username, password, logger: Logger): self.model_serving = ModelServing(tracking_uri=host, username=username, password=password, logger=logger) self.logger = logger - def transform(self, model_name: str, data: pd.DataFrame, model_config: dict) -> dict: + def detect_and_parse_datetime_index(self, data: pd.DataFrame, metadata: dict) -> pd.DataFrame: + """ + Detect and parse datetime index from data. index must be a timestamp like column. + This function must detect the timestamp type (pandas Timestamp or datetime) and convert it to DATETIME_FORMAT_WITH_TZ. + If the index is a string, must be in format DATETIME_FORMAT_WITH_TZ. + If another type or format, must raise an error. + """ + index = data.index + + # Get type of first element of index + index_type = type(index[0]) + + self.logger.custom_info(f"Index type: {index_type}", metadata) + + message = f"Index must be all timestamp like column. Valid formats are: pandas Timestamp, datetime, string in format {DATETIME_FORMAT_WITH_TZ}" + + # Check if all in index are of the same type + if not all(isinstance(i, index_type) for i in index): + raise ValueError( + f"{message}") + + # Check type and converts to DATETIME_FORMAT_WITH_TZ + if index_type == str: + # Validate format of string and return error if not valid + try: + pd.to_datetime(data.index, format=DATETIME_FORMAT_WITH_TZ) + except ValueError: + raise ValueError( + f"{message}") + + elif index_type == datetime or index_type == pd.Timestamp: + data.index = data.index.strftime(DATETIME_FORMAT_WITH_TZ) + else: + raise ValueError( + f"{message}") + + return data + + def transform(self, model_name: str, data: pd.DataFrame, model_config: dict, metadata: dict) -> dict: """ Transform data using a model. @@ -40,6 +80,9 @@ class MLFlowRepository(): """ try: + self.logger.custom_debug( + f"Data received for model transformation: {data.to_csv()}", metadata) + model_retention = model_config.get('retention_minutes', 0) flavor = model_config.get('transform_flavor', 'sklearn') compressed = model_config.get('is_compressed', False) @@ -47,12 +90,20 @@ class MLFlowRepository(): transform_keyword = model_config.get( 'transform_function_keyword', 'predict') + transformed_data = self.model_serving.get_cached_transform( + model_name, data, model_retention, flavor, + compressed, retention_target, transform_keyword + ) + + self.logger.custom_debug( + f"Data received from model transformation: {transformed_data.to_csv()}", metadata) + + transformed_data = self.detect_and_parse_datetime_index( + transformed_data, metadata) + return { 'success': True, - 'content': self.model_serving.get_cached_transform( - model_name, data, model_retention, flavor, - compressed, retention_target, transform_keyword - ).to_dict() + 'content': transformed_data.to_dict() } except Exception as e: @@ -64,7 +115,7 @@ class MLFlowRepository(): } } - def predict(self, model_name: str, data: pd.DataFrame, model_config: dict) -> dict: + def predict(self, model_name: str, data: pd.DataFrame, model_config: dict, metadata: dict) -> dict: """ Predict data using a model. @@ -85,8 +136,8 @@ class MLFlowRepository(): input_index = data.index start_time = datetime.now() - self.logger.debug( - f"Data received for model prediction: {data.to_string()}") + self.logger.custom_debug( + f"Data received for model prediction: {data.to_csv()}", metadata) data = self.model_serving.get_cached_predict( model_name, data, model_retention, flavor, compressed, retention_target @@ -94,8 +145,8 @@ class MLFlowRepository(): end_time = datetime.now() data = pd.DataFrame(data, columns=['prediction']) - self.logger.debug( - f"Data received from model prediction: {data.to_string()}") + self.logger.custom_debug( + f"Data received from model prediction: {data.to_csv()}", metadata) data.index = input_index data['response_time'] = (end_time - start_time).total_seconds() diff --git a/tests/laborious/activities/test_mlflow.py b/tests/laborious/activities/test_mlflow.py index 7aced7c..ed68729 100644 --- a/tests/laborious/activities/test_mlflow.py +++ b/tests/laborious/activities/test_mlflow.py @@ -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 = { diff --git a/tests/laborious/utils/repository/test_model_repository.py b/tests/laborious/utils/repository/test_model_repository.py index 8b3716e..51a3223 100644 --- a/tests/laborious/utils/repository/test_model_repository.py +++ b/tests/laborious/utils/repository/test_model_repository.py @@ -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,