97 lines
2.9 KiB
Python
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")
|