diff --git a/openhcs/mcp/server.py b/openhcs/mcp/server.py index 7b6670dde..90a44d50e 100644 --- a/openhcs/mcp/server.py +++ b/openhcs/mcp/server.py @@ -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 @@ -1977,31 +1978,7 @@ 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( @@ -2009,13 +1986,10 @@ def control_args( 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( @@ -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, diff --git a/tests/integration/test_cellprofiler_official30_zmq.py b/tests/integration/test_cellprofiler_official30_zmq.py index a6657b84b..4489bbb40 100644 --- a/tests/integration/test_cellprofiler_official30_zmq.py +++ b/tests/integration/test_cellprofiler_official30_zmq.py @@ -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 ): @@ -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) @@ -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( @@ -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 @@ -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( diff --git a/tests/unit/agent/test_mcp_server.py b/tests/unit/agent/test_mcp_server.py index 0394ea9d0..e2a984e10 100644 --- a/tests/unit/agent/test_mcp_server.py +++ b/tests/unit/agent/test_mcp_server.py @@ -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" @@ -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