diff --git a/packages/data-designer-engine/src/data_designer/engine/models/telemetry.py b/packages/data-designer-engine/src/data_designer/engine/models/telemetry.py index b2b01384b..b2cf04a45 100644 --- a/packages/data-designer-engine/src/data_designer/engine/models/telemetry.py +++ b/packages/data-designer-engine/src/data_designer/engine/models/telemetry.py @@ -55,14 +55,21 @@ class DeploymentTypeEnum(str, Enum): UNDEFINED = "undefined" -_deployment_type_raw = os.getenv("NEMO_DEPLOYMENT_TYPE", "library").lower() -try: - DEPLOYMENT_TYPE = DeploymentTypeEnum(_deployment_type_raw) -except ValueError: - valid_values = [e.value for e in DeploymentTypeEnum] - raise ValueError( - f"Invalid NEMO_DEPLOYMENT_TYPE: {_deployment_type_raw!r}. Must be one of: {valid_values}" - ) from None +def _normalize_deployment_type(value: str | None) -> DeploymentTypeEnum: + if value is None: + return DeploymentTypeEnum.LIBRARY + + try: + return DeploymentTypeEnum(value.lower()) + except ValueError: + return DeploymentTypeEnum.UNDEFINED + + +def get_nemo_deployment_type() -> DeploymentTypeEnum: + return _normalize_deployment_type(os.getenv("NEMO_DEPLOYMENT_TYPE")) + + +DEPLOYMENT_TYPE = get_nemo_deployment_type() class TaskStatusEnum(str, Enum): diff --git a/packages/data-designer-engine/tests/engine/models/test_telemetry.py b/packages/data-designer-engine/tests/engine/models/test_telemetry.py index f78739097..dc9b52070 100644 --- a/packages/data-designer-engine/tests/engine/models/test_telemetry.py +++ b/packages/data-designer-engine/tests/engine/models/test_telemetry.py @@ -5,6 +5,8 @@ from datetime import datetime, timezone +import pytest + from data_designer.engine.models.telemetry import ( DeploymentTypeEnum, InferenceEvent, @@ -12,6 +14,7 @@ QueuedEvent, TaskStatusEnum, build_payload, + get_nemo_deployment_type, ) @@ -30,3 +33,9 @@ def test_nvidia_internal_deployment_type_uses_schema_version_1_9() -> None: assert payload["eventSchemaVer"] == "1.9" assert payload["events"][0]["parameters"]["deploymentType"] == "nvidia-internal" + + +def test_unrecognized_env_deployment_type_defaults_to_undefined(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setenv("NEMO_DEPLOYMENT_TYPE", "unrecognized") + + assert get_nemo_deployment_type() == DeploymentTypeEnum.UNDEFINED