- Made `date_column` a required field in `TrainModelParams`, ensuring it must be present in the input data. - Updated related documentation in `input-sample.md`, `README.md`, and various test scenarios to reflect the change in requirement. - Adjusted the handling of `date_format` to default to `yyyy-MM-dd HH:mm:ss` if omitted, enhancing usability. - Refined test scenarios to include new examples and ensure compliance with the updated parameter structure. These changes improve the robustness of the model training workflow and clarify the expectations for input data.
249 lines
7.9 KiB
Python
249 lines
7.9 KiB
Python
"""
|
||
End-to-end tests for TrainModel parameter validation paths.
|
||
|
||
Covers scenarios 2.1.x: workflows that must terminate with
|
||
ORCHESTRATOR_VALIDATION_ERROR due to invalid parameter values.
|
||
"""
|
||
|
||
import pytest
|
||
from temporalio.client import WorkflowFailureError
|
||
from temporalio.testing import WorkflowEnvironment
|
||
from temporalio.worker import Worker
|
||
|
||
from e2e.helpers import (
|
||
assert_experiment_error,
|
||
insert_experiment_run,
|
||
load_scenario,
|
||
make_workflow_id,
|
||
start_and_await_workflow,
|
||
)
|
||
from model_manager.workflows.train_model import TrainModel
|
||
|
||
# Base experiment_run ids for validation test scenarios (offset to avoid collision)
|
||
_VALIDATION_ID_BASE = 3000
|
||
|
||
|
||
def _exception_chain_text(exc: BaseException) -> str:
|
||
"""Concatenate messages from an exception __cause__/__context__ chain."""
|
||
parts: list[str] = []
|
||
cur: BaseException | None = exc
|
||
seen: set[int] = set()
|
||
while cur is not None and id(cur) not in seen:
|
||
seen.add(id(cur))
|
||
text = str(cur).strip()
|
||
if text:
|
||
parts.append(text)
|
||
cur = cur.__cause__ or getattr(cur, '__context__', None)
|
||
return ' | '.join(parts).lower()
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
@pytest.mark.integration
|
||
async def test_scenario_2_1_1_train_size_out_of_range(
|
||
temporal_test_env: WorkflowEnvironment,
|
||
temporal_worker: Worker,
|
||
postgres_engine,
|
||
):
|
||
"""Scenario 2.1.1 – train_size=5 violates the 10–100 business rule.
|
||
|
||
Expected: workflow updates status → ORCHESTRATOR_VALIDATION_ERROR
|
||
and error_message references 'train_size'.
|
||
"""
|
||
experiment_run_id = _VALIDATION_ID_BASE + 1
|
||
scenario = load_scenario('01-linear-regression-basic.json')
|
||
scenario = {**scenario, 'experiment_run_id': experiment_run_id, 'train_size': 5}
|
||
insert_experiment_run(postgres_engine, experiment_run_id)
|
||
|
||
with pytest.raises(Exception):
|
||
await start_and_await_workflow(
|
||
temporal_test_env.client,
|
||
TrainModel.run,
|
||
scenario,
|
||
make_workflow_id('test-s2-1-1'),
|
||
)
|
||
|
||
assert_experiment_error(
|
||
postgres_engine,
|
||
experiment_run_id,
|
||
expected_status='ORCHESTRATOR_VALIDATION_ERROR',
|
||
error_substr='train_size',
|
||
)
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
@pytest.mark.integration
|
||
async def test_scenario_2_1_2_empty_variable_columns(
|
||
temporal_test_env: WorkflowEnvironment,
|
||
temporal_worker: Worker,
|
||
postgres_engine,
|
||
):
|
||
"""Scenario 2.1.2 – variable_columns=[] → ORCHESTRATOR_VALIDATION_ERROR."""
|
||
experiment_run_id = _VALIDATION_ID_BASE + 2
|
||
scenario = load_scenario('01-linear-regression-basic.json')
|
||
scenario = {**scenario, 'experiment_run_id': experiment_run_id, 'variable_columns': []}
|
||
insert_experiment_run(postgres_engine, experiment_run_id)
|
||
|
||
with pytest.raises(Exception):
|
||
await start_and_await_workflow(
|
||
temporal_test_env.client,
|
||
TrainModel.run,
|
||
scenario,
|
||
make_workflow_id('test-s2-1-2'),
|
||
)
|
||
|
||
assert_experiment_error(
|
||
postgres_engine,
|
||
experiment_run_id,
|
||
expected_status='ORCHESTRATOR_VALIDATION_ERROR',
|
||
error_substr='variable_columns',
|
||
)
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
@pytest.mark.integration
|
||
async def test_scenario_2_1_3_invalid_date_format(
|
||
temporal_test_env: WorkflowEnvironment,
|
||
temporal_worker: Worker,
|
||
postgres_engine,
|
||
):
|
||
"""Scenario 2.1.3 – date_format='INVALID' is not in the allowed list."""
|
||
experiment_run_id = _VALIDATION_ID_BASE + 3
|
||
scenario = load_scenario('01-linear-regression-basic.json')
|
||
scenario = {**scenario, 'experiment_run_id': experiment_run_id, 'date_format': 'INVALID'}
|
||
insert_experiment_run(postgres_engine, experiment_run_id)
|
||
|
||
with pytest.raises(Exception):
|
||
await start_and_await_workflow(
|
||
temporal_test_env.client,
|
||
TrainModel.run,
|
||
scenario,
|
||
make_workflow_id('test-s2-1-3'),
|
||
)
|
||
|
||
assert_experiment_error(
|
||
postgres_engine,
|
||
experiment_run_id,
|
||
expected_status='ORCHESTRATOR_VALIDATION_ERROR',
|
||
error_substr='date_format',
|
||
)
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
@pytest.mark.integration
|
||
async def test_scenario_2_1_4_whitespace_only_model_name(
|
||
temporal_test_env: WorkflowEnvironment,
|
||
temporal_worker: Worker,
|
||
postgres_engine,
|
||
):
|
||
"""Scenario 2.1.4 – model_name=' ' (whitespace) → ORCHESTRATOR_VALIDATION_ERROR."""
|
||
experiment_run_id = _VALIDATION_ID_BASE + 4
|
||
scenario = load_scenario('01-linear-regression-basic.json')
|
||
scenario = {**scenario, 'experiment_run_id': experiment_run_id, 'model_name': ' '}
|
||
insert_experiment_run(postgres_engine, experiment_run_id)
|
||
|
||
with pytest.raises(Exception):
|
||
await start_and_await_workflow(
|
||
temporal_test_env.client,
|
||
TrainModel.run,
|
||
scenario,
|
||
make_workflow_id('test-s2-1-4'),
|
||
)
|
||
|
||
assert_experiment_error(
|
||
postgres_engine,
|
||
experiment_run_id,
|
||
expected_status='ORCHESTRATOR_VALIDATION_ERROR',
|
||
error_substr='model_name',
|
||
)
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
@pytest.mark.integration
|
||
async def test_scenario_2_1_5_unknown_model_type(
|
||
temporal_test_env: WorkflowEnvironment,
|
||
temporal_worker: Worker,
|
||
postgres_engine,
|
||
):
|
||
"""Scenario 2.1.5 – model_type='totally_unknown' → ORCHESTRATOR_VALIDATION_ERROR.
|
||
|
||
The PluginStore will not find this model in the Gitea repo, causing
|
||
load_model_metadata to fail before validate_train_params is even called.
|
||
"""
|
||
experiment_run_id = _VALIDATION_ID_BASE + 5
|
||
scenario = load_scenario('01-linear-regression-basic.json')
|
||
scenario = {
|
||
**scenario,
|
||
'experiment_run_id': experiment_run_id,
|
||
'model_type': 'totally_unknown_model',
|
||
}
|
||
insert_experiment_run(postgres_engine, experiment_run_id)
|
||
|
||
with pytest.raises(Exception):
|
||
await start_and_await_workflow(
|
||
temporal_test_env.client,
|
||
TrainModel.run,
|
||
scenario,
|
||
make_workflow_id('test-s2-1-5'),
|
||
)
|
||
|
||
assert_experiment_error(
|
||
postgres_engine,
|
||
experiment_run_id,
|
||
expected_status='ORCHESTRATOR_VALIDATION_ERROR',
|
||
error_substr='totally_unknown_model',
|
||
)
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
@pytest.mark.integration
|
||
async def test_scenario_2_1_6_missing_target_variable(
|
||
temporal_test_env: WorkflowEnvironment,
|
||
temporal_worker: Worker,
|
||
postgres_engine,
|
||
):
|
||
"""Scenario 2.1.6 – target_variable='' (empty string) → ORCHESTRATOR_VALIDATION_ERROR."""
|
||
experiment_run_id = _VALIDATION_ID_BASE + 6
|
||
scenario = load_scenario('01-linear-regression-basic.json')
|
||
scenario = {**scenario, 'experiment_run_id': experiment_run_id, 'target_variable': ''}
|
||
insert_experiment_run(postgres_engine, experiment_run_id)
|
||
|
||
with pytest.raises(Exception):
|
||
await start_and_await_workflow(
|
||
temporal_test_env.client,
|
||
TrainModel.run,
|
||
scenario,
|
||
make_workflow_id('test-s2-1-6'),
|
||
)
|
||
|
||
assert_experiment_error(
|
||
postgres_engine,
|
||
experiment_run_id,
|
||
expected_status='ORCHESTRATOR_VALIDATION_ERROR',
|
||
error_substr='target_variable',
|
||
)
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
@pytest.mark.integration
|
||
async def test_scenario_2_1_7_missing_experiment_run_id(
|
||
temporal_test_env: WorkflowEnvironment,
|
||
temporal_worker: Worker,
|
||
):
|
||
"""Scenario 2.1.7 – experiment_run_id missing → workflow raises ValueError immediately.
|
||
|
||
No DB row is inserted because experiment_run_id is mandatory to even
|
||
know which row to update. The workflow should raise before any DB call.
|
||
"""
|
||
scenario = load_scenario('01-linear-regression-basic.json')
|
||
scenario = {k: v for k, v in scenario.items() if k != 'experiment_run_id'}
|
||
|
||
with pytest.raises(WorkflowFailureError) as excinfo:
|
||
await start_and_await_workflow(
|
||
temporal_test_env.client,
|
||
TrainModel.run,
|
||
scenario,
|
||
make_workflow_id('test-s2-1-7'),
|
||
)
|
||
combined = _exception_chain_text(excinfo.value)
|
||
assert 'experiment_run_id' in combined
|