1use std::io::{Read, Write};
42use std::time::Instant;
43
44use prost::Message;
45
46use crate::broker::protocol::{
47 read_frame, validate_frame_envelope, write_frame, Frame, FrameKind, FrameValidationError,
48 HandoffAck, HandoffOffer, PayloadEncoding, PROTOCOL_VERSION,
49};
50use crate::broker::server::deadline_stream::DeadlineStream;
51use crate::broker::server::handoff::handoff_token::HandoffToken;
52use crate::broker::server::handoff::orchestrate::{HandoffDelivery, HandoffDeliveryError};
53use crate::broker::server::handoff::windows::WindowsHandleValue;
54
55pub use crate::broker::protocol::registry::HANDOFF_PAYLOAD_PROTOCOL;
63
64pub fn handoff_offer_frame(offer: &HandoffOffer) -> Frame {
66 let mut payload = Vec::with_capacity(64);
67 offer.encode(&mut payload).expect(
68 "prost encoding HandoffOffer into Vec cannot fail because Vec writes are infallible",
69 );
70 Frame {
71 envelope_version: PROTOCOL_VERSION,
72 kind: FrameKind::Request as i32,
73 payload_protocol: HANDOFF_PAYLOAD_PROTOCOL,
74 payload,
75 request_id: offer.correlation_id,
76 payload_encoding: PayloadEncoding::None as i32,
77 deadline_unix_ms: 0,
78 traceparent: String::new(),
79 tracestate: String::new(),
80 }
81}
82
83#[allow(dead_code)]
84fn handoff_offer_frame_with_attachment(
85 attachment: crate::platform::ipc::HandoffAttachment,
86 token: &HandoffToken,
87 service_name: String,
88 correlation_id: u64,
89) -> Frame {
90 let offer = HandoffOffer {
91 handle_value: 0,
92 token: token.as_bytes().to_vec(),
93 service_name,
94 correlation_id,
95 };
96 let mut payload = Vec::with_capacity(64);
97 payload.push(0x08);
101 attachment.append_unsigned_varint(&mut payload);
102 offer.encode(&mut payload).expect(
103 "prost encoding HandoffOffer into Vec cannot fail because Vec writes are infallible",
104 );
105 Frame {
106 envelope_version: PROTOCOL_VERSION,
107 kind: FrameKind::Request as i32,
108 payload_protocol: HANDOFF_PAYLOAD_PROTOCOL,
109 payload,
110 request_id: correlation_id,
111 payload_encoding: PayloadEncoding::None as i32,
112 deadline_unix_ms: 0,
113 traceparent: String::new(),
114 tracestate: String::new(),
115 }
116}
117
118pub fn handoff_ack_frame(ack: &HandoffAck) -> Frame {
120 let mut payload = Vec::with_capacity(64);
121 ack.encode(&mut payload)
122 .expect("prost encoding HandoffAck into Vec cannot fail because Vec writes are infallible");
123 Frame {
124 envelope_version: PROTOCOL_VERSION,
125 kind: FrameKind::Response as i32,
126 payload_protocol: HANDOFF_PAYLOAD_PROTOCOL,
127 payload,
128 request_id: ack.correlation_id,
129 payload_encoding: PayloadEncoding::None as i32,
130 deadline_unix_ms: 0,
131 traceparent: String::new(),
132 tracestate: String::new(),
133 }
134}
135
136pub fn handoff_ready_frame(ack: &HandoffAck) -> Frame {
149 let mut payload = Vec::with_capacity(64);
150 ack.encode(&mut payload)
151 .expect("prost encoding HandoffAck into Vec cannot fail because Vec writes are infallible");
152 Frame {
153 envelope_version: PROTOCOL_VERSION,
154 kind: FrameKind::Event as i32,
155 payload_protocol: HANDOFF_PAYLOAD_PROTOCOL,
156 payload,
157 request_id: ack.correlation_id,
158 payload_encoding: PayloadEncoding::None as i32,
159 deadline_unix_ms: 0,
160 traceparent: String::new(),
161 tracestate: String::new(),
162 }
163}
164
165pub fn validate_handoff_frame(frame: &Frame, expected_kind: FrameKind) -> Result<(), &'static str> {
170 validate_frame_envelope(frame, expected_kind, HANDOFF_PAYLOAD_PROTOCOL).map_err(|error| {
171 match error {
172 FrameValidationError::EnvelopeVersion { .. } => "envelope_version is not v1",
173 FrameValidationError::Kind { .. } => match expected_kind {
174 FrameKind::Request => "kind is not REQUEST",
175 FrameKind::Event => "kind is not EVENT",
176 _ => "kind is not RESPONSE",
177 },
178 FrameValidationError::PayloadProtocol { .. } => "payload_protocol is not handoff",
179 FrameValidationError::PayloadEncoding { .. } => "payload is compressed",
180 }
181 })
182}
183
184#[derive(Debug)]
193pub struct WireHandoffDelivery<S> {
194 stream: S,
195 service_name: String,
196 correlation_id: u64,
197 configure_ack_bounded_io: Option<fn(&S, bool) -> std::io::Result<()>>,
198 ack_setup_error: Option<String>,
199 io_deadline: Instant,
200 backend_explicitly_rejected: bool,
201}
202
203impl<S> WireHandoffDelivery<S> {
204 pub fn new_preconfigured_nonblocking(
209 stream: S,
210 service_name: impl Into<String>,
211 correlation_id: u64,
212 io_deadline: Instant,
213 ) -> Self {
214 Self {
215 stream,
216 service_name: service_name.into(),
217 correlation_id,
218 configure_ack_bounded_io: None,
219 ack_setup_error: None,
220 io_deadline,
221 backend_explicitly_rejected: false,
222 }
223 }
224
225 #[deprecated(
231 note = "use new_preconfigured_nonblocking, or new_local_socket for production sockets"
232 )]
233 pub fn new(stream: S, service_name: impl Into<String>, correlation_id: u64) -> Self {
234 Self {
235 stream,
236 service_name: service_name.into(),
237 correlation_id,
238 configure_ack_bounded_io: None,
239 ack_setup_error: None,
240 io_deadline: Instant::now() + std::time::Duration::from_secs(30),
241 backend_explicitly_rejected: false,
242 }
243 }
244
245 pub fn correlation_id(&self) -> u64 {
247 self.correlation_id
248 }
249
250 pub fn stream(&self) -> &S {
253 &self.stream
254 }
255
256 pub fn backend_explicitly_rejected(&self) -> bool {
260 self.backend_explicitly_rejected
261 }
262
263 pub fn into_stream(self) -> S {
266 self.stream
267 }
268}
269
270impl WireHandoffDelivery<crate::platform::ipc::Stream> {
271 pub fn new_platform_stream(
274 stream: crate::platform::ipc::Stream,
275 service_name: impl Into<String>,
276 correlation_id: u64,
277 io_deadline: Instant,
278 ) -> Self {
279 let ack_setup_error = stream
280 .set_nonblocking(true)
281 .err()
282 .map(|error| format!("failed to enable bounded HandoffAck reads: {error}"));
283 Self {
284 stream,
285 service_name: service_name.into(),
286 correlation_id,
287 configure_ack_bounded_io: Some(|stream, bounded| stream.set_nonblocking(bounded)),
288 ack_setup_error,
289 io_deadline,
290 backend_explicitly_rejected: false,
291 }
292 }
293
294 #[allow(dead_code)]
296 pub(crate) fn deliver_attachment(
297 &mut self,
298 attachment: crate::platform::ipc::HandoffAttachment,
299 token: &HandoffToken,
300 ) -> Result<(), HandoffDeliveryError> {
301 let frame = handoff_offer_frame_with_attachment(
302 attachment,
303 token,
304 self.service_name.clone(),
305 self.correlation_id,
306 );
307 self.write_offer_frame(frame)
308 }
309}
310
311impl<S: Write> WireHandoffDelivery<S> {
312 fn write_offer_frame(&mut self, frame: Frame) -> Result<(), HandoffDeliveryError> {
313 if let Some(error) = self.ack_setup_error.as_ref() {
314 return Err(HandoffDeliveryError::DeliveryFailed {
315 detail: error.clone(),
316 });
317 }
318 let mut bytes = Vec::with_capacity(64);
319 frame
320 .encode(&mut bytes)
321 .expect("prost encoding Frame into Vec cannot fail because Vec writes are infallible");
322 let mut deadline_stream = DeadlineStream::new(&mut self.stream, self.io_deadline);
323 if let Err(write_error) = write_frame(&mut deadline_stream, &bytes) {
324 let restore_error = self
325 .configure_ack_bounded_io
326 .and_then(|configure| configure(&self.stream, false).err());
327 let detail = match restore_error {
328 Some(restore_error) => format!(
329 "failed to write HandoffOffer frame: {write_error}; additionally failed to \
330 restore blocking mode: {restore_error}"
331 ),
332 None => format!("failed to write HandoffOffer frame: {write_error}"),
333 };
334 return Err(HandoffDeliveryError::DeliveryFailed { detail });
335 }
336 Ok(())
337 }
338}
339
340impl WireHandoffDelivery<interprocess::local_socket::Stream> {
343 pub fn new_local_socket(
344 stream: interprocess::local_socket::Stream,
345 service_name: impl Into<String>,
346 correlation_id: u64,
347 io_deadline: Instant,
348 ) -> Self {
349 use interprocess::local_socket::traits::Stream as _;
350
351 let ack_setup_error = stream
352 .set_nonblocking(true)
353 .err()
354 .map(|error| format!("failed to enable bounded HandoffAck reads: {error}"));
355 Self {
356 stream,
357 service_name: service_name.into(),
358 correlation_id,
359 configure_ack_bounded_io: Some(|stream, bounded| {
360 use interprocess::local_socket::traits::Stream as _;
361 stream.set_nonblocking(bounded)
362 }),
363 ack_setup_error,
364 io_deadline,
365 backend_explicitly_rejected: false,
366 }
367 }
368}
369
370impl<S: Read + Write> HandoffDelivery for WireHandoffDelivery<S> {
371 fn deliver(
372 &mut self,
373 handle: WindowsHandleValue,
374 token: &HandoffToken,
375 ) -> Result<(), HandoffDeliveryError> {
376 let offer = HandoffOffer {
377 handle_value: handle.get() as u64,
378 token: token.as_bytes().to_vec(),
379 service_name: self.service_name.clone(),
380 correlation_id: self.correlation_id,
381 };
382 let frame = handoff_offer_frame(&offer);
383 self.write_offer_frame(frame)
384 }
385
386 fn await_backend_ack(
387 &mut self,
388 token: &HandoffToken,
389 deadline: Instant,
390 ) -> Result<Instant, HandoffDeliveryError> {
391 if let Some(error) = self.ack_setup_error.take() {
392 return Err(ack_not_observed(error));
393 }
394 let read_result = {
395 let mut deadline_stream = DeadlineStream::new(&mut self.stream, deadline);
396 read_frame(&mut deadline_stream)
397 };
398 let restore_result = self
399 .configure_ack_bounded_io
400 .map(|configure| configure(&self.stream, false))
401 .unwrap_or(Ok(()));
402 let bytes = match (read_result, restore_result) {
403 (Ok(bytes), Ok(())) => bytes,
404 (Err(read_error), Ok(())) => {
405 return Err(ack_not_observed(format!(
406 "failed to read HandoffAck frame: {read_error}"
407 )));
408 }
409 (Ok(_), Err(restore_error)) => {
410 return Err(ack_not_observed(format!(
411 "failed to restore blocking HandoffAck stream: {restore_error}"
412 )));
413 }
414 (Err(read_error), Err(restore_error)) => {
415 return Err(ack_not_observed(format!(
416 "failed to read HandoffAck frame: {read_error}; additionally failed to \
417 restore blocking mode: {restore_error}"
418 )));
419 }
420 };
421 let observed_at = Instant::now();
422 let frame = Frame::decode(bytes.as_slice()).map_err(|error| {
423 ack_not_observed(format!("failed to decode HandoffAck Frame: {error}"))
424 })?;
425 validate_handoff_frame(&frame, FrameKind::Response)
426 .map_err(|detail| ack_not_observed(format!("unexpected HandoffAck frame: {detail}")))?;
427 if frame.request_id != self.correlation_id {
428 return Err(ack_not_observed(format!(
429 "HandoffAck frame request_id {} does not match correlation id {}",
430 frame.request_id, self.correlation_id
431 )));
432 }
433 let ack = HandoffAck::decode(frame.payload.as_slice()).map_err(|error| {
434 ack_not_observed(format!("failed to decode HandoffAck payload: {error}"))
435 })?;
436 if ack.correlation_id != self.correlation_id {
437 return Err(ack_not_observed(format!(
438 "HandoffAck correlation id {} does not match offer correlation id {}",
439 ack.correlation_id, self.correlation_id
440 )));
441 }
442 if ack.token != token.as_bytes() {
443 return Err(ack_not_observed(
444 "HandoffAck token echo does not match the offered token".to_string(),
445 ));
446 }
447 if !ack.accepted {
448 self.backend_explicitly_rejected = true;
449 return Err(ack_not_observed(format!(
450 "backend refused the handoff: {}",
451 if ack.error_detail.is_empty() {
452 "no detail provided"
453 } else {
454 ack.error_detail.as_str()
455 }
456 )));
457 }
458 if observed_at > deadline {
459 return Err(ack_not_observed(
460 "backend HandoffAck arrived after the ACK deadline".to_string(),
461 ));
462 }
463 Ok(observed_at)
464 }
465}
466
467fn ack_not_observed(detail: String) -> HandoffDeliveryError {
468 HandoffDeliveryError::AckNotObserved { detail }
469}