feat: add workflow rpc transport params

This commit is contained in:
lda
2026-06-03 06:45:25 +07:00 Verified
parent cdac8c7599
commit c9f4a87cdf
3 changed files with 191 additions and 1 deletions
+33 -1
View File
@@ -1,3 +1,35 @@
from __future__ import annotations
__all__: list[str] = []
from .models import (
CreateDraftFromCapabilityParams,
HealthParams,
InspectCapabilityParams,
InspectRunParams,
ListCapabilitiesParams,
PatchDraftParams,
ReadRunTraceParams,
ResumeRunParams,
SaveArtifactParams,
SaveDeploymentParams,
StartRunParams,
TraceRangeParams,
ValidateDeploymentParams,
ValidateDraftParams,
)
__all__ = [
"CreateDraftFromCapabilityParams",
"HealthParams",
"InspectCapabilityParams",
"InspectRunParams",
"ListCapabilitiesParams",
"PatchDraftParams",
"ReadRunTraceParams",
"ResumeRunParams",
"SaveArtifactParams",
"SaveDeploymentParams",
"StartRunParams",
"TraceRangeParams",
"ValidateDeploymentParams",
"ValidateDraftParams",
]
+100
View File
@@ -0,0 +1,100 @@
from __future__ import annotations
from typing import Any
from pydantic import BaseModel, ConfigDict, Field
from wf_api.models import TraceRange
class RpcParamsModel(BaseModel):
"""Base transport DTO: reject misspelled JSON-RPC params early."""
model_config = ConfigDict(extra="forbid")
class TraceRangeParams(RpcParamsModel):
start: int = Field(default=0, ge=0, description="Zero-based trace offset.")
limit: int = Field(
default=20,
ge=1,
le=100,
description="Maximum trace entries to return; full traces are never implicit.",
)
def to_api_trace_range(self) -> TraceRange:
return TraceRange(start=self.start, limit=self.limit)
class HealthParams(RpcParamsModel):
pass
class ListCapabilitiesParams(RpcParamsModel):
query: str | None = Field(default=None)
source_id: str | None = Field(default=None)
cursor: str | None = Field(default=None)
limit: int = Field(default=50, ge=1, le=200)
class InspectCapabilityParams(RpcParamsModel):
qualified_name: str = Field(min_length=1)
class CreateDraftFromCapabilityParams(RpcParamsModel):
workspace_id: str = Field(min_length=1)
capability_name: str = Field(min_length=1)
name: str | None = None
title: str | None = None
input_schema: dict[str, Any] | None = None
state_schema: dict[str, Any] | None = None
output_schema: dict[str, Any] | None = None
input: list[Any] | None = None
output: list[Any] | None = None
input_map: dict[str, str] | None = None
output_map: dict[str, str] | None = None
error_message_source: Any | None = None
class PatchDraftParams(RpcParamsModel):
draft: dict[str, Any]
patch: list[dict[str, Any]]
class ValidateDraftParams(RpcParamsModel):
draft: dict[str, Any]
class SaveArtifactParams(RpcParamsModel):
artifact: dict[str, Any]
class SaveDeploymentParams(RpcParamsModel):
deployment: dict[str, Any]
class ValidateDeploymentParams(RpcParamsModel):
deployment_id: str = Field(min_length=1)
live_check: bool = False
class StartRunParams(RpcParamsModel):
deployment_id: str = Field(min_length=1)
workflow_input: dict[str, Any] = Field(default_factory=dict)
trace_range: TraceRangeParams | None = None
class InspectRunParams(RpcParamsModel):
run_id: str = Field(min_length=1)
class ReadRunTraceParams(RpcParamsModel):
run_id: str = Field(min_length=1)
trace_range: TraceRangeParams
class ResumeRunParams(RpcParamsModel):
run_id: str = Field(min_length=1)
resume_payload: dict[str, Any] = Field(default_factory=dict)
resume_outcome: str = Field(default="submitted", min_length=1)
trace_range: TraceRangeParams | None = None
@@ -0,0 +1,58 @@
from __future__ import annotations
import pytest
from pydantic import ValidationError
from wf_transport_rpc_http.models import (
InspectCapabilityParams,
ListCapabilitiesParams,
ReadRunTraceParams,
StartRunParams,
TraceRangeParams,
)
def test_trace_range_params_converts_to_api_trace_range() -> None:
trace_range = TraceRangeParams(start=2, limit=5).to_api_trace_range()
assert trace_range.start == 2
assert trace_range.limit == 5
def test_trace_range_params_rejects_invalid_values() -> None:
with pytest.raises(ValidationError):
TraceRangeParams(start=-1, limit=5)
with pytest.raises(ValidationError):
TraceRangeParams(start=0, limit=0)
with pytest.raises(ValidationError):
TraceRangeParams(start=0, limit=101)
def test_capability_params_are_explicit_models() -> None:
listed = ListCapabilitiesParams(query="echo", source_id="wf.std", limit=10)
inspected = InspectCapabilityParams(qualified_name="wf.std.constant")
assert listed.query == "echo"
assert listed.source_id == "wf.std"
assert listed.limit == 10
assert inspected.qualified_name == "wf.std.constant"
def test_run_params_are_explicit_models() -> None:
started = StartRunParams(
deployment_id="demo.default",
workflow_input={"message": "hello"},
trace_range=TraceRangeParams(start=0, limit=3),
)
trace = ReadRunTraceParams(
run_id="run_demo",
trace_range=TraceRangeParams(start=0, limit=1),
)
assert started.deployment_id == "demo.default"
assert started.workflow_input["message"] == "hello"
assert started.trace_range is not None
assert trace.run_id == "run_demo"
assert trace.trace_range.limit == 1