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::sealed::{read_sealed, Sealed, SealedContext};
16use super::{
17 bounded_text, check_payload, entry, fixed, has_fields, identity_signer, names_request,
18 object_refusal, protocol_uint, read_fields, received_frame, text_of, uint, FrameError,
19 RequestType, Rule, VerifiedRequest, CALLER_STREAM_LABEL, MAX_ERROR_CODE_BYTES,
20 MAX_ERROR_TEXT_BYTES, MAX_PROTOCOL_INT, PROTOCOL_VERSION, STREAM_LABEL,
21};
22
23const STREAM_DATA: &str = "stream_data";
24const STREAM_END: &str = "stream_end";
25const STREAM_ERROR: &str = "stream_error";
26const STREAM_REPLY: &str = "stream_reply";
27
28#[derive(Debug, Clone, Copy, PartialEq, Eq)]
31pub enum StreamMode {
32 ServerStream,
33 ClientStream,
34 Bidi,
35}
36
37impl StreamMode {
38 pub fn name(self) -> &'static str {
39 match self {
40 StreamMode::ServerStream => "server_stream",
41 StreamMode::ClientStream => "client_stream",
42 StreamMode::Bidi => "bidi",
43 }
44 }
45
46 pub fn parse(name: &str) -> Option<StreamMode> {
47 match name {
48 "server_stream" => Some(StreamMode::ServerStream),
49 "client_stream" => Some(StreamMode::ClientStream),
50 "bidi" => Some(StreamMode::Bidi),
51 _ => None,
52 }
53 }
54}
55
56#[derive(Debug, Clone, Copy, PartialEq, Eq)]
58pub enum StreamEncoding {
59 Raw,
60 Msgpack,
61}
62
63impl StreamEncoding {
64 fn name(self) -> &'static str {
65 match self {
66 StreamEncoding::Raw => "raw",
67 StreamEncoding::Msgpack => "msgpack",
68 }
69 }
70}
71
72#[derive(Debug, Clone, Copy, PartialEq, Eq)]
74pub enum StreamRole {
75 Send,
76 Both,
77}
78
79impl StreamRole {
80 fn name(self) -> &'static str {
81 match self {
82 StreamRole::Send => "send",
83 StreamRole::Both => "both",
84 }
85 }
86}
87
88#[derive(Debug, Clone, PartialEq)]
90pub enum StreamFields {
91 Data {
94 seq: u64,
95 encoding: StreamEncoding,
96 body: Value,
97 },
98 End { seq: u64, role: StreamRole },
100 Error {
102 seq: u64,
103 code: String,
104 message: String,
105 },
106 Reply { seq: u64, payload: Value },
108 SealedData {
111 seq: u64,
112 encoding: StreamEncoding,
113 sealed: Sealed,
114 },
115 SealedError { seq: u64, sealed: Sealed },
118 SealedReply { seq: u64, sealed: Sealed },
121}
122
123impl StreamFields {
124 pub fn seq(&self) -> u64 {
126 match self {
127 StreamFields::Data { seq, .. }
128 | StreamFields::End { seq, .. }
129 | StreamFields::Error { seq, .. }
130 | StreamFields::Reply { seq, .. }
131 | StreamFields::SealedData { seq, .. }
132 | StreamFields::SealedError { seq, .. }
133 | StreamFields::SealedReply { seq, .. } => *seq,
134 }
135 }
136
137 fn frame_type(&self) -> &'static str {
138 match self {
139 StreamFields::Data { .. } | StreamFields::SealedData { .. } => STREAM_DATA,
140 StreamFields::End { .. } => STREAM_END,
141 StreamFields::Error { .. } | StreamFields::SealedError { .. } => STREAM_ERROR,
142 StreamFields::Reply { .. } | StreamFields::SealedReply { .. } => STREAM_REPLY,
143 }
144 }
145
146 fn sealed(&self) -> Option<&Sealed> {
148 match self {
149 StreamFields::SealedData { sealed, .. }
150 | StreamFields::SealedError { sealed, .. }
151 | StreamFields::SealedReply { sealed, .. } => Some(sealed),
152 _ => None,
153 }
154 }
155}
156
157fn side_context(caller: bool) -> SealedContext {
160 if caller {
161 SealedContext::CallerStream
162 } else {
163 SealedContext::ProviderStream
164 }
165}
166
167#[derive(Debug, Clone, PartialEq)]
170pub struct VerifiedStreamFrame {
171 pub signer: [u8; 32],
172 pub fields: StreamFields,
173}
174
175#[derive(Debug, Clone, PartialEq)]
180pub struct StreamState {
181 open: VerifiedRequest,
182 mode: StreamMode,
183 provider: Side,
184 caller: Side,
185}
186
187#[derive(Debug, Clone, PartialEq, Default)]
188struct Side {
189 next: u64,
190 ended: bool,
191 key: Option<Vec<u8>>,
192 signer: [u8; 32],
193}
194
195pub fn open_stream(open: &VerifiedRequest) -> Result<StreamState, FrameError> {
198 match (open.frame_type, open.mode) {
199 (RequestType::StreamOpen, Some(mode)) => Ok(StreamState {
200 open: open.clone(),
201 mode,
202 provider: Side::default(),
203 caller: Side::default(),
204 }),
205 _ => Err(FrameError::OutOfRange(
206 "a stream opens on a STREAM_OPEN".into(),
207 )),
208 }
209}
210
211pub fn sign_provider_stream(
215 fields: &StreamFields,
216 open: &VerifiedRequest,
217 key: &NodeKey,
218) -> Result<Value, FrameError> {
219 let tbs = stream_build(fields, open, key, false)?;
220 let object = if fields.seq() == 0 {
221 sign_object(STREAM_LABEL, &tbs, key)
222 .map_err(object_refusal)?
223 .to_value()
224 } else {
225 sign_held_object(STREAM_LABEL, &tbs, key)
226 .map_err(object_refusal)?
227 .to_value()
228 };
229 Ok(stream_frame(fields.frame_type(), "stream", object))
230}
231
232pub fn sign_caller_stream(
236 fields: &StreamFields,
237 open: &VerifiedRequest,
238 key: &NodeKey,
239) -> Result<Value, FrameError> {
240 let tbs = stream_build(fields, open, key, true)?;
241 let object = sign_held_object(CALLER_STREAM_LABEL, &tbs, key).map_err(object_refusal)?;
242 Ok(stream_frame(
243 fields.frame_type(),
244 "caller_stream",
245 object.to_value(),
246 ))
247}
248
249fn stream_build(
253 fields: &StreamFields,
254 open: &VerifiedRequest,
255 key: &NodeKey,
256 caller: bool,
257) -> Result<Vec<(Value, Value)>, FrameError> {
258 identity_signer(key)?;
259 let sender = if caller { open.caller } else { open.target };
260 if open.frame_type != RequestType::StreamOpen || open.mode.is_none() || key.key_id() != sender {
261 return Err(FrameError::Unsignable);
262 }
263 if caller {
264 match fields {
265 StreamFields::Reply { .. } | StreamFields::SealedReply { .. } => {
266 return Err(FrameError::NotAllowed("a caller's STREAM_REPLY".into()))
267 }
268 StreamFields::Data { .. } | StreamFields::SealedData { .. }
269 if open.mode == Some(StreamMode::ServerStream) =>
270 {
271 return Err(FrameError::NotAllowed(
272 "a caller's STREAM_DATA in a server_stream".into(),
273 ))
274 }
275 _ => {}
276 }
277 }
278 if fields
279 .sealed()
280 .is_some_and(|s| !s.shaped(side_context(caller)))
281 {
282 return Err(FrameError::SealedShape);
283 }
284 if let StreamFields::Error { code, message, .. } = fields {
285 bounded_text("code", code.as_bytes(), MAX_ERROR_CODE_BYTES)?;
286 bounded_text("message", message.as_bytes(), MAX_ERROR_TEXT_BYTES)?;
287 }
288 match fields {
289 StreamFields::Data {
290 encoding: StreamEncoding::Msgpack,
291 body,
292 ..
293 } => check_payload(body)?,
294 StreamFields::Reply { payload, .. } => check_payload(payload)?,
295 _ => {}
296 }
297 let mut tbs = match fields {
298 StreamFields::Data { encoding, body, .. } => {
299 if *encoding == StreamEncoding::Raw && !matches!(body, Value::Bytes(_)) {
300 return Err(FrameError::OutOfRange(
301 "a raw body that is not a byte string".into(),
302 ));
303 }
304 vec![
305 entry("encoding", Value::text(encoding.name())),
306 entry("body", body.clone()),
307 ]
308 }
309 StreamFields::End { role, .. } => vec![entry("role", Value::text(role.name()))],
310 StreamFields::Error { code, message, .. } => {
311 vec![
312 entry("code", Value::text(code.clone())),
313 entry("message", Value::text(message.clone())),
314 ]
315 }
316 StreamFields::Reply { payload, .. } => vec![entry("payload", payload.clone())],
317 StreamFields::SealedData {
318 encoding, sealed, ..
319 } => vec![
320 entry("encoding", Value::text(encoding.name())),
321 entry("sealed", sealed.value()),
322 ],
323 StreamFields::SealedError { sealed, .. } | StreamFields::SealedReply { sealed, .. } => {
324 vec![entry("sealed", sealed.value())]
325 }
326 };
327 if fields.seq() >= MAX_PROTOCOL_INT {
328 return Err(FrameError::OutOfRange("a seq of 2^53 or more".into()));
329 }
330 tbs.extend([
331 entry("frame_type", Value::text(fields.frame_type())),
332 entry("request_id", Value::Bytes(open.request_id.to_vec())),
333 entry("request_hash", Value::Bytes(open.request_hash.to_vec())),
334 entry("signer", Value::Bytes(key.key_id().to_vec())),
335 entry("seq", uint(fields.seq())),
336 ]);
337 Ok(tbs)
338}
339
340fn stream_frame(frame_type: &str, object_name: &str, object: Value) -> Value {
341 Value::Map(vec![
342 entry("version", Value::Int(i128::from(PROTOCOL_VERSION))),
343 entry("frame_type", Value::text(frame_type)),
344 entry(object_name, object),
345 ])
346}
347
348const PROVIDER_TYPES: &[&str] = &[STREAM_DATA, STREAM_END, STREAM_ERROR, STREAM_REPLY];
349const CALLER_TYPES: &[&str] = &[STREAM_DATA, STREAM_END, STREAM_ERROR];
350
351pub fn verify_provider_stream(
358 frame: &Value,
359 state: &StreamState,
360 profile: Profile,
361) -> Result<(VerifiedStreamFrame, StreamState), FrameError> {
362 let (frame_type, object) =
363 received_frame(frame, "stream", Rule::StreamObject, &[], PROVIDER_TYPES)
364 .ok_or(FrameError::Malformed)?;
365 if state.provider.ended {
366 return Err(FrameError::StreamEnded);
367 }
368 let carries_key = object.get("key").is_some();
369 let Some(held_key) = &state.provider.key else {
370 return provider_first(&frame_type, &object, carries_key, state, profile);
371 };
372 let verified = if carries_key {
373 verify_object(STREAM_LABEL, &object, profile)
374 } else {
375 verify_held_object(STREAM_LABEL, &object, held_key, profile)
376 }
377 .map_err(object_refusal)?;
378 if &verified.key != held_key {
379 return Err(FrameError::KeyIdMismatch);
380 }
381 let (signer, fields, read) =
382 stream_read(&frame_type, &verified, SealedContext::ProviderStream)?;
383 if signer != state.provider.signer {
384 return Err(FrameError::KeyIdMismatch);
385 }
386 if !names_request(&read, &state.open) {
387 return Err(FrameError::RequestMismatch);
388 }
389 if fields.seq() != state.provider.next {
390 return Err(FrameError::SeqMismatch);
391 }
392 if carries_key {
393 return Err(FrameError::Malformed);
394 }
395 let mut next = state.clone();
396 next.provider.next = fields.seq() + 1;
397 next.provider.ended = frame_type == STREAM_END;
398 Ok((VerifiedStreamFrame { signer, fields }, next))
399}
400
401fn provider_first(
402 frame_type: &str,
403 object: &Value,
404 carries_key: bool,
405 state: &StreamState,
406 profile: Profile,
407) -> Result<(VerifiedStreamFrame, StreamState), FrameError> {
408 if !carries_key {
409 return Err(FrameError::SeqMismatch);
410 }
411 let verified = verify_object(STREAM_LABEL, object, profile).map_err(object_refusal)?;
412 let (signer, fields, read) = stream_read(frame_type, &verified, SealedContext::ProviderStream)?;
413 if signer != node_id_of(&verified.key, profile) {
414 return Err(FrameError::KeyIdMismatch);
415 }
416 if !names_request(&read, &state.open) {
417 return Err(FrameError::RequestMismatch);
418 }
419 if signer != state.open.target {
420 return Err(FrameError::NotTheTarget);
421 }
422 if fields.seq() != 0 {
423 return Err(FrameError::SeqMismatch);
424 }
425 let mut next = state.clone();
426 next.provider = Side {
427 next: 1,
428 ended: frame_type == STREAM_END,
429 key: Some(verified.key),
430 signer,
431 };
432 Ok((VerifiedStreamFrame { signer, fields }, next))
433}
434
435pub fn verify_caller_stream(
439 frame: &Value,
440 state: &StreamState,
441 profile: Profile,
442) -> Result<(VerifiedStreamFrame, StreamState), FrameError> {
443 let (frame_type, object) =
444 received_frame(frame, "caller_stream", Rule::HeldObject, &[], CALLER_TYPES)
445 .ok_or(FrameError::Malformed)?;
446 if state.caller.ended {
447 return Err(FrameError::StreamEnded);
448 }
449 let verified = verify_held_object(CALLER_STREAM_LABEL, &object, &state.open.key, profile)
450 .map_err(object_refusal)?;
451 let (signer, fields, read) = stream_read(&frame_type, &verified, SealedContext::CallerStream)?;
452 if frame_type == STREAM_DATA && state.mode == StreamMode::ServerStream {
453 return Err(FrameError::Malformed);
454 }
455 if signer != state.open.caller {
456 return Err(FrameError::KeyIdMismatch);
457 }
458 if !names_request(&read, &state.open) {
459 return Err(FrameError::RequestMismatch);
460 }
461 if fields.seq() != state.caller.next {
462 return Err(FrameError::SeqMismatch);
463 }
464 let mut next = state.clone();
465 next.caller.next = fields.seq() + 1;
466 next.caller.ended = frame_type == STREAM_END;
467 Ok((VerifiedStreamFrame { signer, fields }, next))
468}
469
470fn stream_read(
476 frame_type: &str,
477 verified: &VerifiedObject,
478 context: SealedContext,
479) -> Result<([u8; 32], StreamFields, super::Fields), FrameError> {
480 let types: &'static [&'static str] = match frame_type {
481 STREAM_DATA => &[STREAM_DATA],
482 STREAM_END => &[STREAM_END],
483 STREAM_ERROR => &[STREAM_ERROR],
484 _ => &[STREAM_REPLY],
485 };
486 let table = [
487 ("frame_type", Rule::TextIn(types)),
488 ("alg", Rule::Any),
489 ("request_id", Rule::BytesOf(16)),
490 ("request_hash", Rule::BytesOf(48)),
491 ("signer", Rule::BytesOf(32)),
492 ("seq", Rule::ProtocolUint),
493 ("encoding", Rule::TextIn(&["raw", "msgpack"])),
494 ("body", Rule::Any),
495 ("role", Rule::TextIn(&["send", "both"])),
496 ("code", Rule::TextWithin(MAX_ERROR_CODE_BYTES)),
497 ("message", Rule::TextWithin(MAX_ERROR_TEXT_BYTES)),
498 ("payload", Rule::Any),
499 ("sealed", Rule::Sealed(context)),
500 ];
501 let fields = read_fields(&verified.fields, &table).ok_or(FrameError::Malformed)?;
502 if !has_fields(
503 &fields,
504 &["frame_type", "request_id", "request_hash", "signer", "seq"],
505 ) {
506 return Err(FrameError::Malformed);
507 }
508 let sealed = fields.get("sealed").and_then(|v| read_sealed(v, context));
509 let own: &[&str] = match (frame_type, sealed.is_some()) {
510 (STREAM_DATA, false) => &["encoding", "body"],
511 (STREAM_DATA, true) => &["encoding", "sealed"],
512 (STREAM_END, _) => &["role"],
513 (STREAM_ERROR, false) => &["code", "message"],
514 (_, false) => &["payload"],
515 (_, true) => &["sealed"],
516 };
517 let carried = 5 + own.len() + usize::from(fields.contains_key("alg"));
518 if !has_fields(&fields, own) || fields.len() != carried {
519 return Err(FrameError::Malformed);
520 }
521 let seq = protocol_uint(&fields["seq"]).unwrap_or(0);
522 let parsed = match (frame_type, sealed) {
523 (STREAM_DATA, Some(sealed)) => StreamFields::SealedData {
524 seq,
525 encoding: stream_encoding(&fields),
526 sealed,
527 },
528 (STREAM_ERROR, Some(sealed)) => StreamFields::SealedError { seq, sealed },
529 (STREAM_REPLY, Some(sealed)) => StreamFields::SealedReply { seq, sealed },
530 (STREAM_END, _) => StreamFields::End {
531 seq,
532 role: if text_of(&fields["role"]) == "send" {
533 StreamRole::Send
534 } else {
535 StreamRole::Both
536 },
537 },
538 (_, Some(_)) => return Err(FrameError::Malformed),
539 (STREAM_DATA, None) => {
540 let encoding = stream_encoding(&fields);
541 if encoding == StreamEncoding::Raw && !matches!(fields["body"], Value::Bytes(_)) {
542 return Err(FrameError::Malformed);
543 }
544 StreamFields::Data {
545 seq,
546 encoding,
547 body: fields["body"].clone(),
548 }
549 }
550 (STREAM_ERROR, None) => StreamFields::Error {
551 seq,
552 code: text_of(&fields["code"]),
553 message: text_of(&fields["message"]),
554 },
555 (_, None) => StreamFields::Reply {
556 seq,
557 payload: fields["payload"].clone(),
558 },
559 };
560 Ok((fixed(&fields["signer"]), parsed, fields))
561}
562
563fn stream_encoding(fields: &super::Fields) -> StreamEncoding {
564 if text_of(&fields["encoding"]) == "raw" {
565 StreamEncoding::Raw
566 } else {
567 StreamEncoding::Msgpack
568 }
569}