-
-
Notifications
You must be signed in to change notification settings - Fork 14
Expand file tree
/
Copy pathframes.rs
More file actions
160 lines (143 loc) · 5.88 KB
/
Copy pathframes.rs
File metadata and controls
160 lines (143 loc) · 5.88 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
// SPDX-License-Identifier: BUSL-1.1
//! Native-protocol frame helpers: handshake, raw frame read/write, and a
//! JSON-session `send_sql` convenience for dispatch-routing tests.
use std::time::Duration;
use tokio::io::{AsyncReadExt, AsyncWriteExt};
use tokio::net::TcpStream;
use nodedb_types::protocol::request_fields::RequestFields;
use nodedb_types::protocol::text_fields::TextFields;
use nodedb_types::protocol::{
AuthMethod, FRAME_HEADER_LEN, HELLO_ACK_MAGIC, HELLO_ERROR_MAGIC_U32, HelloAckFrame,
HelloErrorFrame, HelloFrame, NativeRequest, NativeResponse, OpCode,
};
/// Perform the handshake with a custom `HelloFrame`.
/// Returns `(stream, ack_frame)` on success, or the parsed `HelloErrorFrame` via `Err`.
pub async fn do_handshake(
addr: std::net::SocketAddr,
hello: &HelloFrame,
) -> Result<(TcpStream, HelloAckFrame), HelloErrorFrame> {
let mut stream = TcpStream::connect(addr).await.expect("connect");
stream
.write_all(&hello.encode())
.await
.expect("write hello");
stream.flush().await.expect("flush");
let mut magic_buf = [0u8; 4];
stream.read_exact(&mut magic_buf).await.expect("read magic");
let magic = u32::from_be_bytes(magic_buf);
if magic == HELLO_ERROR_MAGIC_U32 {
// Read error code + msg_len + message.
let mut code_buf = [0u8; 1];
stream.read_exact(&mut code_buf).await.expect("read code");
let mut len_buf = [0u8; 1];
stream.read_exact(&mut len_buf).await.expect("read msg_len");
let msg_len = len_buf[0] as usize;
let mut msg = vec![0u8; msg_len];
if msg_len > 0 {
stream.read_exact(&mut msg).await.expect("read msg");
}
// Reassemble the full error frame bytes for HelloErrorFrame::decode.
let mut full = Vec::with_capacity(6 + msg_len);
full.extend_from_slice(b"NDBE");
full.push(code_buf[0]);
full.push(len_buf[0]);
full.extend_from_slice(&msg);
let err_frame = HelloErrorFrame::decode(&full).expect("decode error frame");
return Err(err_frame);
}
assert_eq!(magic, HELLO_ACK_MAGIC, "expected HelloAck magic");
// Read fixed rest: proto_version(2) + capabilities(8) + sv_len(1).
let mut fixed_rest = [0u8; 11];
stream
.read_exact(&mut fixed_rest)
.await
.expect("read fixed");
let sv_len = fixed_rest[10] as usize;
let var_len = sv_len + 1 + 7 * 5;
let mut var_buf = vec![0u8; var_len];
stream.read_exact(&mut var_buf).await.expect("read var");
let mut ack_buf = Vec::with_capacity(4 + 11 + var_len);
ack_buf.extend_from_slice(&magic_buf);
ack_buf.extend_from_slice(&fixed_rest);
ack_buf.extend_from_slice(&var_buf);
let ack = HelloAckFrame::decode(&ack_buf).expect("decode ack");
Ok((stream, ack))
}
/// Write a length-prefixed frame payload to the stream.
pub async fn write_frame(stream: &mut TcpStream, payload: &[u8]) {
let len = (payload.len() as u32).to_be_bytes();
stream.write_all(&len).await.expect("write len");
stream.write_all(payload).await.expect("write payload");
stream.flush().await.expect("flush");
}
/// Read a length-prefixed frame from the stream.
/// Returns `None` on EOF.
pub async fn read_frame(stream: &mut TcpStream) -> Option<Vec<u8>> {
let mut len_buf = [0u8; FRAME_HEADER_LEN];
match stream.read_exact(&mut len_buf).await {
Ok(_) => {}
Err(e) if e.kind() == std::io::ErrorKind::UnexpectedEof => return None,
Err(e) if e.kind() == std::io::ErrorKind::ConnectionReset => return None,
Err(e) => panic!("read_frame error: {e}"),
}
let payload_len = u32::from_be_bytes(len_buf) as usize;
let mut payload = vec![0u8; payload_len];
stream.read_exact(&mut payload).await.expect("read payload");
Some(payload)
}
/// Send any native-protocol request (opcode + `TextFields`) over an
/// established JSON-encoding session and decode the `NativeResponse`.
/// Assumes the session's first frame already selected JSON (see
/// `json_request_gets_json_response`) — callers that open a fresh connection
/// must send one JSON frame before calling this. Shared by [`send_sql`] and
/// any test driving a direct-op opcode (`PointGet`, `RangeScan`,
/// `VectorSearch`, `KvBatchPut`, ...) directly rather than through SQL text.
pub async fn send_request(
stream: &mut TcpStream,
seq: u64,
op: OpCode,
fields: TextFields,
) -> NativeResponse {
let req = NativeRequest {
op,
seq,
fields: RequestFields::Text(fields),
};
let json_bytes = sonic_rs::to_vec(&req).expect("json encode");
write_frame(stream, &json_bytes).await;
let response_payload = tokio::time::timeout(Duration::from_secs(5), read_frame(stream))
.await
.expect("timeout waiting for response")
.expect("response frame");
sonic_rs::from_slice(&response_payload).expect("json decode NativeResponse")
}
/// Authenticate a fresh native connection with an API key. The JSON Auth
/// request also selects JSON framing for the rest of the session.
pub async fn send_api_key_auth(stream: &mut TcpStream, seq: u64, token: String) -> NativeResponse {
send_request(
stream,
seq,
OpCode::Auth,
TextFields {
auth: Some(AuthMethod::ApiKey { token }),
..Default::default()
},
)
.await
}
/// Send a `SHOW`/SQL statement over an established JSON-encoding session and
/// decode the `NativeResponse`. Assumes the session's first frame already
/// selected JSON (see `json_request_gets_json_response`) — callers that open
/// a fresh connection must send one JSON frame before calling this.
pub async fn send_sql(stream: &mut TcpStream, seq: u64, sql: &str) -> NativeResponse {
send_request(
stream,
seq,
OpCode::Sql,
TextFields {
sql: Some(sql.into()),
..Default::default()
},
)
.await
}