1use std::time::Duration;
35
36use serde::{Deserialize, Serialize};
37use tokio::io::{AsyncReadExt, AsyncWriteExt};
38use tokio::net::{TcpListener, TcpStream};
39#[cfg(unix)]
40use tokio::net::{UnixListener, UnixStream};
41
42use super::cross_mob_remote::{RemoteEndpoint, RemoteMobError};
43
44const MAX_CONTROL_PAYLOAD: u32 = 64 * 1024;
48
49pub const DEFAULT_CONTROL_TIMEOUT: Duration = Duration::from_secs(5);
53
54#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
56#[serde(tag = "op", rename_all = "snake_case")]
57pub enum ControlRequest {
58 Wire {
60 remote_member: String,
62 local_peer_spec_address: String,
64 local_comms_name: String,
65 local_peer_id: String,
66 #[serde(default, skip_serializing_if = "Option::is_none")]
70 local_pubkey_b64: Option<String>,
71 },
72 Unwire {
74 remote_member: String,
75 local_peer_spec_address: String,
76 local_comms_name: String,
77 local_peer_id: String,
78 #[serde(default, skip_serializing_if = "Option::is_none")]
79 local_pubkey_b64: Option<String>,
80 },
81 Inject {
84 remote_member: String,
85 content: serde_json::Value,
87 },
88 LookupMember { remote_member: String },
92}
93
94#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
96#[serde(tag = "result", rename_all = "snake_case")]
97pub enum ControlResponse {
98 Ok,
100 Injected { session_id: String },
103 Member { peer_id: String, comms_name: String },
106 Err { code: String, message: String },
110}
111
112enum ControlStream {
118 Tcp(TcpStream),
119 #[cfg(unix)]
120 Uds(UnixStream),
121}
122
123impl ControlStream {
124 async fn write_frame(&mut self, payload: &[u8]) -> Result<(), std::io::Error> {
125 let len = u32::try_from(payload.len()).map_err(|_| {
126 std::io::Error::new(std::io::ErrorKind::InvalidInput, "payload too large")
127 })?;
128 let header = len.to_be_bytes();
129 match self {
130 Self::Tcp(s) => {
131 s.write_all(&header).await?;
132 s.write_all(payload).await?;
133 s.flush().await
134 }
135 #[cfg(unix)]
136 Self::Uds(s) => {
137 s.write_all(&header).await?;
138 s.write_all(payload).await?;
139 s.flush().await
140 }
141 }
142 }
143
144 async fn read_frame(&mut self) -> Result<Vec<u8>, std::io::Error> {
145 let mut header = [0u8; 4];
146 match self {
147 Self::Tcp(s) => s.read_exact(&mut header).await?,
148 #[cfg(unix)]
149 Self::Uds(s) => s.read_exact(&mut header).await?,
150 };
151 let len = u32::from_be_bytes(header);
152 if len > MAX_CONTROL_PAYLOAD {
153 return Err(std::io::Error::new(
154 std::io::ErrorKind::InvalidData,
155 format!("frame too large: {len} bytes"),
156 ));
157 }
158 let mut buf = vec![0u8; len as usize];
159 match self {
160 Self::Tcp(s) => s.read_exact(&mut buf).await?,
161 #[cfg(unix)]
162 Self::Uds(s) => s.read_exact(&mut buf).await?,
163 };
164 Ok(buf)
165 }
166}
167
168pub struct RemoteControlClient;
170
171impl RemoteControlClient {
172 pub async fn send(
179 endpoint: &RemoteEndpoint,
180 request: &ControlRequest,
181 timeout: Duration,
182 ) -> Result<ControlResponse, RemoteMobError> {
183 tokio::time::timeout(timeout, Self::send_inner(endpoint, request))
184 .await
185 .map_err(|_| RemoteMobError::ControlChannelUnavailable {
186 mob_id: String::new(),
187 endpoint: endpoint.comms_address(),
188 operation: "timeout",
189 })?
190 }
191
192 async fn send_inner(
193 endpoint: &RemoteEndpoint,
194 request: &ControlRequest,
195 ) -> Result<ControlResponse, RemoteMobError> {
196 let mut stream = match endpoint {
197 RemoteEndpoint::Tcp(addr) => ControlStream::Tcp(
198 TcpStream::connect(addr)
199 .await
200 .map_err(|err| io_error("connect", endpoint, err))?,
201 ),
202 #[cfg(unix)]
203 RemoteEndpoint::Uds(path) => ControlStream::Uds(
204 UnixStream::connect(std::path::Path::new(path))
205 .await
206 .map_err(|err| io_error("connect", endpoint, err))?,
207 ),
208 #[cfg(not(unix))]
209 RemoteEndpoint::Uds(_) => {
210 return Err(RemoteMobError::UnsupportedTransport {
211 mob_id: String::new(),
212 transport: endpoint.comms_address(),
213 });
214 }
215 };
216 let payload =
217 serde_json::to_vec(request).map_err(|err| encode_error(endpoint, err.to_string()))?;
218 stream
219 .write_frame(&payload)
220 .await
221 .map_err(|err| io_error("write", endpoint, err))?;
222 let response_payload = stream
223 .read_frame()
224 .await
225 .map_err(|err| io_error("read", endpoint, err))?;
226 serde_json::from_slice::<ControlResponse>(&response_payload)
227 .map_err(|err| decode_error(endpoint, err.to_string()))
228 }
229}
230
231fn io_error(stage: &'static str, endpoint: &RemoteEndpoint, err: std::io::Error) -> RemoteMobError {
232 RemoteMobError::ControlChannelUnavailable {
233 mob_id: String::new(),
234 endpoint: endpoint.comms_address(),
235 operation: match stage {
236 "connect" => "connect",
237 "write" => "write",
238 "read" => "read",
239 _ => "io",
240 },
241 }
242 .with_context(err.to_string())
243}
244
245fn encode_error(endpoint: &RemoteEndpoint, message: String) -> RemoteMobError {
246 RemoteMobError::Encode {
247 endpoint: endpoint.comms_address(),
248 message,
249 }
250}
251
252fn decode_error(endpoint: &RemoteEndpoint, message: String) -> RemoteMobError {
253 RemoteMobError::Decode {
254 endpoint: endpoint.comms_address(),
255 message,
256 }
257}
258
259pub trait ControlHandler: Send + Sync + 'static {
266 fn handle(
267 &self,
268 request: ControlRequest,
269 ) -> std::pin::Pin<Box<dyn std::future::Future<Output = ControlResponse> + Send + '_>>;
270}
271
272pub struct MobHandleControlHandler {
276 handle: meerkat_mob::MobHandle,
277}
278
279impl MobHandleControlHandler {
280 pub fn new(handle: meerkat_mob::MobHandle) -> Self {
281 Self { handle }
282 }
283}
284
285impl ControlHandler for MobHandleControlHandler {
286 fn handle(
287 &self,
288 request: ControlRequest,
289 ) -> std::pin::Pin<Box<dyn std::future::Future<Output = ControlResponse> + Send + '_>> {
290 let handle = self.handle.clone();
291 Box::pin(async move {
292 match request {
293 ControlRequest::Wire {
294 remote_member,
295 local_peer_spec_address,
296 local_comms_name,
297 local_peer_id,
298 local_pubkey_b64,
299 } => {
300 handle_wire(
301 &handle,
302 &remote_member,
303 &local_peer_spec_address,
304 &local_comms_name,
305 &local_peer_id,
306 local_pubkey_b64.as_deref(),
307 true,
308 )
309 .await
310 }
311 ControlRequest::Unwire {
312 remote_member,
313 local_peer_spec_address,
314 local_comms_name,
315 local_peer_id,
316 local_pubkey_b64,
317 } => {
318 handle_wire(
319 &handle,
320 &remote_member,
321 &local_peer_spec_address,
322 &local_comms_name,
323 &local_peer_id,
324 local_pubkey_b64.as_deref(),
325 false,
326 )
327 .await
328 }
329 ControlRequest::Inject {
330 remote_member,
331 content,
332 } => handle_inject(&handle, &remote_member, content).await,
333 ControlRequest::LookupMember { remote_member } => {
334 handle_lookup_member(&handle, &remote_member).await
335 }
336 }
337 })
338 }
339}
340
341async fn handle_wire(
342 handle: &meerkat_mob::MobHandle,
343 remote_member: &str,
344 local_peer_spec_address: &str,
345 local_comms_name: &str,
346 local_peer_id: &str,
347 local_pubkey_b64: Option<&str>,
348 wire: bool,
349) -> ControlResponse {
350 let pubkey = match local_pubkey_b64 {
351 Some(s) if !s.is_empty() => match crate::auth::peer_keys::decode_pubkey_b64(s) {
352 Ok(bytes) => Some(bytes),
353 Err(err) => {
354 return ControlResponse::Err {
355 code: "decode".to_string(),
356 message: format!("local_pubkey_b64: {err}"),
357 };
358 }
359 },
360 _ => None,
361 };
362 let spec_result = match pubkey {
363 Some(bytes) => meerkat_core::comms::TrustedPeerDescriptor::unsigned_with_pubkey(
364 local_comms_name,
365 local_peer_id,
366 bytes,
367 local_peer_spec_address,
368 ),
369 None => meerkat_core::comms::TrustedPeerDescriptor::test_only_unsigned(
370 local_comms_name,
371 local_peer_id,
372 local_peer_spec_address,
373 ),
374 };
375 let spec = match spec_result {
376 Ok(spec) => spec,
377 Err(err) => {
378 return ControlResponse::Err {
379 code: "peer_spec".to_string(),
380 message: err,
381 };
382 }
383 };
384 let mid = crate::member_comms_id::mob_member_id(remote_member);
388 let result = if wire {
389 handle
390 .wire(mid, meerkat_mob::PeerTarget::External(spec))
391 .await
392 } else {
393 handle
394 .unwire(mid, meerkat_mob::PeerTarget::External(spec))
395 .await
396 };
397 match result {
398 Ok(()) => ControlResponse::Ok,
399 Err(err) => ControlResponse::Err {
400 code: "mob_error".to_string(),
401 message: err.to_string(),
402 },
403 }
404}
405
406async fn handle_inject(
407 handle: &meerkat_mob::MobHandle,
408 remote_member: &str,
409 content: serde_json::Value,
410) -> ControlResponse {
411 let content_input: meerkat_core::ContentInput = match serde_json::from_value(content) {
412 Ok(c) => c,
413 Err(err) => {
414 return ControlResponse::Err {
415 code: "decode".to_string(),
416 message: format!("content: {err}"),
417 };
418 }
419 };
420 let mid = crate::member_comms_id::mob_member_id(remote_member);
421 let member = match handle.member(&mid).await {
422 Ok(m) => m,
423 Err(err) => {
424 return ControlResponse::Err {
425 code: "unknown_member".to_string(),
426 message: err.to_string(),
427 };
428 }
429 };
430 if let Err(err) = member
431 .send(content_input, meerkat_core::types::HandlingMode::Queue)
432 .await
433 {
434 return ControlResponse::Err {
435 code: "mob_error".to_string(),
436 message: err.to_string(),
437 };
438 }
439 match handle.resolve_bridge_session_id(&mid).await {
440 Some(sid) => ControlResponse::Injected {
441 session_id: sid.to_string(),
442 },
443 None => ControlResponse::Err {
444 code: "no_session".to_string(),
445 message: format!("member '{remote_member}' has no bound bridge session"),
446 },
447 }
448}
449
450async fn handle_lookup_member(
451 handle: &meerkat_mob::MobHandle,
452 remote_member: &str,
453) -> ControlResponse {
454 let mid = crate::member_comms_id::mob_member_id(remote_member);
455 let mob_id = handle.mob_id().to_string();
456 let entry = match handle.get_member(&mid).await {
457 Ok(Some(e)) => e,
458 Ok(None) => {
459 return ControlResponse::Err {
460 code: "unknown_member".to_string(),
461 message: format!("member '{remote_member}' not in mob '{mob_id}'"),
462 };
463 }
464 Err(err) => {
466 return ControlResponse::Err {
467 code: "mob_error".to_string(),
468 message: err.to_string(),
469 };
470 }
471 };
472 let peer_id = match entry.peer_id() {
473 Some(p) => p.to_string(),
474 None => {
475 return ControlResponse::Err {
476 code: "no_comms".to_string(),
477 message: format!("member '{remote_member}' has no comms runtime"),
478 };
479 }
480 };
481 let comms_name = match meerkat_core::MemberCommsName::new(
487 mob_id.as_str(),
488 entry.role.as_str(),
489 mid.as_str(),
490 ) {
491 Ok(name) => name.to_string(),
492 Err(err) => {
493 return ControlResponse::Err {
494 code: "invalid_comms_name".to_string(),
495 message: format!(
496 "member '{remote_member}' in mob '{mob_id}' has an invalid comms name component: {err}"
497 ),
498 };
499 }
500 };
501 ControlResponse::Member {
502 peer_id,
503 comms_name,
504 }
505}
506
507pub async fn serve_tcp_control(listener: TcpListener, handler: std::sync::Arc<dyn ControlHandler>) {
513 loop {
514 let (stream, _peer_addr) = match listener.accept().await {
515 Ok(pair) => pair,
516 Err(err) => {
517 tracing::warn!(error = %err, "control listener accept failed; exiting");
518 return;
519 }
520 };
521 let handler = handler.clone();
522 tokio::spawn(serve_one_tcp(stream, handler));
523 }
524}
525
526#[cfg(unix)]
528pub async fn serve_uds_control(
529 listener: UnixListener,
530 handler: std::sync::Arc<dyn ControlHandler>,
531) {
532 loop {
533 let (stream, _peer_addr) = match listener.accept().await {
534 Ok(pair) => pair,
535 Err(err) => {
536 tracing::warn!(error = %err, "uds control listener accept failed; exiting");
537 return;
538 }
539 };
540 let handler = handler.clone();
541 tokio::spawn(serve_one_uds(stream, handler));
542 }
543}
544
545async fn serve_one_tcp(stream: TcpStream, handler: std::sync::Arc<dyn ControlHandler>) {
546 let mut s = ControlStream::Tcp(stream);
547 serve_one(&mut s, handler).await;
548}
549
550#[cfg(unix)]
551async fn serve_one_uds(stream: UnixStream, handler: std::sync::Arc<dyn ControlHandler>) {
552 let mut s = ControlStream::Uds(stream);
553 serve_one(&mut s, handler).await;
554}
555
556async fn serve_one(stream: &mut ControlStream, handler: std::sync::Arc<dyn ControlHandler>) {
557 let payload = match stream.read_frame().await {
558 Ok(buf) => buf,
559 Err(err) => {
560 tracing::debug!(error = %err, "control listener: read failed");
561 return;
562 }
563 };
564 let request = match serde_json::from_slice::<ControlRequest>(&payload) {
565 Ok(req) => req,
566 Err(err) => {
567 let response = ControlResponse::Err {
568 code: "decode".to_string(),
569 message: err.to_string(),
570 };
571 let response_payload = serde_json::to_vec(&response).unwrap_or_default();
572 let _ = stream.write_frame(&response_payload).await;
573 return;
574 }
575 };
576 let response = handler.handle(request).await;
577 let response_payload = serde_json::to_vec(&response).unwrap_or_default();
578 let _ = stream.write_frame(&response_payload).await;
579}
580
581#[cfg(test)]
582#[allow(clippy::expect_used, clippy::unwrap_used, clippy::panic)]
583mod tests {
584 use super::*;
585 use std::sync::Arc;
586
587 struct EchoHandler;
588
589 impl ControlHandler for EchoHandler {
590 fn handle(
591 &self,
592 request: ControlRequest,
593 ) -> std::pin::Pin<Box<dyn std::future::Future<Output = ControlResponse> + Send + '_>>
594 {
595 Box::pin(async move {
596 match request {
597 ControlRequest::Wire { .. } | ControlRequest::Unwire { .. } => {
598 ControlResponse::Ok
599 }
600 ControlRequest::Inject { remote_member, .. } => ControlResponse::Injected {
601 session_id: format!("session-for-{remote_member}"),
602 },
603 ControlRequest::LookupMember { remote_member } => ControlResponse::Member {
604 peer_id: format!("peer-id-for-{remote_member}"),
605 comms_name: format!("mob/role/{remote_member}"),
606 },
607 }
608 })
609 }
610 }
611
612 #[tokio::test]
613 async fn tcp_round_trip_inject() {
614 let listener = TcpListener::bind("127.0.0.1:0").await.expect("bind");
615 let addr = listener.local_addr().expect("addr");
616 let handler: Arc<dyn ControlHandler> = Arc::new(EchoHandler);
617 let server = tokio::spawn(serve_tcp_control(listener, handler));
618
619 let endpoint = RemoteEndpoint::Tcp(addr.to_string());
620 let request = ControlRequest::Inject {
621 remote_member: "alice".to_string(),
622 content: serde_json::json!({"text": "hello"}),
623 };
624 let response = RemoteControlClient::send(&endpoint, &request, DEFAULT_CONTROL_TIMEOUT)
625 .await
626 .expect("control rpc");
627 assert_eq!(
628 response,
629 ControlResponse::Injected {
630 session_id: "session-for-alice".to_string(),
631 },
632 );
633 server.abort();
634 }
635
636 #[tokio::test]
637 async fn tcp_round_trip_wire() {
638 let listener = TcpListener::bind("127.0.0.1:0").await.expect("bind");
639 let addr = listener.local_addr().expect("addr");
640 let handler: Arc<dyn ControlHandler> = Arc::new(EchoHandler);
641 let server = tokio::spawn(serve_tcp_control(listener, handler));
642
643 let endpoint = RemoteEndpoint::Tcp(addr.to_string());
644 let request = ControlRequest::Wire {
645 remote_member: "bob".to_string(),
646 local_peer_spec_address: "tcp://127.0.0.1:9001".to_string(),
647 local_comms_name: "demo/role/alice".to_string(),
648 local_peer_id: "00000000-0000-4000-8000-000000000001".to_string(),
649 local_pubkey_b64: None,
650 };
651 let response = RemoteControlClient::send(&endpoint, &request, DEFAULT_CONTROL_TIMEOUT)
652 .await
653 .expect("control rpc");
654 assert_eq!(response, ControlResponse::Ok);
655 server.abort();
656 }
657
658 #[tokio::test]
659 async fn malformed_request_returns_decode_error() {
660 let listener = TcpListener::bind("127.0.0.1:0").await.expect("bind");
661 let addr = listener.local_addr().expect("addr");
662 let handler: Arc<dyn ControlHandler> = Arc::new(EchoHandler);
663 let _server = tokio::spawn(serve_tcp_control(listener, handler));
664
665 let mut stream = TcpStream::connect(addr).await.expect("connect");
669 stream
670 .write_all(&u32::to_be_bytes(5))
671 .await
672 .expect("write header");
673 stream.write_all(b"hello").await.expect("write payload");
674 stream.flush().await.expect("flush");
675
676 let mut header = [0u8; 4];
677 stream.read_exact(&mut header).await.expect("read header");
678 let len = u32::from_be_bytes(header) as usize;
679 let mut buf = vec![0u8; len];
680 stream.read_exact(&mut buf).await.expect("read payload");
681 let response: ControlResponse = serde_json::from_slice(&buf).expect("decode response");
682 match response {
683 ControlResponse::Err { code, .. } => assert_eq!(code, "decode"),
684 other => panic!("expected decode error, got {other:?}"),
685 }
686 }
687}