diff --git a/src/autointent/_optimization_config.py b/src/autointent/_optimization_config.py index 547474a92..64767c7bd 100644 --- a/src/autointent/_optimization_config.py +++ b/src/autointent/_optimization_config.py @@ -2,7 +2,7 @@ from typing import TYPE_CHECKING, Any -from pydantic import BaseModel, Field, PositiveInt, field_validator +from pydantic import BaseModel, Field, NonNegativeInt, field_validator from .configs import ( CrossEncoderConfig, @@ -49,7 +49,7 @@ def validate_embedder_config(cls, v: Any) -> EmbedderConfig: # noqa: ANN401 hpo_config: HPOConfig = HPOConfig() - seed: PositiveInt = 42 + seed: NonNegativeInt = 42 @classmethod def from_preset(cls, preset: SearchSpacePreset) -> OptimizationConfig: diff --git a/tests/configs/test_full_config.py b/tests/configs/test_full_config.py index b95b4a63f..02f732204 100644 --- a/tests/configs/test_full_config.py +++ b/tests/configs/test_full_config.py @@ -17,3 +17,13 @@ def test_not_valid_reporting() -> None: with pytest.raises(ValidationError): OptimizationConfig(**config) + + +def test_optimization_config_accepts_seed_zero() -> None: + config = OptimizationConfig(seed=0, search_space=[]) + assert config.seed == 0 + + +def test_optimization_config_rejects_negative_seed() -> None: + with pytest.raises(ValidationError): + OptimizationConfig(seed=-1, search_space=[])