use crate::{
traits::PreimageFetcher, PipeHandle, PreimageKey, PreimageOracleClient, PreimageOracleServer,
};
use alloc::{boxed::Box, vec::Vec};
use anyhow::{bail, Result};
use tracing::trace;
#[derive(Debug, Clone, Copy)]
pub struct OracleReader {
pipe_handle: PipeHandle,
}
impl OracleReader {
pub const fn new(pipe_handle: PipeHandle) -> Self {
Self { pipe_handle }
}
async fn write_key(&self, key: PreimageKey) -> Result<usize> {
let key_bytes: [u8; 32] = key.into();
self.pipe_handle.write(&key_bytes).await?;
let mut length_buffer = [0u8; 8];
self.pipe_handle.read_exact(&mut length_buffer).await?;
Ok(u64::from_be_bytes(length_buffer) as usize)
}
}
#[async_trait::async_trait]
impl PreimageOracleClient for OracleReader {
async fn get(&self, key: PreimageKey) -> Result<Vec<u8>> {
trace!(target: "oracle_client", "Requesting data from preimage oracle. Key {key}");
let length = self.write_key(key).await?;
if length == 0 {
return Ok(Default::default());
}
let mut data_buffer = alloc::vec![0; length];
trace!(target: "oracle_client", "Reading data from preimage oracle. Key {key}");
self.pipe_handle.read_exact(&mut data_buffer).await?;
trace!(target: "oracle_client", "Successfully read data from preimage oracle. Key: {key}");
Ok(data_buffer)
}
async fn get_exact(&self, key: PreimageKey, buf: &mut [u8]) -> Result<()> {
trace!(target: "oracle_client", "Requesting data from preimage oracle. Key {key}");
let length = self.write_key(key).await?;
trace!(target: "oracle_client", "Reading data from preimage oracle. Key {key}");
if buf.len() != length {
bail!("Buffer size {} does not match preimage size {}", buf.len(), length);
}
if length == 0 {
return Ok(());
}
self.pipe_handle.read_exact(buf).await?;
trace!(target: "oracle_client", "Successfully read data from preimage oracle. Key: {key}");
Ok(())
}
}
#[derive(Debug, Clone, Copy)]
pub struct OracleServer {
pipe_handle: PipeHandle,
}
impl OracleServer {
pub fn new(pipe_handle: PipeHandle) -> Self {
Self { pipe_handle }
}
}
#[async_trait::async_trait]
impl PreimageOracleServer for OracleServer {
async fn next_preimage_request<F>(&self, fetcher: &F) -> Result<()>
where
F: PreimageFetcher + Send + Sync,
{
let mut buf = [0u8; 32];
self.pipe_handle.read_exact(&mut buf).await?;
let preimage_key = PreimageKey::try_from(buf)?;
trace!(target: "oracle_server", "Fetching preimage for key {preimage_key}");
let value = fetcher.get_preimage(preimage_key).await?;
let data = [(value.len() as u64).to_be_bytes().as_ref(), value.as_ref()]
.into_iter()
.flatten()
.copied()
.collect::<Vec<_>>();
self.pipe_handle.write(data.as_slice()).await?;
trace!(target: "oracle_server", "Successfully wrote preimage data for key {preimage_key}");
Ok(())
}
}
#[cfg(test)]
mod test {
extern crate std;
use super::*;
use crate::{test_utils::bidirectional_pipe, PreimageKeyType};
use alloc::sync::Arc;
use alloy_primitives::keccak256;
use anyhow::anyhow;
use kona_common::FileDescriptor;
use std::{collections::HashMap, os::fd::AsRawFd};
use tokio::sync::Mutex;
struct TestFetcher {
preimages: Arc<Mutex<HashMap<PreimageKey, Vec<u8>>>>,
}
#[async_trait::async_trait]
impl PreimageFetcher for TestFetcher {
async fn get_preimage(&self, key: PreimageKey) -> Result<Vec<u8>> {
let read_lock = self.preimages.lock().await;
read_lock.get(&key).cloned().ok_or_else(|| anyhow!("Key not found"))
}
}
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
async fn test_oracle_client_and_host() {
const MOCK_DATA_A: &[u8] = b"1234567890";
const MOCK_DATA_B: &[u8] = b"FACADE";
let key_a: PreimageKey =
PreimageKey::new(*keccak256(MOCK_DATA_A), PreimageKeyType::Keccak256);
let key_b: PreimageKey =
PreimageKey::new(*keccak256(MOCK_DATA_B), PreimageKeyType::Keccak256);
let preimages = {
let mut preimages = HashMap::new();
preimages.insert(key_a, MOCK_DATA_A.to_vec());
preimages.insert(key_b, MOCK_DATA_B.to_vec());
Arc::new(Mutex::new(preimages))
};
let preimage_pipe = bidirectional_pipe().unwrap();
let client = tokio::task::spawn(async move {
let oracle_reader = OracleReader::new(PipeHandle::new(
FileDescriptor::Wildcard(preimage_pipe.client.read.as_raw_fd() as usize),
FileDescriptor::Wildcard(preimage_pipe.client.write.as_raw_fd() as usize),
));
let contents_a = oracle_reader.get(key_a).await.unwrap();
let contents_b = oracle_reader.get(key_b).await.unwrap();
(contents_a, contents_b)
});
tokio::task::spawn(async move {
let oracle_server = OracleServer::new(PipeHandle::new(
FileDescriptor::Wildcard(preimage_pipe.host.read.as_raw_fd() as usize),
FileDescriptor::Wildcard(preimage_pipe.host.write.as_raw_fd() as usize),
));
let test_fetcher = TestFetcher { preimages: Arc::clone(&preimages) };
loop {
if oracle_server.next_preimage_request(&test_fetcher).await.is_err() {
break;
}
}
});
let (c,) = tokio::join!(client);
let (contents_a, contents_b) = c.unwrap();
assert_eq!(contents_a, MOCK_DATA_A);
assert_eq!(contents_b, MOCK_DATA_B);
}
}