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

44 lines
1.4 KiB
Python

"""Tests for mock and Azure ML stub training backends."""
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.mock import MockTrainingBackend
@pytest.fixture
def spec() -> TrainingJobSpec:
return TrainingJobSpec(
dataset_name="incident_sft",
dataset_version="v1",
model_name="gpt4o_lora_base",
model_version="v1",
hyperparameters=TrainingHyperparameters(epochs=2),
)
@pytest.mark.asyncio
async def test_mock_submit_and_status_succeeds(spec: TrainingJobSpec) -> None:
backend = MockTrainingBackend()
run_id = uuid4()
job_id = await backend.submit(spec, run_id)
assert job_id == f"mock:{run_id}"
status = await backend.status(job_id)
assert status.status == "succeeded"
assert len(status.artifacts) == 1
@pytest.mark.asyncio
async def test_azure_stub_submit_and_status_succeeds(spec: TrainingJobSpec) -> None:
backend = AzureMLTrainingBackend(settings=Settings())
run_id = uuid4()
job_id = await backend.submit(spec, run_id)
assert job_id.startswith("azureml-stub:")
status = await backend.status(job_id)
assert status.status == "succeeded"
assert "azureml://" in status.artifacts[0].uri