Source code for nvflare.fuel.f3.streaming.byte_streamer

# 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 logging
import threading
import time
from collections import OrderedDict
from concurrent.futures import TimeoutError, as_completed
from dataclasses import dataclass
from typing import Callable, Optional

from nvflare.fuel.f3.cellnet.core_cell import CoreCell
from nvflare.fuel.f3.cellnet.defs import MessageHeaderKey, ReturnCode
from nvflare.fuel.f3.comm_config import CommConfigurator
from nvflare.fuel.f3.message import Message
from nvflare.fuel.f3.mpm import MainProcessMonitor
from nvflare.fuel.f3.stats_pool import StatsPoolManager
from nvflare.fuel.f3.streaming.stream_const import (
    STREAM_ACK_INTERVAL,
    STREAM_ACK_TOPIC,
    STREAM_CHANNEL,
    STREAM_CHUNK_SIZE,
    STREAM_DATA_TOPIC,
    STREAM_ERROR_TOPIC,
    STREAM_RETRY_MAX_PENDING_BYTES,
    STREAM_WINDOW_SIZE,
    StreamDataType,
    StreamHeaderKey,
)
from nvflare.fuel.f3.streaming.stream_types import (
    BlobSizeError,
    Stream,
    StreamError,
    StreamFuture,
    StreamTargetUnreachable,
    StreamTaskSpec,
)
from nvflare.fuel.f3.streaming.stream_utils import (
    ONE_MB,
    CheckedExecutor,
    gen_stream_id,
    stream_stats_category,
    stream_thread_pool,
    wrap_view,
)

STREAM_ACK_WAIT = 300
STREAM_RETRY_WAIT = 5.0
STREAM_RETRY_TIMEOUT = 60.0
STREAM_RETRY_WORKERS = 32
STREAM_RETRY_RESULT_TIMEOUT = 1.0
STREAM_ERROR_CONTEXT_TTL = STREAM_ACK_WAIT
MAX_STREAM_ERROR_CONTEXTS = 10000

STREAM_TYPE_BYTE = "byte"
STREAM_TYPE_BLOB = "blob"
STREAM_TYPE_FILE = "file"

COUNTER_NAME_SENT = "sent"

log = logging.getLogger(__name__)


@dataclass(frozen=True)
class _TxTaskContext:
    cell: CoreCell
    target: str
    channel: str
    topic: str
    req_id: object
    expires_at: float


def _payload_size(payload) -> int:
    if payload is None:
        return 0

    if isinstance(payload, list):
        return sum(len(item) for item in payload)

    return len(payload)


def _snapshot_payload(payload):
    if payload is None:
        return None

    if isinstance(payload, list):
        return [bytes(item) for item in payload]

    return bytes(payload)


