diff --git a/docs/Flash_SDK_Reference.md b/docs/Flash_SDK_Reference.md index dd231a30..2a99403e 100644 --- a/docs/Flash_SDK_Reference.md +++ b/docs/Flash_SDK_Reference.md @@ -32,6 +32,7 @@ Endpoint( template: Optional[PodTemplate] = None, min_cuda_version: Optional[CudaVersion | str] = None, max_concurrency: int = 1, + workers_standby: Optional[int] = None, ) ``` @@ -43,7 +44,8 @@ Endpoint( | `id` | `str` | `None` | Existing endpoint ID. Mutually exclusive with `image`. | | `gpu` | `GpuGroup`, `GpuType`, or list | `None` | GPU type(s). Mutually exclusive with `cpu`. Defaults to `GpuGroup.ANY` if neither is set. | | `cpu` | `str`, `CpuInstanceType`, or list | `None` | CPU instance type(s). Mutually exclusive with `gpu`. | -| `workers` | `int` or `(int, int)` | `(0, 3)` | Worker scaling. `N` = `(0, N)`. `(min, max)` = explicit range. | +| `workers` | `int` or `(int, int)` | `(0, 3)` | Worker scaling. `N` = `(0, N)`. `(min, max)` = explicit range. Standby (active) workers default to `min`, so `(0, N)` scales to zero. | +| `workers_standby` | `int` | `None` | Override the number of active (standby) workers kept warm. Defaults to the `workers` minimum. Must be within the `workers` range. Set to `0` to force scale-to-zero. | | `idle_timeout` | `int` | `60` | Seconds before idle workers scale down. | | `dependencies` | `list[str]` | `None` | Python packages to install (e.g., `["torch", "numpy==1.24"]`). | | `system_dependencies` | `list[str]` | `None` | System packages to install. | @@ -67,6 +69,7 @@ Endpoint( - `id` and `image` are mutually exclusive - `name` or `id` is required - `workers` rejects negative values and `min > max` +- `workers_standby` must be within the `workers` range when set - `max_concurrency` must be >= 1 ### Usage Patterns diff --git a/src/runpod_flash/core/resources/serverless.py b/src/runpod_flash/core/resources/serverless.py index 3dd46e01..a73540fb 100644 --- a/src/runpod_flash/core/resources/serverless.py +++ b/src/runpod_flash/core/resources/serverless.py @@ -286,6 +286,10 @@ class ServerlessResource(DeployableResource): type: Optional[ServerlessType] = ServerlessType.QB workersMax: Optional[int] = DEFAULT_WORKERS_MAX workersMin: Optional[int] = DEFAULT_WORKERS_MIN + # RunPod defaults an omitted workersStandby server-side (to a warm worker), + # which breaks scale-to-zero for workersMin=0. Defaulting to workersMin + # keeps the deployed config matching the requested (min, max) range. + workersStandby: Optional[int] = None workersPFBTarget: Optional[int] = 0 # === Private Attributes === @@ -640,6 +644,33 @@ def validate_worker_range(self): ) return self + @model_validator(mode="after") + def default_and_validate_workers_standby(self): + """Default workersStandby to workersMin and validate explicit values. + + Only validate explicit standby values against bounds when the + (workersMin, workersMax) range itself is valid; invalid ranges raise + in validate_worker_range. + """ + if self.workersStandby is None: + self.workersStandby = self.workersMin + return self + + range_valid = ( + self.workersMin is None + or self.workersMax is None + or self.workersMin <= self.workersMax + ) + if range_valid and not ( + (self.workersMin is None or self.workersStandby >= self.workersMin) + and (self.workersMax is None or self.workersStandby <= self.workersMax) + ): + raise ValueError( + f"workersStandby ({self.workersStandby}) must be between " + f"workersMin ({self.workersMin}) and workersMax ({self.workersMax})" + ) + return self + def _has_cpu_instances(self) -> bool: """Check if endpoint has CPU instances configured. diff --git a/src/runpod_flash/endpoint.py b/src/runpod_flash/endpoint.py index 632369db..ff4698b2 100644 --- a/src/runpod_flash/endpoint.py +++ b/src/runpod_flash/endpoint.py @@ -417,6 +417,7 @@ def __init__( template: Optional[PodTemplate] = None, min_cuda_version: Optional[CudaVersion | str] = CudaVersion.V12_8, max_concurrency: int = 1, + workers_standby: Optional[int] = None, ): if gpu is not None and cpu is not None: raise ValueError( @@ -442,6 +443,16 @@ def __init__( self._cpu = _normalize_cpu(cpu) self._is_cpu = _is_cpu_config(cpu) self._workers_min, self._workers_max = _normalize_workers(workers) + if workers_standby is not None and not ( + self._workers_min <= workers_standby <= self._workers_max + ): + raise ValueError( + f"workers_standby ({workers_standby}) must be between " + f"workers min ({self._workers_min}) and max ({self._workers_max})" + ) + # workers_standby defaults to workers min on the resource config, so + # workers=(0, N) scales to zero as documented. + self.workers_standby = workers_standby self.idle_timeout = idle_timeout self.dependencies = dependencies self.system_dependencies = system_dependencies @@ -570,6 +581,9 @@ def _build_resource_config(self): "scalerValue": self.scaler_value, } + if self.workers_standby is not None: + kwargs["workersStandby"] = self.workers_standby + if self.template is not None: # serialize to dict to avoid pydantic model identity issues # when modules get re-imported across different contexts diff --git a/tests/unit/resources/test_serverless.py b/tests/unit/resources/test_serverless.py index 0db5b1a2..e682700a 100644 --- a/tests/unit/resources/test_serverless.py +++ b/tests/unit/resources/test_serverless.py @@ -377,6 +377,50 @@ def test_workers_min_cannot_exceed_workers_max(self): ): ServerlessResource(name="test", workersMin=5, workersMax=1) + def test_workers_standby_defaults_to_workers_min(self): + """workersStandby follows workersMin so workers=(0, N) scales to zero.""" + serverless = ServerlessResource(name="test", workersMin=0, workersMax=1) + assert serverless.workersStandby == 0 + + def test_workers_standby_defaults_to_nonzero_min(self): + serverless = ServerlessResource(name="test", workersMin=2, workersMax=5) + assert serverless.workersStandby == 2 + + def test_workers_standby_in_deploy_payload(self): + """The saveEndpoint payload carries the requested worker scaling config.""" + serverless = ServerlessResource(name="test", workersMin=0, workersMax=1) + payload = serverless.model_dump( + exclude=serverless._payload_exclude(), exclude_none=True, mode="json" + ) + assert payload["workersMin"] == 0 + assert payload["workersStandby"] == 0 + + def test_workers_standby_explicit_value_respected(self): + serverless = ServerlessResource( + name="test", workersMin=0, workersMax=3, workersStandby=1 + ) + assert serverless.workersStandby == 1 + + def test_workers_standby_above_max_raises(self): + with pytest.raises( + ValueError, + match=r"workersStandby \(4\) must be between workersMin \(0\) " + r"and workersMax \(3\)", + ): + ServerlessResource( + name="test", workersMin=0, workersMax=3, workersStandby=4 + ) + + def test_workers_standby_below_min_raises(self): + with pytest.raises( + ValueError, + match=r"workersStandby \(1\) must be between workersMin \(2\) " + r"and workersMax \(5\)", + ): + ServerlessResource( + name="test", workersMin=2, workersMax=5, workersStandby=1 + ) + @pytest.mark.parametrize("idle_timeout", [0, -1, 3601]) def test_idle_timeout_must_be_between_1_and_3600(self, idle_timeout): with pytest.raises( diff --git a/tests/unit/test_endpoint.py b/tests/unit/test_endpoint.py index 477bc7f8..2388df15 100644 --- a/tests/unit/test_endpoint.py +++ b/tests/unit/test_endpoint.py @@ -99,6 +99,23 @@ def test_workers_default(self): assert ep.workers_min == DEFAULT_WORKERS_MIN assert ep.workers_max == DEFAULT_WORKERS_MAX + def test_workers_standby_defaults_to_none(self): + """workers_standby defers to the resource default (workers min).""" + ep = Endpoint(name="test", workers=(0, 1)) + assert ep.workers_standby is None + + def test_workers_standby_out_of_range_raises(self): + with pytest.raises( + ValueError, + match=r"workers_standby \(4\) must be between workers min \(0\) " + r"and max \(3\)", + ): + Endpoint(name="test", workers=(0, 3), workers_standby=4) + + def test_workers_standby_in_range_ok(self): + ep = Endpoint(name="test", workers=(0, 3), workers_standby=1) + assert ep.workers_standby == 1 + def test_all_params(self): vol = NetworkVolume(name="test-vol", size=50) ep = Endpoint( @@ -378,6 +395,30 @@ def test_config_passes_idle_timeout(self): assert config.workersMin == 1 assert config.workersMax == 5 + @patch.dict(os.environ, {"FLASH_IS_LIVE_PROVISIONING": "true"}) + def test_config_workers_standby_defaults_to_workers_min(self): + """workers=(0, N) must deploy workersStandby=0 so the endpoint scales to zero.""" + ep = Endpoint(name="test", gpu=GpuGroup.ADA_24, workers=(0, 1)) + config = ep._build_resource_config() + assert config.workersMin == 0 + assert config.workersStandby == 0 + + @patch.dict(os.environ, {"FLASH_IS_LIVE_PROVISIONING": "true"}) + def test_config_workers_standby_follows_nonzero_min(self): + ep = Endpoint(name="test", gpu=GpuGroup.ADA_24, workers=(2, 5)) + config = ep._build_resource_config() + assert config.workersMin == 2 + assert config.workersStandby == 2 + + @patch.dict(os.environ, {"FLASH_IS_LIVE_PROVISIONING": "true"}) + def test_config_workers_standby_explicit(self): + ep = Endpoint( + name="test", gpu=GpuGroup.ADA_24, workers=(0, 3), workers_standby=1 + ) + config = ep._build_resource_config() + assert config.workersMin == 0 + assert config.workersStandby == 1 + @patch.dict(os.environ, {"FLASH_IS_LIVE_PROVISIONING": "true"}) def test_config_passes_volume(self): vol = NetworkVolume(name="test-vol", size=50)