1use crate::cbor::Value;
9use crate::node_key::{node_id_of, NodeKey};
10use crate::profile::Profile;
11use crate::signed_object::{
12 sign_held_object, sign_object, verify_held_object, verify_object, VerifiedObject,
13};
14
15use super::{
16 bounded_text, check_payload, entry, fixed, has_fields, identity_signer, names_request,
17 object_refusal, protocol_uint, read_fields, received_frame, text_of, uint, FrameError,
18 RequestType, Rule, VerifiedRequest, CALLER_STREAM_LABEL, MAX_ERROR_CODE_BYTES,
19 MAX_ERROR_TEXT_BYTES, MAX_PROTOCOL_INT, PROTOCOL_VERSION, STREAM_LABEL,
20};
21
22const STREAM_DATA: &str = "stream_data";
23const STREAM_END: &str = "stream_end";
24const STREAM_ERROR: &str = "stream_error";
25const STREAM_REPLY: &str = "stream_reply";
26
27#[derive(Debug, Clone, Copy, PartialEq, Eq)]
30pub enum StreamMode {
31 ServerStream,
32 ClientStream,
33 Bidi,
34}
35
36impl StreamMode {
37 pub fn name(self) -> &'static str {
38 match self {
39 StreamMode::ServerStream => "server_stream",
40 StreamMode::ClientStream => "client_stream",
41 StreamMode::Bidi => "bidi",
42 }
43 }
44
45 pub fn parse(name: &str) -> Option<StreamMode> {
46 match name {
47 "server_stream" => Some(StreamMode::ServerStream),
48 "client_stream" => Some(StreamMode::ClientStream),
49 "bidi" => Some(StreamMode::Bidi),
50 _ => None,
51 }
52 }
53}
54
55#[derive(Debug, Clone, Copy, PartialEq, Eq)]
57pub enum StreamEncoding {
58 Raw,
59 Msgpack,
60}
61
62impl StreamEncoding {
63 fn name(self) -> &'static str {
64 match self {
65 StreamEncoding::Raw => "raw",
66 StreamEncoding::Msgpack => "msgpack",
67 }
68 }
69}
70
71#[derive(Debug, Clone, Copy, PartialEq, Eq)]
73pub enum StreamRole {
74 Send,
75 Both,
76}
77
78impl StreamRole {
79 fn name(self) -> &'static str {
80 match self {
81 StreamRole::Send => "send",
82 StreamRole::Both => "both",
83 }
84 }
85}
86
87#[derive(Debug, Clone, PartialEq)]
89pub enum StreamFields {
90 Data {
93 seq: u64,
94 encoding: StreamEncoding,
95 body: Value,
96 },
97 End { seq: u64, role: StreamRole },
99 Error {
101 seq: u64,
102 code: String,
103 message: String,
104 },
105 Reply { seq: u64, payload: Value },
107}
108
109impl StreamFields {
110 fn seq(&self) -> u64 {
111 match self {
112 StreamFields::Data { seq, .. }
113 | StreamFields::End { seq, .. }
114 | StreamFields::Error { seq, .. }
115 | StreamFields::Reply { seq, .. } => *seq,
116 }
117 }
118
119 fn frame_type(&self) -> &'static str {
120 match self {
121 StreamFields::Data { .. } => STREAM_DATA,
122 StreamFields::End { .. } => STREAM_END,
123 StreamFields::Error { .. } => STREAM_ERROR,
124 StreamFields::Reply { .. } => STREAM_REPLY,
125 }
126 }
127}
128
129#[derive(Debug, Clone, PartialEq)]
132pub struct VerifiedStreamFrame {
133 pub signer: [u8; 32],
134 pub fields: StreamFields,
135}
136
137#[derive(Debug, Clone, PartialEq)]
142pub struct StreamState {
143 open: VerifiedRequest,
144 mode: StreamMode,
145 provider: Side,
146 caller: Side,
147}
148
149#[derive(Debug, Clone, PartialEq, Default)]
150struct Side {
151 next: u64,
152 ended: bool,
153 key: Option<Vec<u8>>,
154 signer: [u8; 32],
155}
156
157pub fn open_stream(open: &VerifiedRequest) -> Result<StreamState, FrameError> {
160 match (open.frame_type, open.mode) {
161 (RequestType::StreamOpen, Some(mode)) => Ok(StreamState {
162 open: open.clone(),
163 mode,
164 provider: Side::default(),
165 caller: Side::default(),
166 }),
167 _ => Err(FrameError::OutOfRange(
168 "a stream opens on a STREAM_OPEN".into(),
169 )),
170 }
171}
172
173pub fn sign_provider_stream(
177 fields: &StreamFields,
178 open: &VerifiedRequest,
179 key: &NodeKey,
180) -> Result<Value, FrameError> {
181 let tbs = stream_build(fields, open, key, false)?;
182 let object = if fields.seq() == 0 {
183 sign_object(STREAM_LABEL, &tbs, key)
184 .map_err(object_refusal)?
185 .to_value()
186 } else {
187 sign_held_object(STREAM_LABEL, &tbs, key)
188 .map_err(object_refusal)?
189 .to_value()
190 };
191 Ok(stream_frame(fields.frame_type(), "stream", object))
192}
193
194pub fn sign_caller_stream(
198 fields: &StreamFields,
199 open: &VerifiedRequest,
200 key: &NodeKey,
201) -> Result<Value, FrameError> {
202 let tbs = stream_build(fields, open, key, true)?;
203 let object = sign_held_object(CALLER_STREAM_LABEL, &tbs, key).map_err(object_refusal)?;
204 Ok(stream_frame(
205 fields.frame_type(),
206 "caller_stream",
207 object.to_value(),
208 ))
209}
210
211fn stream_build(
215 fields: &StreamFields,
216 open: &VerifiedRequest,
217 key: &NodeKey,
218 caller: bool,
219) -> Result<Vec<(Value, Value)>, FrameError> {
220 identity_signer(key)?;
221 let sender = if caller { open.caller } else { open.target };
222 if open.frame_type != RequestType::StreamOpen || open.mode.is_none() || key.key_id() != sender {
223 return Err(FrameError::Unsignable);
224 }
225 if caller {
226 match fields {
227 StreamFields::Reply { .. } => {
228 return Err(FrameError::NotAllowed("a caller's STREAM_REPLY".into()))
229 }
230 StreamFields::Data { .. } if open.mode == Some(StreamMode::ServerStream) => {
231 return Err(FrameError::NotAllowed(
232 "a caller's STREAM_DATA in a server_stream".into(),
233 ))
234 }
235 _ => {}
236 }
237 }
238 if let StreamFields::Error { code, message, .. } = fields {
239 bounded_text("code", code.as_bytes(), MAX_ERROR_CODE_BYTES)?;
240 bounded_text("message", message.as_bytes(), MAX_ERROR_TEXT_BYTES)?;
241 }
242 match fields {
243 StreamFields::Data {
244 encoding: StreamEncoding::Msgpack,
245 body,
246 ..
247 } => check_payload(body)?,
248 StreamFields::Reply { payload, .. } => check_payload(payload)?,
249 _ => {}
250 }
251 let mut tbs = match fields {
252 StreamFields::Data { encoding, body, .. } => {
253 if *encoding == StreamEncoding::Raw && !matches!(body, Value::Bytes(_)) {
254 return Err(FrameError::OutOfRange(
255 "a raw body that is not a byte string".into(),
256 ));
257 }
258 vec![
259 entry("encoding", Value::text(encoding.name())),
260 entry("body", body.clone()),
261 ]
262 }
263 StreamFields::End { role, .. } => vec![entry("role", Value::text(role.name()))],
264 StreamFields::Error { code, message, .. } => {
265 vec![
266 entry("code", Value::text(code.clone())),
267 entry("message", Value::text(message.clone())),
268 ]
269 }
270 StreamFields::Reply { payload, .. } => vec![entry("payload", payload.clone())],
271 };
272 if fields.seq() >= MAX_PROTOCOL_INT {
273 return Err(FrameError::OutOfRange("a seq of 2^53 or more".into()));
274 }
275 tbs.extend([
276 entry("frame_type", Value::text(fields.frame_type())),
277 entry("request_id", Value::Bytes(open.request_id.to_vec())),
278 entry("request_hash", Value::Bytes(open.request_hash.to_vec())),
279 entry("signer", Value::Bytes(key.key_id().to_vec())),
280 entry("seq", uint(fields.seq())),
281 ]);
282 Ok(tbs)
283}
284
285fn stream_frame(frame_type: &str, object_name: &str, object: Value) -> Value {
286 Value::Map(vec![
287 entry("version", Value::Int(i128::from(PROTOCOL_VERSION))),
288 entry("frame_type", Value::text(frame_type)),
289 entry(object_name, object),
290 ])
291}
292
293const PROVIDER_TYPES: &[&str] = &[STREAM_DATA, STREAM_END, STREAM_ERROR, STREAM_REPLY];
294const CALLER_TYPES: &[&str] = &[STREAM_DATA, STREAM_END, STREAM_ERROR];
295
296pub fn verify_provider_stream(
303 frame: &Value,
304 state: &StreamState,
305 profile: Profile,
306) -> Result<(VerifiedStreamFrame, StreamState), FrameError> {
307 let (frame_type, object) =
308 received_frame(frame, "stream", Rule::StreamObject, &[], PROVIDER_TYPES)
309 .ok_or(FrameError::Malformed)?;
310 if state.provider.ended {
311 return Err(FrameError::StreamEnded);
312 }
313 let carries_key = object.get("key").is_some();
314 let Some(held_key) = &state.provider.key else {
315 return provider_first(&frame_type, &object, carries_key, state, profile);
316 };
317 let verified = if carries_key {
318 verify_object(STREAM_LABEL, &object, profile)
319 } else {
320 verify_held_object(STREAM_LABEL, &object, held_key, profile)
321 }
322 .map_err(object_refusal)?;
323 if &verified.key != held_key {
324 return Err(FrameError::KeyIdMismatch);
325 }
326 let (signer, fields, read) = stream_read(&frame_type, &verified)?;
327 if signer != state.provider.signer {
328 return Err(FrameError::KeyIdMismatch);
329 }
330 if !names_request(&read, &state.open) {
331 return Err(FrameError::RequestMismatch);
332 }
333 if fields.seq() != state.provider.next {
334 return Err(FrameError::SeqMismatch);
335 }
336 if carries_key {
337 return Err(FrameError::Malformed);
338 }
339 let mut next = state.clone();
340 next.provider.next = fields.seq() + 1;
341 next.provider.ended = frame_type == STREAM_END;
342 Ok((VerifiedStreamFrame { signer, fields }, next))
343}
344
345fn provider_first(
346 frame_type: &str,
347 object: &Value,
348 carries_key: bool,
349 state: &StreamState,
350 profile: Profile,
351) -> Result<(VerifiedStreamFrame, StreamState), FrameError> {
352 if !carries_key {
353 return Err(FrameError::SeqMismatch);
354 }
355 let verified = verify_object(STREAM_LABEL, object, profile).map_err(object_refusal)?;
356 let (signer, fields, read) = stream_read(frame_type, &verified)?;
357 if signer != node_id_of(&verified.key, profile) {
358 return Err(FrameError::KeyIdMismatch);
359 }
360 if !names_request(&read, &state.open) {
361 return Err(FrameError::RequestMismatch);
362 }
363 if signer != state.open.target {
364 return Err(FrameError::NotTheTarget);
365 }
366 if fields.seq() != 0 {
367 return Err(FrameError::SeqMismatch);
368 }
369 let mut next = state.clone();
370 next.provider = Side {
371 next: 1,
372 ended: frame_type == STREAM_END,
373 key: Some(verified.key),
374 signer,
375 };
376 Ok((VerifiedStreamFrame { signer, fields }, next))
377}
378
379pub fn verify_caller_stream(
383 frame: &Value,
384 state: &StreamState,
385 profile: Profile,
386) -> Result<(VerifiedStreamFrame, StreamState), FrameError> {
387 let (frame_type, object) =
388 received_frame(frame, "caller_stream", Rule::HeldObject, &[], CALLER_TYPES)
389 .ok_or(FrameError::Malformed)?;
390 if state.caller.ended {
391 return Err(FrameError::StreamEnded);
392 }
393 let verified = verify_held_object(CALLER_STREAM_LABEL, &object, &state.open.key, profile)
394 .map_err(object_refusal)?;
395 let (signer, fields, read) = stream_read(&frame_type, &verified)?;
396 if frame_type == STREAM_DATA && state.mode == StreamMode::ServerStream {
397 return Err(FrameError::Malformed);
398 }
399 if signer != state.open.caller {
400 return Err(FrameError::KeyIdMismatch);
401 }
402 if !names_request(&read, &state.open) {
403 return Err(FrameError::RequestMismatch);
404 }
405 if fields.seq() != state.caller.next {
406 return Err(FrameError::SeqMismatch);
407 }
408 let mut next = state.clone();
409 next.caller.next = fields.seq() + 1;
410 next.caller.ended = frame_type == STREAM_END;
411 Ok((VerifiedStreamFrame { signer, fields }, next))
412}
413
414fn stream_read(
418 frame_type: &str,
419 verified: &VerifiedObject,
420) -> Result<([u8; 32], StreamFields, super::Fields), FrameError> {
421 let types: &'static [&'static str] = match frame_type {
422 STREAM_DATA => &[STREAM_DATA],
423 STREAM_END => &[STREAM_END],
424 STREAM_ERROR => &[STREAM_ERROR],
425 _ => &[STREAM_REPLY],
426 };
427 let table = [
428 ("frame_type", Rule::TextIn(types)),
429 ("alg", Rule::Any),
430 ("request_id", Rule::BytesOf(16)),
431 ("request_hash", Rule::BytesOf(48)),
432 ("signer", Rule::BytesOf(32)),
433 ("seq", Rule::ProtocolUint),
434 ("encoding", Rule::TextIn(&["raw", "msgpack"])),
435 ("body", Rule::Any),
436 ("role", Rule::TextIn(&["send", "both"])),
437 ("code", Rule::TextWithin(MAX_ERROR_CODE_BYTES)),
438 ("message", Rule::TextWithin(MAX_ERROR_TEXT_BYTES)),
439 ("payload", Rule::Any),
440 ];
441 let fields = read_fields(&verified.fields, &table).ok_or(FrameError::Malformed)?;
442 if !has_fields(
443 &fields,
444 &["frame_type", "request_id", "request_hash", "signer", "seq"],
445 ) {
446 return Err(FrameError::Malformed);
447 }
448 let own: &[&str] = match frame_type {
449 STREAM_DATA => &["encoding", "body"],
450 STREAM_END => &["role"],
451 STREAM_ERROR => &["code", "message"],
452 _ => &["payload"],
453 };
454 let carried = 5 + own.len() + usize::from(fields.contains_key("alg"));
455 if !has_fields(&fields, own) || fields.len() != carried {
456 return Err(FrameError::Malformed);
457 }
458 let seq = protocol_uint(&fields["seq"]).unwrap_or(0);
459 let parsed = match frame_type {
460 STREAM_DATA => {
461 let encoding = if text_of(&fields["encoding"]) == "raw" {
462 StreamEncoding::Raw
463 } else {
464 StreamEncoding::Msgpack
465 };
466 if encoding == StreamEncoding::Raw && !matches!(fields["body"], Value::Bytes(_)) {
467 return Err(FrameError::Malformed);
468 }
469 StreamFields::Data {
470 seq,
471 encoding,
472 body: fields["body"].clone(),
473 }
474 }
475 STREAM_END => StreamFields::End {
476 seq,
477 role: if text_of(&fields["role"]) == "send" {
478 StreamRole::Send
479 } else {
480 StreamRole::Both
481 },
482 },
483 STREAM_ERROR => StreamFields::Error {
484 seq,
485 code: text_of(&fields["code"]),
486 message: text_of(&fields["message"]),
487 },
488 _ => StreamFields::Reply {
489 seq,
490 payload: fields["payload"].clone(),
491 },
492 };
493 Ok((fixed(&fields["signer"]), parsed, fields))
494}