1use sha2::{Digest, Sha384};
6
7use crate::cbor::Value;
8use crate::node_key::{node_id_of, NodeKey};
9use crate::profile::Profile;
10use crate::signed_object::{sign_object, verify_object};
11
12use super::{
13 bounded_text, check_payload, entry, fixed, has_fields, identity_signer, object_refusal,
14 protocol_uint, read_fields, received_frame, text_of, uint, FrameError, Rule, StreamMode,
15 MAX_PROCEDURE_BYTES, MAX_PROTOCOL_INT, PROTOCOL_VERSION, REQUEST_LABEL,
16};
17
18pub const MAX_PROOFS: usize = 8;
21pub const MAX_PROOFS_BYTES: usize = 256 * 1024;
22
23const CALL: &str = "call";
24const STREAM_OPEN: &str = "stream_open";
25
26#[derive(Debug, Clone, Copy, PartialEq, Eq)]
28pub enum RequestType {
29 Call,
30 StreamOpen,
31}
32
33impl RequestType {
34 fn name(self) -> &'static str {
35 match self {
36 RequestType::Call => CALL,
37 RequestType::StreamOpen => STREAM_OPEN,
38 }
39 }
40}
41
42#[derive(Debug, Clone, PartialEq)]
48pub struct RequestSpec {
49 pub request_id: [u8; 16],
50 pub realm: [u8; 32],
51 pub procedure: String,
52 pub target: [u8; 32],
53 pub deadline: u64,
54 pub payload: Value,
55 pub mode: Option<StreamMode>,
56 pub token: Option<Vec<u8>>,
57 pub proofs: Vec<Vec<u8>>,
58 pub source_route: Option<Vec<u8>>,
59 pub retry_budget: Option<u64>,
60}
61
62#[derive(Debug, Clone, PartialEq)]
66pub struct VerifiedRequest {
67 pub frame_type: RequestType,
68 pub key: Vec<u8>,
69 pub request_hash: [u8; 48],
70 pub caller: [u8; 32],
71 pub request_id: [u8; 16],
72 pub realm: [u8; 32],
73 pub procedure: String,
74 pub target: [u8; 32],
75 pub deadline: u64,
76 pub payload: Value,
77 pub mode: Option<StreamMode>,
78 pub token: Option<Vec<u8>>,
79 pub proofs: Option<Vec<Vec<u8>>>,
80}
81
82pub fn sign_call(spec: &RequestSpec, key: &NodeKey) -> Result<Value, FrameError> {
88 sign_request(RequestType::Call, spec, key)
89}
90
91pub fn sign_stream_open(spec: &RequestSpec, key: &NodeKey) -> Result<Value, FrameError> {
94 sign_request(RequestType::StreamOpen, spec, key)
95}
96
97fn sign_request(
98 frame_type: RequestType,
99 spec: &RequestSpec,
100 key: &NodeKey,
101) -> Result<Value, FrameError> {
102 identity_signer(key)?;
103 bounded_text("procedure", spec.procedure.as_bytes(), MAX_PROCEDURE_BYTES)?;
104 check_payload(&spec.payload)?;
105 if spec.deadline >= MAX_PROTOCOL_INT || spec.retry_budget.is_some_and(|b| b >= MAX_PROTOCOL_INT)
106 {
107 return Err(FrameError::OutOfRange(
108 "a deadline or retry budget of 2^53 or more".into(),
109 ));
110 }
111 match (frame_type, spec.mode) {
112 (RequestType::Call, Some(_)) => {
113 return Err(FrameError::OutOfRange(
114 "a CALL carries no stream mode".into(),
115 ))
116 }
117 (RequestType::StreamOpen, None) => {
118 return Err(FrameError::OutOfRange(
119 "a STREAM_OPEN carries one of the three stream modes".into(),
120 ))
121 }
122 _ => {}
123 }
124 let proofs = proofs_value(&spec.proofs);
125 if !proofs_within_bound(&proofs) {
126 return Err(FrameError::ProofsOutOfBound);
127 }
128 let mut fields = vec![
129 entry("frame_type", Value::text(frame_type.name())),
130 entry("caller", Value::Bytes(key.key_id().to_vec())),
131 entry("request_id", Value::Bytes(spec.request_id.to_vec())),
132 entry("realm", Value::Bytes(spec.realm.to_vec())),
133 entry("procedure", Value::text(spec.procedure.clone())),
134 entry("target", Value::Bytes(spec.target.to_vec())),
135 entry("deadline", uint(spec.deadline)),
136 entry("payload", spec.payload.clone()),
137 ];
138 if let Some(mode) = spec.mode {
139 fields.push(entry("mode", Value::text(mode.name())));
140 }
141 if let Some(token) = &spec.token {
142 fields.push(entry("token", Value::Bytes(token.clone())));
143 }
144 if !spec.proofs.is_empty() {
145 fields.push(entry("proofs", proofs));
146 }
147 let request = sign_object(REQUEST_LABEL, &fields, key).map_err(object_refusal)?;
148 let mut frame = vec![
149 entry("version", Value::Int(i128::from(PROTOCOL_VERSION))),
150 entry("frame_type", Value::text(frame_type.name())),
151 entry("request", request.to_value()),
152 ];
153 if let Some(route) = &spec.source_route {
154 frame.push(entry("source_route", Value::Bytes(route.clone())));
155 }
156 if let Some(budget) = spec.retry_budget {
157 frame.push(entry("retry_budget", uint(budget)));
158 }
159 Ok(Value::Map(frame))
160}
161
162const REQUEST_ROUTES: &[(&str, Rule)] = &[
163 ("source_route", Rule::AnyBytes),
164 ("retry_budget", Rule::ProtocolUint),
165];
166
167pub fn verify_request(frame: &Value, profile: Profile) -> Result<VerifiedRequest, FrameError> {
173 let (frame_type, object) = received_frame(
174 frame,
175 "request",
176 Rule::CarriedObject,
177 REQUEST_ROUTES,
178 &[CALL, STREAM_OPEN],
179 )
180 .ok_or(FrameError::Malformed)?;
181 let frame_type = if frame_type == CALL {
182 RequestType::Call
183 } else {
184 RequestType::StreamOpen
185 };
186 let verified = verify_object(REQUEST_LABEL, &object, profile).map_err(object_refusal)?;
187 let fields =
188 read_fields(&verified.fields, &request_table(frame_type)).ok_or(FrameError::Malformed)?;
189 let has_mode = fields.contains_key("mode");
190 if !has_fields(
191 &fields,
192 &[
193 "frame_type",
194 "caller",
195 "request_id",
196 "realm",
197 "procedure",
198 "target",
199 "deadline",
200 "payload",
201 ],
202 ) || has_mode != (frame_type == RequestType::StreamOpen)
203 {
204 return Err(FrameError::Malformed);
205 }
206 let request = VerifiedRequest {
207 frame_type,
208 request_hash: Sha384::digest(&verified.tbs).into(),
209 caller: fixed(&fields["caller"]),
210 request_id: fixed(&fields["request_id"]),
211 realm: fixed(&fields["realm"]),
212 procedure: text_of(&fields["procedure"]),
213 target: fixed(&fields["target"]),
214 deadline: protocol_uint(&fields["deadline"]).unwrap_or(0),
215 payload: fields["payload"].clone(),
216 mode: fields
217 .get("mode")
218 .and_then(|m| StreamMode::parse(&text_of(m))),
219 token: fields.get("token").map(super::bytes_of),
220 proofs: fields.get("proofs").map(|p| match p {
221 Value::List(items) => items.iter().map(super::bytes_of).collect(),
222 _ => Vec::new(),
223 }),
224 key: verified.key,
225 };
226 if request.caller != node_id_of(&request.key, profile) {
227 return Err(FrameError::KeyIdMismatch);
228 }
229 Ok(request)
230}
231
232fn request_table(frame_type: RequestType) -> Vec<(&'static str, Rule)> {
233 vec![
234 (
235 "frame_type",
236 Rule::TextIn(match frame_type {
237 RequestType::Call => &[CALL],
238 RequestType::StreamOpen => &[STREAM_OPEN],
239 }),
240 ),
241 ("alg", Rule::Any),
242 ("caller", Rule::BytesOf(32)),
243 ("request_id", Rule::BytesOf(16)),
244 ("realm", Rule::BytesOf(32)),
245 ("procedure", Rule::TextWithin(MAX_PROCEDURE_BYTES)),
246 ("target", Rule::BytesOf(32)),
247 ("deadline", Rule::ProtocolUint),
248 ("payload", Rule::Any),
249 (
250 "mode",
251 Rule::TextIn(&["server_stream", "client_stream", "bidi"]),
252 ),
253 ("token", Rule::AnyBytes),
254 ("proofs", Rule::Proofs),
255 ]
256}
257
258pub fn request_fields_accepted(fields: &Value) -> bool {
262 read_fields(fields, &request_table(RequestType::Call)).is_some()
263}
264
265fn proofs_value(proofs: &[Vec<u8>]) -> Value {
266 Value::List(proofs.iter().map(|p| Value::Bytes(p.clone())).collect())
267}
268
269pub(super) fn proofs_within_bound(v: &Value) -> bool {
272 let Value::List(items) = v else {
273 return false;
274 };
275 if items.len() > MAX_PROOFS {
276 return false;
277 }
278 let mut seen = std::collections::HashSet::with_capacity(items.len());
279 let mut total = 0;
280 for item in items {
281 let Value::Bytes(b) = item else {
282 return false;
283 };
284 if !seen.insert(b.as_slice()) {
285 return false;
286 }
287 total += b.len();
288 }
289 total <= MAX_PROOFS_BYTES
290}