use std::{
collections::HashMap,
error::Error,
fmt,
future::Future,
io,
path::{Path, PathBuf},
pin::Pin,
sync::{
atomic::{AtomicBool, AtomicU64, Ordering},
Arc, Mutex, MutexGuard,
},
task::{Context, Poll},
time::Duration,
};
use subc_control::{
CatalogEntry, ClientControlRequest, ClientControlResponse, ConsumerIdentity, PollKind,
};
use subc_protocol::{
AdmissionClass, BindIdentity, ErrorBody, Flags, Frame, FrameBuildError, FrameType, Priority,
RouteTarget, SUBC_LAUNCH_NONCE_ENV, SUBC_MODULE_ID_ENV,
};
use crate::RouteHandle;
use subc_transport::{
authenticate_client, connection_file, read_frame, write_frame, AuthError, ConnectionFileError,
FrameIoError,
};
use tokio::{
io::{AsyncWrite, AsyncWriteExt, BufWriter},
net::{tcp::OwnedReadHalf, TcpStream},
sync::{mpsc, oneshot, Notify, OwnedSemaphorePermit, Semaphore},
task::JoinHandle,
time::{sleep, timeout_at, Instant},
};
use tokio_util::sync::CancellationToken;
const DEFAULT_HANDSHAKE_TIMEOUT: Duration = Duration::from_secs(2);
const DEFAULT_CALL_TIMEOUT: Duration = Duration::from_secs(30);
const DEFAULT_ROUTE_RETRY_DEADLINE: Duration = Duration::from_secs(30);
const DEFAULT_RESTORED_DEBOUNCE: Duration = Duration::from_millis(250);
const EGRESS_BUFFER: usize = 128;
const DEFAULT_ROUTE_WINDOW: usize = 1024;
const DEFAULT_SUBSCRIPTION_EVENT_BUFFER: usize = 128;
const DEFAULT_PUSH_EVENT_BUFFER: usize = 128;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct RetryBackoff {
pub base: Duration,
pub cap: Duration,
pub max_attempts: usize,
}
impl Default for RetryBackoff {
fn default() -> Self {
Self {
base: Duration::from_millis(100),
cap: Duration::from_secs(2),
max_attempts: 6,
}
}
}
impl RetryBackoff {
fn delay_after_attempt(self, attempt: usize) -> Duration {
let mut delay = self.base;
for _ in 1..attempt {
delay = (delay * 2).min(self.cap);
}
delay
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct ConsumerOptions {
pub handshake_timeout: Duration,
pub call_timeout: Duration,
pub reconnect_backoff: RetryBackoff,
pub restored_debounce: Duration,
}
impl Default for ConsumerOptions {
fn default() -> Self {
Self {
handshake_timeout: DEFAULT_HANDSHAKE_TIMEOUT,
call_timeout: DEFAULT_CALL_TIMEOUT,
reconnect_backoff: RetryBackoff::default(),
restored_debounce: DEFAULT_RESTORED_DEBOUNCE,
}
}
}
#[derive(Debug, Clone)]
pub struct CloseRouteOptions {
pub drain: bool,
pub drain_timeout: Duration,
pub consumer_identity: Option<ConsumerIdentity>,
pub consumer_capabilities: Option<Vec<String>>,
}
impl Default for CloseRouteOptions {
fn default() -> Self {
Self {
drain: false,
drain_timeout: DEFAULT_CALL_TIMEOUT,
consumer_identity: None,
consumer_capabilities: None,
}
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct CallOptions {
pub timeout: Duration,
pub priority: Priority,
pub admission_class: AdmissionClass,
pub route_retry: RetryBackoff,
pub route_retry_deadline: Duration,
pub consumer_identity: Option<ConsumerIdentity>,
pub consumer_capabilities: Option<Vec<String>>,
}
impl Default for CallOptions {
fn default() -> Self {
Self {
timeout: DEFAULT_CALL_TIMEOUT,
priority: Priority::Interactive,
admission_class: AdmissionClass::Normal,
route_retry: RetryBackoff::default(),
route_retry_deadline: DEFAULT_ROUTE_RETRY_DEADLINE,
consumer_identity: None,
consumer_capabilities: None,
}
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct SubscribeOptions {
pub priority: Priority,
pub admission_class: AdmissionClass,
pub event_buffer: usize,
pub route_retry: RetryBackoff,
pub route_retry_deadline: Duration,
pub route_open_timeout: Duration,
pub consumer_identity: Option<ConsumerIdentity>,
pub consumer_capabilities: Option<Vec<String>>,
}
impl Default for SubscribeOptions {
fn default() -> Self {
Self {
priority: Priority::Interactive,
admission_class: AdmissionClass::Normal,
event_buffer: DEFAULT_SUBSCRIPTION_EVENT_BUFFER,
route_retry: RetryBackoff::default(),
route_retry_deadline: DEFAULT_ROUTE_RETRY_DEADLINE,
route_open_timeout: DEFAULT_CALL_TIMEOUT,
consumer_identity: None,
consumer_capabilities: None,
}
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum ConnectionState {
Dropped,
Restored { epoch: u64 },
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct RoutePollResult {
pub handle: RouteHandle,
pub status: Option<String>,
pub live: Option<bool>,
}
#[derive(Debug, Clone, serde::Deserialize, PartialEq)]
pub struct CatalogList {
pub generation: u64,
#[serde(default)]
pub modules: Vec<CatalogEntry>,
#[serde(default)]
pub subc_ops: Vec<String>,
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct PushEvent {
pub handle: RouteHandle,
pub body: Vec<u8>,
}
pub struct SubcConsumer {
shared: Arc<Shared>,
}
pub struct Subscription {
events: mpsc::Receiver<Vec<u8>>,
closed: SubscriptionClosed,
cancel: SubscriptionCancel,
}
impl Subscription {
pub fn events(&mut self) -> &mut mpsc::Receiver<Vec<u8>> {
&mut self.events
}
pub fn closed(&mut self) -> &mut SubscriptionClosed {
&mut self.closed
}
pub fn unsubscribe(&self) -> Result<(), CallError> {
self.cancel.unsubscribe()
}
}
impl Drop for Subscription {
fn drop(&mut self) {
let _ = self.cancel.unsubscribe();
}
}
pub struct SubscriptionClosed {
rx: oneshot::Receiver<Result<(), CallError>>,
}
impl Unpin for SubscriptionClosed {}
impl Future for SubscriptionClosed {
type Output = Result<(), CallError>;
fn poll(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Self::Output> {
let this = self.get_mut();
match Pin::new(&mut this.rx).poll(cx) {
Poll::Ready(Ok(result)) => Poll::Ready(result),
Poll::Ready(Err(_)) => Poll::Ready(Err(CallError::outcome_unknown(
"subscription closed result channel dropped",
))),
Poll::Pending => Poll::Pending,
}
}
}
struct SubscriptionCancel {
shared: Arc<Shared>,
key: PendingKey,
priority: Priority,
cancelled: AtomicBool,
}
impl SubscriptionCancel {
fn new(shared: Arc<Shared>, key: PendingKey, priority: Priority) -> Self {
Self {
shared,
key,
priority,
cancelled: AtomicBool::new(false),
}
}
fn unsubscribe(&self) -> Result<(), CallError> {
let handle = RouteHandle::new(self.key.channel, self.key.epoch, self.key.generation);
self.shared.validate_current_handle(handle)?;
if self.cancelled.swap(true, Ordering::AcqRel) {
return Ok(());
}
self.shared
.unsubscribe_subscription(self.key, self.priority)
}
}
impl SubcConsumer {
pub async fn connect(
connection_file: &Path,
opts: ConsumerOptions,
) -> Result<Self, ConsumerError> {
let opened = open_connection(connection_file, opts.handshake_timeout).await?;
let shared = Arc::new(Shared::new(connection_file.to_path_buf(), opts));
shared.install_initial(opened)?;
Ok(Self { shared })
}
pub async fn open_route(
&self,
target: RouteTarget,
identity: BindIdentity,
opts: CallOptions,
) -> Result<RouteHandle, CallError> {
let deadline = Instant::now() + opts.timeout;
let consumer_identity = route_open_consumer_identity(&opts);
let consumer_capabilities = route_open_consumer_capabilities(&opts);
let key = RouteKey::new(
&target,
&identity,
consumer_identity.as_ref(),
consumer_capabilities.as_deref(),
);
let params = RouteOpenParams {
target: &target,
identity: &identity,
consumer_identity: &consumer_identity,
consumer_capabilities: &consumer_capabilities,
};
self.shared
.ensure_route(&key, ¶ms, &opts, deadline)
.await
.map(|route| route.handle)
}
pub async fn open_route_with_admission_facts(
&self,
target: RouteTarget,
identity: BindIdentity,
facts: serde_json::Value,
) -> Result<RouteHandle, CallError> {
let deadline = Instant::now() + self.shared.opts.call_timeout;
let opts = CallOptions::default();
let body = serde_json::to_vec(&ClientControlRequest::RouteOpen {
target,
identity,
consumer_identity: route_open_consumer_identity(&opts),
consumer_capabilities: None,
admission_facts: Some(facts),
})
.map_err(|err| CallError::not_sent(format!("failed to encode route.open: {err}")))?;
let terminal = self.shared.control_call(body, deadline, true).await?;
let TerminalFrame::Response {
generation, body, ..
} = terminal
else {
return Err(CallError::not_sent(
"route.open returned a non-response frame",
));
};
let ClientControlResponse::RouteOpen {
route_channel,
route_epoch,
} = serde_json::from_slice(&body).map_err(|err| {
CallError::not_sent(format!("failed to decode route.open response: {err}"))
})?
else {
return Err(CallError::not_sent(
"route.open returned an unexpected control response",
));
};
let route = RouteState {
handle: RouteHandle::new(route_channel, route_epoch, generation),
sem: Arc::new(Semaphore::new(DEFAULT_ROUTE_WINDOW)),
};
self.shared.install_one_shot_route(route.clone())?;
Ok(route.handle)
}
pub async fn catalog_list(&self) -> Result<CatalogList, CallError> {
let deadline = Instant::now() + self.shared.opts.call_timeout;
let body = serde_json::to_vec(&serde_json::json!({
"op": subc_control::ops::CATALOG_LIST,
}))
.map_err(|err| CallError::not_sent(format!("failed to encode catalog.list: {err}")))?;
loop {
match self
.shared
.control_call(body.clone(), deadline, false)
.await
{
Ok(TerminalFrame::Response { body, .. }) => {
let response =
serde_json::from_slice::<ClientControlResponse>(&body).map_err(|err| {
CallError::not_sent(format!(
"failed to decode catalog.list response: {err}"
))
})?;
let ClientControlResponse::CatalogList {
generation,
modules,
subc_ops,
} = response
else {
return Err(CallError::not_sent(
"catalog.list returned an unexpected control response",
));
};
return Ok(CatalogList {
generation,
modules,
subc_ops,
});
}
Ok(TerminalFrame::Error { body }) => return Err(CallError::Module(body)),
Ok(TerminalFrame::StreamEnd) => {
return Err(CallError::not_sent("catalog.list returned StreamEnd"));
}
Err(err)
if is_retryable_catalog_transport_error(&err) && Instant::now() < deadline =>
{
continue;
}
Err(err) => return Err(err),
}
}
}
pub async fn request(
&self,
handle: &RouteHandle,
body: Vec<u8>,
opts: CallOptions,
) -> Result<Vec<u8>, CallError> {
let deadline = Instant::now() + opts.timeout;
let route = self.shared.route_state(*handle)?;
let permit = match timeout_at(deadline, Arc::clone(&route.sem).acquire_owned()).await {
Ok(Ok(permit)) => permit,
Ok(Err(_)) => return Err(CallError::StaleRouteHandle(*handle)),
Err(_) => {
return Err(CallError::not_sent(
"call deadline elapsed waiting for route flow-control",
))
}
};
let result = self
.shared
.send_request(RequestSend {
expected_handle: Some(*handle),
channel: handle.channel,
epoch: handle.epoch,
body,
priority: opts.priority,
admission_class: opts.admission_class,
deadline,
retain_late_route_open: false,
})
.await;
drop(permit);
match result? {
TerminalFrame::Response { body, .. } => Ok(body),
TerminalFrame::StreamEnd => Ok(Vec::new()),
TerminalFrame::Error { body } => Err(CallError::Module(body)),
}
}
pub async fn subscribe_route(
&self,
handle: &RouteHandle,
body: Vec<u8>,
opts: SubscribeOptions,
) -> Result<Subscription, CallError> {
let deadline = Instant::now() + opts.route_open_timeout;
let route = self.shared.route_state(*handle)?;
let permit = match timeout_at(deadline, Arc::clone(&route.sem).acquire_owned()).await {
Ok(Ok(permit)) => permit,
Ok(Err(_)) => return Err(CallError::StaleRouteHandle(*handle)),
Err(_) => {
return Err(CallError::not_sent(
"subscription deadline elapsed waiting for route flow-control",
))
}
};
self.shared
.send_subscription(SubscriptionSend {
expected_handle: Some(*handle),
channel: handle.channel,
epoch: handle.epoch,
body,
priority: opts.priority,
admission_class: opts.admission_class,
event_buffer: opts.event_buffer,
deadline,
permit,
})
.await
}
pub async fn poll_route(
&self,
handle: &RouteHandle,
kind: PollKind,
timeout: Duration,
) -> Result<RoutePollResult, CallError> {
let deadline = Instant::now() + timeout;
let body = serde_json::to_vec(&ClientControlRequest::RoutePoll {
route_channel: handle.channel,
route_epoch: handle.epoch,
kind,
})
.map_err(|err| CallError::not_sent(format!("failed to encode route.poll: {err}")))?;
let terminal = self
.shared
.send_request(RequestSend {
expected_handle: Some(*handle),
channel: 0,
epoch: 0,
body,
priority: Priority::Interactive,
admission_class: AdmissionClass::Normal,
deadline,
retain_late_route_open: false,
})
.await?;
let TerminalFrame::Response { body, .. } = terminal else {
return Err(CallError::not_sent(
"route.poll returned a non-response frame",
));
};
let ClientControlResponse::RoutePoll {
route_channel,
route_epoch,
status,
live,
} = serde_json::from_slice(&body)
.map_err(|err| CallError::not_sent(format!("failed to decode route.poll: {err}")))?
else {
return Err(CallError::not_sent(
"route.poll returned an unexpected control response",
));
};
if route_channel != handle.channel || route_epoch != handle.epoch {
return Err(CallError::not_sent(
"route.poll response echoed a different route handle",
));
}
self.shared.validate_current_handle(*handle)?;
Ok(RoutePollResult {
handle: *handle,
status,
live,
})
}
pub fn dropped_route_frames(&self) -> u64 {
self.shared.lock_inner().dropped_route_frames
}
pub fn push_events(
&self,
handle: &RouteHandle,
) -> Result<mpsc::Receiver<PushEvent>, CallError> {
self.shared.register_push_events(*handle)
}
pub fn pushes_dropped_no_receiver(&self) -> u64 {
self.shared
.pushes_dropped_no_receiver
.load(Ordering::Relaxed)
}
pub async fn call(
&self,
target: RouteTarget,
identity: BindIdentity,
body: Vec<u8>,
opts: CallOptions,
) -> Result<Vec<u8>, CallError> {
let call_deadline = Instant::now() + opts.timeout;
let mut retried_unknown_channel = false;
let consumer_identity = route_open_consumer_identity(&opts);
let consumer_capabilities = route_open_consumer_capabilities(&opts);
let route_key = RouteKey::new(
&target,
&identity,
consumer_identity.as_ref(),
consumer_capabilities.as_deref(),
);
let route_open = RouteOpenParams {
target: &target,
identity: &identity,
consumer_identity: &consumer_identity,
consumer_capabilities: &consumer_capabilities,
};
loop {
let route = self
.shared
.ensure_route(&route_key, &route_open, &opts, call_deadline)
.await?;
let permit =
match timeout_at(call_deadline, Arc::clone(&route.sem).acquire_owned()).await {
Ok(Ok(permit)) => permit,
Ok(Err(_)) => {
return Err(CallError::not_sent("route flow-control semaphore closed"));
}
Err(_) => {
return Err(CallError::not_sent(
"call deadline elapsed waiting for route flow-control",
));
}
};
if !self.shared.route_is_current(&route_key, &route) {
drop(permit);
self.shared
.sleep_until_retry(call_deadline, opts.route_retry.base)
.await?;
continue;
}
let response = self
.shared
.send_request(RequestSend {
expected_handle: Some(route.handle),
channel: route.handle.channel,
epoch: route.handle.epoch,
body: body.clone(),
priority: opts.priority,
admission_class: opts.admission_class,
deadline: call_deadline,
retain_late_route_open: false,
})
.await;
drop(permit);
match response {
Ok(TerminalFrame::Response { body, .. }) => return Ok(body),
Ok(TerminalFrame::StreamEnd) => return Ok(Vec::new()),
Ok(TerminalFrame::Error { body, .. })
if body.code == "unknown_channel"
&& !retried_unknown_channel
&& Instant::now() < call_deadline =>
{
retried_unknown_channel = true;
self.shared.invalidate_route(&route_key, Some(route.handle));
continue;
}
Ok(TerminalFrame::Error { body, .. }) => return Err(CallError::Module(body)),
Err(err) if err.is_not_sent() && Instant::now() < call_deadline => {
self.shared.invalidate_route(&route_key, Some(route.handle));
self.shared.ensure_connected_for_call(call_deadline).await?;
continue;
}
Err(err) => return Err(err),
}
}
}
pub async fn subscribe(
&self,
target: RouteTarget,
identity: BindIdentity,
body: Vec<u8>,
opts: SubscribeOptions,
) -> Result<Subscription, CallError> {
let open_deadline = Instant::now() + opts.route_open_timeout;
let route_opts = CallOptions {
timeout: opts.route_open_timeout,
priority: opts.priority,
admission_class: opts.admission_class,
route_retry: opts.route_retry,
route_retry_deadline: opts.route_retry_deadline,
consumer_identity: opts.consumer_identity.clone(),
consumer_capabilities: opts.consumer_capabilities.clone(),
};
let consumer_identity = route_open_consumer_identity(&route_opts);
let consumer_capabilities = route_open_consumer_capabilities(&route_opts);
let route_key = RouteKey::new(
&target,
&identity,
consumer_identity.as_ref(),
consumer_capabilities.as_deref(),
);
let route_open = RouteOpenParams {
target: &target,
identity: &identity,
consumer_identity: &consumer_identity,
consumer_capabilities: &consumer_capabilities,
};
loop {
let route = self
.shared
.ensure_route(&route_key, &route_open, &route_opts, open_deadline)
.await?;
let permit =
match timeout_at(open_deadline, Arc::clone(&route.sem).acquire_owned()).await {
Ok(Ok(permit)) => permit,
Ok(Err(_)) => {
return Err(CallError::not_sent("route flow-control semaphore closed"));
}
Err(_) => {
return Err(CallError::not_sent(
"subscription open deadline elapsed waiting for route flow-control",
));
}
};
if !self.shared.route_is_current(&route_key, &route) {
drop(permit);
self.shared
.sleep_until_retry(open_deadline, opts.route_retry.base)
.await?;
continue;
}
match self
.shared
.send_subscription(SubscriptionSend {
expected_handle: Some(route.handle),
channel: route.handle.channel,
epoch: route.handle.epoch,
body: body.clone(),
priority: opts.priority,
admission_class: opts.admission_class,
event_buffer: opts.event_buffer,
deadline: open_deadline,
permit,
})
.await
{
Ok(subscription) => return Ok(subscription),
Err(err) if err.is_not_sent() && Instant::now() < open_deadline => {
self.shared.invalidate_route(&route_key, Some(route.handle));
self.shared.ensure_connected_for_call(open_deadline).await?;
continue;
}
Err(err) => return Err(err),
}
}
}
pub async fn close_route(
&self,
target: RouteTarget,
identity: BindIdentity,
opts: CloseRouteOptions,
) {
let consumer_identity = close_route_consumer_identity(&opts);
let consumer_capabilities = close_route_consumer_capabilities(&opts);
let key = RouteKey::new(
&target,
&identity,
consumer_identity.as_ref(),
consumer_capabilities.as_deref(),
);
self.shared.close_route(&key, &opts).await;
}
pub async fn close_handle(
&self,
handle: &RouteHandle,
opts: CloseRouteOptions,
) -> Result<(), CallError> {
self.shared.close_handle(*handle, &opts).await
}
pub fn current_epoch(&self) -> u64 {
self.shared.lock_inner().epoch
}
pub fn on_connection_state(&self, cb: impl Fn(ConnectionState) + Send + 'static) {
self.shared
.lock_inner()
.callbacks
.push(Arc::new(Mutex::new(Box::new(cb))));
}
pub async fn close(&self) {
self.shared.close_sync("consumer closed");
tokio::task::yield_now().await;
}
}
impl Drop for SubcConsumer {
fn drop(&mut self) {
self.shared.close_sync("consumer dropped");
}
}
#[derive(Debug)]
pub enum ConsumerError {
ConnectionFile {
path: PathBuf,
source: ConnectionFileError,
},
NoEndpoint {
path: PathBuf,
},
Connect {
path: PathBuf,
endpoint: String,
source: io::Error,
},
Auth {
path: PathBuf,
endpoint: String,
source: AuthError,
},
Closed,
}
impl fmt::Display for ConsumerError {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
Self::ConnectionFile { path, source } => write!(
f,
"failed to read subc connection file '{}': {source}",
path.display()
),
Self::NoEndpoint { path } => {
write!(
f,
"subc connection file '{}' has no endpoints",
path.display()
)
}
Self::Connect {
path,
endpoint,
source,
} => write!(
f,
"failed to connect to subc endpoint {endpoint} from '{}': {source}",
path.display()
),
Self::Auth {
path,
endpoint,
source,
} => write!(
f,
"failed to authenticate to subc endpoint {endpoint} from '{}': {source}",
path.display()
),
Self::Closed => write!(f, "consumer closed"),
}
}
}
impl Error for ConsumerError {
fn source(&self) -> Option<&(dyn Error + 'static)> {
match self {
Self::ConnectionFile { source, .. } => Some(source),
Self::Connect { source, .. } => Some(source),
Self::Auth { source, .. } => Some(source),
Self::NoEndpoint { .. } | Self::Closed => None,
}
}
}
#[derive(Debug)]
pub enum CallError {
NotSent(Box<dyn Error + Send + Sync>),
OutcomeUnknown(Box<dyn Error + Send + Sync>),
Module(ErrorBody),
SubscriptionBackpressure(Box<dyn Error + Send + Sync>),
StaleRouteHandle(RouteHandle),
}
impl CallError {
fn not_sent(reason: impl Into<String>) -> Self {
Self::NotSent(Box::new(SimpleError(reason.into())))
}
fn outcome_unknown(reason: impl Into<String>) -> Self {
Self::OutcomeUnknown(Box::new(SimpleError(reason.into())))
}
fn is_not_sent(&self) -> bool {
matches!(self, Self::NotSent(_))
}
fn subscription_backpressure(reason: impl Into<String>) -> Self {
Self::SubscriptionBackpressure(Box::new(SimpleError(reason.into())))
}
}
impl fmt::Display for CallError {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
Self::NotSent(err) => write!(f, "request not sent: {err}"),
Self::OutcomeUnknown(err) => write!(f, "request outcome unknown: {err}"),
Self::Module(body) => write!(f, "module error {}: {}", body.code, body.message),
Self::SubscriptionBackpressure(err) => {
write!(f, "subscription event channel backpressure: {err}")
}
Self::StaleRouteHandle(handle) => write!(f, "stale route handle: {handle:?}"),
}
}
}
impl Error for CallError {
fn source(&self) -> Option<&(dyn Error + 'static)> {
match self {
Self::NotSent(err)
| Self::OutcomeUnknown(err)
| Self::SubscriptionBackpressure(err) => Some(err.as_ref()),
Self::Module(_) | Self::StaleRouteHandle(_) => None,
}
}
}
#[derive(Debug)]
struct SimpleError(String);
impl fmt::Display for SimpleError {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.write_str(&self.0)
}
}
impl Error for SimpleError {}
type Callback = Arc<Mutex<Box<dyn Fn(ConnectionState) + Send + 'static>>>;
type OpeningWaiter = oneshot::Sender<Result<RouteState, SharedCallFailure>>;
struct Opening {
waiters: Vec<OpeningWaiter>,
closed: bool,
}
struct Shared {
connection_file: PathBuf,
opts: ConsumerOptions,
inner: Mutex<Inner>,
notify: Notify,
close_token: CancellationToken,
pushes_dropped_no_receiver: AtomicU64,
}
struct Inner {
generation: u64,
epoch: u64,
next_corr: Option<u64>,
writer: Option<mpsc::Sender<WriteCommand>>,
pending: HashMap<PendingKey, PendingEntry>,
routes: HashMap<RouteKey, RouteState>,
route_by_channel: HashMap<u16, RouteKey>,
one_shot_routes: HashMap<u16, RouteState>,
route_epochs: HashMap<u16, RouteHandle>,
push_event_receivers: HashMap<RouteHandle, mpsc::Sender<PushEvent>>,
dropped_route_frames: u64,
openings: HashMap<RouteKey, Opening>,
callbacks: Vec<Callback>,
closed: bool,
reconnect: ReconnectState,
restored_token: u64,
reader_task: Option<JoinHandle<()>>,
writer_task: Option<JoinHandle<()>>,
}
impl Inner {
fn cache_route(&mut self, key: RouteKey, route: RouteState) -> RouteState {
let cached = self.routes.entry(key.clone()).or_insert(route).clone();
let previous = self
.route_by_channel
.insert(cached.handle.channel, key.clone());
debug_assert!(previous.as_ref().is_none_or(|previous| previous == &key));
self.route_epochs
.insert(cached.handle.channel, cached.handle);
cached
}
fn remove_route(&mut self, key: &RouteKey) -> Option<RouteState> {
let route = self.routes.remove(key)?;
let indexed = self.route_by_channel.remove(&route.handle.channel);
debug_assert_eq!(indexed.as_ref(), Some(key));
Some(route)
}
fn remove_route_by_handle(&mut self, handle: RouteHandle) -> Option<RouteState> {
if let Some(key) = self.route_by_channel.get(&handle.channel).cloned() {
let matches = self
.routes
.get(&key)
.is_some_and(|route| route.handle == handle);
debug_assert!(matches);
if matches {
return self.remove_route(&key);
}
}
self.one_shot_routes
.get(&handle.channel)
.is_some_and(|route| route.handle == handle)
.then(|| self.one_shot_routes.remove(&handle.channel))
.flatten()
}
fn drain_routes(&mut self) -> Vec<RouteState> {
self.route_by_channel.clear();
self.routes
.drain()
.map(|(_, route)| route)
.chain(self.one_shot_routes.drain().map(|(_, route)| route))
.collect()
}
fn close_routes(&mut self) {
self.push_event_receivers.clear();
for route in self.drain_routes() {
route.sem.close();
}
}
}
impl Shared {
fn new(connection_file: PathBuf, opts: ConsumerOptions) -> Self {
Self {
connection_file,
opts,
inner: Mutex::new(Inner {
generation: 1,
epoch: 1,
next_corr: Some(1),
writer: None,
pending: HashMap::new(),
routes: HashMap::new(),
route_by_channel: HashMap::new(),
one_shot_routes: HashMap::new(),
route_epochs: HashMap::new(),
push_event_receivers: HashMap::new(),
dropped_route_frames: 0,
openings: HashMap::new(),
callbacks: Vec::new(),
closed: false,
reconnect: ReconnectState::Idle,
restored_token: 0,
reader_task: None,
writer_task: None,
}),
notify: Notify::new(),
close_token: CancellationToken::new(),
pushes_dropped_no_receiver: AtomicU64::new(0),
}
}
fn lock_inner(&self) -> MutexGuard<'_, Inner> {
self.inner
.lock()
.unwrap_or_else(|poisoned| poisoned.into_inner())
}
fn install_initial(self: &Arc<Self>, opened: OpenedConnection) -> Result<(), ConsumerError> {
self.install_connection(opened, InstallKind::Initial)
.map(|_| ())
}
fn install_reconnected(
self: &Arc<Self>,
opened: OpenedConnection,
) -> Result<(u64, u64), ConsumerError> {
let generation_epoch = self.install_connection(opened, InstallKind::Reconnect)?;
self.notify.notify_waiters();
Ok(generation_epoch)
}
fn install_connection(
self: &Arc<Self>,
opened: OpenedConnection,
kind: InstallKind,
) -> Result<(u64, u64), ConsumerError> {
if self.close_token.is_cancelled() {
return Err(ConsumerError::Closed);
}
let (reader, writer) = opened.stream.into_split();
let (tx, rx) = mpsc::channel(EGRESS_BUFFER);
let (generation, epoch, old_reader, old_writer) = {
let mut inner = self.lock_inner();
if inner.closed {
return Err(ConsumerError::Closed);
}
let (generation, epoch) = match kind {
InstallKind::Initial => (inner.generation, inner.epoch),
InstallKind::Reconnect => {
inner.generation = inner
.generation
.checked_add(1)
.ok_or(ConsumerError::Closed)?;
inner.epoch = inner.epoch.checked_add(1).ok_or(ConsumerError::Closed)?;
(inner.generation, inner.epoch)
}
};
inner.close_routes();
inner.route_epochs.clear();
inner.next_corr = Some(1);
inner.writer = Some(tx);
(
generation,
epoch,
inner.reader_task.take(),
inner.writer_task.take(),
)
};
if let Some(handle) = old_reader {
handle.abort();
}
if let Some(handle) = old_writer {
handle.abort();
}
let reader_shared = Arc::clone(self);
let reader_task = tokio::spawn(async move {
reader_loop(reader_shared, reader, generation).await;
});
let writer_shared = Arc::clone(self);
let writer_task = tokio::spawn(async move {
writer_loop(writer_shared, writer, rx, generation).await;
});
{
let mut inner = self.lock_inner();
if inner.closed || inner.generation != generation {
reader_task.abort();
writer_task.abort();
return Err(ConsumerError::Closed);
}
inner.reader_task = Some(reader_task);
inner.writer_task = Some(writer_task);
}
Ok((generation, epoch))
}
async fn ensure_connected_for_call(
self: &Arc<Self>,
deadline: Instant,
) -> Result<(), CallError> {
loop {
if Instant::now() >= deadline {
return Err(CallError::not_sent(
"call deadline elapsed waiting for reconnection",
));
}
let action = {
let mut inner = self.lock_inner();
if inner.closed {
return Err(CallError::not_sent("consumer closed"));
}
if inner.writer.is_some() {
return Ok(());
}
let generation = inner.generation;
let reconnect_is_live = match &inner.reconnect {
ReconnectState::Idle => false,
ReconnectState::Inline { generation: active } => *active == generation,
ReconnectState::Background {
generation: active,
task,
} => *active == generation && !task.is_finished(),
};
if reconnect_is_live {
EnsureAction::Wait
} else {
let stale_task =
match std::mem::replace(&mut inner.reconnect, ReconnectState::Idle) {
ReconnectState::Background { task, .. } => Some(task),
ReconnectState::Idle | ReconnectState::Inline { .. } => None,
};
inner.reconnect = ReconnectState::Inline { generation };
EnsureAction::Lead {
generation,
stale_task,
}
}
};
match action {
EnsureAction::Wait => {
timeout_at(deadline, self.notify.notified())
.await
.map_err(|_| {
CallError::not_sent("call deadline elapsed waiting for reconnection")
})?;
}
EnsureAction::Lead {
generation,
stale_task,
} => {
if let Some(handle) = stale_task {
handle.abort();
}
let mut guard = InlineReconnectGuard::new(Arc::clone(self), generation);
let result = timeout_at(deadline, self.reconnect_with_retry(generation)).await;
guard.finish();
return match result {
Ok(Ok(())) => Ok(()),
Ok(Err(err)) => Err(CallError::not_sent(err.to_string())),
Err(_) => Err(CallError::not_sent(
"call deadline elapsed waiting for reconnection",
)),
};
}
}
}
}
fn spawn_reconnect(self: &Arc<Self>, generation: u64) -> bool {
let stale_task = {
let mut inner = self.lock_inner();
if inner.closed || inner.writer.is_some() || inner.generation != generation {
return false;
}
let should_spawn = match &inner.reconnect {
ReconnectState::Idle => true,
ReconnectState::Inline { generation: active } => *active < generation,
ReconnectState::Background {
generation: active,
task,
} => *active < generation || (*active == generation && task.is_finished()),
};
if !should_spawn {
return false;
}
let stale_task = match std::mem::replace(&mut inner.reconnect, ReconnectState::Idle) {
ReconnectState::Background { task, .. } => Some(task),
ReconnectState::Idle | ReconnectState::Inline { .. } => None,
};
let shared = Arc::clone(self);
let handle = tokio::spawn(async move {
let _result = shared.reconnect_with_retry(generation).await;
shared.finish_background_reconnect(generation);
});
inner.reconnect = ReconnectState::Background {
generation,
task: handle,
};
stale_task
};
if let Some(handle) = stale_task {
handle.abort();
}
true
}
async fn reconnect_with_retry(
self: &Arc<Self>,
reconnect_generation: u64,
) -> Result<(), ConsumerError> {
let mut last_error: Option<ConsumerError> = None;
for attempt in 1..=self.opts.reconnect_backoff.max_attempts {
if self.close_token.is_cancelled() {
return Err(ConsumerError::Closed);
}
if !self.reconnect_attempt_is_current(reconnect_generation) {
return Ok(());
}
match open_connection(&self.connection_file, self.opts.handshake_timeout).await {
Ok(opened) => {
if !self.reconnect_attempt_is_current(reconnect_generation) {
return Ok(());
}
let (generation, epoch) = self.install_reconnected(opened)?;
self.schedule_restored(generation, epoch);
return Ok(());
}
Err(err) => {
let transient = is_reconnect_transient(&err);
last_error = Some(err);
if !transient || attempt >= self.opts.reconnect_backoff.max_attempts {
break;
}
let delay = self.opts.reconnect_backoff.delay_after_attempt(attempt);
tokio::select! {
() = self.close_token.cancelled() => return Err(ConsumerError::Closed),
() = sleep(delay) => {}
}
}
}
}
Err(last_error.unwrap_or(ConsumerError::Closed))
}
fn reconnect_attempt_is_current(&self, generation: u64) -> bool {
let inner = self.lock_inner();
!inner.closed
&& inner.writer.is_none()
&& inner.generation == generation
&& matches!(
&inner.reconnect,
ReconnectState::Inline { generation: active }
| ReconnectState::Background {
generation: active,
..
} if *active == generation
)
}
fn finish_inline_reconnect(&self, generation: u64) {
let finished = {
let mut inner = self.lock_inner();
if matches!(
&inner.reconnect,
ReconnectState::Inline { generation: active } if *active == generation
) {
inner.reconnect = ReconnectState::Idle;
true
} else {
false
}
};
if finished {
self.notify.notify_waiters();
}
}
fn finish_background_reconnect(&self, generation: u64) {
let completed_task = {
let mut inner = self.lock_inner();
if !matches!(
&inner.reconnect,
ReconnectState::Background { generation: active, .. } if *active == generation
) {
None
} else {
match std::mem::replace(&mut inner.reconnect, ReconnectState::Idle) {
ReconnectState::Background { task, .. } => Some(task),
ReconnectState::Idle | ReconnectState::Inline { .. } => unreachable!(),
}
}
};
if completed_task.is_some() {
drop(completed_task);
self.notify.notify_waiters();
}
}
fn schedule_restored(self: &Arc<Self>, generation: u64, epoch: u64) {
let token = {
let mut inner = self.lock_inner();
inner.restored_token = inner.restored_token.saturating_add(1);
inner.restored_token
};
let shared = Arc::clone(self);
tokio::spawn(async move {
tokio::select! {
() = shared.close_token.cancelled() => {}
() = sleep(shared.opts.restored_debounce) => {
let should_emit = {
let inner = shared.lock_inner();
!inner.closed
&& inner.generation == generation
&& inner.epoch == epoch
&& inner.restored_token == token
&& inner.writer.is_some()
};
if should_emit {
shared.emit_connection_state(ConnectionState::Restored { epoch });
}
}
}
});
}
fn install_one_shot_route(&self, route: RouteState) -> Result<(), CallError> {
let handle = route.handle;
let mut inner = self.lock_inner();
if inner.closed || inner.generation != handle.connection_token() || inner.writer.is_none() {
return Err(CallError::StaleRouteHandle(handle));
}
if inner.route_epochs.contains_key(&handle.channel) {
return Err(CallError::not_sent(
"daemon returned a route channel already in use",
));
}
inner.one_shot_routes.insert(handle.channel, route);
inner.route_epochs.insert(handle.channel, handle);
Ok(())
}
async fn ensure_route(
self: &Arc<Self>,
key: &RouteKey,
route_open: &RouteOpenParams<'_>,
opts: &CallOptions,
call_deadline: Instant,
) -> Result<RouteState, CallError> {
loop {
let action = {
let mut inner = self.lock_inner();
if inner.closed {
return Err(CallError::not_sent("consumer closed"));
}
if let Some(route) = inner.routes.get(key) {
if route.handle.connection_token() == inner.generation && inner.writer.is_some()
{
return Ok(route.clone());
}
}
if let Some(opening) = inner.openings.get_mut(key) {
let (tx, rx) = oneshot::channel();
opening.waiters.push(tx);
RouteOpenAction::Wait(rx)
} else {
inner.openings.insert(
key.clone(),
Opening {
waiters: Vec::new(),
closed: false,
},
);
RouteOpenAction::Lead
}
};
match action {
RouteOpenAction::Wait(rx) => match timeout_at(call_deadline, rx).await {
Ok(Ok(Ok(route))) => return Ok(route),
Ok(Ok(Err(err))) => return Err(err.into_call_error()),
Ok(Err(_)) => continue,
Err(_) => {
return Err(CallError::not_sent(
"call deadline elapsed waiting for route.open",
));
}
},
RouteOpenAction::Lead => {
let mut guard = OpeningGuard::new(Arc::clone(self), key.clone());
let result = self
.open_route_with_retry(key, route_open, opts, call_deadline)
.await
.map_err(SharedCallFailure::from);
guard.finish(result.clone());
return result.map_err(SharedCallFailure::into_call_error);
}
}
}
}
async fn open_route_with_retry(
self: &Arc<Self>,
key: &RouteKey,
route_open: &RouteOpenParams<'_>,
opts: &CallOptions,
call_deadline: Instant,
) -> Result<RouteState, CallError> {
let route_deadline = (Instant::now() + opts.route_retry_deadline).min(call_deadline);
let mut attempt = 0usize;
loop {
attempt = attempt.saturating_add(1);
let body = serde_json::to_vec(&ClientControlRequest::RouteOpen {
target: route_open.target.clone(),
identity: route_open.identity.clone(),
consumer_identity: route_open.consumer_identity.clone(),
consumer_capabilities: route_open.consumer_capabilities.clone(),
admission_facts: None,
})
.map_err(|err| CallError::not_sent(format!("failed to encode route.open: {err}")))?;
match self.control_call(body, route_deadline, true).await {
Ok(TerminalFrame::Response {
generation, body, ..
}) => {
let response =
serde_json::from_slice::<ClientControlResponse>(&body).map_err(|err| {
CallError::not_sent(format!(
"failed to decode route.open response: {err}"
))
})?;
let ClientControlResponse::RouteOpen {
route_channel,
route_epoch,
} = response
else {
return Err(CallError::not_sent(
"route.open returned an unexpected control response",
));
};
let route = RouteState {
handle: RouteHandle::new(route_channel, route_epoch, generation),
sem: Arc::new(Semaphore::new(DEFAULT_ROUTE_WINDOW)),
};
let install = {
let mut inner = self.lock_inner();
if inner.closed {
return Err(CallError::not_sent("consumer closed"));
}
let closed_during_open = inner.openings.get(key).is_some_and(|o| o.closed);
if closed_during_open
|| inner.generation != generation
|| inner.writer.is_none()
{
RouteInstall::Discard {
closed: closed_during_open,
}
} else {
let cached = inner.cache_route(key.clone(), route.clone());
RouteInstall::Cached(cached)
}
};
match install {
RouteInstall::Cached(cached) => return Ok(cached),
RouteInstall::Discard { closed } => {
if closed {
self.send_route_goodbye(route.handle, true);
self.uninstall_route_handle(route.handle);
return Err(CallError::not_sent(
"route was closed before route.open completed",
));
}
}
}
self.sleep_until_retry(route_deadline, opts.route_retry.base)
.await?;
}
Ok(TerminalFrame::Error { body, .. }) => {
if is_retryable_route_open_code(&body.code)
&& attempt < opts.route_retry.max_attempts
&& Instant::now() < route_deadline
{
let delay = opts.route_retry.delay_after_attempt(attempt);
self.sleep_until_retry(route_deadline, delay).await?;
continue;
}
return Err(CallError::not_sent(format!(
"route.open failed for target {}: {} ({})",
key.target_label(),
body.code,
body.message
)));
}
Ok(TerminalFrame::StreamEnd) => {
return Err(CallError::not_sent("route.open returned StreamEnd"));
}
Err(err)
if err.is_not_sent()
&& attempt < opts.route_retry.max_attempts
&& Instant::now() < route_deadline =>
{
let delay = opts.route_retry.delay_after_attempt(attempt);
self.sleep_until_retry(route_deadline, delay).await?;
}
Err(err) => return Err(err),
}
}
#[allow(unreachable_code)]
Err(CallError::not_sent(format!(
"route.open retry deadline elapsed for target {}",
key.target_label()
)))
}
async fn sleep_until_retry(&self, deadline: Instant, delay: Duration) -> Result<(), CallError> {
if Instant::now() >= deadline {
return Err(CallError::not_sent("retry deadline elapsed"));
}
let bounded = delay.min(deadline.saturating_duration_since(Instant::now()));
tokio::select! {
() = self.close_token.cancelled() => Err(CallError::not_sent("consumer closed")),
() = sleep(bounded) => Ok(()),
}
}
async fn control_call(
self: &Arc<Self>,
body: Vec<u8>,
deadline: Instant,
retain_late_route_open: bool,
) -> Result<TerminalFrame, CallError> {
self.ensure_connected_for_call(deadline).await?;
self.send_request(RequestSend {
expected_handle: None,
channel: 0,
epoch: 0,
body,
priority: Priority::Interactive,
admission_class: AdmissionClass::Normal,
deadline,
retain_late_route_open,
})
.await
}
async fn send_request(
self: &Arc<Self>,
request: RequestSend,
) -> Result<TerminalFrame, CallError> {
let RequestSend {
expected_handle,
channel,
epoch,
body,
priority,
admission_class,
deadline,
retain_late_route_open,
} = request;
if Instant::now() >= deadline {
return Err(CallError::not_sent(
"call deadline elapsed before request was sent",
));
}
let (generation, corr, writer) = {
let mut inner = self.lock_inner();
if inner.closed {
return Err(CallError::not_sent("consumer closed"));
}
let generation = inner.generation;
if let Some(expected) = expected_handle {
let route_pair_matches =
channel == 0 || (expected.channel == channel && expected.epoch == epoch);
if expected.connection_token() != generation
|| !route_pair_matches
|| inner.route_epochs.get(&expected.channel) != Some(&expected)
{
return Err(CallError::StaleRouteHandle(expected));
}
}
let Some(writer) = inner.writer.clone() else {
return Err(CallError::not_sent("subc connection is down before send"));
};
let Some(corr) = inner.next_corr else {
drop(inner);
self.handle_generation_drop(
generation,
"channel-0 correlation allocator exhausted".to_string(),
);
return Err(CallError::not_sent(
"correlation allocator exhausted; connection closed",
));
};
inner.next_corr = corr.checked_add(1);
(generation, corr, writer)
};
let frame = Frame::build(
FrameType::Request,
Flags::new(false, priority, false).with_admission_class(admission_class),
channel,
epoch,
corr,
body,
)
.map_err(|err| CallError::not_sent(format!("failed to build request frame: {err}")))?;
let key = PendingKey {
generation,
channel,
epoch,
corr,
};
let (tx, rx) = oneshot::channel();
{
let mut inner = self.lock_inner();
if inner.closed || inner.generation != generation || inner.writer.is_none() {
return Err(CallError::not_sent(
"connection changed before request registration",
));
}
if let Some(expected) = expected_handle {
if inner.route_epochs.get(&expected.channel) != Some(&expected) {
return Err(CallError::StaleRouteHandle(expected));
}
}
let expected_control_handle = (channel == 0 && !retain_late_route_open)
.then_some(expected_handle)
.flatten();
inner.pending.insert(
key,
PendingEntry::unary(tx, retain_late_route_open, expected_control_handle),
);
}
let mut registration =
PendingRegistration::new(Arc::clone(self), key, retain_late_route_open);
match timeout_at(
deadline,
writer.send(WriteCommand {
frame,
pending: Some(key),
}),
)
.await
{
Ok(Ok(())) => {}
Ok(Err(_)) => {
let accepted = registration.remove_pending().unwrap_or(false);
return Err(classify_failure(
accepted,
"writer task closed before accepting request",
));
}
Err(_) => {
let _ = registration.remove_pending();
return Err(CallError::not_sent(
"call deadline elapsed waiting for writer capacity",
));
}
}
tokio::select! {
result = timeout_at(deadline, rx) => match result {
Ok(Ok(result)) => {
registration.disarm();
result.into_call_result()
}
Ok(Err(_)) => {
registration.disarm();
Err(CallError::not_sent("pending response channel closed"))
}
Err(_) => {
let accepted = if retain_late_route_open {
registration.disarm();
self.pending_accepted(key).unwrap_or(false)
} else {
registration.remove_pending().unwrap_or(false)
};
Err(classify_failure(
accepted,
format!("request on channel {channel} timed out at its deadline"),
))
}
},
() = self.close_token.cancelled() => {
let accepted = registration.remove_pending().unwrap_or(false);
Err(classify_failure(accepted, "consumer closed while request was pending"))
}
}
}
async fn send_subscription(
self: &Arc<Self>,
subscription: SubscriptionSend,
) -> Result<Subscription, CallError> {
let SubscriptionSend {
expected_handle,
channel,
epoch,
body,
priority,
admission_class,
event_buffer,
deadline,
permit,
} = subscription;
if Instant::now() >= deadline {
return Err(CallError::not_sent(
"subscription deadline elapsed before request was sent",
));
}
let (generation, corr, writer) = {
let mut inner = self.lock_inner();
if inner.closed {
return Err(CallError::not_sent("consumer closed"));
}
let generation = inner.generation;
if let Some(expected) = expected_handle {
if expected.connection_token() != generation
|| expected.channel != channel
|| expected.epoch != epoch
|| inner.route_epochs.get(&channel) != Some(&expected)
{
return Err(CallError::StaleRouteHandle(expected));
}
}
let Some(writer) = inner.writer.clone() else {
return Err(CallError::not_sent("subc connection is down before send"));
};
let Some(corr) = inner.next_corr else {
drop(inner);
self.handle_generation_drop(
generation,
"correlation allocator exhausted".to_string(),
);
return Err(CallError::not_sent(
"correlation allocator exhausted; connection closed",
));
};
inner.next_corr = corr.checked_add(1);
(generation, corr, writer)
};
let frame = Frame::build(
FrameType::Request,
Flags::new(false, priority, false).with_admission_class(admission_class),
channel,
epoch,
corr,
body,
)
.map_err(|err| CallError::not_sent(format!("failed to build request frame: {err}")))?;
let key = PendingKey {
generation,
channel,
epoch,
corr,
};
let (events_tx, events_rx) = mpsc::channel(event_buffer.max(1));
let (closed_tx, closed_rx) = oneshot::channel();
{
let mut inner = self.lock_inner();
if inner.closed || inner.generation != generation || inner.writer.is_none() {
return Err(CallError::not_sent(
"connection changed before subscription registration",
));
}
if let Some(expected) = expected_handle {
if inner.route_epochs.get(&expected.channel) != Some(&expected) {
return Err(CallError::StaleRouteHandle(expected));
}
}
inner.pending.insert(
key,
PendingEntry::subscription(events_tx, closed_tx, permit, priority),
);
}
let mut registration = PendingRegistration::new(Arc::clone(self), key, false);
match timeout_at(
deadline,
writer.send(WriteCommand {
frame,
pending: Some(key),
}),
)
.await
{
Ok(Ok(())) => {}
Ok(Err(_)) => {
let accepted = registration.remove_pending().unwrap_or(false);
return Err(classify_failure(
accepted,
"writer task closed before accepting subscription request",
));
}
Err(_) => {
let _ = registration.remove_pending();
return Err(CallError::not_sent(
"subscription deadline elapsed waiting for writer capacity",
));
}
}
registration.disarm();
Ok(Subscription {
events: events_rx,
closed: SubscriptionClosed { rx: closed_rx },
cancel: SubscriptionCancel::new(Arc::clone(self), key, priority),
})
}
fn unsubscribe_subscription(
&self,
key: PendingKey,
priority: Priority,
) -> Result<(), CallError> {
let handle = RouteHandle::new(key.channel, key.epoch, key.generation);
self.validate_current_handle(handle)?;
let entry = self.lock_inner().pending.remove(&key);
if let Some(entry) = entry {
entry.settle_subscription_result(Ok(()));
self.send_cancel(handle, key.corr, priority);
}
Ok(())
}
fn route_stream_data(&self, key: PendingKey, body: Vec<u8>) {
let overflow = {
let mut inner = self.lock_inner();
let Some(entry) = inner.pending.get(&key) else {
return;
};
match entry.try_send_stream_data(body) {
Ok(()) | Err(StreamDataDelivery::NotSubscription) => return,
Err(StreamDataDelivery::Full) => {
let priority = entry
.subscription_priority()
.unwrap_or(Priority::Interactive);
let entry = inner.pending.remove(&key);
entry.map(|entry| {
(
entry,
priority,
"subscription event channel filled; reader dropped the stream instead of blocking",
)
})
}
Err(StreamDataDelivery::Closed) => {
let priority = entry
.subscription_priority()
.unwrap_or(Priority::Interactive);
let entry = inner.pending.remove(&key);
entry.map(|entry| {
(
entry,
priority,
"subscription event receiver closed before the stream ended",
)
})
}
}
};
if let Some((entry, priority, reason)) = overflow {
entry.settle_call_error(CallError::subscription_backpressure(reason));
self.send_cancel(
RouteHandle::new(key.channel, key.epoch, key.generation),
key.corr,
priority,
);
}
}
fn register_push_events(
&self,
handle: RouteHandle,
) -> Result<mpsc::Receiver<PushEvent>, CallError> {
let (events_tx, events_rx) = mpsc::channel(DEFAULT_PUSH_EVENT_BUFFER);
let mut inner = self.lock_inner();
if inner.closed
|| inner.generation != handle.connection_token()
|| inner.writer.is_none()
|| inner.route_epochs.get(&handle.channel) != Some(&handle)
{
return Err(CallError::StaleRouteHandle(handle));
}
inner.push_event_receivers.insert(handle, events_tx);
Ok(events_rx)
}
fn route_push(&self, handle: RouteHandle, body: Vec<u8>) {
let should_count_drop = {
let mut inner = self.lock_inner();
if inner.closed
|| inner.generation != handle.connection_token()
|| inner.route_epochs.get(&handle.channel) != Some(&handle)
{
return;
}
match inner.push_event_receivers.get(&handle) {
None => true,
Some(events) => match events.try_send(PushEvent { handle, body }) {
Ok(()) => false,
Err(mpsc::error::TrySendError::Closed(_)) => {
inner.push_event_receivers.remove(&handle);
true
}
Err(mpsc::error::TrySendError::Full(_)) => {
inner.push_event_receivers.remove(&handle);
false
}
},
}
};
if should_count_drop {
self.pushes_dropped_no_receiver
.fetch_add(1, Ordering::Relaxed);
}
}
fn send_cancel(&self, handle: RouteHandle, corr: u64, priority: Priority) {
let writer = {
let inner = self.lock_inner();
if inner.closed
|| inner.generation != handle.connection_token()
|| inner.route_epochs.get(&handle.channel) != Some(&handle)
{
return;
}
inner.writer.clone()
};
let Some(writer) = writer else {
return;
};
let Ok(frame) = Frame::build(
FrameType::Cancel,
Flags::new(false, priority, false),
handle.channel,
handle.epoch,
corr,
Vec::new(),
) else {
return;
};
let _ = writer.try_send(WriteCommand {
frame,
pending: None,
});
}
fn mark_pending_accepted(&self, key: PendingKey) -> bool {
let mut inner = self.lock_inner();
if inner.closed || inner.generation != key.generation {
return false;
}
let Some(entry) = inner.pending.get_mut(&key) else {
return false;
};
entry.accepted = true;
true
}
fn settle_pending(self: &Arc<Self>, key: PendingKey, terminal: PendingTerminal) {
let entry = self.lock_inner().pending.remove(&key);
let Some(entry) = entry else {
return;
};
if entry.retain_late_route_open && entry.completion_is_closed() {
if let PendingTerminal::Response { generation, body } = &terminal {
if let Ok(ClientControlResponse::RouteOpen {
route_channel,
route_epoch,
}) = serde_json::from_slice::<ClientControlResponse>(body)
{
let handle = RouteHandle::new(route_channel, route_epoch, *generation);
self.send_route_goodbye(handle, true);
self.uninstall_route_handle(handle);
}
}
return;
}
entry.settle_terminal(terminal);
}
fn pending_accepted(&self, key: PendingKey) -> Option<bool> {
self.lock_inner()
.pending
.get(&key)
.map(|entry| entry.accepted)
}
fn handle_generation_drop(self: &Arc<Self>, generation: u64, reason: String) {
let (should_emit, pending, openings, callbacks) = {
let mut inner = self.lock_inner();
if inner.closed || inner.generation != generation || inner.writer.is_none() {
return;
}
inner.writer = None;
inner.restored_token = inner.restored_token.saturating_add(1);
inner.close_routes();
inner.route_epochs.clear();
let pending = drain_pending_generation(&mut inner.pending, generation);
let openings = drain_openings(&mut inner.openings);
let callbacks = inner.callbacks.clone();
(true, pending, openings, callbacks)
};
if should_emit {
settle_pending_entries(pending, reason.clone());
fail_openings(openings, SharedCallFailure::not_sent(reason.clone()));
emit_callbacks(callbacks, ConnectionState::Dropped);
self.notify.notify_waiters();
let _ = self.spawn_reconnect(generation);
}
}
fn close_sync(&self, reason: &str) {
let (pending, openings, routes, reader, writer, reconnect) = {
let mut inner = self.lock_inner();
if inner.closed {
return;
}
inner.closed = true;
inner.writer = None;
inner.route_epochs.clear();
inner.push_event_receivers.clear();
self.close_token.cancel();
let reconnect = match std::mem::replace(&mut inner.reconnect, ReconnectState::Idle) {
ReconnectState::Background { task, .. } => Some(task),
ReconnectState::Idle | ReconnectState::Inline { .. } => None,
};
(
inner
.pending
.drain()
.map(|(_, entry)| entry)
.collect::<Vec<_>>(),
inner
.openings
.drain()
.map(|(_, opening)| opening.waiters)
.collect::<Vec<_>>(),
inner.drain_routes(),
inner.reader_task.take(),
inner.writer_task.take(),
reconnect,
)
};
for route in routes {
route.sem.close();
}
if let Some(handle) = reader {
handle.abort();
}
if let Some(handle) = writer {
handle.abort();
}
if let Some(handle) = reconnect {
handle.abort();
}
settle_pending_entries(pending, reason.to_string());
fail_openings(openings, SharedCallFailure::not_sent(reason.to_string()));
self.notify.notify_waiters();
}
fn validate_current_handle(&self, handle: RouteHandle) -> Result<(), CallError> {
let inner = self.lock_inner();
if inner.closed
|| inner.generation != handle.connection_token()
|| inner.writer.is_none()
|| inner.route_epochs.get(&handle.channel) != Some(&handle)
{
Err(CallError::StaleRouteHandle(handle))
} else {
Ok(())
}
}
fn route_state(&self, handle: RouteHandle) -> Result<RouteState, CallError> {
let inner = self.lock_inner();
if inner.closed
|| inner.generation != handle.connection_token()
|| inner.writer.is_none()
|| inner.route_epochs.get(&handle.channel) != Some(&handle)
{
return Err(CallError::StaleRouteHandle(handle));
}
let route = inner
.route_by_channel
.get(&handle.channel)
.and_then(|key| inner.routes.get(key))
.or_else(|| inner.one_shot_routes.get(&handle.channel));
debug_assert!(route.is_none_or(|route| route.handle == handle));
route
.filter(|route| route.handle == handle)
.cloned()
.ok_or(CallError::StaleRouteHandle(handle))
}
fn route_is_current(&self, key: &RouteKey, route: &RouteState) -> bool {
let inner = self.lock_inner();
if inner.closed
|| inner.generation != route.handle.connection_token()
|| inner.writer.is_none()
{
return false;
}
inner.routes.get(key).is_some_and(|cached| {
cached.handle == route.handle && Arc::ptr_eq(&cached.sem, &route.sem)
})
}
fn invalidate_route(&self, key: &RouteKey, expected_handle: Option<RouteHandle>) {
let removed = {
let mut inner = self.lock_inner();
match inner.routes.get(key) {
Some(route) if expected_handle.is_none_or(|expected| expected == route.handle) => {
let removed = inner.remove_route(key);
if let Some(route) = &removed {
if inner.route_epochs.get(&route.handle.channel) == Some(&route.handle) {
inner.route_epochs.remove(&route.handle.channel);
}
inner.push_event_receivers.remove(&route.handle);
}
removed
}
_ => None,
}
};
if let Some(route) = removed {
route.sem.close();
}
}
fn finish_opening(&self, key: &RouteKey, result: Result<RouteState, SharedCallFailure>) {
let opening = self.lock_inner().openings.remove(key);
for waiter in opening.map(|o| o.waiters).unwrap_or_default() {
let _ = waiter.send(result.clone());
}
}
async fn close_handle(
self: &Arc<Self>,
handle: RouteHandle,
opts: &CloseRouteOptions,
) -> Result<(), CallError> {
self.validate_current_handle(handle)?;
let routes = {
let mut inner = self.lock_inner();
inner
.remove_route_by_handle(handle)
.into_iter()
.collect::<Vec<_>>()
};
if opts.drain {
self.drain_channel(handle, opts.drain_timeout).await;
}
for route in routes {
route.sem.close();
}
self.fail_channel_pending(handle, "route closed by close_handle");
self.send_route_goodbye(handle, false);
self.uninstall_route_handle(handle);
Ok(())
}
async fn close_route(self: &Arc<Self>, key: &RouteKey, opts: &CloseRouteOptions) {
let route = {
let mut inner = self.lock_inner();
if let Some(opening) = inner.openings.get_mut(key) {
opening.closed = true;
}
inner.remove_route(key)
};
let Some(route) = route else {
return;
};
if opts.drain {
self.drain_channel(route.handle, opts.drain_timeout).await;
}
route.sem.close();
self.fail_channel_pending(route.handle, "route closed by close_route");
self.send_route_goodbye(route.handle, false);
self.uninstall_route_handle(route.handle);
}
fn uninstall_route_handle(&self, handle: RouteHandle) {
let mut inner = self.lock_inner();
if inner.route_epochs.get(&handle.channel) == Some(&handle) {
inner.route_epochs.remove(&handle.channel);
inner.push_event_receivers.remove(&handle);
}
}
fn fail_channel_pending(&self, handle: RouteHandle, reason: &str) {
let entries = {
let mut inner = self.lock_inner();
drain_pending_handle(&mut inner.pending, handle, true)
};
settle_pending_entries(entries, reason.to_string());
}
async fn drain_channel(&self, handle: RouteHandle, timeout: Duration) {
let deadline = Instant::now() + timeout;
loop {
let has_inflight = {
let inner = self.lock_inner();
inner.pending.iter().any(|(key, entry)| {
key.generation == handle.connection_token()
&& key.channel == handle.channel
&& key.epoch == handle.epoch
&& !entry.is_subscription()
})
};
if !has_inflight || Instant::now() >= deadline {
return;
}
sleep(Duration::from_millis(5)).await;
}
}
fn send_route_goodbye(self: &Arc<Self>, handle: RouteHandle, close_on_failure: bool) -> bool {
let writer = {
let inner = self.lock_inner();
if inner.closed
|| inner.generation != handle.connection_token()
|| inner.route_epochs.get(&handle.channel) != Some(&handle)
{
return false;
}
inner.writer.clone()
};
let Some(writer) = writer else {
return false;
};
let Ok(frame) = Frame::build(
FrameType::Goodbye,
Flags::new(false, Priority::Interactive, false),
handle.channel,
handle.epoch,
0,
Vec::new(),
) else {
return false;
};
if writer
.try_send(WriteCommand {
frame,
pending: None,
})
.is_ok()
{
return true;
}
if close_on_failure {
self.handle_generation_drop(
handle.connection_token(),
"failed to queue late route.open cleanup GOODBYE".to_string(),
);
}
false
}
fn emit_connection_state(&self, state: ConnectionState) {
let callbacks = self.lock_inner().callbacks.clone();
emit_callbacks(callbacks, state);
}
}
#[derive(Clone, Copy)]
enum InstallKind {
Initial,
Reconnect,
}
enum EnsureAction {
Wait,
Lead {
generation: u64,
stale_task: Option<JoinHandle<()>>,
},
}
enum ReconnectState {
Idle,
Inline {
generation: u64,
},
Background {
generation: u64,
task: JoinHandle<()>,
},
}
enum RouteInstall {
Cached(RouteState),
Discard { closed: bool },
}
enum RouteOpenAction {
Wait(oneshot::Receiver<Result<RouteState, SharedCallFailure>>),
Lead,
}
struct RouteOpenParams<'a> {
target: &'a RouteTarget,
identity: &'a BindIdentity,
consumer_identity: &'a Option<ConsumerIdentity>,
consumer_capabilities: &'a Option<Vec<String>>,
}
struct RequestSend {
expected_handle: Option<RouteHandle>,
channel: u16,
epoch: u32,
body: Vec<u8>,
priority: Priority,
admission_class: AdmissionClass,
deadline: Instant,
retain_late_route_open: bool,
}
struct SubscriptionSend {
expected_handle: Option<RouteHandle>,
channel: u16,
epoch: u32,
body: Vec<u8>,
priority: Priority,
admission_class: AdmissionClass,
event_buffer: usize,
deadline: Instant,
permit: OwnedSemaphorePermit,
}
struct OpeningGuard {
shared: Arc<Shared>,
key: RouteKey,
finished: bool,
}
impl OpeningGuard {
fn new(shared: Arc<Shared>, key: RouteKey) -> Self {
Self {
shared,
key,
finished: false,
}
}
fn finish(&mut self, result: Result<RouteState, SharedCallFailure>) {
self.shared.finish_opening(&self.key, result);
self.finished = true;
}
}
impl Drop for OpeningGuard {
fn drop(&mut self) {
if !self.finished {
self.shared.finish_opening(
&self.key,
Err(SharedCallFailure::not_sent(
"route.open future was cancelled",
)),
);
}
}
}
struct InlineReconnectGuard {
shared: Arc<Shared>,
generation: u64,
finished: bool,
}
impl InlineReconnectGuard {
fn new(shared: Arc<Shared>, generation: u64) -> Self {
Self {
shared,
generation,
finished: false,
}
}
fn finish(&mut self) {
self.shared.finish_inline_reconnect(self.generation);
self.finished = true;
}
}
impl Drop for InlineReconnectGuard {
fn drop(&mut self) {
if !self.finished {
self.shared.finish_inline_reconnect(self.generation);
}
}
}
#[derive(Clone)]
struct RouteState {
handle: RouteHandle,
sem: Arc<Semaphore>,
}
#[derive(Debug, Clone, Hash, PartialEq, Eq)]
struct RouteKey {
target: RouteTargetKey,
project_root: PathBuf,
harness: String,
session: String,
consumer_identity: Option<ConsumerIdentityKey>,
consumer_capabilities: Option<ConsumerCapabilitiesKey>,
}
impl RouteKey {
fn new(
target: &RouteTarget,
identity: &BindIdentity,
consumer_identity: Option<&ConsumerIdentity>,
consumer_capabilities: Option<&[String]>,
) -> Self {
Self {
target: RouteTargetKey::from(target),
project_root: identity.project_root.clone(),
harness: identity.harness.clone(),
session: identity.session.clone(),
consumer_identity: consumer_identity.map(ConsumerIdentityKey::from),
consumer_capabilities: consumer_capabilities.map(ConsumerCapabilitiesKey::from_slice),
}
}
fn target_label(&self) -> String {
match &self.target {
RouteTargetKey::ToolProvider { module_id } => format!("tool_provider:{module_id}"),
RouteTargetKey::ManagementSurface { module_id } => {
format!("management_surface:{module_id}")
}
RouteTargetKey::InternalService {
module_id,
service_id,
} => format!("internal_service:{module_id}:{service_id}"),
}
}
}
#[derive(Debug, Clone, Hash, PartialEq, Eq)]
struct ConsumerIdentityKey {
module_id: String,
launch_nonce: String,
}
impl From<&ConsumerIdentity> for ConsumerIdentityKey {
fn from(value: &ConsumerIdentity) -> Self {
Self {
module_id: value.module_id.clone(),
launch_nonce: value.launch_nonce.clone(),
}
}
}
#[derive(Debug, Clone, Hash, PartialEq, Eq)]
struct ConsumerCapabilitiesKey {
values: Vec<String>,
}
impl ConsumerCapabilitiesKey {
fn from_slice(values: &[String]) -> Self {
let mut values = values.to_vec();
values.sort();
values.dedup();
Self { values }
}
}
#[derive(Debug, Clone, Hash, PartialEq, Eq)]
enum RouteTargetKey {
ToolProvider {
module_id: String,
},
ManagementSurface {
module_id: String,
},
InternalService {
module_id: String,
service_id: String,
},
}
impl From<&RouteTarget> for RouteTargetKey {
fn from(value: &RouteTarget) -> Self {
match value {
RouteTarget::ToolProvider { module_id } => Self::ToolProvider {
module_id: module_id.clone(),
},
RouteTarget::ManagementSurface { module_id } => Self::ManagementSurface {
module_id: module_id.clone(),
},
RouteTarget::InternalService {
module_id,
service_id,
} => Self::InternalService {
module_id: module_id.clone(),
service_id: service_id.clone(),
},
}
}
}
#[derive(Debug, Clone)]
struct SharedCallFailure {
kind: FailureKind,
message: String,
}
impl SharedCallFailure {
fn not_sent(message: impl Into<String>) -> Self {
Self {
kind: FailureKind::NotSent,
message: message.into(),
}
}
fn into_call_error(self) -> CallError {
match self.kind {
FailureKind::NotSent => CallError::not_sent(self.message),
FailureKind::OutcomeUnknown => CallError::outcome_unknown(self.message),
}
}
}
impl From<CallError> for SharedCallFailure {
fn from(value: CallError) -> Self {
match value {
CallError::NotSent(err) => Self {
kind: FailureKind::NotSent,
message: err.to_string(),
},
CallError::OutcomeUnknown(err) => Self {
kind: FailureKind::OutcomeUnknown,
message: err.to_string(),
},
CallError::Module(body) => Self {
kind: FailureKind::OutcomeUnknown,
message: format!(
"unexpected module error during route.open: {} ({})",
body.code, body.message
),
},
CallError::SubscriptionBackpressure(err) => Self {
kind: FailureKind::OutcomeUnknown,
message: err.to_string(),
},
CallError::StaleRouteHandle(handle) => Self {
kind: FailureKind::NotSent,
message: format!("stale route handle: {handle:?}"),
},
}
}
}
#[derive(Debug, Clone, Copy)]
enum FailureKind {
NotSent,
OutcomeUnknown,
}
#[derive(Debug, Clone, Copy, Hash, PartialEq, Eq)]
struct PendingKey {
generation: u64,
channel: u16,
epoch: u32,
corr: u64,
}
struct PendingEntry {
accepted: bool,
retain_late_route_open: bool,
expected_control_handle: Option<RouteHandle>,
completion: PendingCompletion,
}
enum PendingCompletion {
Unary(oneshot::Sender<PendingResult>),
Subscription {
events: mpsc::Sender<Vec<u8>>,
closed: oneshot::Sender<Result<(), CallError>>,
_permit: OwnedSemaphorePermit,
priority: Priority,
},
}
enum StreamDataDelivery {
NotSubscription,
Full,
Closed,
}
impl PendingEntry {
fn unary(
tx: oneshot::Sender<PendingResult>,
retain_late_route_open: bool,
expected_control_handle: Option<RouteHandle>,
) -> Self {
Self {
accepted: false,
retain_late_route_open,
expected_control_handle,
completion: PendingCompletion::Unary(tx),
}
}
fn subscription(
events: mpsc::Sender<Vec<u8>>,
closed: oneshot::Sender<Result<(), CallError>>,
permit: OwnedSemaphorePermit,
priority: Priority,
) -> Self {
Self {
accepted: false,
retain_late_route_open: false,
expected_control_handle: None,
completion: PendingCompletion::Subscription {
events,
closed,
_permit: permit,
priority,
},
}
}
fn completion_is_closed(&self) -> bool {
match &self.completion {
PendingCompletion::Unary(tx) => tx.is_closed(),
PendingCompletion::Subscription { closed, .. } => closed.is_closed(),
}
}
fn is_subscription(&self) -> bool {
matches!(&self.completion, PendingCompletion::Subscription { .. })
}
fn subscription_priority(&self) -> Option<Priority> {
match &self.completion {
PendingCompletion::Subscription { priority, .. } => Some(*priority),
PendingCompletion::Unary(_) => None,
}
}
fn try_send_stream_data(&self, body: Vec<u8>) -> Result<(), StreamDataDelivery> {
let PendingCompletion::Subscription { events, .. } = &self.completion else {
return Err(StreamDataDelivery::NotSubscription);
};
events.try_send(body).map_err(|err| match err {
mpsc::error::TrySendError::Full(_) => StreamDataDelivery::Full,
mpsc::error::TrySendError::Closed(_) => StreamDataDelivery::Closed,
})
}
fn settle_terminal(self, terminal: PendingTerminal) {
match self.completion {
PendingCompletion::Unary(tx) => {
let _ = tx.send(PendingResult::Terminal(terminal));
}
PendingCompletion::Subscription { closed, .. } => {
let result = match terminal {
PendingTerminal::Response { .. } | PendingTerminal::StreamEnd => Ok(()),
PendingTerminal::Error { body } => Err(CallError::Module(body)),
};
let _ = closed.send(result);
}
}
}
fn settle_failure(self, reason: String) {
let accepted = self.accepted;
self.settle_call_error(classify_failure(accepted, reason));
}
fn settle_call_error(self, err: CallError) {
match self.completion {
PendingCompletion::Unary(tx) => {
let _ = tx.send(PendingResult::Failure {
accepted: self.accepted,
reason: err.to_string(),
});
}
PendingCompletion::Subscription { closed, .. } => {
let _ = closed.send(Err(err));
}
}
}
fn settle_subscription_result(self, result: Result<(), CallError>) {
match self.completion {
PendingCompletion::Subscription { closed, .. } => {
let _ = closed.send(result);
}
PendingCompletion::Unary(tx) => {
let _ = tx.send(PendingResult::Failure {
accepted: self.accepted,
reason: "subscription cancel matched a unary request".to_string(),
});
}
}
}
}
struct PendingRegistration {
shared: Arc<Shared>,
key: PendingKey,
active: bool,
retain_on_drop: bool,
}
impl PendingRegistration {
fn new(shared: Arc<Shared>, key: PendingKey, retain_on_drop: bool) -> Self {
Self {
shared,
key,
active: true,
retain_on_drop,
}
}
fn remove_pending(&mut self) -> Option<bool> {
if !self.active {
return None;
}
self.active = false;
self.shared
.lock_inner()
.pending
.remove(&self.key)
.map(|entry| entry.accepted)
}
fn disarm(&mut self) {
self.active = false;
}
}
impl Drop for PendingRegistration {
fn drop(&mut self) {
if self.retain_on_drop {
self.disarm();
} else {
let _ = self.remove_pending();
}
}
}
enum PendingResult {
Terminal(PendingTerminal),
Failure { accepted: bool, reason: String },
}
impl PendingResult {
fn into_call_result(self) -> Result<TerminalFrame, CallError> {
match self {
Self::Terminal(terminal) => Ok(terminal.into_terminal_frame()),
Self::Failure { accepted, reason } => Err(classify_failure(accepted, reason)),
}
}
}
enum PendingTerminal {
Response { generation: u64, body: Vec<u8> },
Error { body: ErrorBody },
StreamEnd,
}
impl PendingTerminal {
fn into_terminal_frame(self) -> TerminalFrame {
match self {
Self::Response { generation, body } => TerminalFrame::Response { generation, body },
Self::Error { body } => TerminalFrame::Error { body },
Self::StreamEnd => TerminalFrame::StreamEnd,
}
}
}
#[derive(Debug)]
enum TerminalFrame {
Response { generation: u64, body: Vec<u8> },
Error { body: ErrorBody },
StreamEnd,
}
struct WriteCommand {
frame: Frame,
pending: Option<PendingKey>,
}
struct OpenedConnection {
stream: TcpStream,
}
async fn open_connection(
path: &Path,
deadline: Duration,
) -> Result<OpenedConnection, ConsumerError> {
let conn =
connection_file::read_for_client(path).map_err(|source| ConsumerError::ConnectionFile {
path: path.to_path_buf(),
source,
})?;
let endpoint = conn
.endpoints
.first()
.ok_or_else(|| ConsumerError::NoEndpoint {
path: path.to_path_buf(),
})?;
let endpoint_label = format!("{}:{}", endpoint.host, endpoint.port);
let mut stream = TcpStream::connect(&endpoint_label)
.await
.map_err(|source| ConsumerError::Connect {
path: path.to_path_buf(),
endpoint: endpoint_label.clone(),
source,
})?;
let _ = stream.set_nodelay(true);
authenticate_client(&mut stream, &conn, deadline)
.await
.map_err(|source| ConsumerError::Auth {
path: path.to_path_buf(),
endpoint: endpoint_label,
source,
})?;
Ok(OpenedConnection { stream })
}
async fn reader_loop(shared: Arc<Shared>, mut reader: OwnedReadHalf, generation: u64) {
loop {
match read_frame(&mut reader).await {
Ok(Some(frame)) => {
if !dispatch_frame(&shared, generation, frame).await {
return;
}
}
Ok(None) => {
shared.handle_generation_drop(generation, "subc connection closed".to_string());
return;
}
Err(err) => {
shared.handle_generation_drop(generation, err.to_string());
return;
}
}
}
}
async fn dispatch_frame(shared: &Arc<Shared>, generation: u64, frame: Frame) -> bool {
if !shared.generation_is_current(generation) {
return false;
}
if frame.header.channel != 0
&& !shared.validate_ingress_handle(generation, frame.header.channel, frame.header.epoch)
{
return true;
}
let key = PendingKey {
generation,
channel: frame.header.channel,
epoch: frame.header.epoch,
corr: frame.header.corr,
};
if frame.header.channel == 0 && frame.header.ty == FrameType::Response {
if let Some(expected) = shared.pending_expected_control_handle(key) {
let echoes_expected = matches!(
serde_json::from_slice::<ClientControlResponse>(&frame.body),
Ok(ClientControlResponse::RoutePoll {
route_channel,
route_epoch,
..
}) if route_channel == expected.channel && route_epoch == expected.epoch
);
if !echoes_expected {
shared.count_dropped_route_frame();
return true;
}
}
}
if frame.header.channel == 0
&& frame.header.ty == FrameType::Response
&& shared.pending_expects_route_open(key)
{
if let Ok(ClientControlResponse::RouteOpen {
route_channel,
route_epoch,
}) = serde_json::from_slice::<ClientControlResponse>(&frame.body)
{
shared.install_ingress_handle(RouteHandle::new(route_channel, route_epoch, generation));
}
}
match frame.header.ty {
FrameType::Response => shared.settle_pending(
key,
PendingTerminal::Response {
generation,
body: frame.body,
},
),
FrameType::Error => {
let body =
serde_json::from_slice::<ErrorBody>(&frame.body).unwrap_or_else(|err| ErrorBody {
code: "invalid_error_body".to_string(),
message: err.to_string(),
});
shared.settle_pending(key, PendingTerminal::Error { body });
}
FrameType::StreamEnd => shared.settle_pending(key, PendingTerminal::StreamEnd),
FrameType::StreamData => shared.route_stream_data(key, frame.body),
FrameType::Push => shared.route_push(
RouteHandle::new(frame.header.channel, frame.header.epoch, generation),
frame.body,
),
FrameType::Goodbye if frame.header.channel == 0 => {
shared.handle_generation_drop(generation, "subc sent GOODBYE".to_string());
return false;
}
FrameType::Goodbye => {
let handle = RouteHandle::new(frame.header.channel, frame.header.epoch, generation);
shared.invalidate_routes_for_handle(handle);
let pending = {
let mut inner = shared.lock_inner();
drain_pending_handle(&mut inner.pending, handle, true)
};
settle_pending_entries(pending, "route closed by subc".to_string());
}
FrameType::Ping if frame.header.channel == 0 => {
if let Ok(pong) = Frame::build_with_version(
frame.header.ver,
FrameType::Pong,
frame.header.flags,
0,
0,
frame.header.corr,
Vec::new(),
) {
let writer = shared.lock_inner().writer.clone();
if let Some(writer) = writer {
let _ = writer
.send(WriteCommand {
frame: pong,
pending: None,
})
.await;
}
}
}
_ => {}
}
true
}
impl Shared {
fn generation_is_current(&self, generation: u64) -> bool {
let inner = self.lock_inner();
!inner.closed && inner.generation == generation && inner.writer.is_some()
}
fn pending_expected_control_handle(&self, key: PendingKey) -> Option<RouteHandle> {
self.lock_inner()
.pending
.get(&key)
.and_then(|entry| entry.expected_control_handle)
}
fn count_dropped_route_frame(&self) {
let mut inner = self.lock_inner();
inner.dropped_route_frames = inner.dropped_route_frames.saturating_add(1);
}
fn pending_expects_route_open(&self, key: PendingKey) -> bool {
self.lock_inner()
.pending
.get(&key)
.is_some_and(|entry| entry.retain_late_route_open)
}
fn validate_ingress_handle(&self, generation: u64, channel: u16, epoch: u32) -> bool {
let mut inner = self.lock_inner();
let expected = RouteHandle::new(channel, epoch, generation);
if inner.route_epochs.get(&channel) == Some(&expected) {
true
} else {
inner.dropped_route_frames = inner.dropped_route_frames.saturating_add(1);
false
}
}
fn install_ingress_handle(&self, handle: RouteHandle) {
let mut inner = self.lock_inner();
if !inner.closed && inner.generation == handle.connection_token() && inner.writer.is_some()
{
inner.route_epochs.insert(handle.channel, handle);
}
}
fn invalidate_routes_for_handle(&self, handle: RouteHandle) {
let removed = {
let mut inner = self.lock_inner();
if inner.route_epochs.get(&handle.channel) != Some(&handle) {
return;
}
inner.route_epochs.remove(&handle.channel);
inner.push_event_receivers.remove(&handle);
inner
.remove_route_by_handle(handle)
.into_iter()
.collect::<Vec<_>>()
};
for route in removed {
route.sem.close();
}
}
}
async fn writer_loop<W>(
shared: Arc<Shared>,
writer: W,
mut rx: mpsc::Receiver<WriteCommand>,
generation: u64,
) where
W: AsyncWrite + Unpin,
{
let mut writer = BufWriter::new(writer);
while let Some(command) = rx.recv().await {
if let Some(key) = command.pending {
if !shared.mark_pending_accepted(key) {
continue;
}
}
if let Err(err) = write_frame(&mut writer, &command.frame).await {
shared.handle_generation_drop(generation, err.to_string());
return;
}
while let Ok(command) = rx.try_recv() {
if let Some(key) = command.pending {
if !shared.mark_pending_accepted(key) {
continue;
}
}
if let Err(err) = write_frame(&mut writer, &command.frame).await {
shared.handle_generation_drop(generation, err.to_string());
return;
}
}
if let Err(err) = writer.flush().await.map_err(FrameIoError::Io) {
shared.handle_generation_drop(generation, err.to_string());
return;
}
}
}
fn route_open_consumer_identity(opts: &CallOptions) -> Option<ConsumerIdentity> {
opts.consumer_identity
.clone()
.or_else(consumer_identity_from_env)
}
fn close_route_consumer_identity(opts: &CloseRouteOptions) -> Option<ConsumerIdentity> {
opts.consumer_identity
.clone()
.or_else(consumer_identity_from_env)
}
fn route_open_consumer_capabilities(opts: &CallOptions) -> Option<Vec<String>> {
opts.consumer_capabilities.clone()
}
fn close_route_consumer_capabilities(opts: &CloseRouteOptions) -> Option<Vec<String>> {
opts.consumer_capabilities.clone()
}
fn consumer_identity_from_env() -> Option<ConsumerIdentity> {
let module_id = std::env::var(SUBC_MODULE_ID_ENV)
.ok()
.filter(|value| !value.is_empty())?;
let launch_nonce = std::env::var(SUBC_LAUNCH_NONCE_ENV)
.ok()
.filter(|value| !value.is_empty())?;
Some(ConsumerIdentity {
module_id,
launch_nonce,
})
}
fn classify_failure(accepted: bool, reason: impl Into<String>) -> CallError {
if accepted {
CallError::outcome_unknown(reason)
} else {
CallError::not_sent(reason)
}
}
fn settle_pending_entries(entries: Vec<PendingEntry>, reason: String) {
for entry in entries {
entry.settle_failure(reason.clone());
}
}
fn drain_pending_generation(
pending: &mut HashMap<PendingKey, PendingEntry>,
generation: u64,
) -> Vec<PendingEntry> {
let keys = pending
.keys()
.copied()
.filter(|key| key.generation == generation)
.collect::<Vec<_>>();
keys.into_iter()
.filter_map(|key| pending.remove(&key))
.collect()
}
fn drain_pending_handle(
pending: &mut HashMap<PendingKey, PendingEntry>,
handle: RouteHandle,
include_subscriptions: bool,
) -> Vec<PendingEntry> {
let keys = pending
.iter()
.filter_map(|(key, entry)| {
(key.generation == handle.connection_token()
&& key.channel == handle.channel
&& key.epoch == handle.epoch
&& (include_subscriptions || !entry.is_subscription()))
.then_some(*key)
})
.collect::<Vec<_>>();
keys.into_iter()
.filter_map(|key| pending.remove(&key))
.collect()
}
fn drain_openings(openings: &mut HashMap<RouteKey, Opening>) -> Vec<Vec<OpeningWaiter>> {
openings
.drain()
.map(|(_, opening)| opening.waiters)
.collect()
}
fn fail_openings(openings: Vec<Vec<OpeningWaiter>>, failure: SharedCallFailure) {
for waiters in openings {
for waiter in waiters {
let _ = waiter.send(Err(failure.clone()));
}
}
}
fn emit_callbacks(callbacks: Vec<Callback>, state: ConnectionState) {
for callback in callbacks {
if let Ok(callback) = callback.lock() {
callback(state.clone());
}
}
}
fn is_retryable_route_open_code(code: &str) -> bool {
matches!(
code,
"unknown_module" | "module_reloading" | "target_unavailable" | "module_timeout"
)
}
fn is_retryable_catalog_transport_error(err: &CallError) -> bool {
matches!(err, CallError::NotSent(_) | CallError::OutcomeUnknown(_))
}
fn is_reconnect_transient(err: &ConsumerError) -> bool {
match err {
ConsumerError::Connect { source, .. } => matches!(
source.kind(),
io::ErrorKind::ConnectionRefused
| io::ErrorKind::ConnectionReset
| io::ErrorKind::ConnectionAborted
| io::ErrorKind::TimedOut
| io::ErrorKind::NotConnected
| io::ErrorKind::AddrNotAvailable
),
ConsumerError::ConnectionFile { source, .. } => match source {
ConnectionFileError::Io { source, .. } => source.kind() == io::ErrorKind::NotFound,
_ => false,
},
ConsumerError::Auth { .. } => true,
ConsumerError::NoEndpoint { .. } | ConsumerError::Closed => false,
}
}
impl From<FrameBuildError> for CallError {
fn from(err: FrameBuildError) -> Self {
Self::not_sent(err.to_string())
}
}
#[cfg(test)]
mod tests {
use super::*;
#[derive(Clone)]
struct InstrumentedWriter {
state: Arc<InstrumentedWriterState>,
fail_flush: bool,
}
#[derive(Default)]
struct InstrumentedWriterState {
bytes: Mutex<Vec<u8>>,
flushes: std::sync::atomic::AtomicUsize,
}
impl InstrumentedWriter {
fn new(fail_flush: bool) -> Self {
Self {
state: Arc::new(InstrumentedWriterState::default()),
fail_flush,
}
}
fn bytes(&self) -> Vec<u8> {
self.state
.bytes
.lock()
.unwrap_or_else(|poisoned| poisoned.into_inner())
.clone()
}
fn flush_count(&self) -> usize {
self.state.flushes.load(Ordering::SeqCst)
}
}
impl AsyncWrite for InstrumentedWriter {
fn poll_write(
self: Pin<&mut Self>,
_cx: &mut Context<'_>,
buf: &[u8],
) -> Poll<io::Result<usize>> {
self.state
.bytes
.lock()
.unwrap_or_else(|poisoned| poisoned.into_inner())
.extend_from_slice(buf);
Poll::Ready(Ok(buf.len()))
}
fn poll_flush(self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll<io::Result<()>> {
self.state.flushes.fetch_add(1, Ordering::SeqCst);
if self.fail_flush {
Poll::Ready(Err(io::Error::new(
io::ErrorKind::BrokenPipe,
"instrumented flush failure",
)))
} else {
Poll::Ready(Ok(()))
}
}
fn poll_shutdown(self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll<io::Result<()>> {
Poll::Ready(Ok(()))
}
}
fn writer_test_shared() -> Arc<Shared> {
Arc::new(Shared::new(
PathBuf::from("/tmp/does-not-exist"),
ConsumerOptions {
reconnect_backoff: RetryBackoff {
max_attempts: 1,
..RetryBackoff::default()
},
..ConsumerOptions::default()
},
))
}
#[tokio::test]
async fn writer_batches_ready_frames_into_one_flush() {
const FRAME_COUNT: usize = 8;
let shared = writer_test_shared();
let (live_writer, _live_rx) = mpsc::channel(1);
let (tx, rx) = mpsc::channel(FRAME_COUNT + 1);
let instrumented = InstrumentedWriter::new(false);
let observer = instrumented.clone();
let mut expected = Vec::with_capacity(FRAME_COUNT);
let mut keys = Vec::with_capacity(FRAME_COUNT);
{
let mut inner = shared.lock_inner();
inner.writer = Some(live_writer);
for index in 0..FRAME_COUNT {
let corr = index as u64 + 1;
let frame = response_frame(7, 1, corr, vec![index as u8; index + 1]);
let key = PendingKey {
generation: 1,
channel: 7,
epoch: 1,
corr,
};
let (pending_tx, _pending_rx) = oneshot::channel();
inner
.pending
.insert(key, PendingEntry::unary(pending_tx, false, None));
expected.push(frame.clone());
keys.push(key);
tx.try_send(WriteCommand {
frame,
pending: Some(key),
})
.expect("the burst should fit in the writer queue");
if index == 0 {
tx.try_send(WriteCommand {
frame: response_frame(7, 1, 999, b"skip".to_vec()),
pending: Some(PendingKey {
generation: 1,
channel: 7,
epoch: 1,
corr: 999,
}),
})
.expect("the skipped command should fit in the writer queue");
}
}
}
drop(tx);
writer_loop(Arc::clone(&shared), instrumented, rx, 1).await;
{
let inner = shared.lock_inner();
for key in keys {
assert!(
inner.pending.get(&key).is_some_and(|entry| entry.accepted),
"every written command must be marked accepted"
);
}
}
let mut wire = std::io::Cursor::new(observer.bytes());
for expected_frame in expected {
let actual = read_frame(&mut wire)
.await
.expect("the emitted frame should decode")
.expect("the emitted frame should be present");
assert_eq!(actual, expected_frame);
}
assert!(
read_frame(&mut wire)
.await
.expect("the end of the emitted burst should be clean")
.is_none(),
"the writer must not emit extra frames"
);
let flush_count = observer.flush_count();
assert_eq!(
flush_count, 1,
"a ready burst must be coalesced into one flush"
);
shared.close_sync("test complete");
}
#[tokio::test]
async fn writer_flush_failure_drops_generation_and_preserves_acceptance_classification() {
let shared = writer_test_shared();
let (live_writer, _live_rx) = mpsc::channel(1);
let accepted_key = PendingKey {
generation: 1,
channel: 3,
epoch: 1,
corr: 1,
};
let not_sent_key = PendingKey {
corr: 2,
..accepted_key
};
let (accepted_tx, accepted_rx) = oneshot::channel();
let (not_sent_tx, not_sent_rx) = oneshot::channel();
{
let mut inner = shared.lock_inner();
inner.writer = Some(live_writer);
inner
.pending
.insert(accepted_key, PendingEntry::unary(accepted_tx, false, None));
inner
.pending
.insert(not_sent_key, PendingEntry::unary(not_sent_tx, false, None));
}
let (tx, rx) = mpsc::channel(1);
tx.send(WriteCommand {
frame: response_frame(3, 1, accepted_key.corr, b"accepted".to_vec()),
pending: Some(accepted_key),
})
.await
.unwrap();
drop(tx);
writer_loop(Arc::clone(&shared), InstrumentedWriter::new(true), rx, 1).await;
assert!(
shared.lock_inner().writer.is_none(),
"a flush failure must drop the active generation"
);
let accepted_error = accepted_rx
.await
.expect("the accepted request should be settled")
.into_call_result()
.unwrap_err();
assert!(matches!(accepted_error, CallError::OutcomeUnknown(_)));
let not_sent_error = not_sent_rx
.await
.expect("the unwritten request should be settled")
.into_call_result()
.unwrap_err();
assert!(matches!(not_sent_error, CallError::NotSent(_)));
shared.close_sync("test complete");
}
#[test]
fn reconnect_classifier_treats_auth_failure_as_transient() {
let auth = ConsumerError::Auth {
path: PathBuf::from("/tmp/subc-connection.json"),
endpoint: "127.0.0.1:8757".to_string(),
source: subc_transport::AuthError::InvalidServerProof,
};
assert!(is_reconnect_transient(&auth), "rotation race must retry");
let absent = ConsumerError::ConnectionFile {
path: PathBuf::from("/tmp/subc-connection.json"),
source: ConnectionFileError::Io {
op: "read",
path: PathBuf::from("/tmp/subc-connection.json"),
source: io::Error::new(io::ErrorKind::NotFound, "gone"),
},
};
assert!(is_reconnect_transient(&absent));
}
#[tokio::test]
async fn newer_drop_supersedes_reconnect_and_ignores_stale_completion() {
let shared = Arc::new(Shared::new(
PathBuf::from("/tmp/does-not-exist"),
ConsumerOptions::default(),
));
let stale_task = tokio::spawn(std::future::pending::<()>());
{
let mut inner = shared.lock_inner();
inner.generation = 2;
inner.reconnect = ReconnectState::Background {
generation: 1,
task: stale_task,
};
}
assert!(shared.spawn_reconnect(2));
assert!(matches!(
&shared.lock_inner().reconnect,
ReconnectState::Background { generation, .. } if *generation == 2
));
shared.finish_background_reconnect(1);
assert!(matches!(
&shared.lock_inner().reconnect,
ReconnectState::Background { generation, .. } if *generation == 2
));
shared.close_sync("test complete");
}
#[test]
fn retryable_route_open_codes_are_code_specific() {
for code in [
"unknown_module",
"module_reloading",
"target_unavailable",
"module_timeout",
] {
assert!(is_retryable_route_open_code(code), "{code} should retry");
}
assert!(!is_retryable_route_open_code("invalid_project_root"));
assert!(!is_retryable_route_open_code("route_rejected"));
}
#[tokio::test]
async fn close_route_flips_inflight_opening_so_a_racing_open_discards() {
let shared = Arc::new(Shared::new(
PathBuf::from("/tmp/does-not-exist"),
ConsumerOptions::default(),
));
let key = RouteKey::new(
&RouteTarget::ToolProvider {
module_id: "m".into(),
},
&BindIdentity {
project_root: PathBuf::from("/tmp/p"),
harness: "h".into(),
session: "s".into(),
},
None,
None,
);
shared.lock_inner().openings.insert(
key.clone(),
Opening {
waiters: Vec::new(),
closed: false,
},
);
shared
.close_route(&key, &CloseRouteOptions::default())
.await;
assert!(
shared
.lock_inner()
.openings
.get(&key)
.is_some_and(|o| o.closed),
"close_route must flip the in-flight opening's closed flag (close-beats-reopen)"
);
let absent = RouteKey::new(
&RouteTarget::ToolProvider {
module_id: "absent".into(),
},
&BindIdentity {
project_root: PathBuf::from("/tmp/p"),
harness: "h".into(),
session: "s".into(),
},
None,
None,
);
shared
.close_route(&absent, &CloseRouteOptions::default())
.await;
}
#[test]
fn route_key_is_structured() {
let target = RouteTarget::InternalService {
module_id: "a\0b".into(),
service_id: "svc".into(),
};
let identity = BindIdentity {
project_root: PathBuf::from("/tmp/project"),
harness: "h".into(),
session: "s".into(),
};
let key = RouteKey::new(&target, &identity, None, None);
assert_eq!(key.project_root, PathBuf::from("/tmp/project"));
assert!(matches!(key.target, RouteTargetKey::InternalService { .. }));
}
#[tokio::test]
async fn route_channel_index_tracks_lookup_close_and_generation_drop() {
let shared = writer_test_shared();
let (writer, _rx) = mpsc::channel(32);
let mut expected = Vec::new();
{
let mut inner = shared.lock_inner();
inner.writer = Some(writer);
for channel in 1..=8 {
let key = RouteKey::new(
&RouteTarget::ToolProvider {
module_id: format!("module-{channel}"),
},
&BindIdentity {
project_root: PathBuf::from("/tmp/project"),
harness: "test".into(),
session: format!("session-{channel}"),
},
None,
None,
);
let route = RouteState {
handle: RouteHandle::new(channel, channel.into(), 1),
sem: Arc::new(Semaphore::new(DEFAULT_ROUTE_WINDOW)),
};
inner.cache_route(key.clone(), route.clone());
expected.push((key, route));
}
}
for (key, expected_route) in &expected {
let resolved = shared
.route_state(expected_route.handle)
.expect("an indexed route handle should resolve");
assert_eq!(resolved.handle, expected_route.handle);
assert!(Arc::ptr_eq(&resolved.sem, &expected_route.sem));
let inner = shared.lock_inner();
assert_eq!(
inner.route_by_channel.get(&expected_route.handle.channel),
Some(key)
);
assert_eq!(
inner.route_epochs.get(&expected_route.handle.channel),
Some(&expected_route.handle)
);
assert!(inner
.routes
.get(key)
.is_some_and(|route| route.handle == expected_route.handle));
}
let (closed_key, closed_route) = &expected[3];
shared
.close_route(closed_key, &CloseRouteOptions::default())
.await;
{
let inner = shared.lock_inner();
assert!(!inner.routes.contains_key(closed_key));
assert!(!inner
.route_by_channel
.contains_key(&closed_route.handle.channel));
assert!(!inner
.route_epochs
.contains_key(&closed_route.handle.channel));
}
assert!(matches!(
shared.route_state(closed_route.handle),
Err(CallError::StaleRouteHandle(handle)) if handle == closed_route.handle
));
shared.handle_generation_drop(1, "test generation dropped".into());
{
let inner = shared.lock_inner();
assert!(inner.routes.is_empty());
assert!(inner.route_by_channel.is_empty());
assert!(inner.route_epochs.is_empty());
}
shared.close_sync("test complete");
}
#[tokio::test]
async fn stale_push_is_not_delivered_after_connection_generation_changes() {
let shared = writer_test_shared();
let old_handle = RouteHandle::new(7, 1, 1);
let (writer, _writer_rx) = mpsc::channel(1);
{
let mut inner = shared.lock_inner();
inner.writer = Some(writer);
inner.route_epochs.insert(old_handle.channel, old_handle);
}
let mut pushes = shared
.register_push_events(old_handle)
.expect("the old live route should accept a receiver");
{
let mut inner = shared.lock_inner();
inner.generation = 2;
inner.close_routes();
inner.route_epochs.clear();
}
shared.route_push(old_handle, b"stale".to_vec());
assert!(
pushes.recv().await.is_none(),
"connection teardown must end the old receiver before a stale Push can arrive"
);
shared.close_sync("test complete");
}
#[test]
fn route_key_canonicalizes_consumer_capabilities() {
let target = RouteTarget::ToolProvider {
module_id: "aft".into(),
};
let identity = BindIdentity {
project_root: PathBuf::from("/tmp/project"),
harness: "h".into(),
session: "s".into(),
};
let left = RouteKey::new(
&target,
&identity,
None,
Some(&["sampling".to_string(), "elicitation".to_string()]),
);
let right = RouteKey::new(
&target,
&identity,
None,
Some(&[
"elicitation".to_string(),
"sampling".to_string(),
"sampling".to_string(),
]),
);
assert_eq!(left, right);
}
#[test]
fn drain_pending_channel_can_skip_subscriptions() {
let mut pending = HashMap::new();
let generation = 7;
let channel = 11;
let handle = RouteHandle::new(channel, 3, generation);
let unary_key = PendingKey {
generation,
channel,
epoch: handle.epoch,
corr: 1,
};
let subscription_key = PendingKey {
generation,
channel,
epoch: handle.epoch,
corr: 2,
};
let (unary_tx, _unary_rx) = oneshot::channel();
pending.insert(unary_key, PendingEntry::unary(unary_tx, false, None));
let (events_tx, _events_rx) = mpsc::channel(1);
let (closed_tx, _closed_rx) = oneshot::channel();
let permit = Arc::new(Semaphore::new(1))
.try_acquire_owned()
.expect("test semaphore permit should be available");
pending.insert(
subscription_key,
PendingEntry::subscription(events_tx, closed_tx, permit, Priority::Interactive),
);
let drained = drain_pending_handle(&mut pending, handle, false);
assert_eq!(drained.len(), 1);
assert!(pending.contains_key(&subscription_key));
let drained = drain_pending_handle(&mut pending, handle, true);
assert_eq!(drained.len(), 1);
assert!(pending.is_empty());
}
fn response_frame(channel: u16, epoch: u32, corr: u64, body: Vec<u8>) -> Frame {
Frame::build(
FrameType::Response,
Flags::new(false, Priority::Interactive, false),
channel,
epoch,
corr,
body,
)
.unwrap()
}
#[tokio::test]
async fn stale_epoch_ingress_drops_without_settling_matching_corr() {
let shared = Arc::new(Shared::new(
PathBuf::from("/tmp/does-not-exist"),
ConsumerOptions::default(),
));
let (writer, _rx) = mpsc::channel(4);
let current = RouteHandle::new(9, 2, 1);
let stale_key = PendingKey {
generation: 1,
channel: 9,
epoch: 1,
corr: 77,
};
let key = PendingKey {
generation: 1,
channel: 9,
epoch: 2,
corr: 77,
};
let (stale_tx, mut stale_response) = oneshot::channel();
let (tx, mut response) = oneshot::channel();
{
let mut inner = shared.lock_inner();
inner.writer = Some(writer);
inner.route_epochs.insert(9, current);
inner
.pending
.insert(stale_key, PendingEntry::unary(stale_tx, false, None));
inner
.pending
.insert(key, PendingEntry::unary(tx, false, None));
}
assert!(dispatch_frame(&shared, 1, response_frame(9, 1, 77, b"stale".to_vec())).await);
assert!(matches!(
response.try_recv(),
Err(oneshot::error::TryRecvError::Empty)
));
assert!(matches!(
stale_response.try_recv(),
Err(oneshot::error::TryRecvError::Empty)
));
assert!(shared.lock_inner().pending.contains_key(&stale_key));
assert!(shared.lock_inner().pending.contains_key(&key));
assert_eq!(shared.lock_inner().dropped_route_frames, 1);
assert!(dispatch_frame(&shared, 1, response_frame(9, 2, 77, b"current".to_vec())).await);
let PendingResult::Terminal(PendingTerminal::Response { body, .. }) =
response.await.unwrap()
else {
panic!("current epoch must settle its own pending request");
};
assert_eq!(body, b"current");
assert!(shared.lock_inner().pending.contains_key(&stale_key));
}
#[tokio::test]
async fn route_poll_response_must_echo_expected_handle_before_settling() {
let shared = Arc::new(Shared::new(
PathBuf::from("/tmp/does-not-exist"),
ConsumerOptions::default(),
));
let (writer, _rx) = mpsc::channel(4);
let handle = RouteHandle::new(3, 9, 1);
let key = PendingKey {
generation: 1,
channel: 0,
epoch: 0,
corr: 88,
};
let (tx, mut response) = oneshot::channel();
{
let mut inner = shared.lock_inner();
inner.writer = Some(writer);
inner.route_epochs.insert(handle.channel, handle);
inner
.pending
.insert(key, PendingEntry::unary(tx, false, Some(handle)));
}
let wrong = serde_json::to_vec(&ClientControlResponse::RoutePoll {
route_channel: handle.channel,
route_epoch: handle.epoch + 1,
status: Some("wrong".to_string()),
live: Some(true),
})
.unwrap();
assert!(dispatch_frame(&shared, 1, response_frame(0, 0, key.corr, wrong)).await);
assert!(matches!(
response.try_recv(),
Err(oneshot::error::TryRecvError::Empty)
));
assert!(shared.lock_inner().pending.contains_key(&key));
let correct = serde_json::to_vec(&ClientControlResponse::RoutePoll {
route_channel: handle.channel,
route_epoch: handle.epoch,
status: Some("ready".to_string()),
live: Some(true),
})
.unwrap();
assert!(dispatch_frame(&shared, 1, response_frame(0, 0, key.corr, correct)).await);
assert!(matches!(
response.await.unwrap(),
PendingResult::Terminal(PendingTerminal::Response { .. })
));
}
#[tokio::test]
async fn stale_connection_handle_emits_no_request_cancel_or_goodbye() {
let shared = Arc::new(Shared::new(
PathBuf::from("/tmp/does-not-exist"),
ConsumerOptions::default(),
));
let (writer, mut rx) = mpsc::channel(4);
let stale = RouteHandle::new(4, 1, 1);
let current = RouteHandle::new(4, 1, 2);
{
let mut inner = shared.lock_inner();
inner.generation = 2;
inner.writer = Some(writer);
inner.route_epochs.insert(4, current);
}
let err = shared
.send_request(RequestSend {
expected_handle: Some(stale),
channel: stale.channel,
epoch: stale.epoch,
body: b"request".to_vec(),
priority: Priority::Interactive,
admission_class: AdmissionClass::Normal,
deadline: Instant::now() + Duration::from_millis(10),
retain_late_route_open: false,
})
.await
.unwrap_err();
assert!(matches!(err, CallError::StaleRouteHandle(handle) if handle == stale));
shared.send_cancel(stale, 8, Priority::Interactive);
assert!(!shared.send_route_goodbye(stale, false));
assert!(
rx.try_recv().is_err(),
"stale operations must not queue frames"
);
}
#[tokio::test]
async fn late_route_open_queues_goodbye_and_full_queue_closes_connection() {
let shared = Arc::new(Shared::new(
PathBuf::from("/tmp/does-not-exist"),
ConsumerOptions::default(),
));
let (writer, mut rx) = mpsc::channel(2);
let key = PendingKey {
generation: 1,
channel: 0,
epoch: 0,
corr: 41,
};
let (tx, response) = oneshot::channel();
drop(response);
{
let mut inner = shared.lock_inner();
inner.writer = Some(writer);
inner
.pending
.insert(key, PendingEntry::unary(tx, true, None));
}
let body = serde_json::to_vec(&ClientControlResponse::RouteOpen {
route_channel: 12,
route_epoch: 7,
})
.unwrap();
assert!(dispatch_frame(&shared, 1, response_frame(0, 0, 41, body)).await);
let cleanup = rx.recv().await.unwrap().frame;
assert_eq!(cleanup.header.ty, FrameType::Goodbye);
assert_eq!((cleanup.header.channel, cleanup.header.epoch), (12, 7));
let shared = Arc::new(Shared::new(
PathBuf::from("/tmp/does-not-exist"),
ConsumerOptions {
reconnect_backoff: RetryBackoff {
max_attempts: 1,
..RetryBackoff::default()
},
..ConsumerOptions::default()
},
));
let (writer, _rx) = mpsc::channel(1);
let filler = response_frame(0, 0, 1, Vec::new());
writer
.try_send(WriteCommand {
frame: filler,
pending: None,
})
.unwrap();
let key = PendingKey {
generation: 1,
channel: 0,
epoch: 0,
corr: 42,
};
let (tx, response) = oneshot::channel();
drop(response);
{
let mut inner = shared.lock_inner();
inner.writer = Some(writer);
inner
.pending
.insert(key, PendingEntry::unary(tx, true, None));
}
let body = serde_json::to_vec(&ClientControlResponse::RouteOpen {
route_channel: 13,
route_epoch: 8,
})
.unwrap();
assert!(dispatch_frame(&shared, 1, response_frame(0, 0, 42, body)).await);
assert!(shared.lock_inner().writer.is_none());
}
#[tokio::test]
async fn correlation_allocator_emits_max_once_then_closes_without_reuse() {
let shared = Arc::new(Shared::new(
PathBuf::from("/tmp/does-not-exist"),
ConsumerOptions {
reconnect_backoff: RetryBackoff {
max_attempts: 1,
..RetryBackoff::default()
},
..ConsumerOptions::default()
},
));
let (writer, mut rx) = mpsc::channel(4);
{
let mut inner = shared.lock_inner();
inner.writer = Some(writer);
inner.next_corr = Some(u64::MAX);
}
let request_shared = Arc::clone(&shared);
let request = tokio::spawn(async move {
request_shared
.send_request(RequestSend {
expected_handle: None,
channel: 0,
epoch: 0,
body: Vec::new(),
priority: Priority::Interactive,
admission_class: AdmissionClass::Normal,
deadline: Instant::now() + Duration::from_secs(1),
retain_late_route_open: false,
})
.await
});
let command = rx.recv().await.unwrap();
assert_eq!(command.frame.header.corr, u64::MAX);
assert!(dispatch_frame(&shared, 1, response_frame(0, 0, u64::MAX, Vec::new()),).await);
assert!(request.await.unwrap().is_ok());
let exhausted = shared
.send_request(RequestSend {
expected_handle: None,
channel: 0,
epoch: 0,
body: Vec::new(),
priority: Priority::Interactive,
admission_class: AdmissionClass::Normal,
deadline: Instant::now() + Duration::from_millis(10),
retain_late_route_open: false,
})
.await
.unwrap_err();
assert!(matches!(exhausted, CallError::NotSent(_)));
assert!(rx.try_recv().is_err());
assert!(shared.lock_inner().writer.is_none());
}
#[tokio::test]
async fn managed_call_deadline_bounds_flow_control_wait() {
let shared = Arc::new(Shared::new(
PathBuf::from("/tmp/does-not-exist"),
ConsumerOptions::default(),
));
let (writer, mut rx) = mpsc::channel(4);
let target = RouteTarget::ToolProvider {
module_id: "flow-controlled".to_string(),
};
let identity = BindIdentity {
project_root: PathBuf::from("/tmp/project"),
harness: "test".to_string(),
session: "deadline".to_string(),
};
let consumer_identity = Some(ConsumerIdentity {
module_id: "caller".to_string(),
launch_nonce: "nonce".to_string(),
});
let first_opts = CallOptions {
timeout: Duration::from_secs(1),
consumer_identity: consumer_identity.clone(),
..CallOptions::default()
};
let second_opts = CallOptions {
timeout: Duration::from_millis(25),
consumer_identity,
..CallOptions::default()
};
let key = RouteKey::new(
&target,
&identity,
first_opts.consumer_identity.as_ref(),
None,
);
let handle = RouteHandle::new(5, 3, 1);
{
let mut inner = shared.lock_inner();
inner.writer = Some(writer);
inner.cache_route(
key,
RouteState {
handle,
sem: Arc::new(Semaphore::new(1)),
},
);
}
let first = tokio::spawn({
let consumer = SubcConsumer {
shared: Arc::clone(&shared),
};
let target = target.clone();
let identity = identity.clone();
async move {
consumer
.call(target, identity, b"first".to_vec(), first_opts)
.await
}
});
let first_frame = rx
.recv()
.await
.expect("the first request should enter the fake daemon queue");
assert_eq!(first_frame.frame.header.ty, FrameType::Request);
assert_eq!(first_frame.frame.body, b"first");
assert!(shared.mark_pending_accepted(
first_frame
.pending
.expect("request commands retain their pending key"),
));
let consumer = SubcConsumer {
shared: Arc::clone(&shared),
};
let result = tokio::time::timeout(
Duration::from_millis(250),
consumer.call(target, identity, b"second".to_vec(), second_opts),
)
.await
.expect("a flow-controlled call must finish at its own deadline")
.unwrap_err();
assert!(matches!(result, CallError::NotSent(_)));
assert!(
rx.try_recv().is_err(),
"the timed-out second request must not reach the fake daemon"
);
first.abort();
let _ = first.await;
}
#[tokio::test]
async fn admitted_route_open_emits_one_frame_without_retrying_daemon_errors() {
let shared = Arc::new(Shared::new(
PathBuf::from("/tmp/does-not-exist"),
ConsumerOptions {
call_timeout: Duration::from_secs(1),
..ConsumerOptions::default()
},
));
let (writer, mut rx) = mpsc::channel(4);
{
let mut inner = shared.lock_inner();
inner.writer = Some(writer);
}
let consumer = SubcConsumer {
shared: Arc::clone(&shared),
};
let target = RouteTarget::ToolProvider {
module_id: "admitted-target".to_string(),
};
let identity = BindIdentity {
project_root: PathBuf::from("/tmp/project"),
harness: "test".to_string(),
session: "admitted".to_string(),
};
let task = tokio::spawn(async move {
consumer
.open_route_with_admission_facts(
target,
identity,
serde_json::json!({"schema": 1, "verified_class": "member"}),
)
.await
});
let command = rx.recv().await.expect("one route.open must be queued");
let request: ClientControlRequest = serde_json::from_slice(&command.frame.body).unwrap();
let ClientControlRequest::RouteOpen {
admission_facts, ..
} = request
else {
panic!("expected route.open")
};
assert_eq!(
admission_facts,
Some(serde_json::json!({"schema": 1, "verified_class": "member"}))
);
let error_body = serde_json::to_vec(&ErrorBody {
code: "admission_facts_not_permitted".to_string(),
message: "not permitted".to_string(),
})
.unwrap();
assert!(
dispatch_frame(
&shared,
1,
Frame::build(
FrameType::Error,
Flags::new(false, Priority::Interactive, false),
0,
0,
command.frame.header.corr,
error_body,
)
.unwrap(),
)
.await
);
let result = task.await.unwrap();
assert!(matches!(result, Err(CallError::NotSent(_))));
assert!(rx.try_recv().is_err(), "one-shot route.open must not retry");
}
#[tokio::test]
async fn route_open_waiter_deadline_is_not_sent_without_writing() {
let shared = Arc::new(Shared::new(
PathBuf::from("/tmp/does-not-exist"),
ConsumerOptions::default(),
));
let (writer, mut rx) = mpsc::channel(4);
let target = RouteTarget::ToolProvider {
module_id: "single-flight".to_string(),
};
let identity = BindIdentity {
project_root: PathBuf::from("/tmp/project"),
harness: "test".to_string(),
session: "route-open".to_string(),
};
let opts = CallOptions {
timeout: Duration::from_millis(25),
consumer_identity: Some(ConsumerIdentity {
module_id: "caller".to_string(),
launch_nonce: "nonce".to_string(),
}),
..CallOptions::default()
};
let key = RouteKey::new(&target, &identity, opts.consumer_identity.as_ref(), None);
{
let mut inner = shared.lock_inner();
inner.writer = Some(writer);
inner.openings.insert(
key,
Opening {
waiters: Vec::new(),
closed: false,
},
);
}
let consumer = SubcConsumer { shared };
let result = tokio::time::timeout(
Duration::from_millis(250),
consumer.open_route(target, identity, opts),
)
.await
.expect("a route.open waiter must finish at its own deadline")
.unwrap_err();
assert!(matches!(result, CallError::NotSent(_)));
assert!(
rx.try_recv().is_err(),
"a timed-out route.open waiter must not write a control frame"
);
}
#[tokio::test]
async fn route_poll_deadline_bounds_writer_capacity() {
let shared = Arc::new(Shared::new(
PathBuf::from("/tmp/does-not-exist"),
ConsumerOptions::default(),
));
let (writer, mut rx) = mpsc::channel(1);
let handle = RouteHandle::new(8, 4, 1);
writer
.try_send(WriteCommand {
frame: response_frame(0, 0, 99, Vec::new()),
pending: None,
})
.expect("the fake daemon queue should accept its filler frame");
{
let mut inner = shared.lock_inner();
inner.writer = Some(writer);
inner.route_epochs.insert(handle.channel, handle);
}
let consumer = SubcConsumer { shared };
let result = tokio::time::timeout(
Duration::from_millis(250),
consumer.poll_route(&handle, PollKind::Liveness, Duration::from_millis(25)),
)
.await
.expect("a control request must finish at its own deadline")
.unwrap_err();
assert!(matches!(result, CallError::NotSent(_)));
assert_eq!(
rx.recv()
.await
.expect("the filler must still be the only queued frame")
.frame
.header
.corr,
99
);
assert!(
rx.try_recv().is_err(),
"the timed-out control request must not reach the fake daemon"
);
}
#[tokio::test]
async fn subscription_deadline_bounds_writer_capacity() {
let shared = Arc::new(Shared::new(
PathBuf::from("/tmp/does-not-exist"),
ConsumerOptions::default(),
));
let (writer, mut rx) = mpsc::channel(1);
let handle = RouteHandle::new(9, 2, 1);
let route_sem = Arc::new(Semaphore::new(1));
writer
.try_send(WriteCommand {
frame: response_frame(0, 0, 100, Vec::new()),
pending: None,
})
.expect("the fake daemon queue should accept its filler frame");
{
let mut inner = shared.lock_inner();
inner.writer = Some(writer);
inner.cache_route(
RouteKey::new(
&RouteTarget::ToolProvider {
module_id: "subscriptions".to_string(),
},
&BindIdentity {
project_root: PathBuf::from("/tmp/project"),
harness: "test".to_string(),
session: "subscription".to_string(),
},
None,
None,
),
RouteState {
handle,
sem: Arc::clone(&route_sem),
},
);
}
let consumer = SubcConsumer { shared };
let result = tokio::time::timeout(
Duration::from_millis(250),
consumer.subscribe_route(
&handle,
b"subscribe".to_vec(),
SubscribeOptions {
route_open_timeout: Duration::from_millis(25),
..SubscribeOptions::default()
},
),
)
.await
.expect("a subscription must finish at its route-open deadline");
let result = match result {
Ok(_) => panic!("a subscription blocked before writing must time out"),
Err(err) => err,
};
assert!(matches!(result, CallError::NotSent(_)));
assert_eq!(
rx.recv()
.await
.expect("the filler must still be the only queued frame")
.frame
.header
.corr,
100
);
assert!(
rx.try_recv().is_err(),
"the timed-out subscription must not reach the fake daemon"
);
assert!(
route_sem.try_acquire().is_ok(),
"a pre-write subscription timeout must release its route credit"
);
}
#[test]
fn catalog_list_deserializes_golden_reply_and_ignores_unknown_fields() {
let mut reply: serde_json::Value = serde_json::from_str(include_str!(
"../../subc-control/tests/golden/client_control_response_catalog_list.json"
))
.expect("the catalog.list golden reply must be valid JSON");
reply["future_top_level"] = serde_json::json!(true);
reply["modules"][0]["future_module_field"] = serde_json::json!("ignored");
let catalog: CatalogList =
serde_json::from_value(reply).expect("catalog.list should tolerate additive fields");
assert_eq!(catalog.generation, 7);
assert_eq!(catalog.modules.len(), 1);
assert!(catalog.subc_ops.iter().any(|op| op == "catalog.list"));
let tools = catalog.modules[0]
.roles
.iter()
.find_map(|role| match role {
subc_protocol::manifest::ProviderRole::ToolProvider { tools, .. } => Some(tools),
_ => None,
})
.expect("the golden module must advertise a tool_provider role");
let tool = tools
.first()
.expect("the golden tool_provider role must advertise a tool");
assert!(!tool.name.is_empty());
assert_eq!(
tool.schema.get("type").and_then(serde_json::Value::as_str),
Some("object")
);
assert!(matches!(
tool.execution_mode,
subc_protocol::manifest::ExecutionMode::Pure
));
}
#[tokio::test]
async fn catalog_list_sends_an_unfiltered_channel_zero_request() {
let shared = Arc::new(Shared::new(
PathBuf::from("/tmp/does-not-exist"),
ConsumerOptions::default(),
));
let (writer, mut rx) = mpsc::channel(1);
shared.lock_inner().writer = Some(writer);
let consumer = SubcConsumer {
shared: Arc::clone(&shared),
};
let request = tokio::spawn(async move { consumer.catalog_list().await });
let command = rx
.recv()
.await
.expect("catalog.list must queue a channel-0 request");
assert_eq!(command.frame.header.channel, 0);
let body: serde_json::Value = serde_json::from_slice(&command.frame.body).unwrap();
assert_eq!(body["op"], "catalog.list");
assert!(
body.get("module_id").is_none(),
"catalog.list must request the complete catalog without a module filter"
);
let response = serde_json::to_vec(&ClientControlResponse::CatalogList {
generation: 9,
modules: Vec::new(),
subc_ops: vec!["catalog.list".to_string()],
})
.unwrap();
assert!(
dispatch_frame(
&shared,
1,
response_frame(0, 0, command.frame.header.corr, response),
)
.await
);
let catalog = request.await.unwrap().unwrap();
assert_eq!(catalog.generation, 9);
assert!(catalog.modules.is_empty());
}
#[tokio::test]
async fn catalog_list_deadline_is_not_sent_when_reconnection_stays_down() {
let shared = Arc::new(Shared::new(
PathBuf::from("/tmp/subc-client-rs-catalog-list-unavailable"),
ConsumerOptions {
call_timeout: Duration::from_millis(25),
reconnect_backoff: RetryBackoff {
base: Duration::from_millis(1),
cap: Duration::from_millis(1),
max_attempts: 100,
},
..ConsumerOptions::default()
},
));
let consumer = SubcConsumer { shared };
let result = tokio::time::timeout(Duration::from_millis(250), consumer.catalog_list())
.await
.expect("catalog.list must finish at its configured deadline")
.unwrap_err();
assert!(matches!(result, CallError::NotSent(_)));
}
}