1use crate::error::WireErrorCode;
16use crate::frame::Frame;
17use crate::version::{ProtocolVersion, SupportedVersions};
18
19#[derive(Debug, Clone, Copy, PartialEq, Eq)]
23enum CloseReason {
24 UnsupportedVersion,
27 NonHandshakeFirst,
29 DuplicateHandshake,
31 ServerOnlyKind,
33}
34
35impl CloseReason {
36 const fn phrase(self) -> &'static str {
38 match self {
39 CloseReason::UnsupportedVersion => "a rejected handshake (unsupported version)",
40 CloseReason::NonHandshakeFirst => "a frame before the handshake",
41 CloseReason::DuplicateHandshake => "a duplicate handshake",
42 CloseReason::ServerOnlyKind => "a server-to-client frame on an inbound gate",
43 }
44 }
45}
46
47#[derive(Debug, Clone, Copy, PartialEq, Eq)]
49enum State {
50 AwaitingHandshake,
52 Completed(ProtocolVersion),
54 Closed(CloseReason),
58}
59
60#[derive(Debug, Clone, PartialEq)]
62pub enum HandshakeOutcome {
63 Accepted {
67 ack: Frame,
68 version: ProtocolVersion,
69 },
70 Rejected { error: Frame },
74 Admitted,
78}
79
80#[derive(Debug, Clone, PartialEq)]
90pub struct HandshakeSequenceError {
91 pub error: Box<Frame>,
92}
93
94impl std::fmt::Display for HandshakeSequenceError {
95 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
96 match self.error.as_ref() {
97 Frame::Error { code, message, .. } => {
98 write!(f, "handshake sequence rejected ({code}): {message}")
99 }
100 frame => write!(f, "handshake sequence rejected by {} frame", frame.kind()),
101 }
102 }
103}
104
105impl std::error::Error for HandshakeSequenceError {}
106
107#[derive(Debug)]
114pub struct HandshakeGate {
115 state: State,
116 supported: SupportedVersions,
117}
118
119impl HandshakeGate {
120 pub fn new(supported: SupportedVersions) -> Self {
123 Self {
124 state: State::AwaitingHandshake,
125 supported,
126 }
127 }
128
129 pub fn is_complete(&self) -> bool {
131 matches!(self.state, State::Completed(_))
132 }
133
134 pub fn accepted_version(&self) -> Option<ProtocolVersion> {
136 match self.state {
137 State::Completed(v) => Some(v),
138 _ => None,
139 }
140 }
141
142 pub fn admit(&mut self, frame: &Frame) -> Result<HandshakeOutcome, HandshakeSequenceError> {
161 match (&self.state, frame) {
162 (State::AwaitingHandshake, Frame::Handshake { version }) => {
163 if self.supported.contains(*version) {
164 self.state = State::Completed(*version);
165 Ok(HandshakeOutcome::Accepted {
166 ack: Frame::HandshakeAck { version: *version },
167 version: *version,
168 })
169 } else {
170 self.state = State::Closed(CloseReason::UnsupportedVersion);
171 Ok(HandshakeOutcome::Rejected {
172 error: Frame::Error {
173 id: None,
174 code: WireErrorCode::UnsupportedVersion,
175 message: format!(
176 "unsupported protocol version {version}; server supports [{}, {}]",
177 self.supported.min(),
178 self.supported.max()
179 ),
180 unrecognized_code: None,
181 },
182 })
183 }
184 }
185 (State::AwaitingHandshake, _) => {
186 self.state = State::Closed(CloseReason::NonHandshakeFirst);
187 Err(HandshakeSequenceError {
188 error: Box::new(Frame::Error {
189 id: None,
190 code: WireErrorCode::MalformedFrame,
191 message: format!(
192 "expected \"handshake\" as the first frame, got {:?}",
193 frame.kind()
194 ),
195 unrecognized_code: None,
196 }),
197 })
198 }
199 (State::Completed(_), Frame::Handshake { .. }) => {
200 self.state = State::Closed(CloseReason::DuplicateHandshake);
201 Err(HandshakeSequenceError {
202 error: Box::new(Frame::Error {
203 id: None,
204 code: WireErrorCode::MalformedFrame,
205 message: "handshake already completed on this connection".to_string(),
206 unrecognized_code: None,
207 }),
208 })
209 }
210 (State::Completed(_), _) => {
211 if crate::frame::CLIENT_TO_SERVER_KINDS.contains(&frame.kind()) {
216 Ok(HandshakeOutcome::Admitted)
217 } else {
218 self.state = State::Closed(CloseReason::ServerOnlyKind);
219 Err(HandshakeSequenceError {
220 error: Box::new(Frame::Error {
221 id: None,
222 code: WireErrorCode::MalformedFrame,
223 message: format!(
224 "frame kind {:?} is server-to-client only; a server never accepts it as an inbound frame",
225 frame.kind()
226 ),
227 unrecognized_code: None,
228 }),
229 })
230 }
231 }
232 (State::Closed(reason), _) => Err(HandshakeSequenceError {
233 error: Box::new(Frame::Error {
234 id: None,
235 code: WireErrorCode::MalformedFrame,
236 message: format!("connection already closed by {}", reason.phrase()),
237 unrecognized_code: None,
238 }),
239 }),
240 }
241 }
242}
243
244impl Default for HandshakeGate {
245 fn default() -> Self {
246 Self::new(SupportedVersions::current())
247 }
248}
249
250#[cfg(test)]
251mod tests {
252 use super::*;
253 use crate::version::CURRENT_VERSION;
254
255 #[test]
256 fn accepts_current_version() {
257 let mut gate = HandshakeGate::default();
258 let outcome = gate
259 .admit(&Frame::Handshake {
260 version: CURRENT_VERSION,
261 })
262 .unwrap();
263 assert!(matches!(outcome, HandshakeOutcome::Accepted { .. }));
264 assert!(gate.is_complete());
265 assert_eq!(gate.accepted_version(), Some(CURRENT_VERSION));
266 }
267
268 #[test]
269 fn rejects_unsupported_version() {
270 let mut gate = HandshakeGate::default();
271 let outcome = gate
272 .admit(&Frame::Handshake {
273 version: ProtocolVersion::new(9999),
274 })
275 .unwrap();
276 match outcome {
277 HandshakeOutcome::Rejected { error } => match error {
278 Frame::Error { code, .. } => assert_eq!(code, WireErrorCode::UnsupportedVersion),
279 _ => panic!("expected an error frame"),
280 },
281 other => panic!("expected Rejected, got {other:?}"),
282 }
283 }
284
285 #[test]
286 fn rejects_request_before_handshake() {
287 let mut gate = HandshakeGate::default();
288 let result = gate.admit(&Frame::Cancel {
289 id: crate::frame::OperationId::from("op-1"),
290 });
291 assert!(result.is_err());
292 }
293
294 #[test]
295 fn sequence_error_displays_the_contained_refusal_reason() {
296 let mut gate = HandshakeGate::default();
297 let error = gate
298 .admit(&Frame::Cancel {
299 id: crate::frame::OperationId::from("op-1"),
300 })
301 .unwrap_err();
302 assert!(error.to_string().contains("malformed_frame"));
303 assert!(error.to_string().contains("expected \"handshake\""));
304
305 fn assert_error<T: std::error::Error>() {}
306 assert_error::<HandshakeSequenceError>();
307 }
308
309 #[test]
310 fn admits_ordinary_frames_after_handshake() {
311 let mut gate = HandshakeGate::default();
312 gate.admit(&Frame::Handshake {
313 version: CURRENT_VERSION,
314 })
315 .unwrap();
316 let outcome = gate
317 .admit(&Frame::Cancel {
318 id: crate::frame::OperationId::from("op-1"),
319 })
320 .unwrap();
321 assert_eq!(outcome, HandshakeOutcome::Admitted);
322 }
323
324 #[test]
325 fn rejects_second_handshake() {
326 let mut gate = HandshakeGate::default();
327 gate.admit(&Frame::Handshake {
328 version: CURRENT_VERSION,
329 })
330 .unwrap();
331 let result = gate.admit(&Frame::Handshake {
332 version: CURRENT_VERSION,
333 });
334 assert!(result.is_err());
335 }
336
337 fn completed_gate() -> HandshakeGate {
338 let mut gate = HandshakeGate::default();
339 gate.admit(&Frame::Handshake {
340 version: CURRENT_VERSION,
341 })
342 .unwrap();
343 gate
344 }
345
346 #[test]
347 fn rejects_server_only_kinds_on_the_inbound_gate() {
348 let server_only_frames = [
352 Frame::HandshakeAck {
353 version: CURRENT_VERSION,
354 },
355 Frame::Response {
356 id: crate::frame::OperationId::from("op-1"),
357 result: serde_json::json!({}),
358 },
359 Frame::Error {
360 id: None,
361 code: WireErrorCode::Internal,
362 message: "x".to_string(),
363 unrecognized_code: None,
364 },
365 Frame::SubscribeAck {
366 id: crate::frame::OperationId::from("op-2"),
367 topic: "a.b".to_string(),
368 start_cursor: 1,
369 },
370 Frame::UnsubscribeAck {
371 id: crate::frame::OperationId::from("op-3"),
372 topic: "a.b".to_string(),
373 },
374 Frame::Event {
375 topic: "a.b".to_string(),
376 cursor: 1,
377 occurred_at: "2026-08-04T11:00:00Z".to_string(),
378 payload: serde_json::json!({}),
379 },
380 ];
381 for frame in server_only_frames {
382 let mut gate = completed_gate();
383 assert!(
384 gate.admit(&frame).is_err(),
385 "server-only kind {:?} must be rejected on the inbound gate",
386 frame.kind()
387 );
388 }
389 }
390
391 #[test]
392 fn rejects_a_server_only_kind_with_the_gates_rejection_shape() {
393 let mut gate = completed_gate();
394 let response = Frame::Response {
395 id: crate::frame::OperationId::from("op-1"),
396 result: serde_json::json!({}),
397 };
398 let err = gate.admit(&response).unwrap_err();
399 match err.error.as_ref() {
400 Frame::Error {
401 id, code, message, ..
402 } => {
403 assert_eq!(id, &None);
404 assert_eq!(*code, WireErrorCode::MalformedFrame);
405 assert!(message.contains("server-to-client"), "message: {message}");
406 }
407 other => panic!("expected an error frame, got {other:?}"),
408 }
409 let err = gate
412 .admit(&Frame::Cancel {
413 id: crate::frame::OperationId::from("op-1"),
414 })
415 .unwrap_err();
416 match err.error.as_ref() {
417 Frame::Error { code, .. } => assert_eq!(*code, WireErrorCode::MalformedFrame),
418 other => panic!("expected an error frame, got {other:?}"),
419 }
420 }
421
422 #[test]
423 fn admits_every_client_kind_after_handshake() {
424 let client_frames = [
428 Frame::Request {
429 id: crate::frame::OperationId::from("op-1"),
430 ops: "stats()".to_string(),
431 deadline_ms: None,
432 namespace: None,
433 actor_id: None,
434 visible_namespaces: None,
435 },
436 Frame::Cancel {
437 id: crate::frame::OperationId::from("op-1"),
438 },
439 Frame::Subscribe {
440 id: crate::frame::OperationId::from("op-2"),
441 topic: "a.b".to_string(),
442 resume_cursor: None,
443 },
444 Frame::Unsubscribe {
445 id: crate::frame::OperationId::from("op-3"),
446 topic: "a.b".to_string(),
447 },
448 ];
449 for frame in client_frames {
450 let mut gate = completed_gate();
451 let outcome = gate
452 .admit(&frame)
453 .unwrap_or_else(|_| panic!("kind {:?} should be admitted", frame.kind()));
454 assert_eq!(
455 outcome,
456 HandshakeOutcome::Admitted,
457 "kind {:?}",
458 frame.kind()
459 );
460 }
461 }
462
463 #[test]
464 fn stays_closed_after_a_rejected_handshake() {
465 let mut gate = HandshakeGate::default();
471 let outcome = gate
472 .admit(&Frame::Handshake {
473 version: ProtocolVersion::new(9999),
474 })
475 .unwrap();
476 assert!(matches!(outcome, HandshakeOutcome::Rejected { .. }));
477
478 for frame in [
479 Frame::Handshake {
480 version: CURRENT_VERSION,
481 },
482 Frame::Request {
483 id: crate::frame::OperationId::from("op-1"),
484 ops: "stats()".to_string(),
485 deadline_ms: None,
486 namespace: None,
487 actor_id: None,
488 visible_namespaces: None,
489 },
490 ] {
491 let err = gate.admit(&frame).unwrap_err();
492 match err.error.as_ref() {
493 Frame::Error { code, message, .. } => {
494 assert_eq!(*code, WireErrorCode::MalformedFrame);
495 assert!(
496 message.contains("closed by a rejected handshake"),
497 "message: {message}"
498 );
499 }
500 other => panic!("expected an error frame, got {other:?}"),
501 }
502 }
503 }
504
505 #[test]
506 fn closed_state_reports_a_sequence_violation_reason_not_handshake() {
507 let mut gate = HandshakeGate::default();
510 gate.admit(&Frame::Cancel {
511 id: crate::frame::OperationId::from("op-1"),
512 })
513 .unwrap_err();
514
515 let err = gate
516 .admit(&Frame::Handshake {
517 version: CURRENT_VERSION,
518 })
519 .unwrap_err();
520 match err.error.as_ref() {
521 Frame::Error { code, message, .. } => {
522 assert_eq!(*code, WireErrorCode::MalformedFrame);
523 assert!(
524 message.contains("closed by a frame before the handshake"),
525 "message: {message}"
526 );
527 assert!(
528 !message.contains("handshake failure")
529 && !message.contains("rejected handshake"),
530 "sequence-violation closure must not blame a handshake: {message}"
531 );
532 }
533 other => panic!("expected an error frame, got {other:?}"),
534 }
535 }
536
537 #[test]
538 fn closed_state_reports_a_duplicate_handshake_reason() {
539 let mut gate = completed_gate();
540 gate.admit(&Frame::Handshake {
541 version: CURRENT_VERSION,
542 })
543 .unwrap_err();
544
545 let err = gate
546 .admit(&Frame::Cancel {
547 id: crate::frame::OperationId::from("op-1"),
548 })
549 .unwrap_err();
550 match err.error.as_ref() {
551 Frame::Error { message, .. } => {
552 assert!(
553 message.contains("closed by a duplicate handshake"),
554 "message: {message}"
555 );
556 }
557 other => panic!("expected an error frame, got {other:?}"),
558 }
559 }
560
561 #[test]
562 fn closed_state_reports_a_server_only_kind_reason() {
563 let mut gate = completed_gate();
564 gate.admit(&Frame::Response {
565 id: crate::frame::OperationId::from("op-1"),
566 result: serde_json::json!({}),
567 })
568 .unwrap_err();
569
570 let err = gate
571 .admit(&Frame::Cancel {
572 id: crate::frame::OperationId::from("op-1"),
573 })
574 .unwrap_err();
575 match err.error.as_ref() {
576 Frame::Error { message, .. } => {
577 assert!(
578 message.contains("closed by a server-to-client frame"),
579 "message: {message}"
580 );
581 }
582 other => panic!("expected an error frame, got {other:?}"),
583 }
584 }
585}