"""Azure ML SDK backend with mocked MLClient.""" from unittest.mock import MagicMock, patch from uuid import uuid4 import pytest from forja_core.config import Settings from forja_core.domain.training import TrainingHyperparameters, TrainingJobSpec from forja_core.training.azure_ml import AzureMLTrainingBackend from forja_core.training.azure_ml_sdk import poll_job_status, submit_command_job @pytest.fixture def azure_settings() -> Settings: return Settings( training_backend="azure_ml", azure_ml_subscription_id="sub-123", azure_ml_resource_group="rg-forja", azure_ml_workspace_name="ws-forja", azure_ml_compute="cpu-cluster", ) @pytest.fixture def spec() -> TrainingJobSpec: return TrainingJobSpec( dataset_name="incident_sft", dataset_version="v1", model_name="gpt4o_lora_base", model_version="v1", agent_name="test_agent", hyperparameters=TrainingHyperparameters(), ) @patch("forja_core.training.azure_ml_sdk._ml_client") def test_submit_command_job_returns_azureml_prefix( mock_client_factory: MagicMock, azure_settings: Settings, spec: TrainingJobSpec, ) -> None: client = MagicMock() mock_client_factory.return_value = client created = MagicMock() created.name = "forja-job-abc" client.jobs.create_or_update.return_value = created run_id = uuid4() external_id = submit_command_job(azure_settings, spec, run_id) assert external_id == "azureml:forja-job-abc" client.jobs.create_or_update.assert_called_once() @patch("forja_core.training.azure_ml_sdk._ml_client") def test_poll_job_status_maps_completed( mock_client_factory: MagicMock, azure_settings: Settings, ) -> None: client = MagicMock() mock_client_factory.return_value = client job = MagicMock() job.status = "Completed" client.jobs.get.return_value = job result = poll_job_status(azure_settings, "forja-job-abc") assert result.status == "succeeded" assert len(result.artifacts) == 1 assert "azureml://" in result.artifacts[0].uri @pytest.mark.asyncio async def test_backend_uses_sdk_when_configured( azure_settings: Settings, spec: TrainingJobSpec, ) -> None: backend = AzureMLTrainingBackend(azure_settings) assert backend._use_sdk is True with patch( "forja_core.training.azure_ml_sdk.submit_command_job", return_value="azureml:job-1", ) as mock_submit: job_id = await backend.submit(spec, uuid4()) assert job_id == "azureml:job-1" mock_submit.assert_called_once() with patch( "forja_core.training.azure_ml_sdk.poll_job_status", return_value=MagicMock(status="running", message=None, artifacts=[]), ) as mock_poll: status = await backend.status("azureml:job-1") assert status.status == "running" mock_poll.assert_called_once_with(azure_settings, "job-1")