Skip to content
Merged
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
4 changes: 2 additions & 2 deletions .github/workflows/server-sanitize.yml
Original file line number Diff line number Diff line change
Expand Up @@ -103,7 +103,7 @@ jobs:
source .venv/bin/activate
cd tools/server/tests
export ${{ matrix.extra_args }}
./tests.sh
PYTEST_WORKERS=1 ./tests.sh

- name: Slow tests
id: server_integration_tests_slow
Expand All @@ -112,4 +112,4 @@ jobs:
source .venv/bin/activate
cd tools/server/tests
export ${{ matrix.extra_args }}
SLOW_TESTS=1 ./tests.sh
PYTEST_WORKERS=1 SLOW_TESTS=1 ./tests.sh
18 changes: 16 additions & 2 deletions tools/server/tests/conftest.py
Original file line number Diff line number Diff line change
@@ -1,7 +1,17 @@
import os
import pytest
from filelock import FileLock
from utils import *


@pytest.fixture(scope="session", autouse=True)
def configure_worker_port(request):
worker_id = getattr(request.config, "workerinput", {}).get("workerid", "master")
if worker_id != "master":
worker_num = int(worker_id[2:])
os.environ["PORT"] = str(8080 + worker_num * 10)


# ref: https://stackoverflow.com/questions/22627659/run-code-before-and-after-each-test-in-py-test
@pytest.fixture(autouse=True)
def stop_server_after_each_test():
Expand All @@ -16,6 +26,10 @@ def stop_server_after_each_test():


@pytest.fixture(scope="session", autouse=True)
def load_server_presets():
def load_server_presets(configure_worker_port, tmp_path_factory):
# this will be run once per test session, before any tests
ServerPreset.load_all()

# serialize model downloads across parallel workers.
root_tmp_dir = tmp_path_factory.getbasetemp().parent
with FileLock(str(root_tmp_dir / "load_all.lock")):
ServerPreset.load_all()
2 changes: 2 additions & 0 deletions tools/server/tests/requirements.txt
Original file line number Diff line number Diff line change
@@ -1,5 +1,7 @@
aiohttp~=3.9.3
pytest~=8.3.3
pytest-xdist~=3.6
filelock~=3.16
numpy~=1.26.4
openai~=2.14.0
prometheus-client~=0.20.0
Expand Down
8 changes: 5 additions & 3 deletions tools/server/tests/tests.sh
Original file line number Diff line number Diff line change
Expand Up @@ -6,13 +6,15 @@ cd $SCRIPT_DIR

set -eu

WORKERS="${PYTEST_WORKERS:-auto}"

if [ $# -lt 1 ]
then
if [[ "${SLOW_TESTS:-0}" == 1 ]]; then
pytest --durations=30 -v -x
pytest --durations=30 -v -x -n "${WORKERS}" --dist=worksteal
else
pytest --durations=30 -v -x -m "not slow"
pytest --durations=30 -v -x -n "${WORKERS}" --dist=worksteal -m "not slow"
fi
else
pytest --durations=30 "$@"
pytest --durations=30 -n "${WORKERS}" --dist=worksteal "$@"
fi
3 changes: 0 additions & 3 deletions tools/server/tests/unit/test_compat_anthropic.py
Original file line number Diff line number Diff line change
Expand Up @@ -21,7 +21,6 @@ def create_server():
global server
server = ServerPreset.tinyllama2()
server.model_alias = "tinyllama-2-anthropic"
server.server_port = 8082
server.n_slots = 1
server.n_ctx = 8192
server.n_batch = 2048
Expand All @@ -34,7 +33,6 @@ def vision_server():
server = ServerPreset.tinygemma3()
server.offline = False # Allow downloading the model
server.model_alias = "tinygemma3-anthropic"
server.server_port = 8083 # Different port to avoid conflicts
server.n_slots = 1
return server

Expand Down Expand Up @@ -1015,7 +1013,6 @@ def test_anthropic_thinking_with_reasoning_model(stream):
server.jinja = True
server.n_ctx = 8192
server.n_predict = 1024
server.server_port = 8084
server.start(timeout_seconds=600) # large model needs time to download

if stream:
Expand Down
5 changes: 0 additions & 5 deletions tools/server/tests/unit/test_mcp_servers.py
Original file line number Diff line number Diff line change
Expand Up @@ -37,7 +37,6 @@ def _start_server_with_mcp(mcp_json: str, **kwargs) -> ServerProcess:
srv = ServerPreset.router()
srv.server_tools = "all"
srv.no_ui = True
srv.server_port = 8085 # avoid conflict with load_all() which uses 8080
srv.mcp_servers_json = mcp_json
for k, v in kwargs.items():
setattr(srv, k, v)
Expand Down Expand Up @@ -183,7 +182,6 @@ def test_mcp_tools_not_listed_when_not_configured():
server = ServerPreset.router()
server.server_tools = "all"
server.no_ui = True
server.server_port = 8085
server.start()

try:
Expand Down Expand Up @@ -250,7 +248,6 @@ def test_mcp_tools_via_json_config_file():
server = ServerPreset.router()
server.server_tools = "all"
server.no_ui = True
server.server_port = 8085
server.mcp_servers_config = config_path
server.start()

Expand Down Expand Up @@ -468,7 +465,6 @@ def test_mcp_config_file_errors():
server = ServerPreset.router()
server.server_tools = "all"
server.no_ui = True
server.server_port = 8085
server.mcp_servers_json = "not valid json"
try:
server.start()
Expand All @@ -480,7 +476,6 @@ def test_mcp_config_file_errors():
server = ServerPreset.router()
server.server_tools = "all"
server.no_ui = True
server.server_port = 8085
server.mcp_servers_config = "/nonexistent/path.json"
try:
server.start()
Expand Down
8 changes: 4 additions & 4 deletions tools/server/tests/unit/test_slot_save.py
Original file line number Diff line number Diff line change
Expand Up @@ -10,10 +10,10 @@
server = ServerPreset.tinyllama2()

@pytest.fixture(autouse=True)
def create_server():
def create_server(tmp_path):
global server
server = ServerPreset.tinyllama2()
server.slot_save_path = "./tmp"
server.slot_save_path = str(tmp_path)
server.temperature = 0.0


Expand Down Expand Up @@ -94,7 +94,7 @@ def test_slot_restore_legacy_token_list():
assert res.body["n_saved"] == 84

# rewrite the token payload into a plain token list, as written by servers that predate the packed server_tokens format
path = os.path.join("tmp", "slot_legacy.bin")
path = os.path.join(server.slot_save_path, "slot_legacy.bin")
with open(path, "rb") as f:
data = bytearray(f.read())

Expand Down Expand Up @@ -462,7 +462,7 @@ def test_slot_save_restore_image_payload_larger_than_context(mmproj_server):
})
assert res.status_code == 200

