pagina unica
This commit is contained in:
@@ -0,0 +1,96 @@
|
||||
"""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")
|
||||
Reference in New Issue
Block a user