use crate::{
traits::{HintRouter, HintWriterClient},
HintReaderServer, PipeHandle,
};
use alloc::{boxed::Box, string::String, vec};
use anyhow::Result;
use async_trait::async_trait;
use tracing::{error, trace};
#[derive(Debug, Clone, Copy)]
pub struct HintWriter {
pipe_handle: PipeHandle,
}
impl HintWriter {
pub const fn new(pipe_handle: PipeHandle) -> Self {
Self { pipe_handle }
}
}
#[async_trait]
impl HintWriterClient for HintWriter {
async fn write(&self, hint: &str) -> Result<()> {
let mut hint_bytes = vec![0u8; hint.len() + 4];
hint_bytes[0..4].copy_from_slice(u32::to_be_bytes(hint.len() as u32).as_ref());
hint_bytes[4..].copy_from_slice(hint.as_bytes());
trace!(target: "hint_writer", "Writing hint \"{hint}\"");
self.pipe_handle.write(&hint_bytes).await?;
trace!(target: "hint_writer", "Successfully wrote hint");
let mut hint_ack = [0u8; 1];
self.pipe_handle.read_exact(&mut hint_ack).await?;
trace!(target: "hint_writer", "Received hint acknowledgement");
Ok(())
}
}
#[derive(Debug, Clone, Copy)]
pub struct HintReader {
pipe_handle: PipeHandle,
}
impl HintReader {
pub fn new(pipe_handle: PipeHandle) -> Self {
Self { pipe_handle }
}
}
#[async_trait]
impl HintReaderServer for HintReader {
async fn next_hint<R>(&self, hint_router: &R) -> Result<()>
where
R: HintRouter + Send + Sync,
{
let mut len_buf = [0u8; 4];
self.pipe_handle.read_exact(&mut len_buf).await?;
let len = u32::from_be_bytes(len_buf);
let mut raw_payload = vec![0u8; len as usize];
self.pipe_handle.read_exact(raw_payload.as_mut_slice()).await?;
let payload = String::from_utf8(raw_payload)
.map_err(|e| anyhow::anyhow!("Failed to decode hint payload: {e}"))?;
trace!(target: "hint_reader", "Successfully read hint: \"{payload}\"");
if let Err(e) = hint_router.route_hint(payload).await {
self.pipe_handle.write(&[0x00]).await?;
error!("Failed to route hint: {e}");
anyhow::bail!("Failed to rout hint: {e}");
}
self.pipe_handle.write(&[0x00]).await?;
trace!(target: "hint_reader", "Successfully routed and acknowledged hint");
Ok(())
}
}
#[cfg(test)]
mod test {
extern crate std;
use super::*;
use crate::test_utils::bidirectional_pipe;
use alloc::{sync::Arc, vec::Vec};
use kona_common::FileDescriptor;
use std::os::fd::AsRawFd;
use tokio::sync::Mutex;
struct TestRouter {
incoming_hints: Arc<Mutex<Vec<String>>>,
}
#[async_trait]
impl HintRouter for TestRouter {
async fn route_hint(&self, hint: String) -> Result<()> {
self.incoming_hints.lock().await.push(hint);
Ok(())
}
}
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
async fn test_hint_client_and_host() {
const MOCK_DATA: &str = "test-hint 0xfacade";
let incoming_hints = Arc::new(Mutex::new(Vec::new()));
let hint_pipe = bidirectional_pipe().unwrap();
let client = tokio::task::spawn(async move {
let hint_writer = HintWriter::new(PipeHandle::new(
FileDescriptor::Wildcard(hint_pipe.client.read.as_raw_fd() as usize),
FileDescriptor::Wildcard(hint_pipe.client.write.as_raw_fd() as usize),
));
hint_writer.write(MOCK_DATA).await
});
let host = tokio::task::spawn({
let incoming_hints_ref = Arc::clone(&incoming_hints);
async move {
let router = TestRouter { incoming_hints: incoming_hints_ref };
let hint_reader = HintReader::new(PipeHandle::new(
FileDescriptor::Wildcard(hint_pipe.host.read.as_raw_fd() as usize),
FileDescriptor::Wildcard(hint_pipe.host.write.as_raw_fd() as usize),
));
hint_reader.next_hint(&router).await.unwrap();
}
});
let _ = tokio::join!(client, host);
let mut hints = incoming_hints.lock().await;
assert_eq!(hints.len(), 1);
let h = hints.remove(0);
assert_eq!(h, MOCK_DATA);
}
}