Skip to content
Merged
Show file tree
Hide file tree
Changes from 3 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
22 changes: 10 additions & 12 deletions src/tserve/cli/main.py
Original file line number Diff line number Diff line change
Expand Up @@ -121,11 +121,6 @@ def main(argv: list[str] | None = None) -> int:
------
SystemExit
From ``ArgumentParser.error`` or argparse usage errors.
ValueError
If positional tokens are malformed (empty ``id=spec``,
a bare craft spec), or ``Server`` / ``bootstrap`` reject a
load spec (duplicate id, unknown registry id, non-zip path,
unknown executor).
TypeError
If a model object is not a sktime ``BaseForecaster``.
ImportError
Expand All @@ -150,13 +145,16 @@ def main(argv: list[str] | None = None) -> int:

from tserve.server import Server

server = Server(
model=parse_model(args.positional_model),
models_dir=args.models_dir,
host=args.host,
port=args.port,
log_level=args.log_level,
)
try:
server = Server(
model=parse_model(args.positional_model),
models_dir=args.models_dir,
host=args.host,
port=args.port,
log_level=args.log_level,
)
except ValueError as exc:
parser.error(str(exc))
Comment thread
SiddhantGahankari marked this conversation as resolved.
Outdated
try:
server.run()
except KeyboardInterrupt:
Expand Down
15 changes: 15 additions & 0 deletions src/tserve/cli/tests/test_main.py
Original file line number Diff line number Diff line change
Expand Up @@ -107,3 +107,18 @@ def test_main_version(capsys):

assert excinfo.value.code == 0
assert capsys.readouterr().out.strip() == f"tserve {__version__}"


def test_main_reports_invalid_model_without_traceback(capsys):
with (
patch(
"tserve.server.Server", side_effect=ValueError("unknown model 'missing'")
),
pytest.raises(SystemExit) as excinfo,
):
main(["missing"])

assert excinfo.value.code == 2
err = capsys.readouterr().err
assert "unknown model 'missing'" in err
assert "Traceback" not in err
26 changes: 14 additions & 12 deletions src/tserve/runtime/bootstrap.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
----------
Expand Down Expand Up @@ -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] = {}
Expand All @@ -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 = (
Expand Down
23 changes: 22 additions & 1 deletion src/tserve/runtime/tests/test_bootstrap.py
Original file line number Diff line number Diff line change
Expand Up @@ -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():
Expand Down
Loading