mirror of
https://github.com/ggml-org/llama.cpp.git
synced 2026-09-04 02:37:27 +02:00
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:
@@ -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
|
||||||
|
|||||||
@@ -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()
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -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()
|
||||||
|
|||||||
@@ -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", {
|
||||||
|
|||||||
@@ -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)}")
|
||||||
|
|||||||
Reference in New Issue
Block a user