diff --git a/laborious/utils/repository/model_repository.py b/laborious/utils/repository/model_repository.py index b4fdad5..6a25b8c 100644 --- a/laborious/utils/repository/model_repository.py +++ b/laborious/utils/repository/model_repository.py @@ -170,8 +170,19 @@ class MLFlowRepository(SientiaMonitoring): # Sort by version number to get the latest latest_version = max(stage_versions, key=lambda v: int(v.version)) - run_id = latest_version.source.split('/') - return run_id[2] + source = latest_version.source + if source is None: + raise mlflow.exceptions.MlflowException( + f"Model '{model_name}' version '{latest_version.version}' in stage '{stage}' " + 'has no source URI to resolve run ID.' + ) + parts = source.split('/') + if len(parts) <= 2 or not parts[2]: + raise mlflow.exceptions.MlflowException( + f"Model '{model_name}' version '{latest_version.version}' in stage '{stage}' " + f"has invalid source URI '{source}' for run ID resolution." + ) + return parts[2] def get_next_run_name(self, model_name: str) -> str: """ diff --git a/requirements-light.txt b/requirements-light.txt index 3624685..302fb79 100644 --- a/requirements-light.txt +++ b/requirements-light.txt @@ -4,6 +4,7 @@ sqlalchemy asyncua==1.0.6 redis sientia_do>=1.12.1 +mlflow prometheus-client botocore boto3 diff --git a/tests/laborious/utils/repository/test_model_repository.py b/tests/laborious/utils/repository/test_model_repository.py index 15f8c6b..54a75bb 100644 --- a/tests/laborious/utils/repository/test_model_repository.py +++ b/tests/laborious/utils/repository/test_model_repository.py @@ -142,6 +142,41 @@ def test_get_model_run_id_success(mlflow_repository): assert output == '1' +def test_get_model_run_id_missing_source(mlflow_repository): + mlflow_repository.client.search_registered_models.return_value = [MagicMock(name='test')] + + mlflow_repository.client.search_model_versions.return_value = [ + MagicMock(current_stage='Production', version='1', source='runs/test/0'), + MagicMock(current_stage='Production', version='2', source=None), + ] + + with pytest.raises(mlflow_lib.exceptions.MlflowException) as exc_info: + mlflow_repository.get_model_run_id('test') + + assert ( + str(exc_info.value) + == "Model 'test' version '2' in stage 'Production' has no source URI to resolve run ID." + ) + + +def test_get_model_run_id_invalid_source(mlflow_repository): + mlflow_repository.client.search_registered_models.return_value = [MagicMock(name='test')] + + mlflow_repository.client.search_model_versions.return_value = [ + MagicMock(current_stage='Production', version='1', source='runs/test/0'), + MagicMock(current_stage='Production', version='2', source='runs/test'), + ] + + with pytest.raises(mlflow_lib.exceptions.MlflowException) as exc_info: + mlflow_repository.get_model_run_id('test') + + assert ( + str(exc_info.value) + == "Model 'test' version '2' in stage 'Production' has invalid source URI " + "'runs/test' for run ID resolution." + ) + + 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')