Skip to content
This repository was archived by the owner on Jul 13, 2026. It is now read-only.
Open
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
Empty file added tests/__init__.py
Empty file.
22 changes: 22 additions & 0 deletions tests/conftest.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,22 @@
"""
Pytest configuration for the vast-sdk test suite.

The vastai package imports serverless server code (vastai/serverless/server/lib/backend.py)
which requires pycryptodome (Crypto). That package may not be present in all dev
environments, so we stub it in sys.modules here — before any test module is imported —
so the package initialisation doesn't blow up.
"""
import sys
from unittest.mock import MagicMock

_crypto = MagicMock()
for _mod in (
"Crypto",
"Crypto.Signature",
"Crypto.Signature.pkcs1_15",
"Crypto.Hash",
"Crypto.Hash.SHA256",
"Crypto.PublicKey",
"Crypto.PublicKey.RSA",
):
sys.modules.setdefault(_mod, _crypto)
153 changes: 153 additions & 0 deletions tests/test_template_functions.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,153 @@
"""
Tests for search__templates and create__template return-value behavior.

These tests verify that both functions correctly return data when args.raw=True
(the SDK path) and fall through to display/print when args.raw=False (the CLI path).
"""
import argparse
from unittest.mock import MagicMock, patch

import pytest


def make_args(**overrides):
"""Build a minimal argparse.Namespace sufficient for both tested functions."""
defaults = dict(
raw=True,
query=None,
api_key="test-key",
url="https://console.vast.ai",
explain=False,
curl=False,
retry=3,
# create_template fields
name="test-template",
image="ubuntu:22.04",
image_tag=None,
href=None,
repo=None,
login=None,
env=None,
onstart_cmd=None,
no_default=True, # skip parse_query side-effects
search_params=None,
disk_space=None,
jupyter=False,
direct=False,
jupyter_lab=False,
jupyter_dir=None,
ssh=False,
readme=None,
hide_readme=False,
desc=None,
public=False,
)
defaults.update(overrides)
return argparse.Namespace(**defaults)


def make_get_response(templates):
"""Return a mock HTTP response for search__templates."""
resp = MagicMock()
resp.status_code = 200
resp.headers = {"Content-Type": "application/json"}
resp.json.return_value = {"templates": templates}
resp.raise_for_status.return_value = None
return resp


def make_post_response(success=True, template=None, msg="error"):
"""Return a mock HTTP response for create__template."""
resp = MagicMock()
resp.raise_for_status.return_value = None
if success:
resp.json.return_value = {"success": True, "template": template or {"id": 42, "name": "test-template"}}
else:
resp.json.return_value = {"success": False, "msg": msg}
return resp


# ---------------------------------------------------------------------------
# search__templates
# ---------------------------------------------------------------------------

@patch("vastai.vast.apiurl", return_value="https://console.vast.ai/api/v0/template/")
@patch("vastai.vast.http_get")
def test_search_templates_raw_returns_list(mock_http_get, mock_apiurl):
from vastai.vast import search__templates

mock_http_get.return_value = make_get_response([{"id": 1, "name": "my-template"}])

result = search__templates(make_args())

assert result == [{"id": 1, "name": "my-template"}]


@patch("vastai.vast.apiurl", return_value="https://console.vast.ai/api/v0/template/")
@patch("vastai.vast.http_get")
def test_search_templates_empty_list_returns_empty_list(mock_http_get, mock_apiurl):
"""An empty templates array should return [] not None."""
from vastai.vast import search__templates

mock_http_get.return_value = make_get_response([])

result = search__templates(make_args())

assert result == []
assert result is not None


@patch("vastai.vast.display_table")
@patch("vastai.vast.apiurl", return_value="https://console.vast.ai/api/v0/template/")
@patch("vastai.vast.http_get")
def test_search_templates_not_raw_calls_display_table(mock_http_get, mock_apiurl, mock_display):
from vastai.vast import search__templates

mock_http_get.return_value = make_get_response([{"id": 1}])

result = search__templates(make_args(raw=False))

assert result is None
mock_display.assert_called_once()


# ---------------------------------------------------------------------------
# create__template
# ---------------------------------------------------------------------------

@patch("vastai.vast.apiurl", return_value="https://console.vast.ai/api/v0/template/")
@patch("vastai.vast.http_post")
def test_create_template_raw_returns_template_dict(mock_http_post, mock_apiurl):
from vastai.vast import create__template

expected = {"id": 42, "name": "test-template", "hash": "abc123"}
mock_http_post.return_value = make_post_response(template=expected)

result = create__template(make_args())

assert result == expected


@patch("vastai.vast.apiurl", return_value="https://console.vast.ai/api/v0/template/")
@patch("vastai.vast.http_post")
def test_create_template_not_raw_returns_none(mock_http_post, mock_apiurl):
from vastai.vast import create__template

mock_http_post.return_value = make_post_response()

result = create__template(make_args(raw=False))

assert result is None


@patch("vastai.vast.apiurl", return_value="https://console.vast.ai/api/v0/template/")
@patch("vastai.vast.http_post")
def test_create_template_api_failure_returns_none(mock_http_post, mock_apiurl):
"""When the API returns success=False the function should print the message and return None."""
from vastai.vast import create__template

mock_http_post.return_value = make_post_response(success=False, msg="duplicate name")

result = create__template(make_args())

assert result is None
9 changes: 6 additions & 3 deletions vastai/vast.py
Original file line number Diff line number Diff line change
Expand Up @@ -2751,7 +2751,10 @@ def create__template(args):
try:
rj = r.json()
if rj["success"]:
print(f"New Template: {rj['template']}")
if args.raw:
return rj['template']
else:
print(f"New Template: {rj['template']}")
else:
print(rj['msg'])
except requests.exceptions.JSONDecodeError:
Expand Down Expand Up @@ -4521,8 +4524,8 @@ def search__templates(args):
r.raise_for_status()
elif 'json' in r.headers.get("Content-Type"):
rows = r.json().get('templates', [])
if True: #args.raw:
print(json.dumps(rows, indent=1, sort_keys=True))
if args.raw:
return rows
else:
display_table(rows, displayable_fields)
else:
Expand Down