diff --git a/src/tserve/runtime/bootstrap.py b/src/tserve/runtime/bootstrap.py index 6f8e172..3ae413c 100644 --- a/src/tserve/runtime/bootstrap.py +++ b/src/tserve/runtime/bootstrap.py @@ -89,11 +89,10 @@ def loaded_models(self) -> ModelsResult: def bootstrap(model: list[str | Path | tuple[str, Any]]) -> Runtime: """Resolve, construct, load, warmup, and register each selected model. - For every item: ``resolve_model`` → ``create_executor(info.executor)`` - → ``load`` → ``warmup`` → ``stats.register``. Tuple items pass the - second element (craft spec or object) to ``load``; strings and - paths pass the item itself. Duplicate ``ModelInfo.id`` values raise - before a second load. + Resolve and check all items before creating any executor. Then, for + every item: ``create_executor(info.executor)`` → ``load`` → ``warmup`` + → ``stats.register``. Tuple items pass the second element (craft spec + or object) to ``load``; strings and paths pass the item itself. Parameters ---------- @@ -132,6 +131,15 @@ def bootstrap(model: list[str | Path | tuple[str, Any]]) -> Runtime: tserve.scheduling.scheduler.Scheduler Forecast dispatch is not implemented in this package. """ + resolved: list[tuple[ModelInfo, Any]] = [] + seen: set[str] = set() + for item in model: + info = resolve_model(item) + if info.id in seen: + raise ValueError(f"duplicate model id {info.id!r} in model") + seen.add(info.id) + resolved.append((info, item[1] if isinstance(item, tuple) else item)) + stats = Stats() executors: dict[str, Executor] = {} models: dict[str, ModelInfo] = {} @@ -145,13 +153,7 @@ def bootstrap(model: list[str | Path | tuple[str, Any]]) -> Runtime: f"{paint('is always loaded as a baseline', '2')}" ) - for position, item in enumerate(model, start=1): - info = resolve_model(item) - item = item[1] if isinstance(item, tuple) else item - - if info.id in models: - raise ValueError(f"duplicate model id {info.id!r} in model") - + for position, (info, item) in enumerate(resolved, start=1): label = f"{info.id} via {info.executor}" dots = paint("." * max(3, 40 - len(label)), "2") prefix = ( diff --git a/src/tserve/runtime/tests/test_bootstrap.py b/src/tserve/runtime/tests/test_bootstrap.py index 9b1d394..ce1f7f5 100644 --- a/src/tserve/runtime/tests/test_bootstrap.py +++ b/src/tserve/runtime/tests/test_bootstrap.py @@ -77,10 +77,31 @@ def test_bootstrap_loads_craft(): def test_bootstrap_rejects_duplicate(): with ( patch("tserve.runtime.bootstrap.resolve_model", return_value=_info()), - patch("tserve.runtime.bootstrap.create_executor", return_value=_executor()), + patch("tserve.runtime.bootstrap.create_executor") as create_executor, pytest.raises(ValueError, match="duplicate model id 'naive' in model"), ): bootstrap(["naive", "naive"]) + create_executor.assert_not_called() + + +def test_bootstrap_validates_before_loading(): + executor = _executor() + with ( + patch( + "tserve.runtime.bootstrap.resolve_model", + side_effect=[_info(), ValueError("unknown model 'missing'")], + ) as resolve_model, + patch( + "tserve.runtime.bootstrap.create_executor", return_value=executor + ) as create_executor, + pytest.raises(ValueError, match="unknown model 'missing'"), + ): + bootstrap(["naive", "missing"]) + + assert resolve_model.call_count == 2 + create_executor.assert_not_called() + executor.load.assert_not_called() + executor.warmup.assert_not_called() def test_loaded_models():