Files
forja/tests/unit/test_azure_ml_sdk.py
2026-06-10 18:21:33 +02:00

97 lines
2.9 KiB
Python

"""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")