From 0f033cfb72b8f1db4455cead5430ea65a0668f43 Mon Sep 17 00:00:00 2001 From: Orual Date: Thu, 16 Oct 2025 23:49:43 -0400 Subject: [PATCH] fixed decode trait bound issues --- crates/jacquard-common/src/websocket.rs | 91 --------- crates/jacquard-common/src/xrpc.rs | 2 +- .../jacquard-common/src/xrpc/subscription.rs | 193 +++++++++++++++++- crates/jacquard-oauth/src/client.rs | 4 +- .../jacquard/src/client/credential_session.rs | 4 +- 5 files changed, 187 insertions(+), 107 deletions(-) diff --git a/crates/jacquard-common/src/websocket.rs b/crates/jacquard-common/src/websocket.rs index 28415b190..658b4ca5e 100644 --- a/crates/jacquard-common/src/websocket.rs +++ b/crates/jacquard-common/src/websocket.rs @@ -352,97 +352,6 @@ impl WsStream { } } -/// Extension trait for decoding typed messages from WebSocket streams -pub trait WsStreamExt: Sized { - /// Decode JSON text/binary frames into typed messages - /// - /// Deserializes borrowing from temporary frame bytes, then converts to owned. - fn decode_json(self) -> impl Stream> - where - T: IntoStatic, - for<'de> T: serde::Deserialize<'de>, - T::Output: 'static; - - /// Decode DAG-CBOR binary frames into typed messages - /// - /// Deserializes borrowing from temporary frame bytes, then converts to owned. - fn decode_cbor(self) -> impl Stream> - where - T: IntoStatic, - for<'de> T: serde::Deserialize<'de>, - T::Output: 'static; -} - -impl WsStreamExt for WsStream { - fn decode_json(self) -> impl Stream> - where - T: IntoStatic, - for<'de> T: serde::Deserialize<'de>, - T::Output: 'static, - { - use n0_future::StreamExt as _; - - // Helper to deserialize with concrete lifetime - fn parse_json<'a, T>(bytes: &'a [u8]) -> Result - where - T: serde::Deserialize<'a>, - { - serde_json::from_slice(bytes) - } - - Box::pin(self.into_inner().filter_map(|msg_result| { - match msg_result { - Ok(WsMessage::Text(text)) => Some( - parse_json::(text.as_ref()) - .map(|v| v.into_static()) - .map_err(StreamError::decode), - ), - Ok(WsMessage::Binary(bytes)) => Some( - parse_json::(&bytes) - .map(|v| v.into_static()) - .map_err(StreamError::decode), - ), - Ok(WsMessage::Close(_)) => Some(Err(StreamError::closed())), - Err(e) => Some(Err(e)), - } - })) - } - - fn decode_cbor(self) -> impl Stream> - where - T: IntoStatic, - for<'de> T: serde::Deserialize<'de>, - T::Output: 'static, - { - use n0_future::StreamExt as _; - - // Helper to deserialize with concrete lifetime - fn parse_cbor<'a, T>( - bytes: &'a [u8], - ) -> Result> - where - T: serde::Deserialize<'a>, - { - serde_ipld_dagcbor::from_slice(bytes) - } - - Box::pin(self.into_inner().filter_map(|msg_result| { - match msg_result { - Ok(WsMessage::Binary(bytes)) => Some( - parse_cbor::(&bytes) - .map(|v| v.into_static()) - .map_err(|e| StreamError::decode(crate::error::DecodeError::from(e))), - ), - Ok(WsMessage::Text(_)) => Some(Err(StreamError::wrong_message_format( - "expected binary frame for CBOR, got text", - ))), - Ok(WsMessage::Close(_)) => Some(Err(StreamError::closed())), - Err(e) => Some(Err(e)), - } - })) - } -} - impl fmt::Debug for WsStream { fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { f.debug_struct("WsStream").finish_non_exhaustive() diff --git a/crates/jacquard-common/src/xrpc.rs b/crates/jacquard-common/src/xrpc.rs index 88342c0b0..8cef22e95 100644 --- a/crates/jacquard-common/src/xrpc.rs +++ b/crates/jacquard-common/src/xrpc.rs @@ -23,7 +23,7 @@ pub mod subscription; pub use subscription::{ BasicSubscriptionClient, MessageEncoding, SubscriptionCall, SubscriptionClient, SubscriptionEndpoint, SubscriptionExt, SubscriptionOptions, SubscriptionResp, - TungsteniteSubscriptionClient, XrpcSubscription, + SubscriptionStream, TungsteniteSubscriptionClient, XrpcSubscription, }; use bytes::Bytes; diff --git a/crates/jacquard-common/src/xrpc/subscription.rs b/crates/jacquard-common/src/xrpc/subscription.rs index 409bcf35b..66b46b4c0 100644 --- a/crates/jacquard-common/src/xrpc/subscription.rs +++ b/crates/jacquard-common/src/xrpc/subscription.rs @@ -6,8 +6,10 @@ use serde::{Deserialize, Serialize}; use std::error::Error; use std::future::Future; +use std::marker::PhantomData; use url::Url; +use crate::stream::StreamError; use crate::websocket::{WebSocketClient, WebSocketConnection}; use crate::{CowStr, IntoStatic}; @@ -76,6 +78,168 @@ pub trait XrpcSubscription: Serialize { } } +/// Decode JSON messages from a WebSocket stream +fn decode_json_msg( + msg_result: Result, +) -> Option, StreamError>> +where + for<'a> StreamMessage<'a, S>: IntoStatic>, +{ + use crate::websocket::WsMessage; + + fn parse_msg<'a, S: SubscriptionResp>( + bytes: &'a [u8], + ) -> Result, serde_json::Error> { + serde_json::from_slice(bytes) + } + + match msg_result { + Ok(WsMessage::Text(text)) => Some( + parse_msg::(text.as_ref()) + .map(|v| v.into_static()) + .map_err(StreamError::decode), + ), + Ok(WsMessage::Binary(bytes)) => Some( + parse_msg::(&bytes) + .map(|v| v.into_static()) + .map_err(StreamError::decode), + ), + Ok(WsMessage::Close(_)) => Some(Err(StreamError::closed())), + Err(e) => Some(Err(e)), + } +} + +/// Decode CBOR messages from a WebSocket stream +fn decode_cbor_msg( + msg_result: Result, +) -> Option, StreamError>> +where + for<'a> StreamMessage<'a, S>: IntoStatic>, +{ + use crate::websocket::WsMessage; + + fn parse_cbor<'a, S: SubscriptionResp>( + bytes: &'a [u8], + ) -> Result, serde_ipld_dagcbor::DecodeError> { + serde_ipld_dagcbor::from_slice(bytes) + } + + match msg_result { + Ok(WsMessage::Binary(bytes)) => Some( + parse_cbor::(&bytes) + .map(|v| v.into_static()) + .map_err(|e| StreamError::decode(crate::error::DecodeError::from(e))), + ), + Ok(WsMessage::Text(_)) => Some(Err(StreamError::wrong_message_format( + "expected binary frame for CBOR, got text", + ))), + Ok(WsMessage::Close(_)) => Some(Err(StreamError::closed())), + Err(e) => Some(Err(e)), + } +} + +/// Typed subscription stream wrapping a WebSocket connection. +/// +/// Analogous to `Response` for XRPC but for subscription streams. +/// Automatically decodes messages based on the subscription's encoding format. +pub struct SubscriptionStream { + _marker: PhantomData S>, + connection: WebSocketConnection, +} + +impl SubscriptionStream { + /// Create a new subscription stream from a WebSocket connection. + pub fn new(connection: WebSocketConnection) -> Self { + Self { + _marker: PhantomData, + connection, + } + } + + /// Get a reference to the underlying WebSocket connection. + pub fn connection(&self) -> &WebSocketConnection { + &self.connection + } + + /// Get a mutable reference to the underlying WebSocket connection. + pub fn connection_mut(&mut self) -> &mut WebSocketConnection { + &mut self.connection + } + + /// Split the connection and decode messages into a typed stream. + /// + /// Returns a tuple of (sender, typed message stream). + /// Messages are decoded according to the subscription's ENCODING. + pub fn into_stream( + self, + ) -> ( + crate::websocket::WsSink, + n0_future::stream::Boxed, StreamError>>, + ) + where + for<'a> StreamMessage<'a, S>: IntoStatic>, + { + use n0_future::StreamExt as _; + + let (tx, rx) = self.connection.split(); + + let stream: n0_future::stream::Boxed<_> = match S::ENCODING { + MessageEncoding::Json => { + Box::pin(rx.into_inner().filter_map(|msg| decode_json_msg::(msg))) + } + MessageEncoding::DagCbor => { + Box::pin(rx.into_inner().filter_map(|msg| decode_cbor_msg::(msg))) + } + }; + + (tx, stream) + } + + /// Consume the stream and return the underlying connection. + pub fn into_connection(self) -> WebSocketConnection { + self.connection + } + + /// Tee the stream, keeping the raw stream in self and returning a typed stream. + /// + /// Replaces the internal WebSocket stream with one copy and returns a typed decoded + /// stream. Both streams receive all messages. Useful for observing raw messages + /// while also processing typed messages. + pub fn tee( + &mut self, + ) -> n0_future::stream::Boxed, StreamError>> + where + for<'a> StreamMessage<'a, S>: IntoStatic>, + { + use n0_future::StreamExt as _; + + let rx = self.connection.receiver_mut(); + let (raw_rx, typed_rx_source) = std::mem::replace( + rx, + crate::websocket::WsStream::new(futures::stream::empty()), + ) + .tee(); + + // Put the raw stream back + *rx = raw_rx; + + match S::ENCODING { + MessageEncoding::Json => Box::pin( + typed_rx_source + .into_inner() + .filter_map(|msg| decode_json_msg::(msg)), + ), + MessageEncoding::DagCbor => Box::pin( + typed_rx_source + .into_inner() + .filter_map(|msg| decode_cbor_msg::(msg)), + ), + } + } +} + +type StreamMessage<'a, R> = ::Message<'a>; + /// XRPC subscription endpoint trait (server-side) /// /// Analogous to `XrpcEndpoint` but for WebSocket subscriptions. @@ -163,7 +327,11 @@ impl<'a, C: WebSocketClient> SubscriptionCall<'a, C> { /// /// Builds a WebSocket URL from the base, appends the NSID path, /// encodes query parameters from the subscription type, and connects. - pub async fn subscribe(self, params: &Sub) -> Result + /// Returns a typed SubscriptionStream that automatically decodes messages. + pub async fn subscribe( + self, + params: &Sub, + ) -> Result, C::Error> where Sub: XrpcSubscription, { @@ -185,9 +353,12 @@ impl<'a, C: WebSocketClient> SubscriptionCall<'a, C> { url.set_query(None); } - self.client + let connection = self + .client .connect_with_headers(url, self.opts.headers) - .await + .await?; + + Ok(SubscriptionStream::new(connection)) } } @@ -210,7 +381,7 @@ pub trait SubscriptionClient: WebSocketClient { fn subscribe( &self, params: &Sub, - ) -> impl Future> + ) -> impl Future, Self::Error>> where Sub: XrpcSubscription + Send + Sync, Self: Sync; @@ -220,7 +391,7 @@ pub trait SubscriptionClient: WebSocketClient { fn subscribe( &self, params: &Sub, - ) -> impl Future> + ) -> impl Future, Self::Error>> where Sub: XrpcSubscription + Send + Sync; @@ -230,7 +401,7 @@ pub trait SubscriptionClient: WebSocketClient { &self, params: &Sub, opts: SubscriptionOptions<'_>, - ) -> impl Future> + ) -> impl Future, Self::Error>> where Sub: XrpcSubscription + Send + Sync, Self: Sync; @@ -241,7 +412,7 @@ pub trait SubscriptionClient: WebSocketClient { &self, params: &Sub, opts: SubscriptionOptions<'_>, - ) -> impl Future> + ) -> impl Future, Self::Error>> where Sub: XrpcSubscription + Send + Sync; } @@ -308,7 +479,7 @@ impl SubscriptionClient for BasicSubscriptionClient { async fn subscribe( &self, params: &Sub, - ) -> Result + ) -> Result, Self::Error> where Sub: XrpcSubscription + Send + Sync, Self: Sync, @@ -321,7 +492,7 @@ impl SubscriptionClient for BasicSubscriptionClient { async fn subscribe( &self, params: &Sub, - ) -> Result + ) -> Result, Self::Error> where Sub: XrpcSubscription + Send + Sync, { @@ -334,7 +505,7 @@ impl SubscriptionClient for BasicSubscriptionClient { &self, params: &Sub, opts: SubscriptionOptions<'_>, - ) -> Result + ) -> Result, Self::Error> where Sub: XrpcSubscription + Send + Sync, Self: Sync, @@ -351,7 +522,7 @@ impl SubscriptionClient for BasicSubscriptionClient { &self, params: &Sub, opts: SubscriptionOptions<'_>, - ) -> Result + ) -> Result, Self::Error> where Sub: XrpcSubscription + Send + Sync, { diff --git a/crates/jacquard-oauth/src/client.rs b/crates/jacquard-oauth/src/client.rs index e6ddef2e4..179d66880 100644 --- a/crates/jacquard-oauth/src/client.rs +++ b/crates/jacquard-oauth/src/client.rs @@ -618,7 +618,7 @@ where async fn subscribe( &self, params: &Sub, - ) -> std::result::Result + ) -> std::result::Result, Self::Error> where Sub: XrpcSubscription + Send + Sync, { @@ -630,7 +630,7 @@ where &self, params: &Sub, opts: jacquard_common::xrpc::SubscriptionOptions<'_>, - ) -> std::result::Result + ) -> std::result::Result, Self::Error> where Sub: XrpcSubscription + Send + Sync, { diff --git a/crates/jacquard/src/client/credential_session.rs b/crates/jacquard/src/client/credential_session.rs index a7abf66f4..b651cd426 100644 --- a/crates/jacquard/src/client/credential_session.rs +++ b/crates/jacquard/src/client/credential_session.rs @@ -607,7 +607,7 @@ where async fn subscribe( &self, params: &Sub, - ) -> Result + ) -> Result, Self::Error> where Sub: XrpcSubscription + Send + Sync, { @@ -619,7 +619,7 @@ where &self, params: &Sub, opts: jacquard_common::xrpc::SubscriptionOptions<'_>, - ) -> Result + ) -> Result, Self::Error> where Sub: XrpcSubscription + Send + Sync, { -- 2.51.2