kona_preimage/
oracle.rs

1use crate::{
2    PreimageKey, PreimageOracleClient, PreimageOracleServer,
3    errors::{PreimageOracleError, PreimageOracleResult},
4    traits::{Channel, PreimageFetcher},
5};
6use alloc::{boxed::Box, vec::Vec};
7
8/// An [OracleReader] is a high-level interface to the preimage oracle channel.
9#[derive(Debug, Clone, Copy)]
10pub struct OracleReader<C> {
11    channel: C,
12}
13
14impl<C> OracleReader<C>
15where
16    C: Channel,
17{
18    /// Create a new [OracleReader] from a [Channel].
19    pub const fn new(channel: C) -> Self {
20        Self { channel }
21    }
22
23    /// Set the preimage key for the global oracle reader. This will overwrite any existing key, and
24    /// block until the host has prepared the preimage and responded with the length of the
25    /// preimage.
26    async fn write_key(&self, key: PreimageKey) -> PreimageOracleResult<usize> {
27        // Write the key to the host so that it can prepare the preimage.
28        let key_bytes: [u8; 32] = key.into();
29        self.channel.write(&key_bytes).await?;
30
31        // Read the length prefix and reset the cursor.
32        let mut length_buffer = [0u8; 8];
33        self.channel.read_exact(&mut length_buffer).await?;
34        Ok(u64::from_be_bytes(length_buffer) as usize)
35    }
36}
37
38#[async_trait::async_trait]
39impl<C> PreimageOracleClient for OracleReader<C>
40where
41    C: Channel + Send + Sync,
42{
43    /// Get the data corresponding to the currently set key from the host. Return the data in a new
44    /// heap allocated `Vec<u8>`
45    async fn get(&self, key: PreimageKey) -> PreimageOracleResult<Vec<u8>> {
46        trace!(target: "oracle_client", "Requesting data from preimage oracle. Key {key}");
47
48        let length = self.write_key(key).await?;
49
50        if length == 0 {
51            return Ok(Default::default());
52        }
53
54        let mut data_buffer = alloc::vec![0; length];
55
56        trace!(target: "oracle_client", "Reading data from preimage oracle. Key {key}");
57
58        // Grab a read lock on the preimage channel to read the data.
59        self.channel.read_exact(&mut data_buffer).await?;
60
61        trace!(target: "oracle_client", "Successfully read data from preimage oracle. Key: {key}");
62
63        Ok(data_buffer)
64    }
65
66    /// Get the data corresponding to the currently set key from the host. Write the data into the
67    /// provided buffer
68    async fn get_exact(&self, key: PreimageKey, buf: &mut [u8]) -> PreimageOracleResult<()> {
69        trace!(target: "oracle_client", "Requesting data from preimage oracle. Key {key}");
70
71        // Write the key to the host and read the length of the preimage.
72        let length = self.write_key(key).await?;
73
74        trace!(target: "oracle_client", "Reading data from preimage oracle. Key {key}");
75
76        // Ensure the buffer is the correct size.
77        if buf.len() != length {
78            return Err(PreimageOracleError::BufferLengthMismatch(length, buf.len()));
79        }
80
81        if length == 0 {
82            return Ok(());
83        }
84
85        self.channel.read_exact(buf).await?;
86
87        trace!(target: "oracle_client", "Successfully read data from preimage oracle. Key: {key}");
88
89        Ok(())
90    }
91}
92
93/// An [OracleServer] is a router for the host to serve data back to the client [OracleReader].
94#[derive(Debug, Clone, Copy)]
95pub struct OracleServer<C> {
96    channel: C,
97}
98
99impl<C> OracleServer<C>
100where
101    C: Channel,
102{
103    /// Create a new [OracleServer] from a [Channel].
104    pub const fn new(chanel: C) -> Self {
105        Self { channel: chanel }
106    }
107}
108
109#[async_trait::async_trait]
110impl<C> PreimageOracleServer for OracleServer<C>
111where
112    C: Channel + Send + Sync,
113{
114    async fn next_preimage_request<F>(&self, fetcher: &F) -> Result<(), PreimageOracleError>
115    where
116        F: PreimageFetcher + Send + Sync,
117    {
118        // Read the preimage request from the client, and throw early if there isn't is any.
119        let mut buf = [0u8; 32];
120        self.channel.read_exact(&mut buf).await?;
121        let preimage_key = PreimageKey::try_from(buf)?;
122
123        trace!(target: "oracle_server", "Fetching preimage for key {preimage_key}");
124
125        // Fetch the preimage value from the preimage getter.
126        let value = fetcher.get_preimage(preimage_key).await?;
127
128        // Write the length as a big-endian u64 followed by the data.
129        self.channel.write(value.len().to_be_bytes().as_ref()).await?;
130        self.channel.write(value.as_ref()).await?;
131
132        trace!(target: "oracle_server", "Successfully wrote preimage data for key {preimage_key}");
133
134        Ok(())
135    }
136}
137
138#[cfg(test)]
139mod test {
140    use super::*;
141    use crate::{PreimageKeyType, native_channel::BidirectionalChannel};
142    use alloc::sync::Arc;
143    use alloy_primitives::keccak256;
144    use std::collections::HashMap;
145    use tokio::sync::Mutex;
146
147    struct TestFetcher {
148        preimages: Arc<Mutex<HashMap<PreimageKey, Vec<u8>>>>,
149    }
150
151    #[async_trait::async_trait]
152    impl PreimageFetcher for TestFetcher {
153        async fn get_preimage(&self, key: PreimageKey) -> PreimageOracleResult<Vec<u8>> {
154            let read_lock = self.preimages.lock().await;
155            read_lock.get(&key).cloned().ok_or(PreimageOracleError::KeyNotFound)
156        }
157    }
158
159    #[tokio::test(flavor = "multi_thread", worker_threads = 2)]
160    async fn test_oracle_reader_get_exact() {
161        const MOCK_DATA_A: &[u8] = b"1234567890";
162        const MOCK_DATA_B: &[u8] = b"FACADE";
163        let key_a: PreimageKey =
164            PreimageKey::new(*keccak256(MOCK_DATA_A), PreimageKeyType::Keccak256);
165        let key_b: PreimageKey =
166            PreimageKey::new(*keccak256(MOCK_DATA_B), PreimageKeyType::Keccak256);
167
168        let preimages = {
169            let mut preimages = HashMap::default();
170            preimages.insert(key_a, MOCK_DATA_A.to_vec());
171            preimages.insert(key_b, MOCK_DATA_B.to_vec());
172            Arc::new(Mutex::new(preimages))
173        };
174
175        let preimage_channel = BidirectionalChannel::new().unwrap();
176
177        let client = tokio::task::spawn(async move {
178            let oracle_reader = OracleReader::new(preimage_channel.client);
179            let mut contents_a = [0u8; 10];
180            let mut contents_b = [0u8; 6];
181            oracle_reader.get_exact(key_a, &mut contents_a).await.unwrap();
182            oracle_reader.get_exact(key_b, &mut contents_b).await.unwrap();
183
184            (contents_a, contents_b)
185        });
186        tokio::task::spawn(async move {
187            let oracle_server = OracleServer::new(preimage_channel.host);
188            let test_fetcher = TestFetcher { preimages: Arc::clone(&preimages) };
189
190            loop {
191                match oracle_server.next_preimage_request(&test_fetcher).await {
192                    Err(PreimageOracleError::IOError(_)) => break,
193                    Err(e) => panic!("Unexpected error: {:?}", e),
194                    Ok(_) => {}
195                }
196            }
197        });
198
199        let (c,) = tokio::join!(client);
200        let (contents_a, contents_b) = c.unwrap();
201        assert_eq!(contents_a, MOCK_DATA_A);
202        assert_eq!(contents_b, MOCK_DATA_B);
203    }
204
205    #[tokio::test(flavor = "multi_thread", worker_threads = 2)]
206    async fn test_oracle_client_and_host() {
207        const MOCK_DATA_A: &[u8] = b"1234567890";
208        const MOCK_DATA_B: &[u8] = b"FACADE";
209        let key_a: PreimageKey =
210            PreimageKey::new(*keccak256(MOCK_DATA_A), PreimageKeyType::Keccak256);
211        let key_b: PreimageKey =
212            PreimageKey::new(*keccak256(MOCK_DATA_B), PreimageKeyType::Keccak256);
213
214        let preimages = {
215            let mut preimages = HashMap::default();
216            preimages.insert(key_a, MOCK_DATA_A.to_vec());
217            preimages.insert(key_b, MOCK_DATA_B.to_vec());
218            Arc::new(Mutex::new(preimages))
219        };
220
221        let preimage_channel = BidirectionalChannel::new().unwrap();
222
223        let client = tokio::task::spawn(async move {
224            let oracle_reader = OracleReader::new(preimage_channel.client);
225            let contents_a = oracle_reader.get(key_a).await.unwrap();
226            let contents_b = oracle_reader.get(key_b).await.unwrap();
227
228            (contents_a, contents_b)
229        });
230        tokio::task::spawn(async move {
231            let oracle_server = OracleServer::new(preimage_channel.host);
232            let test_fetcher = TestFetcher { preimages: Arc::clone(&preimages) };
233
234            loop {
235                match oracle_server.next_preimage_request(&test_fetcher).await {
236                    Err(PreimageOracleError::IOError(_)) => break,
237                    Err(e) => panic!("Unexpected error: {:?}", e),
238                    Ok(_) => {}
239                }
240            }
241        });
242
243        let (c,) = tokio::join!(client);
244        let (contents_a, contents_b) = c.unwrap();
245        assert_eq!(contents_a, MOCK_DATA_A);
246        assert_eq!(contents_b, MOCK_DATA_B);
247    }
248}