11use async_trait:: async_trait;
2+ use bytes:: Bytes ;
23use dashmap:: DashMap ;
34use futures:: { SinkExt , StreamExt } ;
45use mtw_codec:: json:: JsonCodec ;
56use mtw_codec:: MtwCodec ;
67use mtw_core:: MtwError ;
8+ use mtw_protocol:: frame:: { Frame , FrameType } ;
79use mtw_protocol:: {
810 ConnId , ConnMetadata , DisconnectReason , MsgType , MtwMessage , Payload , TransportEvent ,
911} ;
@@ -25,6 +27,8 @@ pub struct WebSocketTransport {
2527 ping_interval : u64 ,
2628 /// Active connections: conn_id -> sender
2729 connections : Arc < DashMap < ConnId , WsSender > > ,
30+ /// Connections using binary frame protocol (vs JSON text)
31+ binary_connections : Arc < DashMap < ConnId , ( ) > > ,
2832 /// Event channel
2933 event_tx : mpsc:: UnboundedSender < TransportEvent > ,
3034 event_rx : Option < mpsc:: UnboundedReceiver < TransportEvent > > ,
@@ -41,6 +45,7 @@ impl WebSocketTransport {
4145 path : path. into ( ) ,
4246 ping_interval,
4347 connections : Arc :: new ( DashMap :: new ( ) ) ,
48+ binary_connections : Arc :: new ( DashMap :: new ( ) ) ,
4449 event_tx,
4550 event_rx : Some ( event_rx) ,
4651 codec : Arc :: new ( JsonCodec ) ,
@@ -53,6 +58,7 @@ impl WebSocketTransport {
5358 stream : TcpStream ,
5459 addr : SocketAddr ,
5560 connections : Arc < DashMap < ConnId , WsSender > > ,
61+ binary_connections : Arc < DashMap < ConnId , ( ) > > ,
5662 event_tx : mpsc:: UnboundedSender < TransportEvent > ,
5763 codec : Arc < dyn MtwCodec > ,
5864 ping_interval : u64 ,
@@ -143,10 +149,51 @@ impl WebSocketTransport {
143149 }
144150 }
145151 Some ( Ok ( WsMessage :: Binary ( data) ) ) => {
146- let _ = event_tx. send( TransportEvent :: Binary (
147- conn_id. clone( ) ,
148- data. to_vec( ) ,
149- ) ) ;
152+ let bytes = Bytes :: from( data. to_vec( ) ) ;
153+ match Frame :: decode( bytes. clone( ) ) {
154+ Ok ( ( FrameType :: Json , payload) ) => {
155+ // MTW binary frame containing a JSON message
156+ binary_connections. insert( conn_id. clone( ) , ( ) ) ;
157+ match serde_json:: from_slice:: <MtwMessage >( & payload) {
158+ Ok ( mtw_msg) => {
159+ let _ = event_tx. send( TransportEvent :: Message (
160+ conn_id. clone( ) ,
161+ mtw_msg,
162+ ) ) ;
163+ }
164+ Err ( e) => {
165+ let _ = event_tx. send( TransportEvent :: Error (
166+ conn_id. clone( ) ,
167+ format!( "frame JSON decode error: {}" , e) ,
168+ ) ) ;
169+ }
170+ }
171+ }
172+ Ok ( ( FrameType :: Ping , _) ) => {
173+ // MTW Ping frame — respond with MTW Pong
174+ if let Some ( sender) = connections. get( & conn_id) {
175+ let pong = Frame :: encode_pong( ) ;
176+ let _ = sender. send( WsMessage :: Binary ( pong. to_vec( ) . into( ) ) ) ;
177+ }
178+ }
179+ Ok ( ( FrameType :: Pong , _) ) => {
180+ // MTW Pong — connection is alive
181+ }
182+ Ok ( ( FrameType :: Binary , payload) ) => {
183+ // Raw binary data (audio, 3D, etc.)
184+ let _ = event_tx. send( TransportEvent :: Binary (
185+ conn_id. clone( ) ,
186+ payload. to_vec( ) ,
187+ ) ) ;
188+ }
189+ Err ( _) => {
190+ // Not an MTW frame — treat as raw binary
191+ let _ = event_tx. send( TransportEvent :: Binary (
192+ conn_id. clone( ) ,
193+ bytes. to_vec( ) ,
194+ ) ) ;
195+ }
196+ }
150197 }
151198 Some ( Ok ( WsMessage :: Pong ( _) ) ) => {
152199 // Connection is alive
@@ -173,6 +220,7 @@ impl WebSocketTransport {
173220 ping_handle. abort ( ) ;
174221 write_handle. abort ( ) ;
175222 connections. remove ( & conn_id) ;
223+ binary_connections. remove ( & conn_id) ;
176224
177225 let _ = event_tx. send ( TransportEvent :: Disconnected (
178226 conn_id. clone ( ) ,
@@ -198,6 +246,7 @@ impl MtwTransport for WebSocketTransport {
198246 self . shutdown_tx = Some ( shutdown_tx. clone ( ) ) ;
199247
200248 let connections = self . connections . clone ( ) ;
249+ let binary_connections = self . binary_connections . clone ( ) ;
201250 let event_tx = self . event_tx . clone ( ) ;
202251 let codec = self . codec . clone ( ) ;
203252 let ping_interval = self . ping_interval ;
@@ -214,6 +263,7 @@ impl MtwTransport for WebSocketTransport {
214263 stream,
215264 addr,
216265 connections. clone ( ) ,
266+ binary_connections. clone ( ) ,
217267 event_tx. clone ( ) ,
218268 codec. clone ( ) ,
219269 ping_interval,
@@ -231,8 +281,16 @@ impl MtwTransport for WebSocketTransport {
231281 }
232282
233283 async fn send ( & self , conn_id : & ConnId , msg : MtwMessage ) -> Result < ( ) , MtwError > {
234- let encoded = self . codec . encode ( & msg) ?;
235- let ws_msg = WsMessage :: Text ( String :: from_utf8_lossy ( & encoded) . into ( ) ) ;
284+ let ws_msg = if self . binary_connections . contains_key ( conn_id) {
285+ // Client speaks MTW binary protocol — send as binary frame
286+ let frame = Frame :: encode_message ( & msg)
287+ . map_err ( |e| MtwError :: Transport ( format ! ( "frame encode error: {}" , e) ) ) ?;
288+ WsMessage :: Binary ( frame. to_vec ( ) . into ( ) )
289+ } else {
290+ // Client speaks JSON text — send as plain JSON
291+ let encoded = self . codec . encode ( & msg) ?;
292+ WsMessage :: Text ( String :: from_utf8_lossy ( & encoded) . into ( ) )
293+ } ;
236294
237295 if let Some ( sender) = self . connections . get ( conn_id) {
238296 sender
0 commit comments