use core::cell::Cell;
use core::fmt;
use core::future::Future;
use core::marker::PhantomData;
use core::{any, mem};
use alloc::boxed::Box;
use alloc::string::{String, ToString};
use alloc::sync::Arc;
use alloc::vec::Vec;
use std::collections::HashMap;
use std::collections::hash_map::Entry;
use std::sync::Mutex;
use std::sync::atomic::{AtomicBool, AtomicU8, AtomicU32, Ordering};
use bytes::Bytes;
use rand::prelude::*;
use rand::rngs::SmallRng;
use slab::Slab;
use tokio::sync::{mpsc, oneshot, watch};
use tokio::time::{Duration, Instant};
use crate::api::{self, ChannelId, DecodeBody, Event, Format, MessageId};
use crate::format;
const INITIAL_TIMEOUT: Duration = Duration::from_millis(250);
const MAX_TIMEOUT: Duration = Duration::from_millis(4000);
const MAX_FUZZ: u64 = 50;
const DEFAULT_SEED: u64 = 0xdeadbeef;
#[non_exhaustive]
pub struct EmptyBody;
#[non_exhaustive]
pub struct EmptyCallback;
#[cfg_attr(not(feature = "tungstenite029"), allow(dead_code))]
pub(crate) enum Message {
Text,
Binary(Bytes),
Ping,
Pong,
Close,
}
pub(crate) mod sealed_socket {
pub trait Sealed {}
}
pub(crate) trait SocketImpl
where
Self: 'static + Send + Sized + self::sealed_socket::Sealed,
{
#[doc(hidden)]
type Error;
#[doc(hidden)]
fn recv(&mut self) -> impl Future<Output = Option<Result<Message, Self::Error>>> + Send + '_;
#[doc(hidden)]
fn send(&mut self, data: &[u8]) -> impl Future<Output = Result<(), Self::Error>> + Send + '_;
#[doc(hidden)]
fn close(&mut self) -> impl Future<Output = Result<(), Self::Error>> + Send + '_;
}
pub(crate) mod sealed_client {
pub trait Sealed {}
}
pub trait ClientImpl
where
Self: 'static + Copy + Sized + self::sealed_client::Sealed,
{
#[doc(hidden)]
type Error: 'static + Send + Sync + core::error::Error;
#[doc(hidden)]
#[allow(private_bounds)]
type Socket: SocketImpl<Error = Self::Error>;
#[doc(hidden)]
fn connect(url: &str) -> impl Future<Output = Result<Self::Socket, Self::Error>> + Send;
}
pub fn connect<T>(url: impl AsRef<str>) -> ServiceBuilder<T, EmptyCallback>
where
T: ClientImpl,
{
ServiceBuilder {
url: url.as_ref().to_string(),
on_error: EmptyCallback,
reconnect: true,
seed: DEFAULT_SEED,
format: Format::DEFAULT,
_marker: PhantomData,
}
}
#[derive(Debug, PartialEq, Eq, Clone, Copy)]
#[non_exhaustive]
pub enum State {
Open,
Closed,
}
impl State {
#[inline]
pub fn is_open(&self) -> bool {
matches!(self, Self::Open)
}
}
pub trait Callback<I>
where
Self: 'static + Send + Sync,
{
fn call(&self, input: I);
}
impl<I> Callback<I> for EmptyCallback {
#[inline]
fn call(&self, _: I) {}
}
impl<F, I> Callback<I> for F
where
F: 'static + Send + Sync + Fn(I),
{
#[inline]
fn call(&self, input: I) {
self(input)
}
}
#[derive(Debug)]
pub struct Error {
kind: ErrorKind,
}
impl Error {
#[inline]
const fn new(kind: ErrorKind) -> Self {
Self { kind }
}
#[inline]
pub fn is_empty_packet(&self) -> bool {
matches!(self.kind, ErrorKind::EmptyPacket)
}
#[inline]
pub fn is_not_connected(&self) -> bool {
matches!(self.kind, ErrorKind::NotConnected)
}
#[inline]
pub fn as_server_error(&self) -> Option<&str> {
match &self.kind {
ErrorKind::Server(message) => Some(message),
_ => None,
}
}
#[inline]
pub fn message(message: impl fmt::Display) -> Self {
Self::new(ErrorKind::Message(message.to_string()))
}
#[inline]
fn server(message: impl fmt::Display) -> Self {
Self::new(ErrorKind::Server(message.to_string()))
}
#[inline]
fn transport<E>(error: E) -> Self
where
E: 'static + Send + Sync + core::error::Error,
{
Self::new(ErrorKind::Transport(Box::new(error)))
}
#[inline]
fn decode_response_header(error: format::Error) -> Self {
Self::new(ErrorKind::DecodeResponseHeader(error))
}
#[inline]
fn decode_error_message(error: format::Error) -> Self {
Self::new(ErrorKind::DecodeErrorMessage(error))
}
#[inline]
fn decode_packet(error: format::Error) -> Self {
Self::new(ErrorKind::DecodePacket(error))
}
#[inline]
fn encoding_header(error: format::Error) -> Self {
Self::new(ErrorKind::EncodingHeader(error))
}
#[inline]
fn encoding_body(error: format::Error) -> Self {
Self::new(ErrorKind::EncodingBody(error))
}
}
#[derive(Debug)]
enum ErrorKind {
EmptyPacket,
NotConnected,
Message(String),
Server(String),
Transport(Box<dyn core::error::Error + Send + Sync>),
DecodeResponseHeader(format::Error),
DecodeErrorMessage(format::Error),
DecodePacket(format::Error),
EncodingHeader(format::Error),
EncodingBody(format::Error),
}
impl fmt::Display for Error {
#[inline]
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match &self.kind {
ErrorKind::EmptyPacket => write!(f, "Packet is empty"),
ErrorKind::NotConnected => write!(f, "Client is not connected"),
ErrorKind::Message(message) => write!(f, "{message}"),
ErrorKind::Server(message) => write!(f, "Server error: {message}"),
ErrorKind::Transport(..) => write!(f, "Error in underlying transport"),
ErrorKind::DecodeResponseHeader(..) => {
write!(f, "Encoding error when decoding response header")
}
ErrorKind::DecodeErrorMessage(..) => {
write!(f, "Encoding error when decoding error response")
}
ErrorKind::DecodePacket(..) => write!(f, "Encoding error when decoding packet"),
ErrorKind::EncodingHeader(..) => write!(f, "Encoding error when encoding header"),
ErrorKind::EncodingBody(..) => write!(f, "Encoding error when encoding body"),
}
}
}
impl core::error::Error for Error {
#[inline]
fn source(&self) -> Option<&(dyn core::error::Error + 'static)> {
match &self.kind {
ErrorKind::Transport(error) => Some(&**error),
ErrorKind::DecodeResponseHeader(error) => Some(error),
ErrorKind::DecodeErrorMessage(error) => Some(error),
ErrorKind::DecodePacket(error) => Some(error),
ErrorKind::EncodingHeader(error) => Some(error),
ErrorKind::EncodingBody(error) => Some(error),
_ => None,
}
}
}
type Result<T, E = Error> = core::result::Result<T, E>;
type Broadcasts = HashMap<MessageId, Slab<mpsc::UnboundedSender<Result<RawPacket>>>>;
enum Command {
Send {
serial: u32,
data: Vec<u8>,
pending: Pending,
},
Disconnect { channel: ChannelId },
Close,
}
enum Pending {
Negotiate { format: Format },
Request {
id: MessageId,
reply: oneshot::Sender<Result<RawPacket>>,
},
Channel {
reply: oneshot::Sender<Result<ChannelId>>,
},
}
impl Pending {
#[inline]
fn error(self, error: Error) {
match self {
Pending::Negotiate { .. } => {
tracing::debug!("Format negotiation failed: {error}");
}
Pending::Request { reply, .. } => {
_ = reply.send(Err(error));
}
Pending::Channel { reply } => {
_ = reply.send(Err(error));
}
}
}
}
struct Shared {
tx: mpsc::UnboundedSender<Command>,
serial: AtomicU32,
state: watch::Sender<State>,
broadcasts: Mutex<Broadcasts>,
gone: AtomicBool,
format: AtomicU8,
}
impl Shared {
#[inline]
fn next_serial(&self) -> u32 {
self.serial.fetch_add(1, Ordering::Relaxed)
}
#[inline]
fn is_open(&self) -> bool {
self.state.borrow().is_open()
}
#[inline]
fn format(&self) -> Format {
Format::from_u8(self.format.load(Ordering::Acquire)).unwrap_or(Format::DEFAULT)
}
#[inline]
fn set_format(&self, format: Format) {
self.format.store(format.to_u8(), Ordering::Release);
}
#[inline]
fn is_gone(&self) -> bool {
self.gone.load(Ordering::Acquire)
}
fn set_gone(&self) {
self.gone.store(true, Ordering::Release);
self.state.send_modify(|state| *state = State::Closed);
}
#[inline]
fn send(&self, command: Command) -> Result<()> {
if self.tx.send(command).is_err() {
return Err(Error::message("Client service is down"));
}
Ok(())
}
}
pub struct ServiceBuilder<T, E> {
url: String,
on_error: E,
reconnect: bool,
seed: u64,
format: Format,
_marker: PhantomData<T>,
}
impl<T, E> ServiceBuilder<T, E>
where
T: ClientImpl,
E: Callback<Error>,
{
#[inline]
pub fn on_error<U>(self, on_error: U) -> ServiceBuilder<T, U>
where
U: Callback<Error>,
{
ServiceBuilder {
url: self.url,
on_error,
reconnect: self.reconnect,
seed: self.seed,
format: self.format,
_marker: self._marker,
}
}
#[inline]
pub fn format(mut self, format: Format) -> Self {
self.format = format;
self
}
#[inline]
pub fn reconnect(mut self, reconnect: bool) -> Self {
self.reconnect = reconnect;
self
}
#[inline]
pub fn seed(mut self, seed: u64) -> Self {
self.seed = seed;
self
}
pub fn build(self) -> Service<T> {
let (tx, rx) = mpsc::unbounded_channel();
let (state, _) = watch::channel(State::Closed);
let shared = Arc::new(Shared {
tx,
serial: AtomicU32::new(0),
state,
broadcasts: Mutex::new(Broadcasts::new()),
gone: AtomicBool::new(false),
format: AtomicU8::new(self.format.to_u8()),
});
Service {
handle: Handle {
shared: shared.clone(),
},
shared,
rx,
url: self.url,
on_error: Box::new(self.on_error),
reconnect: self.reconnect,
socket: None,
pending: HashMap::new(),
timeout: INITIAL_TIMEOUT,
next_attempt: Some(Instant::now()),
rng: SmallRng::seed_from_u64(self.seed),
closed: false,
requested: self.format,
}
}
}
pub struct Service<T>
where
T: ClientImpl,
{
handle: Handle,
shared: Arc<Shared>,
rx: mpsc::UnboundedReceiver<Command>,
url: String,
on_error: Box<dyn Callback<Error>>,
reconnect: bool,
socket: Option<T::Socket>,
pending: HashMap<u32, Pending>,
timeout: Duration,
next_attempt: Option<Instant>,
rng: SmallRng,
closed: bool,
requested: Format,
}
enum Output<E> {
Message(Option<Result<Message, E>>),
Command(Option<Command>),
Connect,
}
impl<T> Service<T>
where
T: ClientImpl,
{
#[inline]
pub fn handle(&self) -> &Handle {
&self.handle
}
pub async fn run(&mut self) -> Result<()> {
while !self.closed {
let output = {
let rx = &mut self.rx;
match &mut self.socket {
Some(socket) => {
tokio::select! {
message = socket.recv() => Output::Message(message),
command = rx.recv() => Output::Command(command),
}
}
None => match self.next_attempt {
Some(deadline) => {
tokio::select! {
_ = tokio::time::sleep_until(deadline) => Output::Connect,
command = rx.recv() => Output::Command(command),
}
}
None => Output::Command(rx.recv().await),
},
}
};
match output {
Output::Connect => {
self.connect().await;
}
Output::Command(command) => {
let Some(command) = command else {
break;
};
self.command(command).await;
}
Output::Message(message) => {
let Some(message) = message else {
tracing::debug!("Connection closed by server");
self.disconnect().await;
continue;
};
let message = match message {
Ok(message) => message,
Err(error) => {
self.on_error.call(Error::transport(error));
self.disconnect().await;
continue;
}
};
match message {
Message::Binary(bytes) => match self.message(bytes) {
Ok(Post::Negotiate) => self.send_negotiate().await,
Ok(Post::None) => {}
Err(error) => self.on_error.call(error),
},
Message::Text => {
self.on_error
.call(Error::message("Unsupported text message"));
self.disconnect().await;
}
Message::Ping | Message::Pong => {}
Message::Close => {
tracing::debug!("Close message received");
self.disconnect().await;
}
}
}
}
}
self.shutdown().await;
Ok(())
}
async fn connect(&mut self) {
tracing::debug!(url = self.url.as_str(), "Connecting");
match T::connect(&self.url).await {
Ok(socket) => {
tracing::debug!("Connection established");
self.socket = Some(socket);
self.next_attempt = None;
self.timeout = INITIAL_TIMEOUT;
}
Err(error) => {
self.on_error.call(Error::transport(error));
self.schedule_reconnect();
}
}
}
async fn disconnect(&mut self) {
if let Some(mut socket) = self.socket.take() {
_ = socket.close().await;
}
self.emit_state(State::Closed);
self.close_pending(|| Error::message("Connection closed"));
self.schedule_reconnect();
}
async fn shutdown(&mut self) {
if let Some(mut socket) = self.socket.take() {
_ = socket.close().await;
}
self.emit_state(State::Closed);
self.close_pending(|| Error::message("Client service closed"));
}
fn schedule_reconnect(&mut self) {
if !self.reconnect {
tracing::debug!("Reconnecting is disabled, closing service");
self.closed = true;
return;
}
let fuzz = self.rng.random_range(0..=MAX_FUZZ);
let timeout = self
.timeout
.saturating_add(Duration::from_millis(fuzz))
.min(MAX_TIMEOUT);
self.timeout = self.timeout.saturating_mul(2).min(MAX_TIMEOUT);
self.next_attempt = Some(Instant::now() + timeout);
tracing::debug!(?timeout, "Scheduling reconnect");
}
fn close_pending(&mut self, error: impl Fn() -> Error) {
for (_, pending) in self.pending.drain() {
pending.error(error());
}
}
fn emit_state(&mut self, state: State) {
self.shared.state.send_if_modified(|current| {
if *current == state {
return false;
}
*current = state;
true
});
}
async fn command(&mut self, command: Command) {
match command {
Command::Send {
serial,
data,
pending,
} => {
let Some(socket) = self.socket.as_mut() else {
pending.error(Error::new(ErrorKind::NotConnected));
return;
};
if let Err(error) = socket.send(&data).await {
pending.error(Error::transport(error));
self.disconnect().await;
return;
}
if let Some(existing) = self.pending.insert(serial, pending) {
existing.error(Error::message("Request cancelled"));
}
}
Command::Disconnect { channel } => {
if let Err(error) = self.send_disconnect(channel).await {
self.on_error.call(error);
}
}
Command::Close => {
self.closed = true;
}
}
}
async fn send_disconnect(&mut self, channel: ChannelId) -> Result<()> {
let Some(socket) = self.socket.as_mut() else {
return Ok(());
};
let mut data = Vec::new();
let header = api::RequestHeader {
serial: 0,
id: MessageId::DISCONNECT.get(),
format: 0,
channel,
};
format::encode_envelope(&mut data, &header).map_err(Error::encoding_header)?;
tracing::debug!(?channel, "Sending disconnect");
if let Err(error) = socket.send(&data).await {
let error = Error::transport(error);
self.disconnect().await;
return Err(error);
}
Ok(())
}
fn dispatch(&self, id: MessageId, value: impl Fn() -> Result<RawPacket>) {
let broadcasts = self
.shared
.broadcasts
.lock()
.unwrap_or_else(|e| e.into_inner());
let Some(slots) = broadcasts.get(&id) else {
return;
};
for (_, tx) in slots.iter() {
_ = tx.send(value());
}
}
fn body_format(header: &api::ResponseHeader) -> Result<Format> {
let Some(format) = Format::from_u8(header.format) else {
return Err(Error::message(format_args!(
"Server used unknown format id {} for a message body",
header.format
)));
};
Ok(format)
}
fn message(&mut self, bytes: Bytes) -> Result<Post> {
let mut at = 0;
let header: api::ResponseHeader =
format::decode_envelope(&bytes, &mut at).map_err(Error::decode_response_header)?;
if let Some(broadcast) = MessageId::new(header.broadcast) {
tracing::debug!(?header, "Got broadcast");
if broadcast == MessageId::SERVER_HELLO {
tracing::debug!("Server hello, negotiating format");
return Ok(Post::Negotiate);
}
if let Some(id) = MessageId::new(header.error) {
let error = match id {
MessageId::ERROR_MESSAGE => Self::body_format(&header)?
.decode(&bytes, &mut at)
.map_err(Error::decode_error_message)?,
_ => api::ErrorMessage {
message: "Unsupported broadcast",
},
};
self.dispatch(broadcast, || Err(Error::server(error.message)));
return Ok(Post::None);
}
let format = Self::body_format(&header)?;
let packet = RawPacket {
id: broadcast,
buf: bytes,
at: Cell::new(at),
format,
channel: header.channel,
};
self.dispatch(broadcast, || Ok(packet.clone()));
return Ok(Post::None);
}
tracing::debug!(?header, "Got response");
let Some(pending) = self.pending.remove(&header.serial) else {
tracing::trace!(?header.serial, "Got message with unknown serial");
return Ok(Post::None);
};
if let Some(id) = MessageId::new(header.error) {
let error = match id {
MessageId::ERROR_MESSAGE => Self::body_format(&header)?
.decode(&bytes, &mut at)
.map_err(Error::decode_error_message)?,
_ => api::ErrorMessage {
message: "Unsupported request",
},
};
match pending {
Pending::Negotiate { format } => {
self.on_error.call(Error::message(format_args!(
"Server rejected format `{format}` ({}), falling back to `{}`",
error.message,
Format::DEFAULT
)));
self.shared.set_format(Format::DEFAULT);
self.emit_state(State::Open);
}
pending => {
pending.error(Error::server(error.message));
}
}
return Ok(Post::None);
}
match pending {
Pending::Negotiate { format } => {
let accepted = Format::from_u8(header.format).unwrap_or(format);
tracing::debug!(?accepted, "Format negotiated");
self.shared.set_format(accepted);
self.emit_state(State::Open);
}
Pending::Channel { reply } => {
_ = reply.send(Ok(header.channel));
}
Pending::Request { id, reply } => {
let format = Self::body_format(&header)?;
let packet = RawPacket {
id,
buf: bytes,
at: Cell::new(at),
format,
channel: header.channel,
};
_ = reply.send(Ok(packet));
}
}
Ok(Post::None)
}
async fn send_negotiate(&mut self) {
let format = self.requested;
let serial = self.shared.next_serial();
let header = api::RequestHeader {
serial,
id: MessageId::NEGOTIATE.get(),
format: format.to_u8(),
channel: ChannelId::NONE,
};
let mut data = Vec::new();
if let Err(error) = format::encode_envelope(&mut data, &header) {
self.on_error.call(Error::encoding_header(error));
return;
}
let Some(socket) = self.socket.as_mut() else {
return;
};
tracing::debug!(?format, "Requesting format");
if let Err(error) = socket.send(&data).await {
self.on_error.call(Error::transport(error));
self.disconnect().await;
return;
}
self.pending.insert(serial, Pending::Negotiate { format });
}
}
enum Post {
None,
Negotiate,
}
impl<T> Drop for Service<T>
where
T: ClientImpl,
{
fn drop(&mut self) {
self.shared.set_gone();
for (_, pending) in self.pending.drain() {
pending.error(Error::message("Client service closed"));
}
self.shared
.broadcasts
.lock()
.unwrap_or_else(|e| e.into_inner())
.clear();
}
}
impl<T> fmt::Debug for Service<T>
where
T: ClientImpl,
{
#[inline]
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("Service")
.field("url", &self.url)
.field("state", &*self.shared.state.borrow())
.finish_non_exhaustive()
}
}
#[derive(Clone)]
pub struct Handle {
shared: Arc<Shared>,
}
impl Handle {
#[inline]
pub fn state(&self) -> State {
*self.shared.state.borrow()
}
#[inline]
pub fn is_open(&self) -> bool {
self.shared.is_open()
}
#[inline]
pub fn format(&self) -> Format {
self.shared.format()
}
#[inline]
pub fn on_state_change(&self) -> StateListener {
StateListener {
rx: self.shared.state.subscribe(),
shared: self.shared.clone(),
}
}
pub async fn wait_until_open(&self) -> Result<()> {
let mut listener = self.on_state_change();
if listener.wait_until(State::Open).await {
return Ok(());
}
Err(Error::message("Client service is down"))
}
pub async fn channel(&self) -> Result<Channel> {
if !self.shared.is_open() {
return Err(Error::new(ErrorKind::NotConnected));
}
let serial = self.shared.next_serial();
let header = api::RequestHeader {
serial,
id: MessageId::CONNECT.get(),
format: 0,
channel: ChannelId::NONE,
};
let mut data = Vec::new();
format::encode_envelope(&mut data, &header).map_err(Error::encoding_header)?;
let (reply, rx) = oneshot::channel();
self.shared.send(Command::Send {
serial,
data,
pending: Pending::Channel { reply },
})?;
let Ok(result) = rx.await else {
return Err(Error::message("Client service is down"));
};
Ok(Channel {
shared: self.shared.clone(),
id: result?,
})
}
#[inline]
pub fn request(&self) -> RequestBuilder<'_, EmptyBody> {
RequestBuilder {
shared: &self.shared,
channel: ChannelId::NONE,
body: EmptyBody,
}
}
pub fn on_broadcast<T>(&self) -> Listener<T>
where
T: api::Broadcast,
{
let (tx, rx) = mpsc::unbounded_channel();
let index = {
let mut broadcasts = self
.shared
.broadcasts
.lock()
.unwrap_or_else(|e| e.into_inner());
broadcasts.entry(T::ID).or_default().insert(tx)
};
Listener {
shared: Some(self.shared.clone()),
id: T::ID,
index,
rx,
_marker: PhantomData,
}
}
#[inline]
pub fn close(&self) {
_ = self.shared.tx.send(Command::Close);
}
}
impl PartialEq for Handle {
#[inline]
fn eq(&self, other: &Self) -> bool {
Arc::ptr_eq(&self.shared, &other.shared)
}
}
impl Eq for Handle {}
impl fmt::Debug for Handle {
#[inline]
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("Handle")
.field("state", &*self.shared.state.borrow())
.finish_non_exhaustive()
}
}
pub struct Channel {
shared: Arc<Shared>,
id: ChannelId,
}
impl Channel {
#[inline]
pub fn id(&self) -> ChannelId {
self.id
}
#[inline]
pub fn handle(&self) -> Handle {
Handle {
shared: self.shared.clone(),
}
}
#[inline]
pub fn request(&self) -> RequestBuilder<'_, EmptyBody> {
RequestBuilder {
shared: &self.shared,
channel: self.id,
body: EmptyBody,
}
}
}
impl Drop for Channel {
#[inline]
fn drop(&mut self) {
_ = self.shared.send(Command::Disconnect { channel: self.id });
}
}
impl fmt::Debug for Channel {
#[inline]
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("Channel")
.field("id", &self.id)
.field("state", &*self.shared.state.borrow())
.finish_non_exhaustive()
}
}
pub struct RequestBuilder<'a, B> {
shared: &'a Arc<Shared>,
channel: ChannelId,
body: B,
}
impl<'a, B> RequestBuilder<'a, B> {
#[inline]
pub fn body<U>(self, body: U) -> RequestBuilder<'a, U>
where
U: api::Request,
{
RequestBuilder {
shared: self.shared,
channel: self.channel,
body,
}
}
}
impl<B> RequestBuilder<'_, B>
where
B: api::Request,
{
pub async fn send(self) -> Result<Packet<B::Endpoint>> {
Ok(Packet::new(self.send_raw().await?))
}
pub async fn send_raw(self) -> Result<RawPacket> {
let id = <B::Endpoint as api::Endpoint>::ID;
if !self.shared.is_open() {
return Err(Error::new(ErrorKind::NotConnected));
}
let serial = self.shared.next_serial();
let format = self.shared.format();
let header = api::RequestHeader {
serial,
id: id.get(),
format: format.to_u8(),
channel: self.channel,
};
let mut data = Vec::new();
format::encode_envelope(&mut data, &header).map_err(Error::encoding_header)?;
format
.encode(&mut data, &self.body)
.map_err(Error::encoding_body)?;
tracing::debug!(serial, ?id, ?format, len = data.len(), "Sending request");
let (reply, rx) = oneshot::channel();
self.shared.send(Command::Send {
serial,
data,
pending: Pending::Request { id, reply },
})?;
let Ok(result) = rx.await else {
return Err(Error::message("Client service is down"));
};
result
}
}
pub struct Listener<T> {
shared: Option<Arc<Shared>>,
id: MessageId,
index: usize,
rx: mpsc::UnboundedReceiver<Result<RawPacket>>,
_marker: PhantomData<T>,
}
impl<T> Listener<T> {
#[inline]
pub async fn recv_raw(&mut self) -> Option<Result<RawPacket>> {
self.rx.recv().await
}
#[inline]
pub async fn recv(&mut self) -> Option<Result<Packet<T>>> {
Some(match self.rx.recv().await? {
Ok(packet) => Ok(Packet::new(packet)),
Err(error) => Err(error),
})
}
pub fn clear(&mut self) {
let Some(shared) = self.shared.take() else {
return;
};
let index = mem::take(&mut self.index);
let mut broadcasts = shared.broadcasts.lock().unwrap_or_else(|e| e.into_inner());
let Entry::Occupied(mut e) = broadcasts.entry(self.id) else {
return;
};
_ = e.get_mut().try_remove(index);
if e.get().is_empty() {
e.remove();
}
}
}
impl<T> Drop for Listener<T> {
#[inline]
fn drop(&mut self) {
self.clear();
}
}
impl<T> fmt::Debug for Listener<T> {
#[inline]
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("Listener")
.field("type", &any::type_name::<T>())
.field("id", &self.id)
.finish_non_exhaustive()
}
}
pub struct StateListener {
rx: watch::Receiver<State>,
shared: Arc<Shared>,
}
impl StateListener {
#[inline]
pub fn state(&self) -> State {
*self.rx.borrow()
}
#[inline]
pub async fn changed(&mut self) -> Option<State> {
if self.shared.is_gone() {
return None;
}
self.rx.changed().await.ok()?;
if self.shared.is_gone() {
return None;
}
Some(*self.rx.borrow_and_update())
}
pub async fn wait_until(&mut self, state: State) -> bool {
loop {
if *self.rx.borrow_and_update() == state {
return true;
}
if self.shared.is_gone() {
return false;
}
if self.rx.changed().await.is_err() {
return false;
}
}
}
}
impl Clone for StateListener {
#[inline]
fn clone(&self) -> Self {
Self {
rx: self.rx.clone(),
shared: self.shared.clone(),
}
}
}
impl fmt::Debug for StateListener {
#[inline]
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("StateListener")
.field("state", &*self.rx.borrow())
.finish()
}
}
#[derive(Clone)]
pub struct RawPacket {
id: MessageId,
buf: Bytes,
at: Cell<usize>,
format: Format,
channel: ChannelId,
}
impl RawPacket {
#[inline]
pub const fn empty() -> Self {
Self {
id: MessageId::EMPTY,
buf: Bytes::new(),
at: Cell::new(0),
format: Format::DEFAULT,
channel: ChannelId::NONE,
}
}
#[inline]
pub fn format(&self) -> Format {
self.format
}
#[inline]
pub fn channel(&self) -> ChannelId {
self.channel
}
pub fn decode<'this, T>(&'this self) -> Result<T>
where
T: DecodeBody<'this>,
{
if self.id == MessageId::EMPTY {
return Err(Error::new(ErrorKind::EmptyPacket));
}
let mut at = self.at.get();
match self.format.decode(&self.buf, &mut at) {
Ok(value) => {
self.at.set(at);
Ok(value)
}
Err(error) => {
self.at.set(self.len());
Err(Error::decode_packet(error))
}
}
}
#[inline]
pub fn as_slice(&self) -> &[u8] {
&self.buf
}
#[inline]
pub fn remaining(&self) -> usize {
self.buf.len().saturating_sub(self.at.get())
}
#[inline]
pub fn len(&self) -> usize {
self.buf.len()
}
#[inline]
pub fn is_empty(&self) -> bool {
self.at.get() >= self.len()
}
#[inline]
pub fn id(&self) -> MessageId {
self.id
}
}
impl fmt::Debug for RawPacket {
#[inline]
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("RawPacket")
.field("id", &self.id)
.field("remaining", &self.remaining())
.finish()
}
}
pub struct Packet<T> {
raw: RawPacket,
_marker: PhantomData<T>,
}
impl<T> Packet<T> {
#[inline]
pub const fn empty() -> Self {
Self {
raw: RawPacket::empty(),
_marker: PhantomData,
}
}
#[inline]
pub fn new(raw: RawPacket) -> Self {
Self {
raw,
_marker: PhantomData,
}
}
#[inline]
pub fn channel(&self) -> ChannelId {
self.raw.channel()
}
#[inline]
pub fn format(&self) -> Format {
self.raw.format()
}
#[inline]
pub fn into_raw(self) -> RawPacket {
self.raw
}
#[inline]
pub fn remaining(&self) -> usize {
self.raw.remaining()
}
#[inline]
pub fn is_empty(&self) -> bool {
self.raw.is_empty()
}
#[inline]
pub fn id(&self) -> MessageId {
self.raw.id()
}
}
impl<T> Packet<T>
where
T: api::Decodable,
{
#[inline]
pub fn decode(&self) -> Result<T::Type<'_>> {
self.decode_any()
}
#[inline]
pub fn decode_any<'de, R>(&'de self) -> Result<R>
where
R: DecodeBody<'de>,
{
self.raw.decode()
}
}
impl<T> Packet<T>
where
T: api::Endpoint,
{
#[inline]
pub fn decode_response(&self) -> Result<T::Response<'_>> {
self.decode_any_response()
}
#[inline]
pub fn decode_any_response<'de, R>(&'de self) -> Result<R>
where
R: DecodeBody<'de>,
{
self.raw.decode()
}
}
impl<T> Packet<T>
where
T: api::Broadcast,
{
#[inline]
pub fn decode_event<'de>(&'de self) -> Result<T::Event<'de>>
where
T: api::BroadcastWithEvent,
{
self.decode_event_any()
}
#[inline]
pub fn decode_event_any<'de, E>(&'de self) -> Result<E>
where
E: Event<Broadcast = T> + DecodeBody<'de>,
{
self.raw.decode()
}
}
impl<T> Clone for Packet<T> {
#[inline]
fn clone(&self) -> Self {
Self {
raw: self.raw.clone(),
_marker: PhantomData,
}
}
}
impl<T> fmt::Debug for Packet<T> {
#[inline]
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("Packet")
.field("type", &any::type_name::<T>())
.field("remaining", &self.remaining())
.finish()
}
}