path = os.path.join("tmp", "mm_slot_large_payload.bin")
path = os.path.join(server.slot_save_path, "mm_slot_large_payload.bin")
with open(path, "rb") as f:
data = bytearray(f.read())
payload_size = struct.unpack_from("=I", data, STATE_FILE_HEADER_SIZE - 4)[0]
Expand Down
1 change: 0 additions & 1 deletion tools/server/tests/unit/test_tool_call.py
Original file line number Diff line number Diff line change
Expand Up @@ -21,7 +21,6 @@ def create_server():
global server
server = ServerPreset.tinyllama2()
server.model_alias = "tinyllama-2-tool-call"
server.server_port = 8081
server.n_slots = 1
server.n_ctx = 8192
server.n_batch = 2048
Expand Down
12 changes: 6 additions & 6 deletions tools/server/tests/unit/test_tools_builtin.py
Original file line number Diff line number Diff line change
Expand Up @@ -64,11 +64,11 @@ def test_tools_builtin_read_file():
assert "def test_tools_builtin_read_file" in text


def test_tools_builtin_write_then_edit_file():
def test_tools_builtin_write_then_edit_file(tmp_path):
global server
server.start()

log_path = os.path.join(PROJECT_ROOT, "test.log")
log_path = str(tmp_path / "test.log")
try:
write_res = call_tool("write_file", {"path": log_path, "content": "line1\nline2\nline3\n"})
assert write_res["result"] == "file written successfully"
Expand All @@ -93,11 +93,11 @@ def test_tools_builtin_write_then_edit_file():
os.remove(log_path)


def test_tools_builtin_edit_file_rejects_non_unique_old_text():
def test_tools_builtin_edit_file_rejects_non_unique_old_text(tmp_path):
global server
server.start()

log_path = os.path.join(PROJECT_ROOT, "test.log")
log_path = str(tmp_path / "test.log")
try:
call_tool("write_file", {"path": log_path, "content": "dup\ndup\n"})
err = call_tool_expect_error("edit_file", {
Expand Down Expand Up @@ -275,11 +275,11 @@ def test_tools_builtin_docker_runtime_cleans_up_spawned_container():
assert leftover.returncode != 0, f"container {container_id} was not cleaned up after server exit"


def test_tools_builtin_edit_file_rejects_overlapping_edits():
def test_tools_builtin_edit_file_rejects_overlapping_edits(tmp_path):
global server
server.start()

log_path = os.path.join(PROJECT_ROOT, "test.log")
log_path = str(tmp_path / "test.log")
try:
call_tool("write_file", {"path": log_path, "content": "line1\nline2\n"})
err = call_tool_expect_error("edit_file", {
Expand Down
1 change: 1 addition & 0 deletions tools/server/tests/utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -294,6 +294,7 @@ def start(self, timeout_seconds: int = DEFAULT_HTTP_TIMEOUT) -> None:
server_args.append("--backend_sampling")
if self.gcp_compat:
env["AIP_MODE"] = "PREDICTION"
env["AIP_HTTP_PORT"] = str(self.server_port)

args = [str(arg) for arg in [server_path, *server_args]]
print(f"tests: starting server with: {' '.join(args)}")
Expand Down
Loading