# 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,
)