1pub mod lease;
20
21use postcard_rpc::header::VarSeqKind;
22use postcard_rpc::host_client::{HostClient, WireRx, WireSpawn, WireTx};
23use postcard_rpc::standard_icd::{WireError, ERROR_PATH};
24use serde::{Deserialize, Serialize};
25use std::io::{Read, Write};
26use tokio::io::{AsyncRead, AsyncReadExt, AsyncWrite, AsyncWriteExt};
27use tokio::net::TcpStream;
28
29pub const NET_VERSION: u8 = 0;
31
32pub const MAX_FRAME: usize = 256 * 1024;
35
36#[derive(Serialize, Deserialize, Debug, Clone, Copy, PartialEq, Eq)]
38pub enum Role {
39 Node,
40 Lease,
41}
42
43#[derive(Serialize, Deserialize, Debug)]
44struct Hello {
45 version: u8,
46 role: Role,
47 token: String,
48}
49
50#[derive(Serialize, Deserialize, Debug)]
51enum HelloReply {
52 Ok { rig_name: String },
53 Denied,
54}
55
56pub async fn write_frame<W: AsyncWrite + Unpin>(w: &mut W, data: &[u8]) -> std::io::Result<()> {
59 debug_assert!(data.len() <= MAX_FRAME);
60 w.write_all(&(data.len() as u32).to_le_bytes()).await?;
61 w.write_all(data).await?;
62 w.flush().await
63}
64
65pub async fn read_frame<R: AsyncRead + Unpin>(r: &mut R) -> std::io::Result<Vec<u8>> {
66 let mut len = [0u8; 4];
67 r.read_exact(&mut len).await?;
68 let len = u32::from_le_bytes(len) as usize;
69 if len > MAX_FRAME {
70 return Err(std::io::Error::other(format!("frame of {len} bytes exceeds MAX_FRAME")));
71 }
72 let mut buf = vec![0u8; len];
73 r.read_exact(&mut buf).await?;
74 Ok(buf)
75}
76
77pub fn write_frame_sync<W: Write>(w: &mut W, data: &[u8]) -> std::io::Result<()> {
78 debug_assert!(data.len() <= MAX_FRAME);
79 w.write_all(&(data.len() as u32).to_le_bytes())?;
80 w.write_all(data)?;
81 w.flush()
82}
83
84pub fn read_frame_sync<R: Read>(r: &mut R) -> std::io::Result<Vec<u8>> {
85 let mut len = [0u8; 4];
86 r.read_exact(&mut len)?;
87 let len = u32::from_le_bytes(len) as usize;
88 if len > MAX_FRAME {
89 return Err(std::io::Error::other(format!("frame of {len} bytes exceeds MAX_FRAME")));
90 }
91 let mut buf = vec![0u8; len];
92 r.read_exact(&mut buf)?;
93 Ok(buf)
94}
95
96fn token_eq(a: &str, b: &str) -> bool {
99 let (a, b) = (a.as_bytes(), b.as_bytes());
100 if a.len() != b.len() {
101 return false;
102 }
103 a.iter().zip(b).fold(0u8, |acc, (x, y)| acc | (x ^ y)) == 0
104}
105
106async fn handshake_client(stream: &mut TcpStream, role: Role, token: &str) -> anyhow::Result<String> {
109 let hello = Hello { version: NET_VERSION, role, token: token.to_owned() };
110 write_frame(stream, &postcard::to_stdvec(&hello)?).await?;
111 let reply: HelloReply = postcard::from_bytes(&read_frame(stream).await?)?;
112 match reply {
113 HelloReply::Ok { rig_name } => Ok(rig_name),
114 HelloReply::Denied => anyhow::bail!("rig daemon denied the handshake (bad token?)"),
115 }
116}
117
118pub async fn connect_node(addr: &str, token: &str) -> anyhow::Result<HostClient<WireError>> {
122 let mut stream = TcpStream::connect(addr)
123 .await
124 .map_err(|e| anyhow::anyhow!("connecting to node at {addr}: {e}"))?;
125 stream.set_nodelay(true)?;
126 handshake_client(&mut stream, Role::Node, token)
127 .await
128 .map_err(|e| anyhow::anyhow!("node handshake with {addr}: {e}"))?;
129 let (rx, tx) = stream.into_split();
130 Ok(HostClient::new_with_wire(
131 TcpWireTx(tx),
132 TcpWireRx(rx),
133 TokioSpawn,
134 VarSeqKind::Seq2,
135 ERROR_PATH,
136 8,
137 ))
138}
139
140struct TcpWireTx(tokio::net::tcp::OwnedWriteHalf);
141
142impl WireTx for TcpWireTx {
143 type Error = std::io::Error;
144 async fn send(&mut self, data: Vec<u8>) -> Result<(), Self::Error> {
145 write_frame(&mut self.0, &data).await
146 }
147}
148
149struct TcpWireRx(tokio::net::tcp::OwnedReadHalf);
150
151impl WireRx for TcpWireRx {
152 type Error = std::io::Error;
153 async fn receive(&mut self) -> Result<Vec<u8>, Self::Error> {
154 read_frame(&mut self.0).await
155 }
156}
157
158struct TokioSpawn;
159
160impl WireSpawn for TokioSpawn {
161 fn spawn(&mut self, fut: impl std::future::Future<Output = ()> + Send + 'static) {
162 tokio::spawn(fut);
163 }
164}
165
166pub async fn handshake_server(
173 stream: &mut TcpStream,
174 token: &str,
175 rig_name: &str,
176) -> anyhow::Result<Role> {
177 stream.set_nodelay(true)?;
178 let hello: Hello = postcard::from_bytes(&read_frame(stream).await?)?;
179 if hello.version != NET_VERSION || !token_eq(&hello.token, token) {
180 write_frame(stream, &postcard::to_stdvec(&HelloReply::Denied)?).await?;
181 anyhow::bail!(
182 "handshake denied: version {} (want {NET_VERSION}), token {}",
183 hello.version,
184 if token_eq(&hello.token, token) { "ok" } else { "mismatch" },
185 );
186 }
187 let reply = HelloReply::Ok { rig_name: rig_name.to_owned() };
188 write_frame(stream, &postcard::to_stdvec(&reply)?).await?;
189 Ok(hello.role)
190}
191
192pub(crate) fn handshake_client_sync(
194 stream: &mut std::net::TcpStream,
195 role: Role,
196 token: &str,
197) -> anyhow::Result<()> {
198 let hello = Hello { version: NET_VERSION, role, token: token.to_owned() };
199 write_frame_sync(stream, &postcard::to_stdvec(&hello)?)?;
200 let reply: HelloReply = postcard::from_bytes(&read_frame_sync(stream)?)?;
201 match reply {
202 HelloReply::Ok { .. } => Ok(()),
203 HelloReply::Denied => anyhow::bail!("rig daemon denied the handshake (bad token?)"),
204 }
205}