[docs] class ReliableRetryScheduler: def __init__(self): self.tasks = {} self.cv = threading.Condition() self.thread = None self.stopped = False self.generation = 0 self.retry_task_pool = CheckedExecutor(STREAM_RETRY_WORKERS, "stm_retry") # task -> dispatch timestamp, used to detect retry dispatches stuck in transport sends self.inflight_tasks = {} self.stalled_tasks = set()
[docs] def register(self, task): with self.cv: if self.stopped: return self.tasks[task.sid] = task self.generation += 1 if not self.thread or not self.thread.is_alive(): self.thread = threading.Thread(target=self._run, name="stm_retry", daemon=True) self.thread.start() self.cv.notify()
[docs] def unregister(self, task): with self.cv: registered = self.tasks.get(task.sid) if registered is task: self.tasks.pop(task.sid, None) self.inflight_tasks.pop(task, None) self.stalled_tasks.discard(task) self.generation += 1 self.cv.notify()
[docs] def wakeup(self): with self.cv: self.generation += 1 self.cv.notify()
[docs] def shutdown(self): with self.cv: self.stopped = True self.generation += 1 self.cv.notify() thread = self.thread if thread and thread.is_alive() and thread is not threading.current_thread(): thread.join(timeout=1.0) # Cell transport must remain alive until every already-admitted retry # finishes. Standalone trainer teardown stops the Cell immediately after # this scheduler and the shared streaming executors have drained. self.retry_task_pool.shutdown(wait=True)
def _finish_inflight(self, task): with self.cv: self.inflight_tasks.pop(task, None) self.stalled_tasks.discard(task) self.cv.notify() def _run(self): while True: with self.cv: if self.stopped: return now = time.monotonic() tasks = [task for task in self.tasks.values() if task not in self.inflight_tasks] for task in tasks: self.inflight_tasks[task] = now stalled = [ (task, now - start) for task, start in self.inflight_tasks.items() if task not in self.stalled_tasks and now - start > task.retry_timeout ] self.stalled_tasks.update(task for task, _elapsed in stalled) generation = self.generation for task, elapsed in stalled: log.error(f"{task} retry dispatch has not returned after {elapsed:.1f} seconds, retries are stalled") next_wait = None futures = {} completed_futures = set() for task in tasks: future = self.retry_task_pool.submit(task.retry_task) if future is None: self._finish_inflight(task) continue futures[future] = task try: for future in as_completed(futures, timeout=STREAM_RETRY_RESULT_TIMEOUT): completed_futures.add(future) task = futures[future] self._finish_inflight(task) wait_time = future.result() if wait_time is not None: next_wait = wait_time if next_wait is None else min(next_wait, wait_time) except TimeoutError: next_wait = ( STREAM_RETRY_RESULT_TIMEOUT if next_wait is None else min(next_wait, STREAM_RETRY_RESULT_TIMEOUT) ) for future, task in futures.items(): if future not in completed_futures: future.add_done_callback(lambda _future, retry_task=task: self._finish_inflight(retry_task)) with self.cv: if self.stopped: return if self.generation == generation: self.cv.wait(timeout=next_wait)
reliable_retry_scheduler = ReliableRetryScheduler() MainProcessMonitor.add_cleanup_cb(reliable_retry_scheduler.shutdown)
[docs] class TxTask(StreamTaskSpec): def __init__( self, cell: CoreCell, chunk_size: int, channel: str, topic: str, target: str, headers: dict, stream: Stream, reliable: Optional[bool], secure: bool, optional: bool, ): self.cell = cell self.chunk_size = chunk_size self.sid = gen_stream_id() self.buffer = wrap_view(bytearray(chunk_size)) # Optimization to send the original buffer without copying self.direct_buf: Optional[bytes] = None self.buffer_size = 0 self.channel = channel self.topic = topic self.target = target self.headers = headers self.stream = stream self.stream_future = None self.task_future = None self.ack_waiter = threading.Event() self.seq = 0 self.seq_ack = -1 self.offset = 0 self.offset_ack = 0 self.secure = secure self.optional = optional self.stopped = False self.stopping = False self.send_lock = threading.RLock() self.stream_future = StreamFuture(self.sid, task_handle=self) self.stream_future.set_size(stream.get_size()) config = CommConfigurator() self.reliable = config.get_streaming_reliable(False) if reliable is None else reliable self.window_size = config.get_streaming_window_size(STREAM_WINDOW_SIZE) self.ack_interval = config.get_streaming_ack_interval(STREAM_ACK_INTERVAL) if self.ack_interval > self.window_size: log.warning( f"{self} streaming_ack_interval {self.ack_interval} exceeds streaming_window_size " f"{self.window_size}; using {self.window_size}" ) self.ack_interval = self.window_size self.ack_wait = config.get_streaming_ack_wait(STREAM_ACK_WAIT) self.ack_progress_timeout = config.get_streaming_ack_progress_timeout(60.0) # Guard against zero/negative config to avoid wait(0) busy-spin loops. self.ack_progress_check_interval = max(0.01, config.get_streaming_ack_progress_check_interval(5.0)) self.last_ack_progress_ts = time.monotonic() self.retry_wait = max(0.01, config.get_streaming_retry_wait(STREAM_RETRY_WAIT)) self.retry_timeout = max(0.01, config.get_streaming_retry_timeout(STREAM_RETRY_TIMEOUT)) retry_max_pending_default = max(STREAM_RETRY_MAX_PENDING_BYTES, 2 * self.window_size) self.retry_max_pending_bytes = config.get_streaming_retry_max_pending_bytes(retry_max_pending_default) if self.reliable: self.pending_messages = {} self.pending_send_errors = {} self.pending_message_bytes = 0 self.retry_lock = threading.RLock() reliable_retry_scheduler.register(self) else: self.pending_messages = None self.pending_send_errors = None self.pending_message_bytes = 0 self.retry_lock = None def __str__(self): return f"Tx[SID:{self.sid} to {self.target} for {self.channel}/{self.topic}]" @staticmethod def _new_send_error(msg: str, send_error=None) -> StreamError: if send_error == ReturnCode.TARGET_UNREACHABLE: return StreamTargetUnreachable(msg) return StreamError(msg) def _new_pending_error(self, msg: str, seq=None) -> StreamError: if not self.reliable: return StreamError(msg) with self.retry_lock: if seq is not None: return self._new_send_error(msg, self.pending_send_errors.get(seq)) if self.pending_send_errors and all( error == ReturnCode.TARGET_UNREACHABLE for error in self.pending_send_errors.values() ): return StreamTargetUnreachable(msg) return StreamError(msg)
[docs] def send_loop(self): """Read/send loop to transmit the whole stream with flow control""" while not self.stopped: if self.buffer_size == self.chunk_size: read_size = self.chunk_size else: read_size = self.chunk_size - self.buffer_size buf = self.stream.read(read_size) if not buf: # End of Stream if not self.send_pending_buffer(final=True): return self.stop() return # Flow control window = self.offset - self.offset_ack # It may take several ACKs to clear up the window. # Keep the historical strict comparison: a zero window provides # stop-and-wait behavior by allowing the first frame to be sent. # RxTask includes the possible boundary frame in its buffer sizing. while window > self.window_size: log.debug(f"{self} window size {window} exceeds limit: {self.window_size}") wait_start = time.monotonic() while window > self.window_size: if self.stopped: return now = time.monotonic() if now - self.last_ack_progress_ts >= self.ack_progress_timeout: self.stop( self._new_pending_error( f"{self} ACK made no progress for {self.ack_progress_timeout} seconds" ) ) return elapsed = now - wait_start if elapsed >= self.ack_wait: self.stop(self._new_pending_error(f"{self} ACK timeouts after {self.ack_wait} seconds")) return self.ack_waiter.clear() wait_timeout = min(self.ack_progress_check_interval, self.ack_wait - elapsed) self.ack_waiter.wait(timeout=wait_timeout) window = self.offset - self.offset_ack size = len(buf) if size > read_size: raise StreamError(f"{self} Stream returns invalid size: {size} (requested {read_size})") # A full pending buffer is sent only after a non-empty lookahead read. # This avoids an empty final frame when the stream size is an exact # multiple of chunk_size while ensuring all non-final frames are full. if self.buffer_size == self.chunk_size: if not self.send_pending_buffer(): return if size == self.chunk_size: self.direct_buf = buf else: self.buffer[self.buffer_size : self.buffer_size + size] = buf self.buffer_size += size
[docs] def send_pending_buffer(self, final=False): if self.buffer_size == 0: payload = bytes(0) elif self.buffer_size == self.chunk_size: if self.direct_buf: payload = self.direct_buf else: payload = self.buffer else: payload = self.buffer[0 : self.buffer_size] if self.reliable: payload = _snapshot_payload(payload) message = Message(None, payload) if self.headers: message.add_headers(self.headers) stream_headers = { StreamHeaderKey.CHANNEL: self.channel, StreamHeaderKey.TOPIC: self.topic, StreamHeaderKey.SIZE: self.stream.get_size(), StreamHeaderKey.STREAM_ID: self.sid, StreamHeaderKey.DATA_TYPE: StreamDataType.FINAL if final else StreamDataType.CHUNK, StreamHeaderKey.SEQUENCE: self.seq, StreamHeaderKey.OFFSET: self.offset, StreamHeaderKey.RELIABLE: self.reliable, StreamHeaderKey.OPTIONAL: self.optional, # Repeat the buffer-sizing parameters because ConnManager may process a # later frame before sequence 0. Older receivers ignore these headers # after the first frame, so this is wire-compatible. StreamHeaderKey.CHUNK_SIZE: self.chunk_size, StreamHeaderKey.WINDOW_SIZE: self.window_size, } if self.seq == 0: stream_headers[StreamHeaderKey.ACK_INTERVAL] = self.ack_interval if self.reliable: stream_headers[StreamHeaderKey.RETRY_WAIT] = self.retry_wait stream_headers[StreamHeaderKey.RETRY_TIMEOUT] = self.retry_timeout message.add_headers(stream_headers) if self.reliable: errors = None over_limit_error = None with self.send_lock: curr_time = time.monotonic() with self.retry_lock: if self.stopped: return False pending_message_size = _payload_size(message.payload) self.pending_messages[self.seq] = None, curr_time, message self.pending_send_errors[self.seq] = None self.pending_message_bytes += pending_message_size if self.retry_max_pending_bytes > 0 and self.pending_message_bytes > self.retry_max_pending_bytes: self.pending_messages.pop(self.seq, None) self.pending_send_errors.pop(self.seq, None) self.pending_message_bytes -= pending_message_size msg = ( f"{self} has too many retry messages " f"({self.pending_message_bytes + pending_message_size} > {self.retry_max_pending_bytes})" ) over_limit_error = StreamError(msg) if not over_limit_error: reliable_retry_scheduler.wakeup() errors = self.cell.fire_and_forget( STREAM_CHANNEL, STREAM_DATA_TOPIC, self.target, message, secure=self.secure, optional=self.optional, ) if over_limit_error: log.error(str(over_limit_error)) self.stop(over_limit_error) return False else: errors = self.cell.fire_and_forget( STREAM_CHANNEL, STREAM_DATA_TOPIC, self.target, message, secure=self.secure, optional=self.optional ) errors = errors or {} error = errors.get(self.target) if self.reliable: with self.retry_lock: if self.seq in self.pending_messages: self.pending_send_errors[self.seq] = error if error: msg = f"{self} Message sending error to target {self.target}: {error}" if self.reliable: log_fn = log.debug if self.optional and error == ReturnCode.TARGET_UNREACHABLE else log.error log_fn(f"{msg}, will retry in {self.retry_wait} seconds") else: self.stop(self._new_send_error(msg, error)) return False # Update state self.seq += 1 self.offset += self.buffer_size self.buffer_size = 0 self.direct_buf = None # Update future self.stream_future.set_progress(self.offset) return True
[docs] def stop(self, error: Optional[StreamError] = None, notify=True): if self.reliable: if error: with self.send_lock: if not self._prepare_reliable_stop(error): return elif not self._prepare_reliable_stop(error): return reliable_retry_scheduler.unregister(self) else: if self.stopped: return self.stopped = True self.remove_task() if not self.ack_waiter.is_set(): self.ack_waiter.set() if self.task_future: self.task_future.cancel() if not error: # Result is the number of bytes streamed if self.stream_future: self.stream_future.set_result(self.offset) return # Error handling log.debug(f"{self} Stream error: {error}") if self.stream_future: self.stream_future.set_exception(error) if notify: message = Message(None, None) if self.headers: message.add_headers(self.headers) message.add_headers( { StreamHeaderKey.STREAM_ID: self.sid, StreamHeaderKey.DATA_TYPE: StreamDataType.ERROR, StreamHeaderKey.OFFSET: self.offset, StreamHeaderKey.ERROR_MSG: str(error), } ) try: self.cell.fire_and_forget( STREAM_CHANNEL, STREAM_DATA_TOPIC, self.target, message, secure=self.secure, optional=True ) except Exception as ex: log.error(f"{self} failed to notify stream error to target {self.target}: {ex}")
def _prepare_reliable_stop(self, error: Optional[StreamError]) -> bool: with self.retry_lock: if self.stopped: return False if not error and self.pending_messages: self.stopping = True reliable_retry_scheduler.wakeup() if not self.ack_waiter.is_set(): self.ack_waiter.set() return False self.stopped = True self.stopping = False if error: self.pending_messages.clear() self.pending_send_errors.clear() self.pending_message_bytes = 0 return True
[docs] def handle_ack(self, message: Message): origin = message.get_header(MessageHeaderKey.ORIGIN) ack_seq = message.get_header(StreamHeaderKey.SEQUENCE, None) offset = message.get_header(StreamHeaderKey.OFFSET, None) error = message.get_header(StreamHeaderKey.ERROR_MSG, None) if error: error_type = message.get_header(StreamHeaderKey.ERROR_TYPE) error_class = BlobSizeError if error_type == BlobSizeError.__name__ else StreamError self.stop(error_class(f"{self} Received error from {origin}: {error}"), notify=False) return if self.reliable and ack_seq is None: self.stop(StreamError(f"{self} receiving end at {origin} doesn't support reliable streaming"), notify=True) return if self.reliable: should_stop = False ack_progressed = False with self.retry_lock: if offset is not None and offset > self.offset_ack: self.offset_ack = offset ack_progressed = True if ack_seq is not None and ack_seq > self.seq_ack: self.seq_ack = ack_seq ack_progressed = True if ack_progressed: self.last_ack_progress_ts = time.monotonic() if self.pending_messages and ack_seq is not None: for seq in list(self.pending_messages): if seq <= ack_seq: _retry_start_time, _last_retry, message = self.pending_messages.pop(seq) self.pending_send_errors.pop(seq, None) self.pending_message_bytes -= _payload_size(message.payload) should_stop = self.stopping and not self.pending_messages if should_stop: self.stop() elif offset is not None and offset > self.offset_ack: self.offset_ack = offset self.last_ack_progress_ts = time.monotonic() if not self.ack_waiter.is_set(): self.ack_waiter.set()
[docs] def start_task_thread(self, task_handler: Callable): self.task_future = stream_thread_pool.submit(task_handler, self)
[docs] def cancel(self): self.stop(error=StreamError("cancelled"))
[docs] def retry_task(self) -> Optional[float]: try: return self._retry_task() except Exception as ex: msg = f"{self} retry thread ended due to error: {ex}" log.error(msg) self.stop(StreamError(msg), notify=True) return None
def _retry_task(self) -> Optional[float]: should_stop = False next_wait = None messages_to_retry = [] retry_next_wait = None retry_error = None with self.retry_lock: if self.stopped: return None if not self.pending_messages: should_stop = self.stopping else: curr_time = time.monotonic() for seq, value in list(self.pending_messages.items()): retry_start_time, last_retry, message = value wait_time = self.retry_wait - (curr_time - last_retry) remaining_retry_timeout = self.retry_timeout if retry_start_time is not None: retry_time = curr_time - retry_start_time if retry_time > self.retry_timeout: msg = f"{self} seq {seq} retry failed after {retry_time:.2f} seconds from first retry" retry_error = self._new_pending_error(msg, seq) log_fn = ( log.debug if self.optional and isinstance(retry_error, StreamTargetUnreachable) else log.error ) log_fn(msg) break remaining_retry_timeout = self.retry_timeout - retry_time wait_time = min(wait_time, remaining_retry_timeout) if wait_time <= 0: retry_start_time = curr_time if retry_start_time is None else retry_start_time messages_to_retry.append((seq, message)) self.pending_messages[seq] = retry_start_time, curr_time, message after_retry_wait = min(self.retry_wait, remaining_retry_timeout) retry_next_wait = ( after_retry_wait if retry_next_wait is None else min(retry_next_wait, after_retry_wait) ) else: next_wait = wait_time if next_wait is None else min(next_wait, wait_time) if retry_error: self.stop(error=retry_error) return None if should_stop: self.stop() return None if messages_to_retry: # Hold send_lock so stop(error) cannot clear pending state and notify the receiver # while a retry send is still in flight, which would deliver a ghost chunk. with self.send_lock: with self.retry_lock: if self.stopped: return None for seq, message in messages_to_retry: errors = self.cell.fire_and_forget( STREAM_CHANNEL, STREAM_DATA_TOPIC, self.target, message, secure=self.secure, optional=self.optional, ) errors = errors or {} error = errors.get(self.target) with self.retry_lock: if seq in self.pending_messages: self.pending_send_errors[seq] = error if error: log_fn = log.debug if self.optional and error == ReturnCode.TARGET_UNREACHABLE else log.error log_fn( f"{self} message retry error for target {self.target} seq {seq}: " f"{error}, will retry again in {self.retry_wait} seconds" ) next_wait = retry_next_wait if next_wait is None else min(next_wait, retry_next_wait) return next_wait
[docs] def remove_task(self): with ByteStreamer.map_lock: ByteStreamer.tx_task_map.pop(self.sid, None) ByteStreamer._retain_error_context(self) log.debug(f"{self} is removed")
[docs] class ByteStreamer: tx_task_map = {} # Contexts all have the same TTL, so insertion order is expiry order. error_context_map = OrderedDict() map_lock = threading.Lock() sent_stream_counter_pool = StatsPoolManager.add_counter_pool( name="Sent_Stream_Counters", description="Counters of sent streams", counter_names=[COUNTER_NAME_SENT], ) sent_stream_size_pool = StatsPoolManager.add_msg_size_pool("Sent_Stream_Sizes", "Sizes of streams sent (MBs)") def __init__(self, cell: CoreCell): self.cell = cell self.error_callbacks = [] self.cell.add_error_handler(STREAM_CHANNEL, STREAM_DATA_TOPIC, self._forward_error_handler) self.cell.register_request_cb(channel=STREAM_CHANNEL, topic=STREAM_ACK_TOPIC, cb=self._ack_handler) self.cell.register_request_cb(channel=STREAM_CHANNEL, topic=STREAM_ERROR_TOPIC, cb=self._error_handler) self.chunk_size = CommConfigurator().get_streaming_chunk_size(STREAM_CHUNK_SIZE) def _forward_error_handler(self, message: Message, error: str): """Report a downstream routing failure to the original stream sender.""" if not message.get_header(MessageHeaderKey.OPTIONAL, False): # Required reliable streams own their retry policy. A transient downstream # routing failure must not bypass retry_timeout by settling the sender early. return sender = message.get_header(MessageHeaderKey.ORIGIN) failed_destination = message.get_header(MessageHeaderKey.DESTINATION) if not sender or not failed_destination: return error_class = StreamTargetUnreachable if error == ReturnCode.TARGET_UNREACHABLE else StreamError headers = { StreamHeaderKey.STREAM_ID: message.get_header(StreamHeaderKey.STREAM_ID), StreamHeaderKey.DATA_TYPE: StreamDataType.ERROR, StreamHeaderKey.ERROR_MSG: f"stream forwarding to {failed_destination} failed: {error}", StreamHeaderKey.ERROR_TYPE: error_class.__name__, StreamHeaderKey.FAILED_DESTINATION: failed_destination, StreamHeaderKey.CHANNEL: message.get_header(StreamHeaderKey.CHANNEL), StreamHeaderKey.TOPIC: message.get_header(StreamHeaderKey.TOPIC), } req_id = message.get_header(StreamHeaderKey.STREAM_REQ_ID) if req_id: headers[StreamHeaderKey.STREAM_REQ_ID] = req_id errors = self.cell.fire_and_forget(STREAM_CHANNEL, STREAM_ERROR_TOPIC, sender, Message(headers), optional=True) send_error = (errors or {}).get(sender) if send_error: log.debug( f"failed to report stream routing error: stream_id={headers[StreamHeaderKey.STREAM_ID]} " f"sender={sender} failed_destination={failed_destination}: {send_error}" )
[docs] def register_error_callback(self, callback: Callable): if not callable(callback): raise StreamError(f"specified stream error callback {type(callback)} is not callable") self.error_callbacks.append(callback)
def _notify_error_callbacks(self, message: Message): for callback in self.error_callbacks: try: callback(message) except Exception as ex: log.error(f"stream error callback {callback} failed: {ex}") @classmethod def _retain_error_context(cls, task: TxTask): now = time.monotonic() cls._purge_error_contexts(now) context = _TxTaskContext( cell=task.cell, target=task.target, channel=task.channel, topic=task.topic, req_id=(task.headers or {}).get(StreamHeaderKey.STREAM_REQ_ID), expires_at=now + STREAM_ERROR_CONTEXT_TTL, ) cls.error_context_map.pop(task.sid, None) cls.error_context_map[task.sid] = context while len(cls.error_context_map) > MAX_STREAM_ERROR_CONTEXTS: cls.error_context_map.popitem(last=False) @classmethod def _purge_error_contexts(cls, now: float): while cls.error_context_map: sid = next(iter(cls.error_context_map)) context = cls.error_context_map[sid] if context.expires_at > now: break cls.error_context_map.pop(sid, None) @staticmethod def _matches_error_context(message: Message, context: _TxTaskContext) -> bool: origin = message.get_header(MessageHeaderKey.ORIGIN) failed_destination = message.get_header(StreamHeaderKey.FAILED_DESTINATION, origin) return ( failed_destination == context.target and (origin == context.target or message.get_header(StreamHeaderKey.FAILED_DESTINATION) == context.target) and message.get_header(StreamHeaderKey.CHANNEL) == context.channel and message.get_header(StreamHeaderKey.TOPIC) == context.topic and message.get_header(StreamHeaderKey.STREAM_REQ_ID) == context.req_id )
[docs] def get_chunk_size(self): return self.chunk_size
[docs] @classmethod def shutdown(cls): """Cancel every process-owned outgoing stream before F3 executors stop.""" with cls.map_lock: tasks = tuple(cls.tx_task_map.values()) for task in tasks: task.stop(StreamError("streaming shutdown"), notify=False)
[docs] def send( self, channel: str, topic: str, target: str, headers: dict, stream: Stream, stream_type=STREAM_TYPE_BYTE, secure=False, optional=False, reliable: Optional[bool] = None, ) -> StreamFuture: tx_task = TxTask( self.cell, self.chunk_size, channel, topic, target, headers, stream, reliable, secure, optional ) with ByteStreamer.map_lock: ByteStreamer.error_context_map.pop(tx_task.sid, None) ByteStreamer.tx_task_map[tx_task.sid] = tx_task tx_task.start_task_thread(self._transmit_task) fqcn = self.cell.my_info.fqcn ByteStreamer.sent_stream_counter_pool.increment( category=stream_stats_category(fqcn, channel, topic, stream_type), counter_name=COUNTER_NAME_SENT ) ByteStreamer.sent_stream_size_pool.record_value( category=stream_stats_category(fqcn, channel, topic, stream_type), value=stream.get_size() / ONE_MB ) return tx_task.stream_future
@staticmethod def _transmit_task(task: TxTask): try: task.send_loop() except Exception as ex: msg = f"{task} Error while sending: {ex}" if task.optional: log.debug(msg) else: log.error(msg) task.stop(StreamError(msg), True) @staticmethod def _ack_handler(message: Message): sid = message.get_header(StreamHeaderKey.STREAM_ID) with ByteStreamer.map_lock: tx_task = ByteStreamer.tx_task_map.get(sid, None) if not tx_task: origin = message.get_header(MessageHeaderKey.ORIGIN) offset = message.get_header(StreamHeaderKey.OFFSET, None) seq = message.get_header(StreamHeaderKey.SEQUENCE, None) # Last few ACKs always arrive late so this is normal log.debug(f"ACK for stream {sid} received late from {origin} with offset {offset} seq {seq}") return tx_task.handle_ack(message) def _error_handler(self, message: Message): sid = message.get_header(StreamHeaderKey.STREAM_ID) origin = message.get_header(MessageHeaderKey.ORIGIN) channel = message.get_header(StreamHeaderKey.CHANNEL) topic = message.get_header(StreamHeaderKey.TOPIC) error = message.get_header(StreamHeaderKey.ERROR_MSG, "stream rejected by receiver") error_type = message.get_header(StreamHeaderKey.ERROR_TYPE) error_classes = { BlobSizeError.__name__: BlobSizeError, StreamTargetUnreachable.__name__: StreamTargetUnreachable, } error_class = error_classes.get(error_type, StreamError) sender = self.cell.my_info.fqcn failed_destination = message.get_header(StreamHeaderKey.FAILED_DESTINATION, origin) with ByteStreamer.map_lock: tx_task = ByteStreamer.tx_task_map.get(sid) if tx_task and tx_task.cell is not self.cell: tx_task = None ByteStreamer._purge_error_contexts(time.monotonic()) context = ByteStreamer.error_context_map.get(sid) if context and context.cell is not self.cell: context = None if not tx_task: if not context or not self._matches_error_context(message, context): log.warning( f"Ignored uncorrelated stream error: stream_id={sid} channel={channel} topic={topic} " f"sender={sender} failed_destination={failed_destination}: {error}" ) return log.warning( f"Late stream error: stream_id={sid} channel={channel} topic={topic} " f"sender={sender} failed_destination={failed_destination}: {error}" ) self._notify_error_callbacks(message) return active_context = _TxTaskContext( cell=tx_task.cell, target=tx_task.target, channel=tx_task.channel, topic=tx_task.topic, req_id=(tx_task.headers or {}).get(StreamHeaderKey.STREAM_REQ_ID), expires_at=0, ) if not self._matches_error_context(message, active_context): log.warning( f"Ignored stream error with unexpected context: stream_id={sid} channel={channel} topic={topic} " f"sender={sender} expected_destination={tx_task.target} failed_destination={failed_destination}" ) return self._notify_error_callbacks(message) tx_task.stop( error_class( f"Stream rejected: stream_id={sid} channel={tx_task.channel} topic={tx_task.topic} " f"sender={sender} failed_destination={failed_destination}: {error}" ), notify=False, )