Source code for nvflare.fuel.f3.cellnet.cell

# Copyright (c) 2023, NVIDIA CORPORATION.  All rights reserved.
#
# 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.

import concurrent.futures
import copy
import os
import threading
import uuid
from typing import Dict, List, Union

from nvflare.apis.fl_constant import ServerCommandNames
from nvflare.apis.signal import Signal
from nvflare.fuel.f3.cellnet.core_cell import CoreCell, TargetMessage
from nvflare.fuel.f3.cellnet.defs import CellChannel, MessageHeaderKey, MessagePropKey, MessageType, ReturnCode
from nvflare.fuel.f3.cellnet.fqcn import FQCN
from nvflare.fuel.f3.cellnet.utils import decode_payload, encode_payload, make_reply
from nvflare.fuel.f3.message import Message
from nvflare.fuel.f3.stream_cell import StreamCell
from nvflare.fuel.f3.streaming.stream_const import StreamHeaderKey
from nvflare.fuel.f3.streaming.stream_types import BlobSizeError, StreamFuture, StreamTargetUnreachable
from nvflare.fuel.utils.fobs import FOBSContextKey
from nvflare.fuel.utils.log_utils import get_obj_logger
from nvflare.fuel.utils.waiter_utils import WaiterRC, conditional_wait
from nvflare.security.logging import secure_format_exception

CHANNELS_TO_EXCLUDE = (
    CellChannel.CLIENT_MAIN,
    CellChannel.SERVER_MAIN,
    CellChannel.SERVER_PARENT_LISTENER,
    CellChannel.CLIENT_COMMAND,
    CellChannel.CLIENT_SUB_WORKER_COMMAND,
    CellChannel.MULTI_PROCESS_EXECUTOR,
    CellChannel.SIMULATOR_RUNNER,
    CellChannel.RETURN_ONLY,
)


def _is_stream_channel(channel: str) -> bool:
    if channel is None or channel == "":
        return False
    elif channel in CHANNELS_TO_EXCLUDE:
        return False
    # if not excluded, all channels supporting streaming capabilities
    return True


def _is_server_job_cell(my_info) -> bool:
    """Return True only for server cells owned by one job run.

    Parent server cell FQCN is "server"; server job cells start with
    "server.<job_id>" and may have nested child cells.
    The fail-fast path must only exit a job process, never the parent server.
    """
    fqcn = getattr(my_info, "fqcn", "")
    if not fqcn:
        return False

    path = FQCN.split(fqcn)
    return len(path) >= 2 and path[0] == FQCN.ROOT_SERVER


