From 9e31ab679b58ddf0923c5f9730e0dcb5a5091faa Mon Sep 17 00:00:00 2001 From: Bruno Domingues Date: Fri, 17 Oct 2025 09:11:54 -0300 Subject: [PATCH] SIENTIAPDE-1255: Implement MLFlow artifact management and model persistence This commit introduces a new model_repository.py to handle MLFlow artifact generation and model persistence. It also updates the README to reflect this change and modifies training_repository.py to separate training and MLFlow operations. --- README.md | 6 +- .../utils/repository/model_repository.py | 446 +--------------- .../utils/repository/test_model_repository.py | 482 +----------------- 3 files changed, 18 insertions(+), 916 deletions(-) diff --git a/README.md b/README.md index ee62cfe..c0f319c 100644 --- a/README.md +++ b/README.md @@ -166,8 +166,9 @@ The Model Manager system uses a Temporal-based workflow architecture with clear #### **Data Services (`model_manager/utils/`)** - **Connectors Config**: Environment variable-based configuration management -- **Repository**: Data access layer for training operations +- **Repository**: Data access layer for training and MLFlow operations - `training_repository.py`: Training business logic and operations + - `model_repository.py`: MLFlow artifact generation and model persistence - **Models**: Data models and schemas - `train_model_params.py`: Training parameters model - `train_model_result.py`: Training result model @@ -959,7 +960,8 @@ model_manager/ │ │ └── experiment_status.py # Experiment status enum │ └── repository/ # Data access layer │ ├── __init__.py -│ └── training_repository.py # Training business logic +│ ├── training_repository.py # Training business logic +│ └── model_repository.py # MLFlow artifact management ├── metrics.py # Prometheus metrics definitions └── __init__.py ``` diff --git a/model_manager/utils/repository/model_repository.py b/model_manager/utils/repository/model_repository.py index 0630b71..6c74a5e 100644 --- a/model_manager/utils/repository/model_repository.py +++ b/model_manager/utils/repository/model_repository.py @@ -1,28 +1,23 @@ """ -Model Monitoring Repository +MLFlow Repository -This module contains the ModelMonitoringRepository class, -which is responsible for handling the communication with the Model Monitoring API. +This module contains the MLFlowRepository class, which is responsible for +handling model training artifacts and MLFlow operations for the Model Manager system. -It includes the methods that are used to answer ModelMonitoringService -requests using the Model Monitoring API functions. - -By Monitoring we mean the evaluation of the performance of models, the generation of reports. +It includes methods for generating training reports, managing artifacts, +and logging model runs to MLFlow. """ import shutil -import traceback from datetime import datetime -from os import makedirs, path, remove +from os import makedirs, path -import mlflow import numpy as np import pandas as pd from sientia.ModelServing import ModelServing # type: ignore[import-untyped] from sientia.reports import Reports # type: ignore[import-untyped] from sientia_do.observability.logger import Logger -from sientia_do.temporal.constants import DATETIME_FORMAT_WITH_TZ from model_manager.utils.models.train_model_result import TrainModelResult @@ -34,434 +29,7 @@ class MLFlowRepository: ) self.logger = logger - 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 is str: - # Validate format of string and return error if not valid - try: - pd.to_datetime(data.index, format=DATETIME_FORMAT_WITH_TZ) - except ValueError as e: - raise ValueError(f'{message}') from e - - elif index_type == datetime or index_type == pd.Timestamp: - data.index = data.index.strftime(DATETIME_FORMAT_WITH_TZ) # type: ignore[attr-defined] - 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. - - Parameters: - - model_name (str): The name of the model to use for transformation. - - data (pandas.DataFrame): The data to transform. - - model_retention (int): The number of minutes to keep the model. - - Returns: - - dict: A dictionary containing the transformed data. - """ - - 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) - retention_target = model_config.get('retention_target', 'model') - 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': transformed_data.to_dict()} - - except Exception as e: # noqa: BLE001 - return { - 'success': False, - 'content': {'message': str(e), 'traceback': traceback.format_exc()}, - } - - def predict( - self, model_name: str, data: pd.DataFrame, model_config: dict, metadata: dict - ) -> dict: - """ - Predict data using a model. - - Parameters: - - model_name (str): The name of the model to use for prediction. - - data (pandas.DataFrame): The data to predict. - - model_retention (int): The number of minutes to keep the model. - - Returns: - - dict: A dictionary containing the predicted data. - """ - try: - model_retention = model_config.get('retention_minutes', 0) - flavor = model_config.get('predict_flavor', 'pyfunc') - compressed = model_config.get('is_compressed', False) - retention_target = model_config.get('retention_target', 'model') - - input_index = data.index - start_time = datetime.now() - - 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 - ) - - end_time = datetime.now() - data = pd.DataFrame(data, columns=['prediction']) - 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() - - return {'success': True, 'content': data.to_dict()} - - except Exception as e: # noqa: BLE001 - return { - 'success': False, - 'content': {'message': str(e), 'traceback': traceback.format_exc()}, - } - - def get_experiment_by_run_id(self, run_id: str) -> dict: - # Get the run information using the run_id - run = mlflow.get_run(run_id) - - # Extract the experiment ID from the run - experiment_id = run.info.experiment_id - - # Get the experiment details using the experiment ID - experiment = mlflow.get_experiment(experiment_id) - experiment_name = experiment.name - return experiment_name - - def get_next_run_name(self, model_name: str) -> str: - """ - Generate the next run name for a specific MLFlow model. - - This method calculates the next sequential run number for a model - by searching existing runs and incrementing the count. It ensures - unique run names for model training and retraining operations. - - Args: - model_name (str): The name of the MLFlow model - - Returns: - str: The next run name in format 'model_name-run_number' - """ - runs = mlflow.search_runs(experiment_names=[model_name], order_by=['start_time desc']) - next_run_number = len(runs) + 1 - return f'{model_name}-{next_run_number}' - - def create_model_experiment(self, model_name: str, data: pd.DataFrame) -> tuple: - """ - Create a new MLFlow experiment for model retraining. - - This method sets up the complete environment for model retraining by: - 1. Loading the current production prediction model - 2. Loading the current production transformation model - 3. Fitting the transformation model with new data - 4. Preparing data for prediction model retraining - 5. Setting up the MLFlow experiment context - - Args: - model_name (str): Name of the MLFlow model to retrain - data (pd.DataFrame): Training data for model retraining - - Returns: - tuple: (prediction_model, data_model, experiment) - - prediction_model: Loaded prediction model for retraining - - data_model: Fitted transformation model - - experiment: MLFlow experiment name - """ - # load predictor model - predictor_uri = f'models:/{model_name}/production' - # load transform model - latest_production_id = self.model_serving.get_model_info(model_name) # type: ignore[no-any-return] - transform_uri = self.model_serving.get_model_uri(latest_production_id, prediction=False) - # load - data_model = mlflow.sklearn.load_model(transform_uri) - prediction_model = mlflow.sklearn.load_model(predictor_uri) - data_model = data_model.fit(data) - treated_data = data_model.predict(data) - - target_name = data_model.target_variable - y = data[target_name] - treated_data = pd.merge(treated_data, y, left_index=True, right_index=True) - prediction_model = prediction_model.fit(treated_data) - experiment = self.get_experiment_by_run_id(latest_production_id) - mlflow.set_experiment(experiment) - - return prediction_model, data_model, experiment - - def perform_model_retrain( - self, prediction_model, data_model, experiment: str, model_name: str, data: pd.DataFrame - ): - """ - Execute the complete model retraining process in MLFlow. - - This method performs the actual model retraining by: - 1. Starting a new MLFlow run with descriptive metadata - 2. Logging model parameters and hyperparameters - 3. Retraining both prediction and transformation models - 4. Logging training data as artifacts - 5. Saving retrained models to MLFlow registry - - Args: - prediction_model: MLFlow prediction model to retrain - data_model: MLFlow transformation model to retrain - experiment (str): MLFlow experiment name for the retraining - model_name (str): Name of the model being retrained - data (pd.DataFrame): Training data used for retraining - - Returns: - tuple: (status_message, experiment_name) - - status_message (str): Success confirmation message - - experiment_name (str): Name of the experiment - """ - pred_model_atributes = vars(prediction_model) # load class attributes - data_model_atributes = vars(data_model) # load class attributes - experiment_description = f'Retrain model {model_name} with new data' - current_run_name = self.get_next_run_name(experiment) - with mlflow.start_run( - run_name=current_run_name, description=experiment_description - ) as _run: - # update transfomation model - # fixed parameters - for name_atribute, val_atribute in pred_model_atributes.items(): - if name_atribute != 'model': - mlflow.log_param(name_atribute, val_atribute) - # update prediction model - for name_atribute, val_atribute in data_model_atributes.items(): - if name_atribute != 'model': - mlflow.log_param(name_atribute, val_atribute) - # dynamic parameters, including model itself - mlflow.sklearn.log_model(data_model, 'data_model') - - makedirs('temp', exist_ok=True) - - file_path = f'temp/raw_data_{model_name}.csv' - data.to_csv(file_path, index=True) - - # log the data raw - mlflow.log_artifact(file_path) - - # dynamic parameters, including model itself - mlflow.sklearn.log_model(prediction_model, 'prediction_model') - mlflow.log_param('retrain', True) - - # clear temp file - if path.exists(file_path): - remove(file_path) - - return 'Model retrained successfully', experiment - - def retrain_model(self, data: pd.DataFrame, model_name: str) -> tuple: - """ - Orchestrate the complete model retraining workflow. - - This method coordinates the entire model retraining process by: - 1. Creating the MLFlow experiment environment - 2. Loading existing production models - 3. Executing the retraining process - 4. Returning comprehensive retraining results - - Args: - data (pd.DataFrame): Training data for model retraining - model_name (str): Name of the MLFlow model to retrain - - Returns: - tuple: (status_message, experiment_name) - - status_message (str): Retraining operation status - - experiment_name (str): MLFlow experiment identifier - """ - prediction_model, data_model, experiment = self.create_model_experiment(model_name, data) - retrain_result = self.perform_model_retrain( - prediction_model, data_model, experiment, model_name, data - ) - return retrain_result - - def get_experiment(self, experiment_name: str) -> int: - """ - Retrieve MLFlow experiment ID by experiment name. - - This method searches for an MLFlow experiment by name and - returns its unique identifier. It provides error handling - for non-existent experiments. - - Args: - experiment_name (str): Name of the MLFlow experiment - - Returns: - int: MLFlow experiment ID - - Raises: - ValueError: If the experiment name is not found - """ - experiment = mlflow.get_experiment_by_name(experiment_name) - if experiment is None: - raise ValueError(f'Experiment {experiment_name} not found') - - return experiment.experiment_id # type: ignore[no-any-return] - - def get_experiment_last_run(self, experiment_id: int) -> str: - """ - Retrieve the most recent retraining run ID for an experiment. - - This method searches for the latest run in an MLFlow experiment - that has been marked as a retraining run. It filters runs by - the 'retrain' parameter and orders them by completion time. - - Args: - experiment_id (int): MLFlow experiment ID - - Returns: - str: MLFlow run ID of the most recent retraining run - - Raises: - ValueError: If runs data is not in expected DataFrame format - """ - runs = mlflow.search_runs( - experiment_ids=[experiment_id], - filter_string='', # Sem filtro no MLflow ainda - output_format='pandas', - ) - - if not isinstance(runs, pd.DataFrame): - raise ValueError('Runs is not a pandas DataFrame') - - # Filtrar apenas as runs onde params.retrain == True - filtered_runs = runs[runs['params.retrain'] == 'True'] - - # Converter a coluna 'end_time' para datetime - filtered_runs['end_time'] = pd.to_datetime(filtered_runs['end_time']) - - # Ordenar o DataFrame de forma descendente pela coluna 'end_time' - filtered_runs = filtered_runs.sort_values(by='end_time', ascending=False) - - # Pegar a última run_id do DataFrame filtrado e ordenado - latest_run_id = filtered_runs.iloc[0]['run_id'] - - return latest_run_id - - def update_production_model_by_run_id(self, run_id: str, model_name: str) -> dict: - """ - Update production model with a specific MLFlow run. - - This method promotes a model from a specific MLFlow run to - production stage. It handles model registration, versioning, - and stage transitions with proper error handling. - - Args: - run_id (str): MLFlow run ID containing the model to promote - model_name (str): Name of the MLFlow model - - Returns: - dict: Model update metadata containing: - - model_name (str): Name of the updated model - - version (str): New model version number - - mlflow_run_id (str): Source run ID - - Update Process: - 1. Registers the model from the specified run - 2. Retrieves the latest model version - 3. Transitions the model to 'Production' stage - 4. Archives existing production versions - """ - # Registrar o modelo - # Aqui estamos assumindo que você já tem um modelo salvo, caso contrário você precisará treiná-lo e salvá-lo primeiro. - # Se o modelo já está registrado, você pode usar o método register_model() ou pyfunc.load_model() para isso. - mlflow.register_model(f'runs:/{run_id}/prediction_model', model_name) - - # Colocar a versão do modelo em produção - # Depois de registrar o modelo, precisamos pegar a versão mais recente do modelo e movê-lo para o estágio 'Production' - client = mlflow.tracking.MlflowClient() - - # Obter a versão mais recente registrada do modelo - model_versions = client.get_registered_model(model_name).latest_versions - - if not isinstance(model_versions, list): - raise ValueError('Model versions is not a list') - - max_version = max(model_versions, key=lambda x: int(x.version)).version - - # Mover a versão mais recente do modelo para o estágio de 'Production' - client.transition_model_version_stage( - name=model_name, version=max_version, stage='Production', archive_existing_versions=True - ) - - return {'model_name': model_name, 'version': max_version, 'mlflow_run_id': run_id} - - def update_production_model(self, experiment: str, model_name: str) -> dict: - """ - Update production model using the latest retraining run. - - This method orchestrates the complete production model update - process by identifying the most recent retraining run and - promoting it to production stage. - - Args: - experiment (str): MLFlow experiment name - model_name (str): Name of the MLFlow model - - Returns: - dict: Complete model update metadata containing: - - model_name (str): Name of the updated model - - version (str): New model version number - - mlflow_run_id (str): Source run ID - - mlflow_experiment_id (int): Experiment ID - """ - experiment_id = self.get_experiment(experiment) - run_id = self.get_experiment_last_run(experiment_id) - metadata = self.update_production_model_by_run_id(run_id, model_name) - - metadata['mlflow_experiment_id'] = experiment_id - - return metadata - - def get_next_run_name_new(self, experiment_name: str) -> str: + def get_next_run_name(self, experiment_name: str) -> str: """ Generates the next run name for a given experiment. diff --git a/tests/utils/repository/test_model_repository.py b/tests/utils/repository/test_model_repository.py index f77c8db..c18a631 100644 --- a/tests/utils/repository/test_model_repository.py +++ b/tests/utils/repository/test_model_repository.py @@ -1,9 +1,8 @@ -from datetime import UTC, datetime -from unittest.mock import ANY, MagicMock, call, patch +from unittest.mock import MagicMock, patch import numpy as np import pytest -from pandas import DataFrame, Timestamp +from pandas import DataFrame from model_manager.utils.repository.model_repository import MLFlowRepository @@ -32,477 +31,10 @@ metadata = { } -class Any: - pass - - -invalid_cases = [ - ({'value': {'2024-01-01 12:00:00': 1, 2024: 2}}), - ({'value': {'2024-01-01': 1, '2024-01-02': 2}}), - ({'value': {Any(): 1, Any(): 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=UTC): 1, - datetime(2025, 1, 2, 12, 0, 0, tzinfo=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=UTC): 1, - Timestamp(2026, 1, 2, 12, 0, 0, tzinfo=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 = MagicMock() - model_name = 'model' - - 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, 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.detect_and_parse_datetime_index.return_value.to_dict.return_value, - } - - -def test_transform_error(mlflow_repository): - data = MagicMock() - model_name = 'model' - - mlflow_repository.model_serving.get_cached_transform.side_effect = Exception('error') - - output = mlflow_repository.transform(model_name, data, {}, metadata['metadata']) - - mlflow_repository.model_serving.get_cached_transform.assert_called_once_with( - model_name, data, 0, 'sklearn', False, 'model', 'predict' - ) - - assert output == {'success': False, 'content': {'message': 'error', 'traceback': ANY}} - - -def test_predict_success(mlflow_repository): - data = DataFrame({'feat_1': {'index_1': 2, 'index_2': 3}}) - model_name = 'model' - mlflow_repository.model_serving.get_cached_predict.return_value = np.array([2, 3]) - - output = mlflow_repository.predict(model_name, data, {}, metadata['metadata']) - - mlflow_repository.model_serving.get_cached_predict.assert_called_once_with( - model_name, data, 0, 'pyfunc', False, 'model' - ) - - assert output['success'] is True - assert output['content'] == { - 'prediction': {'index_1': 2, 'index_2': 3}, - 'response_time': {'index_1': ANY, 'index_2': ANY}, - } - - -def test_predict_error(mlflow_repository): - data = DataFrame({'feat_1': {'index_1': 2, 'index_2': 3}}) - model_name = 'model' - - mlflow_repository.model_serving.get_cached_predict = MagicMock(side_effect=Exception('error')) - - output = mlflow_repository.predict(model_name, data, {}, metadata['metadata']) - - mlflow_repository.model_serving.get_cached_predict.assert_called_once_with( - model_name, data, 0, 'pyfunc', False, 'model' - ) - - assert output == {'success': False, 'content': {'message': 'error', 'traceback': ANY}} - - -@patch('model_manager.utils.repository.model_repository.mlflow') -def test_get_experiment_by_run_id(mlflow, mlflow_repository): - mlflow.get_run.return_value = MagicMock( - info=MagicMock( - experiment_id='0', - ) - ) - mlflow.get_experiment.return_value = MagicMock() - mlflow.get_experiment.return_value.name = 'test' - - output = mlflow_repository.get_experiment_by_run_id('0') - assert output == 'test' - mlflow.get_run.assert_called_once_with('0') - mlflow.get_experiment.assert_called_once_with('0') - - -@patch('model_manager.utils.repository.model_repository.mlflow') -def test_get_next_run_name(mlflow, mlflow_repository): - mlflow.search_runs.return_value = [1, 2, 3] - output = mlflow_repository.get_next_run_name('run') - assert output == 'run-4' - mlflow.search_runs.assert_called_once_with( - experiment_names=['run'], - order_by=['start_time desc'], - ) - - -@patch('model_manager.utils.repository.model_repository.mlflow') -def test_get_experiment_success(mlflow, mlflow_repository): - mlflow.get_experiment_by_name.return_value = MagicMock(experiment_id='0') - - output = mlflow_repository.get_experiment('test') - - assert output == '0' - - -@patch('model_manager.utils.repository.model_repository.mlflow') -def test_get_experiment_error(mlflow, mlflow_repository): - mlflow.get_experiment_by_name.return_value = None - - try: - mlflow_repository.get_experiment('test') - except ValueError as e: - assert str(e) == 'Experiment test not found' - else: - raise AssertionError('Expected exception') - - -@patch('model_manager.utils.repository.model_repository.mlflow') -def test_get_experiment_last_run(mlflow, mlflow_repository): - mlflow.search_runs.return_value = DataFrame( - { - 'params.retrain': ['True', 'False', 'True', 'False'], - 'end_time': ['2021-01-01', '2021-01-02', '2021-01-03', '2021-01-04'], - 'run_id': ['0', '1', '2', '3'], - } - ) - - output = mlflow_repository.get_experiment_last_run(0) - - mlflow.search_runs.assert_called_once_with( - experiment_ids=[0], - filter_string='', - output_format='pandas', - ) - - assert output == '2' - - -@patch('model_manager.utils.repository.model_repository.mlflow') -def test_get_experiment_last_run_error(mlflow, mlflow_repository): - mlflow.search_runs.return_value = [] - - try: - mlflow_repository.get_experiment_last_run(0) - except ValueError as e: - assert str(e) == 'Runs is not a pandas DataFrame' - else: - raise AssertionError('Expected exception') - - -@patch('model_manager.utils.repository.model_repository.mlflow.sklearn') -@patch('model_manager.utils.repository.model_repository.mlflow.set_experiment') -def test_create_model_experiment(set_experiment, sklearn, mlflow_repository): - mlflow_repository.model_serving.get_model_info = MagicMock(return_value='0') - mlflow_repository.model_serving.get_model_uri = MagicMock(return_value='test') - mlflow_repository.get_experiment_by_run_id = MagicMock() - - data_model_mock = MagicMock() - prediction_model_mock = MagicMock() - - sklearn.load_model.side_effect = [data_model_mock, prediction_model_mock] - - data_model_mock.fit.return_value = data_model_mock - data_model_mock.predict.return_value = DataFrame( - { - 'x': [10, 20, 30], - } - ) - data_model_mock.target_variable = 'y' - - prediction_model_mock.fit.return_value = prediction_model_mock - - data = DataFrame({'x': [1, 2, 3], 'y': [4, 5, 6]}) - - output = mlflow_repository.create_model_experiment('test', data) - - mlflow_repository.model_serving.get_model_info.assert_called_once_with('test') - mlflow_repository.model_serving.get_model_uri.assert_called_once_with('0', prediction=False) - - sklearn.load_model.assert_has_calls( - [ - call(mlflow_repository.model_serving.get_model_uri.return_value), - call('models:/test/production'), - ] - ) - assert sklearn.load_model.call_count == 2 - - data_model_mock.fit.assert_called_once_with(data) - data_model_mock.predict.assert_called_once_with(data) - - fit_args = prediction_model_mock.fit.call_args[0][0] - assert fit_args.equals( - DataFrame( - { - 'x': [10, 20, 30], - 'y': [4, 5, 6], - } - ) - ) - - mlflow_repository.get_experiment_by_run_id.assert_called_once_with('0') - - set_experiment.assert_called_once_with(mlflow_repository.get_experiment_by_run_id.return_value) - - assert output == ( - prediction_model_mock, - data_model_mock, - mlflow_repository.get_experiment_by_run_id.return_value, - ) - - -@patch('model_manager.utils.repository.model_repository.path.exists') -@patch('model_manager.utils.repository.model_repository.remove') -@patch('model_manager.utils.repository.model_repository.mlflow.start_run') -@patch('model_manager.utils.repository.model_repository.mlflow.log_param') -@patch('model_manager.utils.repository.model_repository.mlflow.sklearn.log_model') -@patch('model_manager.utils.repository.model_repository.mlflow.log_artifact') -def test_perform_model_retrain( - log_artifact, log_model, log_param, start_run, mock_remove, mock_path_exists, mlflow_repository -): - # Create mock models with attributes to test the for loops (lines 268-274) - prediction_model_mock = MagicMock() - prediction_model_mock.__dict__ = {'model': 'pred_model', 'param1': 'value1', 'param2': 'value2'} - - data_model_mock = MagicMock() - data_model_mock.__dict__ = {'model': 'data_model', 'param3': 'value3', 'param4': 'value4'} - - experiment = 'test' - model_name = 'test' - data = MagicMock() - - mlflow_repository.get_next_run_name = MagicMock(return_value='test-1') - run = MagicMock() - start_run.__enter__.return_value = run - mock_path_exists.return_value = True - - output = mlflow_repository.perform_model_retrain( - prediction_model_mock, data_model_mock, experiment, model_name, data - ) - - mlflow_repository.get_next_run_name.assert_called_once_with(experiment) - start_run.assert_called_once_with( - run_name='test-1', description='Retrain model test with new data' - ) - - log_model.assert_has_calls( - [ - call(data_model_mock, 'data_model'), - call(prediction_model_mock, 'prediction_model'), - ] - ) - - data.to_csv.assert_called_once_with('temp/raw_data_test.csv', index=True) - - log_artifact.assert_called_once_with('temp/raw_data_test.csv') - - # Verify that model attributes were logged (excluding 'model' key) - log_param.assert_has_calls( - [ - call('param1', 'value1'), # from prediction_model - call('param2', 'value2'), # from prediction_model - call('param3', 'value3'), # from data_model - call('param4', 'value4'), # from data_model - call('retrain', True), - ], - any_order=True, - ) - - # Verify temp file cleanup - mock_path_exists.assert_called_once_with('temp/raw_data_test.csv') - mock_remove.assert_called_once_with('temp/raw_data_test.csv') - - assert output == ('Model retrained successfully', experiment) - - -@patch('model_manager.utils.repository.model_repository.path.exists') -@patch('model_manager.utils.repository.model_repository.remove') -@patch('model_manager.utils.repository.model_repository.mlflow.start_run') -@patch('model_manager.utils.repository.model_repository.mlflow.log_param') -@patch('model_manager.utils.repository.model_repository.mlflow.sklearn.log_model') -@patch('model_manager.utils.repository.model_repository.mlflow.log_artifact') -def test_perform_model_retrain_file_not_exists( - log_artifact, log_model, log_param, start_run, mock_remove, mock_path_exists, mlflow_repository -): - """Test perform_model_retrain when temp file doesn't exist (line 291->294 branch).""" - prediction_model_mock = MagicMock() - prediction_model_mock.__dict__ = {'model': 'pred_model'} - - data_model_mock = MagicMock() - data_model_mock.__dict__ = {'model': 'data_model'} - - experiment = 'test' - model_name = 'test' - data = MagicMock() - - mlflow_repository.get_next_run_name = MagicMock(return_value='test-1') - run = MagicMock() - start_run.__enter__.return_value = run - mock_path_exists.return_value = False # File doesn't exist - - output = mlflow_repository.perform_model_retrain( - prediction_model_mock, data_model_mock, experiment, model_name, data - ) - - # Verify temp file cleanup was checked but not executed - mock_path_exists.assert_called_once_with('temp/raw_data_test.csv') - mock_remove.assert_not_called() # Should not be called when file doesn't exist - - assert output == ('Model retrained successfully', experiment) - - -def test_retrain_model(mlflow_repository): - data = MagicMock() - model_name = 'test' - - mlflow_repository.create_model_experiment = MagicMock( - return_value=('data_model', 'prediction_model', '0') - ) - - mlflow_repository.perform_model_retrain = MagicMock(return_value='Model retrained successfully') - - output = mlflow_repository.retrain_model(data, model_name) - - mlflow_repository.create_model_experiment.assert_called_once_with(model_name, data) - - mlflow_repository.perform_model_retrain.assert_called_once_with( - 'data_model', 'prediction_model', '0', model_name, data - ) - - assert output == 'Model retrained successfully' - - -@patch('model_manager.utils.repository.model_repository.mlflow') -def test_update_production_model_by_run_id(mlflow, mlflow_repository): - client_mock = MagicMock() - mlflow.tracking.MlflowClient.return_value = client_mock - - client_mock.get_registered_model.return_value = MagicMock( - latest_versions=[ - MagicMock(version='1'), - MagicMock(version='2'), - MagicMock(version='3'), - ] - ) - output = mlflow_repository.update_production_model_by_run_id('0', 'test') - - mlflow.register_model.assert_called_once_with( - 'runs:/0/prediction_model', - 'test', - ) - - mlflow.tracking.MlflowClient.assert_called_once() - client_mock.get_registered_model.assert_called_once_with('test') - client_mock.transition_model_version_stage.assert_called_once_with( - name='test', - version='3', - stage='Production', - archive_existing_versions=True, - ) - - assert output == { - 'model_name': 'test', - 'version': '3', - 'mlflow_run_id': '0', - } - - -@patch('model_manager.utils.repository.model_repository.mlflow') -def test_update_production_model_by_run_id_error(mlflow, mlflow_repository): - mlflow.tracking.MlflowClient.return_value = MagicMock( - get_registered_model=MagicMock(return_value=MagicMock(latest_versions={})) - ) - - try: - mlflow_repository.update_production_model_by_run_id('0', 'test') - except Exception as e: # noqa: BLE001 - assert str(e) == 'Model versions is not a list' - else: - raise AssertionError('Expected exception') - - -def test_update_production_model(mlflow_repository): - connector = mlflow_repository - - with patch.object(connector, 'get_experiment', return_value='0') as get_experiment: - with patch.object( - connector, 'get_experiment_last_run', return_value='2' - ) as get_experiment_last_run: - with patch.object( - connector, - 'update_production_model_by_run_id', - return_value={'model_name': 'test', 'version': '3', 'mlflow_run_id': '0'}, - ) as update_production_model_by_run_id: - output = connector.update_production_model('0', 'test') - - get_experiment.assert_called_once_with('0') - get_experiment_last_run.assert_called_once_with('0') - update_production_model_by_run_id.assert_called_once_with('2', 'test') - - assert output == { - 'model_name': 'test', - 'version': '3', - 'mlflow_run_id': '0', - 'mlflow_experiment_id': '0', - } - - # ========== Tests for Model Artifact Generation Methods ========== -def test_get_next_run_name_new(mlflow_repository): +def test_get_next_run_name(mlflow_repository): """Test get_next_run_name generates correct run name based on existing runs.""" mlflow_repository.model_serving.search_runs_by_name.return_value = [ MagicMock(), @@ -510,7 +42,7 @@ def test_get_next_run_name_new(mlflow_repository): MagicMock(), ] - result = mlflow_repository.get_next_run_name_new('test_experiment') + result = mlflow_repository.get_next_run_name('test_experiment') mlflow_repository.model_serving.search_runs_by_name.assert_called_once_with( experiment_names=['test_experiment'], order_by=['start_time desc'] @@ -518,13 +50,13 @@ def test_get_next_run_name_new(mlflow_repository): assert result == 'test_experiment-4' -def test_get_next_run_name_new_first_run(mlflow_repository): +def test_get_next_run_name_first_run(mlflow_repository): """Test get_next_run_name for first run (no existing runs).""" mlflow_repository.model_serving.search_runs_by_name.return_value = [] - result = mlflow_repository.get_next_run_name_new('new_experiment') + result = mlflow_repository.get_next_run_name('test_experiment') - assert result == 'new_experiment-1' + assert result == 'test_experiment-1' @patch('model_manager.utils.repository.model_repository.path')