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::sealed::{clear_or_sealed, sealed_field, Sealed, SealedContext};
13use super::{
14 bounded_text, check_payload, entry, fixed, has_fields, identity_signer, object_refusal,
15 protocol_uint, read_fields, received_frame, text_of, uint, FrameError, Rule, StreamMode,
16 MAX_PROCEDURE_BYTES, MAX_PROTOCOL_INT, PROTOCOL_VERSION, REQUEST_LABEL,
17};
18
19pub const MAX_PROOFS: usize = 8;
22pub const MAX_PROOFS_BYTES: usize = 256 * 1024;
23
24const CALL: &str = "call";
25const STREAM_OPEN: &str = "stream_open";
26
27#[derive(Debug, Clone, Copy, PartialEq, Eq)]
29pub enum RequestType {
30 Call,
31 StreamOpen,
32}
33
34impl RequestType {
35 fn name(self) -> &'static str {
36 match self {
37 RequestType::Call => CALL,
38 RequestType::StreamOpen => STREAM_OPEN,
39 }
40 }
41}
42
43#[derive(Debug, Clone, PartialEq)]
50pub struct RequestSpec {
51 pub request_id: [u8; 16],
52 pub realm: [u8; 32],
53 pub procedure: String,
54 pub target: [u8; 32],
55 pub deadline: u64,
56 pub payload: Value,
57 pub sealed: Option<Sealed>,
58 pub mode: Option<StreamMode>,
59 pub token: Option<Vec<u8>>,
60 pub proofs: Vec<Vec<u8>>,
61 pub source_route: Option<Vec<u8>>,
62 pub retry_budget: Option<u64>,
63}
64
65#[derive(Debug, Clone, PartialEq)]
70pub struct VerifiedRequest {
71 pub frame_type: RequestType,
72 pub key: Vec<u8>,
73 pub request_hash: [u8; 48],
74 pub caller: [u8; 32],
75 pub request_id: [u8; 16],
76 pub realm: [u8; 32],
77 pub procedure: String,
78 pub target: [u8; 32],
79 pub deadline: u64,
80 pub payload: Value,
81 pub sealed: Option<Sealed>,
82 pub mode: Option<StreamMode>,
83 pub token: Option<Vec<u8>>,
84 pub proofs: Option<Vec<Vec<u8>>>,
85}
86
87pub fn sign_call(spec: &RequestSpec, key: &NodeKey) -> Result<Value, FrameError> {
94 sign_request(RequestType::Call, spec, key)
95}
96
97pub fn sign_stream_open(spec: &RequestSpec, key: &NodeKey) -> Result<Value, FrameError> {
100 sign_request(RequestType::StreamOpen, spec, key)
101}
102
103fn sign_request(
104 frame_type: RequestType,
105 spec: &RequestSpec,
106 key: &NodeKey,
107) -> Result<Value, FrameError> {
108 let proofs = check_request(frame_type, spec, key)?;
109 let fields = request_fields(frame_type, spec, key, proofs);
110 let request = sign_object(REQUEST_LABEL, &fields, key).map_err(object_refusal)?;
111 Ok(request_frame(frame_type, spec, request.to_value()))
112}
113
114fn check_request(
117 frame_type: RequestType,
118 spec: &RequestSpec,
119 key: &NodeKey,
120) -> Result<Value, FrameError> {
121 identity_signer(key)?;
122 bounded_text("procedure", spec.procedure.as_bytes(), MAX_PROCEDURE_BYTES)?;
123 match &spec.sealed {
124 Some(sealed) if !sealed.shaped(SealedContext::Request) => {
125 return Err(FrameError::SealedShape)
126 }
127 Some(_) => {}
128 None => check_payload(&spec.payload)?,
129 }
130 if spec.deadline >= MAX_PROTOCOL_INT || spec.retry_budget.is_some_and(|b| b >= MAX_PROTOCOL_INT)
131 {
132 return Err(FrameError::OutOfRange(
133 "a deadline or retry budget of 2^53 or more".into(),
134 ));
135 }
136 match (frame_type, spec.mode) {
137 (RequestType::Call, Some(_)) => {
138 return Err(FrameError::OutOfRange(
139 "a CALL carries no stream mode".into(),
140 ))
141 }
142 (RequestType::StreamOpen, None) => {
143 return Err(FrameError::OutOfRange(
144 "a STREAM_OPEN carries one of the three stream modes".into(),
145 ))
146 }
147 _ => {}
148 }
149 let proofs = proofs_value(&spec.proofs);
150 if !proofs_within_bound(&proofs) {
151 return Err(FrameError::ProofsOutOfBound);
152 }
153 Ok(proofs)
154}
155
156fn request_fields(
159 frame_type: RequestType,
160 spec: &RequestSpec,
161 key: &NodeKey,
162 proofs: Value,
163) -> Vec<(Value, Value)> {
164 let mut fields = vec![
165 entry("frame_type", Value::text(frame_type.name())),
166 entry("caller", Value::Bytes(key.key_id().to_vec())),
167 entry("request_id", Value::Bytes(spec.request_id.to_vec())),
168 entry("realm", Value::Bytes(spec.realm.to_vec())),
169 entry("procedure", Value::text(spec.procedure.clone())),
170 entry("target", Value::Bytes(spec.target.to_vec())),
171 entry("deadline", uint(spec.deadline)),
172 match &spec.sealed {
173 Some(sealed) => entry("sealed", sealed.value()),
174 None => entry("payload", spec.payload.clone()),
175 },
176 ];
177 if let Some(mode) = spec.mode {
178 fields.push(entry("mode", Value::text(mode.name())));
179 }
180 if let Some(token) = &spec.token {
181 fields.push(entry("token", Value::Bytes(token.clone())));
182 }
183 if !spec.proofs.is_empty() {
184 fields.push(entry("proofs", proofs));
185 }
186 fields
187}
188
189fn request_frame(frame_type: RequestType, spec: &RequestSpec, request: Value) -> Value {
192 let mut frame = vec![
193 entry("version", Value::Int(i128::from(PROTOCOL_VERSION))),
194 entry("frame_type", Value::text(frame_type.name())),
195 entry("request", request),
196 ];
197 if let Some(route) = &spec.source_route {
198 frame.push(entry("source_route", Value::Bytes(route.clone())));
199 }
200 if let Some(budget) = spec.retry_budget {
201 frame.push(entry("retry_budget", uint(budget)));
202 }
203 Value::Map(frame)
204}
205
206const REQUEST_ROUTES: &[(&str, Rule)] = &[
207 ("source_route", Rule::AnyBytes),
208 ("retry_budget", Rule::ProtocolUint),
209];
210
211pub fn verify_request(frame: &Value, profile: Profile) -> Result<VerifiedRequest, FrameError> {
217 let (frame_type, object) = received_frame(
218 frame,
219 "request",
220 Rule::CarriedObject,
221 REQUEST_ROUTES,
222 &[CALL, STREAM_OPEN],
223 )
224 .ok_or(FrameError::Malformed)?;
225 let frame_type = if frame_type == CALL {
226 RequestType::Call
227 } else {
228 RequestType::StreamOpen
229 };
230 let verified = verify_object(REQUEST_LABEL, &object, profile).map_err(object_refusal)?;
231 let fields =
232 read_fields(&verified.fields, &request_table(frame_type)).ok_or(FrameError::Malformed)?;
233 let has_mode = fields.contains_key("mode");
234 if !has_fields(
235 &fields,
236 &[
237 "frame_type",
238 "caller",
239 "request_id",
240 "realm",
241 "procedure",
242 "target",
243 "deadline",
244 ],
245 ) || !clear_or_sealed(&fields, "payload")
246 || has_mode != (frame_type == RequestType::StreamOpen)
247 {
248 return Err(FrameError::Malformed);
249 }
250 let request = VerifiedRequest {
251 frame_type,
252 request_hash: Sha384::digest(&verified.tbs).into(),
253 caller: fixed(&fields["caller"]),
254 request_id: fixed(&fields["request_id"]),
255 realm: fixed(&fields["realm"]),
256 procedure: text_of(&fields["procedure"]),
257 target: fixed(&fields["target"]),
258 deadline: protocol_uint(&fields["deadline"]).unwrap_or(0),
259 payload: fields.get("payload").cloned().unwrap_or(Value::Null),
260 sealed: sealed_field(&fields, SealedContext::Request),
261 mode: fields
262 .get("mode")
263 .and_then(|m| StreamMode::parse(&text_of(m))),
264 token: fields.get("token").map(super::bytes_of),
265 proofs: fields.get("proofs").map(|p| match p {
266 Value::List(items) => items.iter().map(super::bytes_of).collect(),
267 _ => Vec::new(),
268 }),
269 key: verified.key,
270 };
271 if request.caller != node_id_of(&request.key, profile) {
272 return Err(FrameError::KeyIdMismatch);
273 }
274 Ok(request)
275}
276
277fn request_table(frame_type: RequestType) -> Vec<(&'static str, Rule)> {
278 vec![
279 (
280 "frame_type",
281 Rule::TextIn(match frame_type {
282 RequestType::Call => &[CALL],
283 RequestType::StreamOpen => &[STREAM_OPEN],
284 }),
285 ),
286 ("alg", Rule::Any),
287 ("caller", Rule::BytesOf(32)),
288 ("request_id", Rule::BytesOf(16)),
289 ("realm", Rule::BytesOf(32)),
290 ("procedure", Rule::TextWithin(MAX_PROCEDURE_BYTES)),
291 ("target", Rule::BytesOf(32)),
292 ("deadline", Rule::ProtocolUint),
293 ("payload", Rule::Any),
294 ("sealed", Rule::Sealed(SealedContext::Request)),
295 (
296 "mode",
297 Rule::TextIn(&["server_stream", "client_stream", "bidi"]),
298 ),
299 ("token", Rule::AnyBytes),
300 ("proofs", Rule::Proofs),
301 ]
302}
303
304pub fn request_fields_accepted(fields: &Value) -> bool {
308 read_fields(fields, &request_table(RequestType::Call)).is_some()
309}
310
311fn proofs_value(proofs: &[Vec<u8>]) -> Value {
312 Value::List(proofs.iter().map(|p| Value::Bytes(p.clone())).collect())
313}
314
315pub(super) fn proofs_within_bound(v: &Value) -> bool {
318 let Value::List(items) = v else {
319 return false;
320 };
321 if items.len() > MAX_PROOFS {
322 return false;
323 }
324 let mut seen = std::collections::HashSet::with_capacity(items.len());
325 let mut total = 0;
326 for item in items {
327 let Value::Bytes(b) = item else {
328 return false;
329 };
330 if !seen.insert(b.as_slice()) {
331 return false;
332 }
333 total += b.len();
334 }
335 total <= MAX_PROOFS_BYTES
336}