From 801553a41a0fe4e513a95c075c1f8fafa89ebb3d Mon Sep 17 00:00:00 2001 From: zhsu Date: Wed, 1 Jul 2026 17:02:12 +0800 Subject: [PATCH 1/2] Support SPU link_desc config --- mplang/backends/spu_state.py | 62 +++++++++++++++++--- mplang/dialects/spu.py | 68 +++++++++++++++++++--- tests/backends/test_spu_state.py | 97 ++++++++++++++++++++++++++++++++ tests/dialects/test_spu.py | 58 +++++++++++++++++++ 4 files changed, 267 insertions(+), 18 deletions(-) create mode 100644 tests/backends/test_spu_state.py diff --git a/mplang/backends/spu_state.py b/mplang/backends/spu_state.py index 77ccb0e3..6e4046b1 100644 --- a/mplang/backends/spu_state.py +++ b/mplang/backends/spu_state.py @@ -22,7 +22,7 @@ import logging import os -from dataclasses import dataclass, field +from dataclasses import dataclass, field, replace from typing import TYPE_CHECKING, Any import spu.api as spu_api @@ -117,19 +117,53 @@ def __init__( # Optional shared infrastructure (for per-request isolation via link.spawn) self._infra = infra self._brpc_config = brpc_config or BrpcLinkConfig() - # Key: (local_rank, world_size, protocol, field, link_mode, spu_endpoints, brpc_config) - # ``brpc_config`` participates only when link_mode == "brpc" (else None) so that - # callers passing custom configs don't silently reuse the first config's link. + # Key: (local_rank, world_size, protocol, field, fxp_fraction_bits, + # link_mode, spu_endpoints, brpc_config). ``brpc_config`` participates + # only when link_mode == "brpc" (else None) so callers passing custom + # configs don't silently reuse the first config's link. # Value: (Runtime, Io) self._runtimes: dict[ tuple[ - int, int, str, str, str, tuple[str, ...] | None, BrpcLinkConfig | None + int, + int, + str, + str, + int, + str, + tuple[str, ...] | None, + BrpcLinkConfig | None, ], tuple[spu_api.Runtime, spu_api.Io], ] = {} # Local template link cache (used when no WorkerInfra is provided) self._template_links: dict[tuple, libspu.link.Context] = {} + def _effective_brpc_config(self, config: spu.SPUConfig) -> BrpcLinkConfig: + """Merge per-SPU link_desc over this state's default brpc config.""" + link_desc = config.link_desc + if link_desc is None: + return self._brpc_config + + overrides: dict[str, Any] = {} + if link_desc.brpc_channel_protocol is not None: + overrides["protocol"] = link_desc.brpc_channel_protocol + if link_desc.brpc_channel_connection_type is not None: + overrides["connection_type"] = link_desc.brpc_channel_connection_type + if link_desc.recv_timeout_ms is not None: + overrides["recv_timeout_ms"] = link_desc.recv_timeout_ms + if link_desc.http_max_payload_size is not None: + overrides["http_max_payload_size"] = link_desc.http_max_payload_size + if link_desc.http_timeout_ms is not None: + overrides["http_timeout_ms"] = link_desc.http_timeout_ms + if link_desc.connect_retry_times is not None: + overrides["connect_retry_times"] = link_desc.connect_retry_times + if link_desc.connect_retry_interval_ms is not None: + overrides["connect_retry_interval_ms"] = link_desc.connect_retry_interval_ms + + if not overrides: + return self._brpc_config + return replace(self._brpc_config, **overrides) + def _get_template_link( self, cache_key: tuple, @@ -138,6 +172,7 @@ def _get_template_link( communicator: object | None, parties: list[int] | None, spu_endpoints: list[str] | None, + brpc_config: BrpcLinkConfig | None = None, ) -> libspu.link.Context: """Get or create a template link for the given configuration. @@ -147,7 +182,8 @@ def _get_template_link( def _create() -> libspu.link.Context: if spu_endpoints: - return self._create_brpc_link(local_rank, spu_endpoints) + cfg = brpc_config or self._brpc_config + return self._create_brpc_link(local_rank, spu_endpoints, cfg) elif communicator is not None: if parties is None: raise ValueError("parties required when using communicator") @@ -203,14 +239,18 @@ def get_or_create( else: link_mode = "mem" + brpc_config = ( + self._effective_brpc_config(config) if link_mode == "brpc" else None + ) cache_key = ( local_rank, spu_world_size, config.protocol, config.field, + config.fxp_fraction_bits, link_mode, tuple(spu_endpoints) if spu_endpoints else None, - self._brpc_config if link_mode == "brpc" else None, + brpc_config, ) if cache_key in self._runtimes: @@ -224,6 +264,7 @@ def get_or_create( communicator, parties, spu_endpoints, + brpc_config, ) link = template_link.spawn() @@ -290,10 +331,13 @@ def _create_channels_link( return libspu.link.create_with_channels(desc, local_rank, channels) def _create_brpc_link( - self, local_rank: int, spu_endpoints: list[str] + self, + local_rank: int, + spu_endpoints: list[str], + brpc_config: BrpcLinkConfig | None = None, ) -> libspu.link.Context: """Create BRPC link for distributed execution.""" - cfg = self._brpc_config + cfg = brpc_config or self._brpc_config desc = libspu.link.Desc() # type: ignore desc.recv_timeout_ms = cfg.recv_timeout_ms desc.http_max_payload_size = cfg.http_max_payload_size diff --git a/mplang/dialects/spu.py b/mplang/dialects/spu.py index e19d6912..37860f08 100644 --- a/mplang/dialects/spu.py +++ b/mplang/dialects/spu.py @@ -94,6 +94,46 @@ def secure_add(x, y): # ============================================================================== +@dataclass(frozen=True) +class SPULinkDesc: + """SPU brpc link configuration compatible with SecretFlow link_desc.""" + + recv_timeout_ms: int | None = None + http_timeout_ms: int | None = None + http_max_payload_size: int | None = None + brpc_channel_protocol: str | None = None + brpc_channel_connection_type: str | None = None + connect_retry_times: int | None = None + connect_retry_interval_ms: int | None = None + + @classmethod + def from_dict(cls, d: dict[str, Any] | SPULinkDesc) -> SPULinkDesc: + if isinstance(d, SPULinkDesc): + return d + return cls( + recv_timeout_ms=d.get("recv_timeout_ms"), + http_timeout_ms=d.get("http_timeout_ms"), + http_max_payload_size=d.get("http_max_payload_size"), + brpc_channel_protocol=d.get("brpc_channel_protocol") or d.get("protocol"), + brpc_channel_connection_type=d.get("brpc_channel_connection_type") + or d.get("connection_type"), + connect_retry_times=d.get("connect_retry_times"), + connect_retry_interval_ms=d.get("connect_retry_interval_ms"), + ) + + def to_json(self) -> dict[str, Any]: + data = { + "recv_timeout_ms": self.recv_timeout_ms, + "http_timeout_ms": self.http_timeout_ms, + "http_max_payload_size": self.http_max_payload_size, + "brpc_channel_protocol": self.brpc_channel_protocol, + "brpc_channel_connection_type": self.brpc_channel_connection_type, + "connect_retry_times": self.connect_retry_times, + "connect_retry_interval_ms": self.connect_retry_interval_ms, + } + return {key: value for key, value in data.items() if value is not None} + + @serde.register_class @dataclass(frozen=True) class SPUConfig: @@ -103,37 +143,47 @@ class SPUConfig: protocol: SPU protocol (e.g., "SEMI2K", "ABY3"). field: SPU field type (e.g., "FM64", "FM128"). fxp_fraction_bits: Fixed-point fraction bits. + link_desc: Optional brpc link configuration. """ protocol: str = "SEMI2K" field: str = "FM128" fxp_fraction_bits: int = 18 + link_desc: SPULinkDesc | None = None @classmethod def from_dict(cls, d: dict[str, Any]) -> SPUConfig: + runtime_config = d.get("runtime_config") or {} + if not isinstance(runtime_config, dict): + raise TypeError( + f"SPUConfig.runtime_config must be a dict, got {runtime_config!r}" + ) + link_desc = d.get("link_desc") return cls( - protocol=d.get("protocol", "SEMI2K"), - field=d.get("field", "FM128"), - fxp_fraction_bits=d.get("fxp_fraction_bits", 18), + protocol=runtime_config.get("protocol", d.get("protocol", "SEMI2K")), + field=runtime_config.get("field", d.get("field", "FM128")), + fxp_fraction_bits=runtime_config.get( + "fxp_fraction_bits", d.get("fxp_fraction_bits", 18) + ), + link_desc=SPULinkDesc.from_dict(link_desc) if link_desc else None, ) # --- Serde methods --- _serde_kind: ClassVar[str] = "spu.SPUConfig" def to_json(self) -> dict[str, Any]: - return { + data: dict[str, Any] = { "protocol": self.protocol, "field": self.field, "fxp_fraction_bits": self.fxp_fraction_bits, } + if self.link_desc is not None: + data["link_desc"] = self.link_desc.to_json() + return data @classmethod def from_json(cls, data: dict[str, Any]) -> SPUConfig: - return cls( - protocol=data["protocol"], - field=data["field"], - fxp_fraction_bits=data["fxp_fraction_bits"], - ) + return cls.from_dict(data) # ============================================================================== diff --git a/tests/backends/test_spu_state.py b/tests/backends/test_spu_state.py new file mode 100644 index 00000000..09b94a58 --- /dev/null +++ b/tests/backends/test_spu_state.py @@ -0,0 +1,97 @@ +# Copyright 2026 Ant Group Co., Ltd. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +"""Tests for SPU dialect state.""" + +from __future__ import annotations + +from types import SimpleNamespace + +from mplang.backends import spu_state +from mplang.backends.spu_state import BrpcLinkConfig, SPUState +from mplang.dialects import spu + + +def test_effective_brpc_config_merges_spu_link_desc() -> None: + base = BrpcLinkConfig( + protocol="h2", + connection_type="single", + recv_timeout_ms=1000, + http_max_payload_size=1024, + http_timeout_ms=None, + connect_retry_times=3, + connect_retry_interval_ms=4, + ) + state = SPUState(brpc_config=base) + config = spu.SPUConfig( + link_desc=spu.SPULinkDesc( + brpc_channel_protocol="http", + brpc_channel_connection_type="pooled", + recv_timeout_ms=7200000, + http_timeout_ms=7200000, + ) + ) + + effective = state._effective_brpc_config(config) + + assert effective.protocol == "http" + assert effective.connection_type == "pooled" + assert effective.recv_timeout_ms == 7200000 + assert effective.http_timeout_ms == 7200000 + assert effective.http_max_payload_size == 1024 + assert effective.connect_retry_times == 3 + assert effective.connect_retry_interval_ms == 4 + + +def test_create_brpc_link_applies_explicit_config(monkeypatch) -> None: + class FakeDesc: + def __init__(self) -> None: + self.parties: list[tuple[str, str]] = [] + + def add_party(self, name: str, endpoint: str) -> None: + self.parties.append((name, endpoint)) + + captured = {} + + def fake_create_brpc(desc: FakeDesc, local_rank: int) -> str: + captured["desc"] = desc + captured["local_rank"] = local_rank + return "fake-link" + + fake_link = SimpleNamespace(Desc=FakeDesc, create_brpc=fake_create_brpc) + monkeypatch.setattr(spu_state.libspu, "link", fake_link) + + cfg = BrpcLinkConfig( + protocol="http", + connection_type="pooled", + recv_timeout_ms=7200000, + http_max_payload_size=67108864, + http_timeout_ms=7200000, + connect_retry_times=9, + connect_retry_interval_ms=10, + ) + + link = SPUState()._create_brpc_link(1, ["127.0.0.1:9000", "127.0.0.1:9001"], cfg) + + assert link == "fake-link" + assert captured["local_rank"] == 1 + desc = captured["desc"] + assert desc.recv_timeout_ms == 7200000 + assert desc.http_max_payload_size == 67108864 + assert desc.http_timeout_ms == 7200000 + assert desc.brpc_channel_protocol == "http" + assert desc.brpc_channel_connection_type == "pooled" + assert desc.connect_retry_times == 9 + assert desc.connect_retry_interval_ms == 10 + assert desc.parties == [("P0", "127.0.0.1:9000"), ("P1", "127.0.0.1:9001")] diff --git a/tests/dialects/test_spu.py b/tests/dialects/test_spu.py index 3289a5fc..870fc446 100644 --- a/tests/dialects/test_spu.py +++ b/tests/dialects/test_spu.py @@ -20,11 +20,69 @@ import mplang.edsl as el import mplang.edsl.typing as elt from mplang.dialects import simp, spu +from mplang.edsl import serde def test_spu_config(): config = spu.SPUConfig() assert config.protocol == "SEMI2K" + assert config.link_desc is None + + +def test_spu_config_from_secretflow_dict(): + config = spu.SPUConfig.from_dict({ + "runtime_config": { + "protocol": "ABY3", + "field": "FM64", + "fxp_fraction_bits": 20, + }, + "link_desc": { + "connect_retry_times": 7200000, + "connect_retry_interval_ms": 7200000, + "brpc_channel_protocol": "http", + "brpc_channel_connection_type": "pooled", + "recv_timeout_ms": 7200000, + "http_timeout_ms": 7200000, + "http_max_payload_size": 67108864, + }, + }) + + assert config.protocol == "ABY3" + assert config.field == "FM64" + assert config.fxp_fraction_bits == 20 + assert config.link_desc is not None + assert config.link_desc.brpc_channel_protocol == "http" + assert config.link_desc.brpc_channel_connection_type == "pooled" + assert config.link_desc.recv_timeout_ms == 7200000 + assert config.link_desc.http_timeout_ms == 7200000 + assert config.link_desc.http_max_payload_size == 67108864 + assert config.link_desc.connect_retry_times == 7200000 + assert config.link_desc.connect_retry_interval_ms == 7200000 + + +def test_spu_config_json_roundtrip_keeps_link_desc(): + config = spu.SPUConfig( + protocol="SEMI2K", + field="FM128", + fxp_fraction_bits=18, + link_desc=spu.SPULinkDesc( + recv_timeout_ms=1234, + http_timeout_ms=5678, + brpc_channel_protocol="h2", + brpc_channel_connection_type="pooled", + ), + ) + + payload = serde.to_json(config) + result = serde.from_json(payload) + + assert result == config + assert payload["link_desc"] == { + "recv_timeout_ms": 1234, + "http_timeout_ms": 5678, + "brpc_channel_protocol": "h2", + "brpc_channel_connection_type": "pooled", + } def test_encrypt_decrypt_flow(): From 174f10d12f601d9d28735f1993e7c3f546a2f924 Mon Sep 17 00:00:00 2001 From: zhsu Date: Thu, 2 Jul 2026 10:06:30 +0800 Subject: [PATCH 2/2] Handle malformed SPU config inputs --- mplang/dialects/spu.py | 20 +++++++++++++++----- tests/dialects/test_spu.py | 10 ++++++++++ 2 files changed, 25 insertions(+), 5 deletions(-) diff --git a/mplang/dialects/spu.py b/mplang/dialects/spu.py index 37860f08..61f02601 100644 --- a/mplang/dialects/spu.py +++ b/mplang/dialects/spu.py @@ -110,6 +110,8 @@ class SPULinkDesc: def from_dict(cls, d: dict[str, Any] | SPULinkDesc) -> SPULinkDesc: if isinstance(d, SPULinkDesc): return d + if not isinstance(d, dict): + raise TypeError(f"SPULinkDesc.from_dict expects dict, got {type(d)!r}") return cls( recv_timeout_ms=d.get("recv_timeout_ms"), http_timeout_ms=d.get("http_timeout_ms"), @@ -153,18 +155,26 @@ class SPUConfig: @classmethod def from_dict(cls, d: dict[str, Any]) -> SPUConfig: + if not isinstance(d, dict): + raise TypeError(f"SPUConfig.from_dict expects dict, got {type(d)!r}") runtime_config = d.get("runtime_config") or {} if not isinstance(runtime_config, dict): raise TypeError( f"SPUConfig.runtime_config must be a dict, got {runtime_config!r}" ) link_desc = d.get("link_desc") + protocol = cast( + str, runtime_config.get("protocol", d.get("protocol", "SEMI2K")) + ) + field = cast(str, runtime_config.get("field", d.get("field", "FM128"))) + fxp_fraction_bits = cast( + int, + runtime_config.get("fxp_fraction_bits", d.get("fxp_fraction_bits", 18)), + ) return cls( - protocol=runtime_config.get("protocol", d.get("protocol", "SEMI2K")), - field=runtime_config.get("field", d.get("field", "FM128")), - fxp_fraction_bits=runtime_config.get( - "fxp_fraction_bits", d.get("fxp_fraction_bits", 18) - ), + protocol=protocol, + field=field, + fxp_fraction_bits=fxp_fraction_bits, link_desc=SPULinkDesc.from_dict(link_desc) if link_desc else None, ) diff --git a/tests/dialects/test_spu.py b/tests/dialects/test_spu.py index 870fc446..66850852 100644 --- a/tests/dialects/test_spu.py +++ b/tests/dialects/test_spu.py @@ -29,6 +29,16 @@ def test_spu_config(): assert config.link_desc is None +def test_spu_config_from_dict_rejects_non_dict(): + with pytest.raises(TypeError, match=r"SPUConfig\.from_dict expects dict"): + spu.SPUConfig.from_dict("bad") + + +def test_spu_link_desc_from_dict_rejects_non_dict(): + with pytest.raises(TypeError, match=r"SPULinkDesc\.from_dict expects dict"): + spu.SPULinkDesc.from_dict([("recv_timeout_ms", 1)]) + + def test_spu_config_from_secretflow_dict(): config = spu.SPUConfig.from_dict({ "runtime_config": {