feat: 完成节点协议与 SDK
This commit is contained in:
@@ -0,0 +1,158 @@
|
||||
import json
|
||||
|
||||
import pytest
|
||||
|
||||
from wov_sdk.models import (
|
||||
HealthResponse,
|
||||
InvokeRequest,
|
||||
InvokeResponse,
|
||||
NodeManifest,
|
||||
ProgressEvent,
|
||||
WorkflowDefinition,
|
||||
WorkflowEdge,
|
||||
WorkflowNode,
|
||||
)
|
||||
|
||||
|
||||
def valid_manifest() -> NodeManifest:
|
||||
return NodeManifest(
|
||||
id="echo",
|
||||
name="Echo",
|
||||
version="1.0.0",
|
||||
capability="echo",
|
||||
command=["python", "-m", "echo"],
|
||||
repo_dir="wov-node-echo",
|
||||
env={"PORT": "0"},
|
||||
input_schema={"text": "string"},
|
||||
output_schema={"text": "string"},
|
||||
max_concurrency=2,
|
||||
idle_ttl_seconds=15,
|
||||
health_timeout_seconds=5,
|
||||
keep_warm=True,
|
||||
)
|
||||
|
||||
|
||||
def test_manifest_round_trip() -> None:
|
||||
manifest = valid_manifest()
|
||||
restored = NodeManifest.from_dict(manifest.to_dict())
|
||||
assert restored == manifest
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("field", "value"),
|
||||
[
|
||||
("id", ""),
|
||||
("name", ""),
|
||||
("version", ""),
|
||||
("capability", ""),
|
||||
("repo_dir", ""),
|
||||
("command", []),
|
||||
("max_concurrency", 0),
|
||||
("idle_ttl_seconds", -1),
|
||||
("health_timeout_seconds", 0),
|
||||
],
|
||||
)
|
||||
def test_manifest_validation(field: str, value: object) -> None:
|
||||
manifest = valid_manifest()
|
||||
setattr(manifest, field, value)
|
||||
with pytest.raises(ValueError):
|
||||
manifest.validate()
|
||||
|
||||
|
||||
def test_manifest_load(tmp_path) -> None:
|
||||
path = tmp_path / "node.manifest.json"
|
||||
path.write_text(json.dumps(valid_manifest().to_dict()), encoding="utf-8")
|
||||
loaded = NodeManifest.load(str(path))
|
||||
assert loaded.id == "echo"
|
||||
|
||||
|
||||
def test_invoke_request_round_trip() -> None:
|
||||
request = InvokeRequest(
|
||||
run_id="run_1",
|
||||
node_instance_id="ni_1",
|
||||
inputs={"text": "hello"},
|
||||
params={"temperature": 0.2},
|
||||
output_dir="out",
|
||||
)
|
||||
restored = InvokeRequest.from_dict(request.to_dict())
|
||||
assert restored == request
|
||||
|
||||
|
||||
def test_invoke_response_round_trip() -> None:
|
||||
response = InvokeResponse(status="completed", outputs={"text": "hello"})
|
||||
restored = InvokeResponse.from_dict(response.to_dict())
|
||||
assert restored == response
|
||||
|
||||
|
||||
def test_health_and_progress_serialization() -> None:
|
||||
health = HealthResponse(status="ok", node_id="echo", version="1.0.0")
|
||||
assert health.to_dict() == {
|
||||
"status": "ok",
|
||||
"node_id": "echo",
|
||||
"version": "1.0.0",
|
||||
}
|
||||
|
||||
progress = ProgressEvent(run_id="run_1", node_id="echo", progress=0.5, message="half")
|
||||
assert progress.to_dict() == {
|
||||
"run_id": "run_1",
|
||||
"node_id": "echo",
|
||||
"progress": 0.5,
|
||||
"message": "half",
|
||||
}
|
||||
|
||||
|
||||
def test_workflow_node_and_edge_round_trip() -> None:
|
||||
node = WorkflowNode(
|
||||
id="asr",
|
||||
node_type="faster-whisper",
|
||||
params={"language": "ja"},
|
||||
inputs={"audio_uri": "extract.audio_uri"},
|
||||
)
|
||||
edge = WorkflowEdge(from_node="extract", to_node="asr")
|
||||
assert WorkflowNode.from_dict(node.to_dict()) == node
|
||||
assert WorkflowEdge.from_dict(edge.to_dict()) == edge
|
||||
assert edge.to_dict() == {"from": "extract", "to": "asr"}
|
||||
|
||||
assert node.to_dict()["inputs"] == {"audio_uri": "extract.audio_uri"}
|
||||
|
||||
|
||||
def test_workflow_definition_round_trip_and_validation() -> None:
|
||||
definition = WorkflowDefinition(
|
||||
name="demo",
|
||||
version=1,
|
||||
nodes=[
|
||||
WorkflowNode(id="extract", node_type="ffmpeg"),
|
||||
WorkflowNode(id="asr", node_type="whisper"),
|
||||
],
|
||||
edges=[WorkflowEdge(from_node="extract", to_node="asr")],
|
||||
entry_inputs={"video_uri": "file"},
|
||||
final_outputs={"srt": "asr.srt_uri"},
|
||||
)
|
||||
restored = WorkflowDefinition.from_dict(definition.to_dict())
|
||||
assert restored == definition
|
||||
restored.validate()
|
||||
|
||||
|
||||
def test_workflow_definition_invalid() -> None:
|
||||
with pytest.raises(ValueError):
|
||||
WorkflowDefinition(name="", version=1).validate()
|
||||
|
||||
with pytest.raises(ValueError):
|
||||
WorkflowDefinition(name="demo", version=0).validate()
|
||||
|
||||
duplicate = WorkflowDefinition(
|
||||
name="demo",
|
||||
version=1,
|
||||
nodes=[WorkflowNode(id="a", node_type="x"), WorkflowNode(id="a", node_type="y")],
|
||||
)
|
||||
with pytest.raises(ValueError):
|
||||
duplicate.validate()
|
||||
|
||||
unknown_edge = WorkflowDefinition(
|
||||
name="demo",
|
||||
version=1,
|
||||
nodes=[WorkflowNode(id="a", node_type="x")],
|
||||
edges=[WorkflowEdge(from_node="a", to_node="missing")],
|
||||
)
|
||||
with pytest.raises(ValueError):
|
||||
unknown_edge.validate()
|
||||
@@ -0,0 +1,145 @@
|
||||
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:
|
||||
return NodeManifest(
|
||||
id="echo",
|
||||
name="Echo",
|
||||
version="1.0.0",
|
||||
capability="echo",
|
||||
command=["python", "-m", "echo"],
|
||||
)
|
||||
|
||||
|
||||
def test_create_node_server_validates_manifest() -> None:
|
||||
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]:
|
||||
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]:
|
||||
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:
|
||||
def handler(request: InvokeRequest) -> InvokeResponse:
|
||||
return InvokeResponse(
|
||||
status="completed",
|
||||
outputs={"text": request.inputs.get("text", "")},
|
||||
)
|
||||
|
||||
server = create_node_server(valid_manifest(), handler)
|
||||
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:
|
||||
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:
|
||||
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
|
||||
Reference in New Issue
Block a user