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 stream_checks(fields, open, key, caller)?;
259 let mut tbs = stream_body(fields)?;
260 if fields.seq() >= MAX_PROTOCOL_INT {
261 return Err(FrameError::OutOfRange("a seq of 2^53 or more".into()));
262 }
263 tbs.extend([
264 entry("frame_type", Value::text(fields.frame_type())),
265 entry("request_id", Value::Bytes(open.request_id.to_vec())),
266 entry("request_hash", Value::Bytes(open.request_hash.to_vec())),
267 entry("signer", Value::Bytes(key.key_id().to_vec())),
268 entry("seq", uint(fields.seq())),
269 ]);
270 Ok(tbs)
271}
272
273fn stream_checks(
277 fields: &StreamFields,
278 open: &VerifiedRequest,
279 key: &NodeKey,
280 caller: bool,
281) -> Result<(), FrameError> {
282 identity_signer(key)?;
283 let sender = if caller { open.caller } else { open.target };
284 if open.frame_type != RequestType::StreamOpen || open.mode.is_none() || key.key_id() != sender {
285 return Err(FrameError::Unsignable);
286 }
287 if caller {
288 caller_sends(fields, open)?;
289 }
290 if fields
291 .sealed()
292 .is_some_and(|s| !s.shaped(side_context(caller)))
293 {
294 return Err(FrameError::SealedShape);
295 }
296 if let StreamFields::Error { code, message, .. } = fields {
297 bounded_text("code", code.as_bytes(), MAX_ERROR_CODE_BYTES)?;
298 bounded_text("message", message.as_bytes(), MAX_ERROR_TEXT_BYTES)?;
299 }
300 match fields {
301 StreamFields::Data {
302 encoding: StreamEncoding::Msgpack,
303 body,
304 ..
305 } => check_payload(body)?,
306 StreamFields::Reply { payload, .. } => check_payload(payload)?,
307 _ => {}
308 }
309 Ok(())
310}
311
312fn caller_sends(fields: &StreamFields, open: &VerifiedRequest) -> Result<(), FrameError> {
315 match fields {
316 StreamFields::Reply { .. } | StreamFields::SealedReply { .. } => {
317 Err(FrameError::NotAllowed("a caller's STREAM_REPLY".into()))
318 }
319 StreamFields::Data { .. } | StreamFields::SealedData { .. }
320 if open.mode == Some(StreamMode::ServerStream) =>
321 {
322 Err(FrameError::NotAllowed(
323 "a caller's STREAM_DATA in a server_stream".into(),
324 ))
325 }
326 _ => Ok(()),
327 }
328}
329
330fn stream_body(fields: &StreamFields) -> Result<Vec<(Value, Value)>, FrameError> {
333 let tbs = match fields {
334 StreamFields::Data { encoding, body, .. } => {
335 if *encoding == StreamEncoding::Raw && !matches!(body, Value::Bytes(_)) {
336 return Err(FrameError::OutOfRange(
337 "a raw body that is not a byte string".into(),
338 ));
339 }
340 vec![
341 entry("encoding", Value::text(encoding.name())),
342 entry("body", body.clone()),
343 ]
344 }
345 StreamFields::End { role, .. } => vec![entry("role", Value::text(role.name()))],
346 StreamFields::Error { code, message, .. } => {
347 vec![
348 entry("code", Value::text(code.clone())),
349 entry("message", Value::text(message.clone())),
350 ]
351 }
352 StreamFields::Reply { payload, .. } => vec![entry("payload", payload.clone())],
353 StreamFields::SealedData {
354 encoding, sealed, ..
355 } => vec![
356 entry("encoding", Value::text(encoding.name())),
357 entry("sealed", sealed.value()),
358 ],
359 StreamFields::SealedError { sealed, .. } | StreamFields::SealedReply { sealed, .. } => {
360 vec![entry("sealed", sealed.value())]
361 }
362 };
363 Ok(tbs)
364}
365
366fn stream_frame(frame_type: &str, object_name: &str, object: Value) -> Value {
367 Value::Map(vec![
368 entry("version", Value::Int(i128::from(PROTOCOL_VERSION))),
369 entry("frame_type", Value::text(frame_type)),
370 entry(object_name, object),
371 ])
372}
373
374const PROVIDER_TYPES: &[&str] = &[STREAM_DATA, STREAM_END, STREAM_ERROR, STREAM_REPLY];
375const CALLER_TYPES: &[&str] = &[STREAM_DATA, STREAM_END, STREAM_ERROR];
376
377pub fn verify_provider_stream(
384 frame: &Value,
385 state: &StreamState,
386 profile: Profile,
387) -> Result<(VerifiedStreamFrame, StreamState), FrameError> {
388 let (frame_type, object) =
389 received_frame(frame, "stream", Rule::StreamObject, &[], PROVIDER_TYPES)
390 .ok_or(FrameError::Malformed)?;
391 if state.provider.ended {
392 return Err(FrameError::StreamEnded);
393 }
394 let carries_key = object.get("key").is_some();
395 let Some(held_key) = &state.provider.key else {
396 return provider_first(&frame_type, &object, carries_key, state, profile);
397 };
398 let verified = if carries_key {
399 verify_object(STREAM_LABEL, &object, profile)
400 } else {
401 verify_held_object(STREAM_LABEL, &object, held_key, profile)
402 }
403 .map_err(object_refusal)?;
404 if &verified.key != held_key {
405 return Err(FrameError::KeyIdMismatch);
406 }
407 let (signer, fields, read) =
408 stream_read(&frame_type, &verified, SealedContext::ProviderStream)?;
409 if signer != state.provider.signer {
410 return Err(FrameError::KeyIdMismatch);
411 }
412 if !names_request(&read, &state.open) {
413 return Err(FrameError::RequestMismatch);
414 }
415 if fields.seq() != state.provider.next {
416 return Err(FrameError::SeqMismatch);
417 }
418 if carries_key {
419 return Err(FrameError::Malformed);
420 }
421 let mut next = state.clone();
422 next.provider.next = fields.seq() + 1;
423 next.provider.ended = frame_type == STREAM_END;
424 Ok((VerifiedStreamFrame { signer, fields }, next))
425}
426
427fn provider_first(
428 frame_type: &str,
429 object: &Value,
430 carries_key: bool,
431 state: &StreamState,
432 profile: Profile,
433) -> Result<(VerifiedStreamFrame, StreamState), FrameError> {
434 if !carries_key {
435 return Err(FrameError::SeqMismatch);
436 }
437 let verified = verify_object(STREAM_LABEL, object, profile).map_err(object_refusal)?;
438 let (signer, fields, read) = stream_read(frame_type, &verified, SealedContext::ProviderStream)?;
439 if signer != node_id_of(&verified.key, profile) {
440 return Err(FrameError::KeyIdMismatch);
441 }
442 if !names_request(&read, &state.open) {
443 return Err(FrameError::RequestMismatch);
444 }
445 if signer != state.open.target {
446 return Err(FrameError::NotTheTarget);
447 }
448 if fields.seq() != 0 {
449 return Err(FrameError::SeqMismatch);
450 }
451 let mut next = state.clone();
452 next.provider = Side {
453 next: 1,
454 ended: frame_type == STREAM_END,
455 key: Some(verified.key),
456 signer,
457 };
458 Ok((VerifiedStreamFrame { signer, fields }, next))
459}
460
461pub fn verify_caller_stream(
465 frame: &Value,
466 state: &StreamState,
467 profile: Profile,
468) -> Result<(VerifiedStreamFrame, StreamState), FrameError> {
469 let (frame_type, object) =
470 received_frame(frame, "caller_stream", Rule::HeldObject, &[], CALLER_TYPES)
471 .ok_or(FrameError::Malformed)?;
472 if state.caller.ended {
473 return Err(FrameError::StreamEnded);
474 }
475 let verified = verify_held_object(CALLER_STREAM_LABEL, &object, &state.open.key, profile)
476 .map_err(object_refusal)?;
477 let (signer, fields, read) = stream_read(&frame_type, &verified, SealedContext::CallerStream)?;
478 if frame_type == STREAM_DATA && state.mode == StreamMode::ServerStream {
479 return Err(FrameError::Malformed);
480 }
481 if signer != state.open.caller {
482 return Err(FrameError::KeyIdMismatch);
483 }
484 if !names_request(&read, &state.open) {
485 return Err(FrameError::RequestMismatch);
486 }
487 if fields.seq() != state.caller.next {
488 return Err(FrameError::SeqMismatch);
489 }
490 let mut next = state.clone();
491 next.caller.next = fields.seq() + 1;
492 next.caller.ended = frame_type == STREAM_END;
493 Ok((VerifiedStreamFrame { signer, fields }, next))
494}
495
496fn stream_read(
502 frame_type: &str,
503 verified: &VerifiedObject,
504 context: SealedContext,
505) -> Result<([u8; 32], StreamFields, super::Fields), FrameError> {
506 let fields = stream_table_read(frame_type, verified, context)?;
507 let sealed = fields.get("sealed").and_then(|v| read_sealed(v, context));
508 let own: &[&str] = match (frame_type, sealed.is_some()) {
509 (STREAM_DATA, false) => &["encoding", "body"],
510 (STREAM_DATA, true) => &["encoding", "sealed"],
511 (STREAM_END, _) => &["role"],
512 (STREAM_ERROR, false) => &["code", "message"],
513 (_, false) => &["payload"],
514 (_, true) => &["sealed"],
515 };
516 let carried = 5 + own.len() + usize::from(fields.contains_key("alg"));
517 if !has_fields(&fields, own) || fields.len() != carried {
518 return Err(FrameError::Malformed);
519 }
520 let seq = protocol_uint(&fields["seq"]).unwrap_or(0);
521 let parsed = stream_parse(frame_type, seq, sealed, &fields)?;
522 Ok((fixed(&fields["signer"]), parsed, fields))
523}
524
525fn stream_table_read(
528 frame_type: &str,
529 verified: &VerifiedObject,
530 context: SealedContext,
531) -> Result<super::Fields, FrameError> {
532 let types: &'static [&'static str] = match frame_type {
533 STREAM_DATA => &[STREAM_DATA],
534 STREAM_END => &[STREAM_END],
535 STREAM_ERROR => &[STREAM_ERROR],
536 _ => &[STREAM_REPLY],
537 };
538 let table = [
539 ("frame_type", Rule::TextIn(types)),
540 ("alg", Rule::Any),
541 ("request_id", Rule::BytesOf(16)),
542 ("request_hash", Rule::BytesOf(48)),
543 ("signer", Rule::BytesOf(32)),
544 ("seq", Rule::ProtocolUint),
545 ("encoding", Rule::TextIn(&["raw", "msgpack"])),
546 ("body", Rule::Any),
547 ("role", Rule::TextIn(&["send", "both"])),
548 ("code", Rule::TextWithin(MAX_ERROR_CODE_BYTES)),
549 ("message", Rule::TextWithin(MAX_ERROR_TEXT_BYTES)),
550 ("payload", Rule::Any),
551 ("sealed", Rule::Sealed(context)),
552 ];
553 let fields = read_fields(&verified.fields, &table).ok_or(FrameError::Malformed)?;
554 if !has_fields(
555 &fields,
556 &["frame_type", "request_id", "request_hash", "signer", "seq"],
557 ) {
558 return Err(FrameError::Malformed);
559 }
560 Ok(fields)
561}
562
563fn stream_parse(
566 frame_type: &str,
567 seq: u64,
568 sealed: Option<Sealed>,
569 fields: &super::Fields,
570) -> Result<StreamFields, FrameError> {
571 let parsed = match (frame_type, sealed) {
572 (STREAM_DATA, Some(sealed)) => StreamFields::SealedData {
573 seq,
574 encoding: stream_encoding(fields),
575 sealed,
576 },
577 (STREAM_ERROR, Some(sealed)) => StreamFields::SealedError { seq, sealed },
578 (STREAM_REPLY, Some(sealed)) => StreamFields::SealedReply { seq, sealed },
579 (STREAM_END, _) => StreamFields::End {
580 seq,
581 role: if text_of(&fields["role"]) == "send" {
582 StreamRole::Send
583 } else {
584 StreamRole::Both
585 },
586 },
587 (_, Some(_)) => return Err(FrameError::Malformed),
588 (STREAM_DATA, None) => {
589 let encoding = stream_encoding(fields);
590 if encoding == StreamEncoding::Raw && !matches!(fields["body"], Value::Bytes(_)) {
591 return Err(FrameError::Malformed);
592 }
593 StreamFields::Data {
594 seq,
595 encoding,
596 body: fields["body"].clone(),
597 }
598 }
599 (STREAM_ERROR, None) => StreamFields::Error {
600 seq,
601 code: text_of(&fields["code"]),
602 message: text_of(&fields["message"]),
603 },
604 (_, None) => StreamFields::Reply {
605 seq,
606 payload: fields["payload"].clone(),
607 },
608 };
609 Ok(parsed)
610}
611
612fn stream_encoding(fields: &super::Fields) -> StreamEncoding {
613 if text_of(&fields["encoding"]) == "raw" {
614 StreamEncoding::Raw
615 } else {
616 StreamEncoding::Msgpack
617 }
618}