[docs] class SimpleWaiter: def __init__(self, req_id, target, result): super().__init__() self.req_id = req_id self.target = target self.result = result self.receiving_future = None self.stream_error = None self.in_receiving = threading.Event()
[docs] class Adapter: def __init__(self, cb, my_info, cell): self.cb = cb self.my_info = my_info self.cell = cell self.logger = get_obj_logger(self)
[docs] def call(self, future, *args, **kwargs): # this will be called by StreamCell upon receiving the first byte of blob headers = future.headers stream_req_id = headers.get(StreamHeaderKey.STREAM_REQ_ID, "") origin = headers.get(MessageHeaderKey.ORIGIN, None) result = future.result() self.logger.debug(f"{stream_req_id=}: {headers=}, incoming data={result}") request = Message(headers, result) # PASS_THROUGH can be requested per-message, per-channel, or for one # exact (channel, topic) route. Any source activates LazyDownloadRef # decode so tensors are not downloaded at this hop. channel = request.get_header(StreamHeaderKey.CHANNEL) topic = request.get_header(StreamHeaderKey.TOPIC) passthrough = bool(request.get_header(MessageHeaderKey.PASS_THROUGH, False)) if channel in self.cell.decode_pass_through_channels: passthrough = True if (channel, topic) in self.cell.decode_pass_through_topics: passthrough = True decode_ctx = self.cell.get_fobs_context(props={FOBSContextKey.PASS_THROUGH: passthrough}) try: decode_payload(request, StreamHeaderKey.PAYLOAD_ENCODING, fobs_ctx=decode_ctx) except Exception as ex: if ( channel == CellChannel.SERVER_COMMAND and topic == ServerCommandNames.SUBMIT_UPDATE and _is_server_job_cell(self.my_info) ): self.logger.critical( f"failed to decode streamed submit_update in server job cell {self.my_info.fqcn}; " f"exiting job process to fail the job: {secure_format_exception(ex)}" ) os._exit(1) raise request.set_header(MessageHeaderKey.CHANNEL, channel) request.set_header(MessageHeaderKey.TOPIC, topic) self.logger.debug(f"Call back on {stream_req_id=}: {channel=}, {topic=}") req_id = request.get_header(MessageHeaderKey.REQ_ID, "") secure = request.get_header(MessageHeaderKey.SECURE, False) self.logger.debug(f"{stream_req_id=}: on {channel=}, {topic=}") response = self.cb(request, *args, **kwargs) if isinstance(response, concurrent.futures.Future): response.add_done_callback( lambda done: self._send_async_response(done, stream_req_id, req_id, channel, topic, origin, secure) ) return self._send_response(response, stream_req_id, req_id, channel, topic, origin, secure)
def _send_async_response(self, response_future, *reply_args): try: response = response_future.result() except Exception as ex: self.logger.error(f"async request callback failed: {secure_format_exception(ex)}") response = make_reply(ReturnCode.PROCESS_EXCEPTION) self._send_response(response, *reply_args) def _handle_reply_stream_done(self, reply_future, transport_optional=False): error = reply_future.exception() if not error: return if isinstance(error, BlobSizeError) and _is_server_job_cell(self.my_info): self.logger.critical( f"streamed response from server job cell {self.my_info.fqcn} was rejected as too large; " f"exiting job process to fail the job: {secure_format_exception(error)}" ) os._exit(1) if transport_optional and isinstance(error, StreamTargetUnreachable): # The response transport is optional because it can outlive its requester or job cell. This callback # cannot determine remote waiter state; request-level failures surface at the requester through its # correlated stream error or timeout, so an unreachable optional response is diagnostic only. self.logger.debug( f"streamed response from {self.my_info.fqcn} was not delivered: {secure_format_exception(error)}" ) return self.logger.error( f"streamed response from {self.my_info.fqcn} failed asynchronously: {secure_format_exception(error)}" ) def _send_response(self, response, stream_req_id, req_id, channel, topic, origin, secure): self.logger.debug(f"response available: {stream_req_id=}: on {channel=}, {topic=}") if not stream_req_id: # no need to reply! self.logger.debug("Do not send reply because there is no stream_req_id!") return if not isinstance(response, Message): self.logger.error(f"request callback must return Message but got {type(response)}") response = make_reply(ReturnCode.PROCESS_EXCEPTION) response.add_headers( { MessageHeaderKey.REQ_ID: req_id, MessageHeaderKey.MSG_TYPE: MessageType.REPLY, StreamHeaderKey.STREAM_REQ_ID: stream_req_id, } ) encode_payload(response, StreamHeaderKey.PAYLOAD_ENCODING, fobs_ctx=self.cell.get_fobs_context()) self.logger.debug(f"sending: {stream_req_id=}: {response.headers=}, target={origin}") try: # Response production is asynchronous and can outlive the request timeout. Mark its transport frames # optional so a stale destination is dropped without ERROR-level routing noise. Reliable delivery and # receiver-side stream errors still settle reply_future and the original request waiter when present. transport_optional = True reply_future = self.cell.send_blob( CellChannel.RETURN_ONLY, f"{channel}:{topic}", origin, response, secure, optional=transport_optional, reliable=True, ) except BlobSizeError as ex: if _is_server_job_cell(self.my_info): self.logger.critical( f"streamed response from server job cell {self.my_info.fqcn} is too large; " f"exiting job process to fail the job: {secure_format_exception(ex)}" ) os._exit(1) raise reply_future.add_done_callback(self._handle_reply_stream_done, reply_future, transport_optional) self.logger.debug(f"Done sending: {stream_req_id=}: {reply_future=}")
[docs] class Cell(StreamCell): def __init__(self, *args, **kwargs): self.core_cell = CoreCell(*args, **kwargs) super().__init__(self.core_cell) self.requests_dict = dict() self.logger = get_obj_logger(self) self.byte_streamer.register_error_callback(self._process_stream_error) self.register_blob_cb(CellChannel.RETURN_ONLY, "*", self._process_reply) # this should be one-time registration self.core_cell.update_fobs_context({FOBSContextKey.CELL: self}) self.decode_pass_through_channels: set = set() # per-channel opt-in for receiver-side PASS_THROUGH self.decode_pass_through_topics: set = set() # exact (channel, topic) receiver-side opt-in
[docs] def update_fobs_context(self, props: dict): self.core_cell.update_fobs_context(props)
[docs] def get_fobs_context(self, props: dict = None): """Return a new copy of the fobs context. If props is specified, they will be set into the context. Returns: a new copy of the fobs context """ ctx = self.core_cell.get_fobs_context() if props: ctx.update(props) return ctx
def __getattr__(self, func): """ This method is called when Python cannot find an invoked method "x" of this class. Method "x" is one of the message sending methods (send_request, broadcast_request, etc.) In this method, we decide which method should be used instead, based on the "channel" of the message. - If the channel is stream channel, use the method "_x" of this class. - Otherwise, user the method "x" of the CoreCell. """ def method(*args, **kwargs): self.logger.debug(f"__getattr__: {args=}, {kwargs=}") if _is_stream_channel(kwargs.get("channel")): self.logger.debug(f"calling cell {func}") return getattr(self, f"_{func}")(*args, **kwargs) if not hasattr(self.core_cell, func): raise AttributeError(f"'{func}' not in core_cell.") self.logger.debug(f"calling core_cell {func}") return getattr(self.core_cell, func)(*args, **kwargs) return method def _broadcast_request( self, channel: str, topic: str, targets: Union[str, List[str]], request: Message, timeout=None, secure=False, optional=False, abort_signal: Signal = None, ) -> Dict[str, Message]: """ Send a message over a channel to specified destination cell(s), and wait for reply Args: channel: channel for the message topic: topic of the message targets: FQCN of the destination cell(s) request: message to be sent timeout: how long to wait for replies secure: End-end encryption optional: whether the message is optional abort_signal: signal to abort the message Returns: a dict of: cell_id => reply message """ self.logger.debug(f"broadcast: {channel=}, {topic=}, {targets=}, {timeout=}") if isinstance(targets, str): targets = [targets] target_argument = {} fixed_dict = dict(channel=channel, topic=topic, timeout=timeout, secure=secure, optional=optional) results = dict() future_to_target = {} # Encode the request now so each target thread won't need to do it again. # For a direct broadcast, the routing targets are also the tensor download # consumers, so the transaction can track their exact identities. With # PASS_THROUGH, each target may forward the refs to another Cell (for # example, an external trainer), and only the final consumer count is known. # Pinning the transaction to the first-hop identities would prevent it from # completing after those downstream consumers finish downloading. pass_through = bool(request.get_header(MessageHeaderKey.PASS_THROUGH, False)) receiver_ids = None if pass_through else targets self._encode_message( request, abort_signal, num_receivers=len(targets), receiver_ids=receiver_ids, ) with concurrent.futures.ThreadPoolExecutor(max_workers=len(targets)) as executor: self.logger.debug(f"broadcast to {targets=}") for t in targets: req = Message(copy.deepcopy(request.headers), request.payload) target_argument["request"] = TargetMessage(t, channel, topic, req).message target_argument["target"] = t target_argument["abort_signal"] = abort_signal target_argument.update(fixed_dict) f = executor.submit(self._send_one_request, **target_argument) future_to_target[f] = t self.logger.debug(f"submitted to {t} with {target_argument.keys()=}") for future in concurrent.futures.as_completed(future_to_target): target = future_to_target[future] self.logger.debug(f"{target} completed") try: data = future.result() except Exception as exc: self.logger.warning(f"{target} raises {exc}") results[target] = make_reply(ReturnCode.TIMEOUT) else: results[target] = data self.logger.debug(f"{target=}: {data=}") self.logger.debug("About to return from broadcast_request") return results def _fire_and_forget( self, channel: str, topic: str, targets: Union[str, List[str]], message: Message, secure=False, optional=False, ) -> Dict[str, str]: """ Send a message over a channel to specified destination cell(s), and do not wait for replies. Args: channel: channel for the message topic: topic of the message targets: one or more destination cell IDs. None means all. message: message to be sent secure: End-end encryption if True optional: whether the message is optional Returns: None """ encode_payload(message, encoding_key=StreamHeaderKey.PAYLOAD_ENCODING, fobs_ctx=self.get_fobs_context()) if isinstance(targets, str): targets = [targets] result = {} futures = {} for target in targets: future = self.send_blob( channel=channel, topic=topic, target=target, message=message, secure=secure, optional=optional ) futures[target] = future result[target] = "" message.set_prop(MessagePropKey.FUTURES, futures) return result def _get_result(self, req_id): waiter = self.requests_dict.pop(req_id) return waiter.result def _check_error(self, future): if future.error: # must return a negative number return -1 else: return WaiterRC.OK def _future_wait(self, future, timeout, abort_signal: Signal): # future could have an error! last_progress = 0 while True: rc = conditional_wait(future.waiter, timeout, abort_signal, condition_cb=self._check_error, future=future) if rc == WaiterRC.IS_SET: # waiter has been set! break elif rc == WaiterRC.TIMEOUT: # timed out: check whether any progress has been made during this time current_progress = future.get_progress() if last_progress == current_progress: # no progress in timeout secs: consider this to be a failure return False else: # good progress self.logger.debug(f"{current_progress=}") last_progress = current_progress else: # error condition: aborted or future error return False if future.error: return False else: return True def _encode_message( self, msg: Message, abort_signal, num_receivers=1, receiver_ids=None, fobs_ctx_props: dict = None, ) -> int: try: # Keep call-scoped FOBS behavior out of the Cell's shared context. This is # especially important when a trainer logs from one thread while another # serializes a result: temporarily mutating the CoreCell context would let # transaction-lifecycle callbacks leak into the log encode. props = dict(fobs_ctx_props or {}) props.update( { FOBSContextKey.ABORT_SIGNAL: abort_signal, FOBSContextKey.NUM_RECEIVERS: num_receivers, } ) if receiver_ids is not None: props[FOBSContextKey.RECEIVER_IDS] = receiver_ids return encode_payload(msg, StreamHeaderKey.PAYLOAD_ENCODING, fobs_ctx=self.get_fobs_context(props)) except BaseException as exc: self.logger.error(f"Can't encode {msg=} {exc=}") raise exc def _send_request( self, channel, target, topic, request, timeout=10.0, secure=False, optional=False, abort_signal: Signal = None, progress_wait_cb=None, num_receivers=1, receiver_ids=None, fobs_ctx_props: dict = None, ): """Stream one request to the target Args: channel: message channel name target: FQCN of the target cell topic: topic of the message request: request message timeout: how long to wait secure: is P2P security to be applied optional: is the message optional abort_signal: signal to abort the message fobs_ctx_props: optional call-scoped FOBS context properties used only while serializing this request Returns: reply data """ self._encode_message( request, abort_signal, num_receivers=num_receivers, receiver_ids=receiver_ids, fobs_ctx_props=fobs_ctx_props, ) return self._send_one_request( channel, target, topic, request, timeout, secure, optional, abort_signal, progress_wait_cb ) def _send_one_request( self, channel, target, topic, request, timeout=10.0, secure=False, optional=False, abort_signal=None, progress_wait_cb=None, ): req_id = str(uuid.uuid4()) request.add_headers({StreamHeaderKey.STREAM_REQ_ID: req_id}) # this future can be used to check sending progress, but not for checking return blob self.logger.debug(f"{req_id=}, {channel=}, {topic=}, {target=}, {timeout=}: send_request about to send_blob") waiter = SimpleWaiter(req_id=req_id, target=target, result=make_reply(ReturnCode.TIMEOUT)) self.requests_dict[req_id] = waiter try: future = self.send_blob( channel=channel, topic=topic, target=target, message=request, secure=secure, optional=optional ) self.logger.debug(f"{req_id=}: Waiting starts") # Three stages, sending, waiting for receiving first byte, receiving # sending with progress timeout self.logger.debug(f"{req_id=}: entering sending wait {timeout=}") sending_complete = self._future_wait(future, timeout, abort_signal) if not sending_complete: stream_error = waiter.stream_error if future.error and not stream_error: waiter.result = make_reply(ReturnCode.COMM_ERROR, error=str(future.error)) self.logger.debug(f"{req_id=}: sending failed with stream error: {future.error}") elif stream_error: self.logger.debug(f"{req_id=}: receiver rejected stream: {stream_error}") else: self.logger.debug(f"{req_id=}: sending timeout {timeout=}") return self._get_result(req_id) self.logger.debug(f"{req_id=}: sending complete") # waiting for receiving first byte self.logger.debug(f"{req_id=}: entering remote process wait {timeout=}") # With progress-aware waits, timeout is the remote-process idle polling interval. # The progress callback owns the total idle budget and decides whether to keep waiting. while True: waiter_rc = conditional_wait(waiter.in_receiving, timeout, abort_signal) if waiter_rc == WaiterRC.IS_SET: break if waiter_rc == WaiterRC.TIMEOUT and progress_wait_cb: try: if progress_wait_cb(): self.logger.debug( f"{req_id=}: remote processing still has transfer progress; continue waiting" ) continue except Exception as ex: self.logger.warning( f"{req_id=}: progress_wait_cb raised, treating as no-progress: " f"{secure_format_exception(ex)}" ) self.logger.debug(f"{req_id=}: remote processing timeout {timeout=} {waiter_rc=}") return self._get_result(req_id) self.logger.debug(f"{req_id=}: in receiving") if waiter.stream_error: self.logger.debug(f"{req_id=}: returning stream error reply") return self._get_result(req_id) # receiving with progress timeout r_future = waiter.receiving_future self.logger.debug(f"{req_id=}: entering receiving wait {timeout=}") while True: receiving_complete = self._future_wait(r_future, timeout, abort_signal) if receiving_complete: break if abort_signal and abort_signal.triggered: self.logger.info(f"{req_id=}: receiving aborted {timeout=}") return self._get_result(req_id) if getattr(r_future, "error", None): self.logger.info(f"{req_id=}: receiving failed with stream error") return self._get_result(req_id) if progress_wait_cb: # The same bounded progress policy applies after the first byte: a streamed reply body can still # be moving even though the peer has already started responding. try: if progress_wait_cb(): self.logger.debug(f"{req_id=}: receiving still has transfer progress; continue waiting") continue except Exception as ex: self.logger.warning( f"{req_id=}: progress_wait_cb raised during receiving, treating as no-progress: " f"{secure_format_exception(ex)}" ) self.logger.info(f"{req_id=}: receiving timeout {timeout=}") return self._get_result(req_id) self.logger.debug(f"{req_id=}: receiving complete") waiter.result = Message(r_future.headers, r_future.result()) pt = bool(waiter.result.get_header(MessageHeaderKey.PASS_THROUGH, False)) if channel in self.decode_pass_through_channels: pt = True if (channel, topic) in self.decode_pass_through_topics: pt = True decode_payload( waiter.result, encoding_key=StreamHeaderKey.PAYLOAD_ENCODING, fobs_ctx=self.get_fobs_context( props={ FOBSContextKey.ABORT_SIGNAL: abort_signal, FOBSContextKey.PASS_THROUGH: pt, } ), ) self.logger.debug(f"{req_id=}: return result {waiter.result=}") return self._get_result(req_id) except Exception as ex: self.logger.error(f"exception sending request: {secure_format_exception(ex)}") self.requests_dict.pop(req_id, None) raise ex def _process_reply(self, future: StreamFuture): headers = future.headers req_id = headers.get(StreamHeaderKey.STREAM_REQ_ID, -1) self.logger.debug(f"{req_id=}: _process_reply") try: waiter = self.requests_dict[req_id] except KeyError as e: self.logger.warning(f"Receiving unknown {req_id=}, discarded: {e}") return waiter.receiving_future = future waiter.in_receiving.set() def _process_stream_error(self, message: Message): req_id = message.get_header(StreamHeaderKey.STREAM_REQ_ID) error = message.get_header(StreamHeaderKey.ERROR_MSG, "stream rejected by receiver") error_type = message.get_header(StreamHeaderKey.ERROR_TYPE) if req_id: waiter = self.requests_dict.get(req_id) if waiter: origin = message.get_header(MessageHeaderKey.ORIGIN) failed_destination = message.get_header(StreamHeaderKey.FAILED_DESTINATION, origin) forwarded_error = message.get_header(StreamHeaderKey.FAILED_DESTINATION) == waiter.target if origin != waiter.target and not forwarded_error: self.logger.warning( f"ignored stream error for {req_id=} with unexpected failed destination {failed_destination}; " f"expected {waiter.target}" ) return waiter.stream_error = error waiter.result = make_reply(ReturnCode.COMM_ERROR, error=error) waiter.in_receiving.set() return self.logger.debug(f"stream error for completed or unknown {req_id=}") if error_type == BlobSizeError.__name__ and _is_server_job_cell(self.core_cell.my_info): self.logger.critical( f"streamed message from server job cell {self.core_cell.my_info.fqcn} was rejected as too large; " f"exiting job process to fail the job: {error}" ) os._exit(1) def _register_request_cb(self, channel: str, topic: str, cb, *args, **kwargs): """ Register a callback for handling request. The CB must follow request_cb_signature. Args: channel: the channel of the request topic: topic of the request cb: *args: **kwargs: Returns: """ if not callable(cb): raise ValueError(f"specified request_cb {type(cb)} is not callable") # always register with core_cell since some requests (e.g. broadcast_multi_requests) will directly go # through the core_cell, even if the channel may be a stream channel (e.g. aux channel). self.core_cell.register_request_cb(channel, topic, cb, *args, **kwargs) if _is_stream_channel(channel): self.logger.info(f"Register blob CB for {channel=}, {topic=}") adapter = Adapter(cb, self.core_cell.my_info, self) self.register_blob_cb(channel, topic, adapter.call, *args, **kwargs)