server : use pytest-xdist for server tests (#28298)

* server : use pytest-xdist for server tests

This commit adds pytest-xdist to the server tests. This is pytest
plugin that distributes test execution across multiple CPU cores.

Assisted-by: pi:llama.cpp/qwen3.8-27B

Refs: https://github.com/ggml-org/llama.cpp/pull/26734#issuecomment-5220707042

* remove server_base_port and BASE_PORT

* use worksteal and pytest builting tmp_path
This commit is contained in:
Daniel Bevenius
2026-09-03 15:04:30 +02:00
committed by GitHub
parent de8656bd94
commit 42f0225fea
10 changed files with 36 additions and 26 deletions
+2 -2
View File
@@ -103,7 +103,7 @@ jobs:
source .venv/bin/activate source .venv/bin/activate
cd tools/server/tests cd tools/server/tests
export ${{ matrix.extra_args }} export ${{ matrix.extra_args }}
./tests.sh PYTEST_WORKERS=1 ./tests.sh
- name: Slow tests - name: Slow tests
id: server_integration_tests_slow id: server_integration_tests_slow
@@ -112,4 +112,4 @@ jobs:
source .venv/bin/activate source .venv/bin/activate
cd tools/server/tests cd tools/server/tests
export ${{ matrix.extra_args }} export ${{ matrix.extra_args }}
SLOW_TESTS=1 ./tests.sh PYTEST_WORKERS=1 SLOW_TESTS=1 ./tests.sh
+16 -2
View File
@@ -1,7 +1,17 @@
import os
import pytest import pytest
from filelock import FileLock
from utils import * 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 # ref: https://stackoverflow.com/questions/22627659/run-code-before-and-after-each-test-in-py-test
@pytest.fixture(autouse=True) @pytest.fixture(autouse=True)
def stop_server_after_each_test(): def stop_server_after_each_test():
@@ -16,6 +26,10 @@ def stop_server_after_each_test():
@pytest.fixture(scope="session", autouse=True) @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 # 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
View File
@@ -1,5 +1,7 @@
aiohttp~=3.9.3 aiohttp~=3.9.3
pytest~=8.3.3 pytest~=8.3.3
pytest-xdist~=3.6
filelock~=3.16
numpy~=1.26.4 numpy~=1.26.4
openai~=2.14.0 openai~=2.14.0
prometheus-client~=0.20.0 prometheus-client~=0.20.0
+5 -3
View File
@@ -6,13 +6,15 @@ cd $SCRIPT_DIR
set -eu set -eu
WORKERS="${PYTEST_WORKERS:-auto}"
if [ $# -lt 1 ] if [ $# -lt 1 ]
then then
if [[ "${SLOW_TESTS:-0}" == 1 ]]; then if [[ "${SLOW_TESTS:-0}" == 1 ]]; then
pytest --durations=30 -v -x pytest --durations=30 -v -x -n "${WORKERS}" --dist=worksteal
else else
pytest --durations=30 -v -x -m "not slow" pytest --durations=30 -v -x -n "${WORKERS}" --dist=worksteal -m "not slow"
fi fi
else else
pytest --durations=30 "$@" pytest --durations=30 -n "${WORKERS}" --dist=worksteal "$@"
fi fi
@@ -21,7 +21,6 @@ def create_server():
global server global server
server = ServerPreset.tinyllama2() server = ServerPreset.tinyllama2()
server.model_alias = "tinyllama-2-anthropic" server.model_alias = "tinyllama-2-anthropic"
server.server_port = 8082
server.n_slots = 1 server.n_slots = 1
server.n_ctx = 8192 server.n_ctx = 8192
server.n_batch = 2048 server.n_batch = 2048
@@ -34,7 +33,6 @@ def vision_server():
server = ServerPreset.tinygemma3() server = ServerPreset.tinygemma3()
server.offline = False # Allow downloading the model server.offline = False # Allow downloading the model
server.model_alias = "tinygemma3-anthropic" server.model_alias = "tinygemma3-anthropic"
server.server_port = 8083 # Different port to avoid conflicts
server.n_slots = 1 server.n_slots = 1
return server return server
@@ -1015,7 +1013,6 @@ def test_anthropic_thinking_with_reasoning_model(stream):
server.jinja = True server.jinja = True
server.n_ctx = 8192 server.n_ctx = 8192
server.n_predict = 1024 server.n_predict = 1024
server.server_port = 8084
server.start(timeout_seconds=600) # large model needs time to download server.start(timeout_seconds=600) # large model needs time to download
if stream: if stream:
@@ -37,7 +37,6 @@ def _start_server_with_mcp(mcp_json: str, **kwargs) -> ServerProcess:
srv = ServerPreset.router() srv = ServerPreset.router()
srv.server_tools = "all" srv.server_tools = "all"
srv.no_ui = True srv.no_ui = True
srv.server_port = 8085 # avoid conflict with load_all() which uses 8080
srv.mcp_servers_json = mcp_json srv.mcp_servers_json = mcp_json
for k, v in kwargs.items(): for k, v in kwargs.items():
setattr(srv, k, v) setattr(srv, k, v)
@@ -183,7 +182,6 @@ def test_mcp_tools_not_listed_when_not_configured():
server = ServerPreset.router() server = ServerPreset.router()
server.server_tools = "all" server.server_tools = "all"
server.no_ui = True server.no_ui = True
server.server_port = 8085
server.start() server.start()
try: try:
@@ -250,7 +248,6 @@ def test_mcp_tools_via_json_config_file():
server = ServerPreset.router() server = ServerPreset.router()
server.server_tools = "all" server.server_tools = "all"
server.no_ui = True server.no_ui = True
server.server_port = 8085
server.mcp_servers_config = config_path server.mcp_servers_config = config_path
server.start() server.start()
@@ -468,7 +465,6 @@ def test_mcp_config_file_errors():
server = ServerPreset.router() server = ServerPreset.router()
server.server_tools = "all" server.server_tools = "all"
server.no_ui = True server.no_ui = True
server.server_port = 8085
server.mcp_servers_json = "not valid json" server.mcp_servers_json = "not valid json"
try: try:
server.start() server.start()
@@ -480,7 +476,6 @@ def test_mcp_config_file_errors():
server = ServerPreset.router() server = ServerPreset.router()
server.server_tools = "all" server.server_tools = "all"
server.no_ui = True server.no_ui = True
server.server_port = 8085
server.mcp_servers_config = "/nonexistent/path.json" server.mcp_servers_config = "/nonexistent/path.json"
try: try:
server.start() server.start()
+4 -4
View File
@@ -10,10 +10,10 @@ STATE_FILE_HEADER_SIZE = 12
server = ServerPreset.tinyllama2() server = ServerPreset.tinyllama2()
@pytest.fixture(autouse=True) @pytest.fixture(autouse=True)
def create_server(): def create_server(tmp_path):
global server global server
server = ServerPreset.tinyllama2() server = ServerPreset.tinyllama2()
server.slot_save_path = "./tmp" server.slot_save_path = str(tmp_path)
server.temperature = 0.0 server.temperature = 0.0
@@ -94,7 +94,7 @@ def test_slot_restore_legacy_token_list():
assert res.body["n_saved"] == 84 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 # 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: with open(path, "rb") as f:
data = bytearray(f.read()) data = bytearray(f.read())
@@ -462,7 +462,7 @@ def test_slot_save_restore_image_payload_larger_than_context(mmproj_server):
}) })
assert res.status_code == 200 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: with open(path, "rb") as f:
data = bytearray(f.read()) data = bytearray(f.read())
payload_size = struct.unpack_from("=I", data, STATE_FILE_HEADER_SIZE - 4)[0] payload_size = struct.unpack_from("=I", data, STATE_FILE_HEADER_SIZE - 4)[0]
@@ -21,7 +21,6 @@ def create_server():
global server global server
server = ServerPreset.tinyllama2() server = ServerPreset.tinyllama2()
server.model_alias = "tinyllama-2-tool-call" server.model_alias = "tinyllama-2-tool-call"
server.server_port = 8081
server.n_slots = 1 server.n_slots = 1
server.n_ctx = 8192 server.n_ctx = 8192
server.n_batch = 2048 server.n_batch = 2048
@@ -64,11 +64,11 @@ def test_tools_builtin_read_file():
assert "def test_tools_builtin_read_file" in text 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 global server
server.start() server.start()
log_path = os.path.join(PROJECT_ROOT, "test.log") log_path = str(tmp_path / "test.log")
try: try:
write_res = call_tool("write_file", {"path": log_path, "content": "line1\nline2\nline3\n"}) write_res = call_tool("write_file", {"path": log_path, "content": "line1\nline2\nline3\n"})
assert write_res["result"] == "file written successfully" assert write_res["result"] == "file written successfully"
@@ -93,11 +93,11 @@ def test_tools_builtin_write_then_edit_file():
os.remove(log_path) 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 global server
server.start() server.start()
log_path = os.path.join(PROJECT_ROOT, "test.log") log_path = str(tmp_path / "test.log")
try: try:
call_tool("write_file", {"path": log_path, "content": "dup\ndup\n"}) call_tool("write_file", {"path": log_path, "content": "dup\ndup\n"})
err = call_tool_expect_error("edit_file", { err = call_tool_expect_error("edit_file", {
@@ -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" 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 global server
server.start() server.start()
log_path = os.path.join(PROJECT_ROOT, "test.log") log_path = str(tmp_path / "test.log")
try: try:
call_tool("write_file", {"path": log_path, "content": "line1\nline2\n"}) call_tool("write_file", {"path": log_path, "content": "line1\nline2\n"})
err = call_tool_expect_error("edit_file", { err = call_tool_expect_error("edit_file", {
+1
View File
@@ -294,6 +294,7 @@ class ServerProcess:
server_args.append("--backend_sampling") server_args.append("--backend_sampling")
if self.gcp_compat: if self.gcp_compat:
env["AIP_MODE"] = "PREDICTION" env["AIP_MODE"] = "PREDICTION"
env["AIP_HTTP_PORT"] = str(self.server_port)
args = [str(arg) for arg in [server_path, *server_args]] args = [str(arg) for arg in [server_path, *server_args]]
print(f"tests: starting server with: {' '.join(args)}") print(f"tests: starting server with: {' '.join(args)}")