Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
64 changes: 28 additions & 36 deletions openhcs/mcp/server.py
Original file line number Diff line number Diff line change
Expand Up @@ -34,6 +34,7 @@
from metaclass_registry import AutoRegisterMeta
from pydantic import Field as PydanticField
from pydantic import WithJsonSchema
from python_introspect import dataclass_from_mapping
from zmqruntime.config import TransportMode

from openhcs.agent.authoring_contexts import AuthoringContextDeclaration
Expand Down Expand Up @@ -1977,45 +1978,18 @@ def connection_parameters(
cls,
timeout_policy: type[McpControlTimeoutPolicy] = McpViewerTimeoutPolicy,
) -> tuple[Parameter, ...]:
return (
Parameter(
"port",
Parameter.KEYWORD_ONLY,
annotation=int,
),
Parameter(
"host",
Parameter.KEYWORD_ONLY,
default="localhost",
annotation=str,
),
Parameter(
"transport_mode",
Parameter.KEYWORD_ONLY,
default=None,
annotation=str | None,
),
Parameter(
"timeout_ms",
Parameter.KEYWORD_ONLY,
default=None,
annotation=_timeout_parameter_annotation(timeout_policy),
),
)
return McpViewerConnectionToolFields.signature_parameters(timeout_policy)

@classmethod
def control_args(
cls,
arguments: Mapping[str, JsonValue],
timeout_policy: type[McpControlTimeoutPolicy] = McpViewerTimeoutPolicy,
) -> "McpViewerConnectionToolArgs":
return McpViewerConnectionToolArgs.from_fields(
port=arguments["port"],
host=arguments["host"],
transport_mode=arguments["transport_mode"],
timeout_ms=arguments["timeout_ms"],
timeout_policy=timeout_policy,
)
return dataclass_from_mapping(
McpViewerConnectionToolFields,
arguments,
).to_control_args(timeout_policy)

@staticmethod
def option_parameters(
Expand Down Expand Up @@ -2664,14 +2638,32 @@ def _json_object_or_empty(value: dict | None) -> dict:
return dict(value)


@dataclass(frozen=True, slots=True)
@dataclass(frozen=True, slots=True, kw_only=True)
class McpViewerConnectionToolFields:
"""Raw MCP viewer connection arguments before policy resolution."""

port: int
host: str
transport_mode: TransportMode | None
timeout_ms: int | None
host: str = "localhost"
transport_mode: TransportMode | None = None
timeout_ms: int | None = None

@classmethod
def signature_parameters(
cls,
timeout_policy: type[McpControlTimeoutPolicy],
) -> tuple[Parameter, ...]:
"""Project the public tool signature from the declared field types."""
annotations = get_type_hints(cls)
return tuple(
parameter.replace(
annotation=(
_timeout_parameter_annotation(timeout_policy)
if parameter.name == "timeout_ms"
else annotations[parameter.name]
)
)
for parameter in inspect_signature(cls).parameters.values()
)

def to_control_args(
self,
Expand Down
12 changes: 6 additions & 6 deletions tests/integration/test_cellprofiler_official30_zmq.py
Original file line number Diff line number Diff line change
Expand Up @@ -193,12 +193,12 @@ def _free_zmq_port_pair(excluded: set[int]) -> int:

def _registered_streaming_config_kwargs(
viewer_specs: Sequence[StreamingViewerConfigSpec],
ports_by_viewer: Mapping[str, int],
ports_by_viewer: Mapping[ViewerType, int],
) -> dict[str, StreamingConfig]:
"""Project one viewer selection through registered config owners."""

selected_specs = {spec.registry_key: spec for spec in viewer_specs}
selected_viewers = {spec.viewer_type.value for spec in viewer_specs}
selected_viewers = {spec.viewer_type for spec in viewer_specs}
if len(selected_specs) != len(viewer_specs) or len(selected_viewers) != len(
viewer_specs
):
Expand All @@ -221,7 +221,7 @@ def _registered_streaming_config_kwargs(
if enabled:
init_kwargs.update(
host="127.0.0.1",
port=ports_by_viewer[spec.viewer_type.value],
port=ports_by_viewer[spec.viewer_type],
transport_mode=TransportMode.TCP,
)
config_kwargs[spec.registry_key] = config_type(**init_kwargs)
Expand Down Expand Up @@ -523,7 +523,7 @@ def test_official30_fiji_variants_project_registered_viewer_configs() -> None:

for viewer_specs in _OFFICIAL30_FIJI_VIEWER_VARIANTS:
ports_by_viewer = {
spec.viewer_type.value: 20_000 + index * 100
spec.viewer_type: 20_000 + index * 100
for index, spec in enumerate(viewer_specs)
}
configs = _registered_streaming_config_kwargs(
Expand All @@ -544,7 +544,7 @@ def test_official30_fiji_variants_project_registered_viewer_configs() -> None:
)
assert {
config.viewer_type for config in configs.values() if config.enabled
} == {spec.viewer_type.value for spec in viewer_specs}
} == {spec.viewer_type for spec in viewer_specs}
assert all(config.persistent is config.enabled for config in configs.values())
assert {
config.viewer_type: config.port
Expand Down Expand Up @@ -703,7 +703,7 @@ def test_official30_persistent_fiji_variants_isolated_per_case(
excluded_ports: set[int] = set()
execution_port = _free_zmq_port_pair(excluded_ports)
ports_by_viewer = {
spec.viewer_type.value: _free_zmq_port_pair(excluded_ports)
spec.viewer_type: _free_zmq_port_pair(excluded_ports)
for spec in viewer_specs
}
streaming_configs = _registered_streaming_config_kwargs(
Expand Down
28 changes: 25 additions & 3 deletions tests/unit/agent/test_mcp_server.py
Original file line number Diff line number Diff line change
Expand Up @@ -855,11 +855,17 @@ def test_mcp_tool_descriptions_expose_debugging_result_contracts():
assert "visible_route_keys" in isolate_properties
assert "selected_route_key" in isolate_properties
assert "axis_indices" in isolate_properties
viewer_validation_properties = schemas["openhcs_validate_viewer_window_state"][
"properties"
]
viewer_validation_schema = schemas["openhcs_validate_viewer_window_state"]
viewer_validation_properties = viewer_validation_schema["properties"]
assert "route_key" in viewer_validation_properties
assert "include_state" in viewer_validation_properties
assert viewer_validation_properties["transport_mode"]["anyOf"][0] == {
"$ref": "#/$defs/TransportMode"
}
assert viewer_validation_schema["$defs"]["TransportMode"]["enum"] == [
"tcp",
"ipc",
]
assert "compact_actions" in schemas["openhcs_ui_get_widget_tree"]["properties"]
assert (
"maximum_item_model_nodes"
Expand Down Expand Up @@ -14817,6 +14823,22 @@ def test_mcp_viewer_connection_fields_project_timeout_policy():
)


def test_mcp_viewer_connection_tool_fields_parse_nominal_transport_from_wire():
control_args = server.McpViewerRequestToolBindingABC.control_args(
{
"port": 5555,
"host": "127.0.0.1",
"transport_mode": "tcp",
"timeout_ms": 2000,
},
server.McpViewerCommandTimeoutPolicy,
)

assert control_args.connection.transport_mode is TransportMode.TCP
assert control_args.connection.host == "127.0.0.1"
assert control_args.timeout_ms == 2000


def test_mcp_viewer_mutation_tools_default_to_command_timeout():
if importlib.util.find_spec("mcp") is None:
return
Expand Down
Loading