Files
wov-sdk/tests/test_server.py
T

160 lines
5.1 KiB
Python

"""wov_sdk.server 的协议测试。
通过真实 HTTP 请求验证 /health 与 /invoke 端点的正常、404 和 500 分支,
并验证 run_node 的启动输出格式。
"""
import json
import threading
import urllib.error
import urllib.request
import pytest
from wov_sdk.models import InvokeRequest, InvokeResponse, NodeManifest
from wov_sdk.server import NodeHTTPServer, create_node_server, run_node
def valid_manifest() -> NodeManifest:
"""构造一个合法的最小 Echo 节点 manifest。"""
return NodeManifest(
id="echo",
name="Echo",
version="1.0.0",
capability="echo",
command=["python", "-m", "echo"],
)
def test_create_node_server_validates_manifest() -> None:
"""验证空 manifest 在创建服务器前就会被校验拒绝。"""
with pytest.raises(ValueError):
create_node_server(
NodeManifest(
id="",
name="",
version="",
capability="",
command=[],
),
lambda request: InvokeResponse(status="completed"),
)
def _post(server: NodeHTTPServer, path: str, body: dict | None = None) -> tuple[int, dict]:
"""向节点服务器发送 POST 请求并返回状态码与 JSON 响应。"""
data = json.dumps(body or {}).encode("utf-8")
request = urllib.request.Request(
f"http://127.0.0.1:{server.server_address[1]}{path}",
data=data,
headers={"Content-Type": "application/json"},
method="POST",
)
try:
with urllib.request.urlopen(request, timeout=5) as response:
return response.status, json.loads(response.read().decode("utf-8"))
except urllib.error.HTTPError as exc:
return exc.code, json.loads(exc.read().decode("utf-8"))
def _get(server: NodeHTTPServer, path: str) -> tuple[int, dict]:
"""向节点服务器发送 GET 请求并返回状态码与 JSON 响应。"""
try:
with urllib.request.urlopen(
f"http://127.0.0.1:{server.server_address[1]}{path}", timeout=5
) as response:
return response.status, json.loads(response.read().decode("utf-8"))
except urllib.error.HTTPError as exc:
return exc.code, json.loads(exc.read().decode("utf-8"))
def test_http_protocol() -> None:
"""验证健康检查、正常调用、404 与非法请求体等协议分支。"""
def handler(request: InvokeRequest) -> InvokeResponse:
return InvokeResponse(
status="completed",
outputs={"text": request.inputs.get("text", "")},
)
server = create_node_server(valid_manifest(), handler)
# 在独立线程中运行服务器,主线程执行 HTTP 客户端断言。
thread = threading.Thread(target=server.serve_forever, daemon=True)
thread.start()
try:
status, health = _get(server, "/health")
assert status == 200
assert health["node_id"] == "echo"
status, invoked = _post(
server,
"/invoke",
{
"run_id": "run_1",
"node_instance_id": "ni_1",
"inputs": {"text": "hello"},
"output_dir": "out",
},
)
assert status == 200
assert invoked["outputs"]["text"] == "hello"
status, not_found = _get(server, "/missing")
assert status == 404
assert not_found["status"] == "error"
status, not_found = _post(server, "/missing", {"run_id": "r"})
assert status == 404
status, bad_body = _post(server, "/invoke", {})
assert status == 500
assert bad_body["status"] == "failed"
finally:
server.shutdown()
server.server_close()
thread.join(timeout=5)
def test_http_protocol_handler_error() -> None:
"""验证业务处理函数抛异常时返回 500 与 failed 状态。"""
def failing_handler(request: InvokeRequest) -> InvokeResponse:
raise RuntimeError("boom")
server = create_node_server(valid_manifest(), failing_handler)
thread = threading.Thread(target=server.serve_forever, daemon=True)
thread.start()
try:
status, payload = _post(
server,
"/invoke",
{
"run_id": "run_1",
"node_instance_id": "ni_1",
"inputs": {},
"output_dir": "out",
},
)
assert status == 500
assert payload["status"] == "failed"
assert "boom" in payload["error"]
finally:
server.shutdown()
server.server_close()
thread.join(timeout=5)
def test_run_node(monkeypatch, capsys) -> None:
"""验证 run_node 打印 WOV_NODE_READY 端口行并进入服务循环。"""
called = []
def fake_serve_forever(server: NodeHTTPServer) -> None:
called.append(server.server_address[1])
server.server_close()
monkeypatch.setattr(NodeHTTPServer, "serve_forever", fake_serve_forever)
run_node(
valid_manifest(),
lambda request: InvokeResponse(status="completed"),
)
assert len(called) == 1
assert "WOV_NODE_READY" in capsys.readouterr().out