Skip to content
Draft
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
5 changes: 4 additions & 1 deletion docs/Flash_SDK_Reference.md
Original file line number Diff line number Diff line change
Expand Up @@ -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,
)
```

Expand All @@ -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. |
Expand All @@ -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
Expand Down
31 changes: 31 additions & 0 deletions src/runpod_flash/core/resources/serverless.py
Original file line number Diff line number Diff line change
Expand Up @@ -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 ===
Expand Down Expand Up @@ -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.

Expand Down
14 changes: 14 additions & 0 deletions src/runpod_flash/endpoint.py
Original file line number Diff line number Diff line change
Expand Up @@ -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(
Expand All @@ -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
Expand Down Expand Up @@ -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
Expand Down
44 changes: 44 additions & 0 deletions tests/unit/resources/test_serverless.py
Original file line number Diff line number Diff line change
Expand Up @@ -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(
Expand Down
41 changes: 41 additions & 0 deletions tests/unit/test_endpoint.py
Original file line number Diff line number Diff line change
Expand Up @@ -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(
Expand Down Expand Up @@ -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)
Expand Down
Loading