1use std::fmt;
8use std::path::PathBuf;
9use std::rc::Rc;
10
11use crate::async_driver::DriverCancel;
12use crate::async_driver::LocalAsyncDriver;
13#[cfg(feature = "tokio")]
14use crate::async_driver::TokioLocalDriver;
15use crate::protocol::GraphError;
16
17#[cfg(feature = "tokio-websocket")]
18use futures_util::{SinkExt, StreamExt};
19
20#[derive(Debug, Clone, PartialEq, Eq)]
22pub struct ProcessCommand {
23 pub program: String,
25 pub args: Vec<String>,
27 pub cwd: Option<PathBuf>,
29 pub env: Vec<(String, String)>,
31}
32
33impl ProcessCommand {
34 pub fn new(program: impl Into<String>) -> Self {
36 Self {
37 program: program.into(),
38 args: Vec::new(),
39 cwd: None,
40 env: Vec::new(),
41 }
42 }
43
44 pub fn args<I, S>(mut self, args: I) -> Self
46 where
47 I: IntoIterator<Item = S>,
48 S: Into<String>,
49 {
50 self.args = args.into_iter().map(Into::into).collect();
51 self
52 }
53
54 pub fn cwd(mut self, cwd: impl Into<PathBuf>) -> Self {
56 self.cwd = Some(cwd.into());
57 self
58 }
59
60 pub fn env(mut self, key: impl Into<String>, value: impl Into<String>) -> Self {
62 self.env.push((key.into(), value.into()));
63 self
64 }
65}
66
67#[derive(Debug, Clone, PartialEq, Eq)]
69pub struct ProcessResult {
70 pub stdout: String,
72 pub stderr: String,
74 pub exit_code: Option<i32>,
76 pub signal: Option<String>,
78}
79
80#[derive(Debug, Clone, PartialEq, Eq)]
81pub struct HttpRequest {
83 pub method: String,
85 pub url: String,
87 pub headers: Vec<(String, String)>,
89 pub body: Vec<u8>,
91}
92
93impl HttpRequest {
94 pub fn new(method: impl Into<String>, url: impl Into<String>) -> Self {
96 Self {
97 method: method.into(),
98 url: url.into(),
99 headers: Vec::new(),
100 body: Vec::new(),
101 }
102 }
103
104 pub fn get(url: impl Into<String>) -> Self {
106 Self::new("GET", url)
107 }
108
109 pub fn header(mut self, key: impl Into<String>, value: impl Into<String>) -> Self {
111 self.headers.push((key.into(), value.into()));
112 self
113 }
114
115 pub fn body(mut self, body: impl Into<Vec<u8>>) -> Self {
117 self.body = body.into();
118 self
119 }
120}
121
122#[derive(Debug, Clone, PartialEq, Eq)]
123pub struct HttpResponse {
125 pub status: u16,
127 pub headers: Vec<(String, String)>,
129 pub body: Vec<u8>,
131}
132
133#[derive(Debug, Clone, PartialEq, Eq)]
134pub struct HttpStreamHead {
136 pub status: u16,
138 pub headers: Vec<(String, String)>,
140}
141
142#[derive(Debug, Clone, PartialEq, Eq)]
143pub struct SseRequest {
145 pub url: String,
147 pub headers: Vec<(String, String)>,
149}
150
151impl SseRequest {
152 pub fn new(url: impl Into<String>) -> Self {
154 Self {
155 url: url.into(),
156 headers: Vec::new(),
157 }
158 }
159
160 pub fn header(mut self, key: impl Into<String>, value: impl Into<String>) -> Self {
162 self.headers.push((key.into(), value.into()));
163 self
164 }
165}
166
167#[derive(Debug, Clone, PartialEq, Eq)]
168pub struct SseEvent {
170 pub event: Option<String>,
172 pub data: String,
174 pub id: Option<String>,
176 pub retry_ms: Option<u64>,
178}
179
180#[derive(Debug, Clone, PartialEq, Eq)]
181pub struct WebSocketRequest {
183 pub url: String,
185 pub headers: Vec<(String, String)>,
187}
188
189impl WebSocketRequest {
190 pub fn new(url: impl Into<String>) -> Self {
192 Self {
193 url: url.into(),
194 headers: Vec::new(),
195 }
196 }
197
198 pub fn header(mut self, key: impl Into<String>, value: impl Into<String>) -> Self {
200 self.headers.push((key.into(), value.into()));
201 self
202 }
203}
204
205#[derive(Debug, Clone, PartialEq, Eq)]
206pub enum WebSocketEvent {
208 Open,
210 Text(String),
212 Binary(Vec<u8>),
214 Close {
216 code: Option<u16>,
218 reason: Option<String>,
220 },
221}
222
223#[derive(Debug, Clone, PartialEq, Eq)]
224pub struct WebSocketSend {
226 pub data: Vec<u8>,
228 frame_kind: WebSocketFrameKind,
229}
230
231#[derive(Debug, Clone, Copy, PartialEq, Eq)]
232enum WebSocketFrameKind {
233 Text,
234 Binary,
235}
236
237impl WebSocketSend {
238 pub fn text(data: impl Into<String>) -> Self {
240 Self {
241 data: data.into().into_bytes(),
242 frame_kind: WebSocketFrameKind::Text,
243 }
244 }
245
246 pub fn binary(data: impl Into<Vec<u8>>) -> Self {
248 Self {
249 data: data.into(),
250 frame_kind: WebSocketFrameKind::Binary,
251 }
252 }
253}
254
255#[derive(Debug, Clone, Copy, PartialEq, Eq)]
256pub struct WebSocketSendResult {
258 pub sent: bool,
260}
261
262#[derive(Debug, Clone, PartialEq, Eq)]
263pub struct WebhookRegistration {
265 pub id: String,
267 pub method: Option<String>,
269 pub path: Option<String>,
271}
272
273impl WebhookRegistration {
274 pub fn new(id: impl Into<String>) -> Self {
276 Self {
277 id: id.into(),
278 method: None,
279 path: None,
280 }
281 }
282
283 pub fn method(mut self, method: impl Into<String>) -> Self {
285 self.method = Some(method.into());
286 self
287 }
288
289 pub fn path(mut self, path: impl Into<String>) -> Self {
291 self.path = Some(path.into());
292 self
293 }
294}
295
296#[derive(Debug, Clone, PartialEq, Eq)]
297pub struct WebhookEvent {
299 pub registration_id: String,
301 pub method: String,
303 pub path: String,
305 pub headers: Vec<(String, String)>,
307 pub query: Vec<(String, String)>,
309 pub body: Vec<u8>,
311}
312
313pub trait LocalProcessDriver {
315 fn run(
317 &self,
318 command: ProcessCommand,
319 callback: Box<dyn FnOnce(Result<ProcessResult, GraphError>)>,
320 ) -> DriverCancel;
321}
322
323pub trait LocalHttpDriver {
325 fn request(
327 &self,
328 request: HttpRequest,
329 callback: Box<dyn FnOnce(Result<HttpResponse, GraphError>)>,
330 ) -> DriverCancel;
331}
332
333pub enum HttpStreamDriverEvent {
335 Head(HttpStreamHead),
337 Chunk(Vec<u8>),
339 Error(GraphError),
341 Complete,
343}
344
345pub trait LocalHttpStreamDriver {
347 fn stream(
349 &self,
350 request: HttpRequest,
351 callback: Rc<dyn Fn(HttpStreamDriverEvent)>,
352 ) -> DriverCancel;
353}
354
355pub enum SseDriverEvent {
357 Event(SseEvent),
359 Error(GraphError),
361 Complete,
363}
364
365pub trait LocalSseDriver {
367 fn connect(&self, request: SseRequest, callback: Rc<dyn Fn(SseDriverEvent)>) -> DriverCancel;
369}
370
371pub enum WebSocketDriverEvent {
373 Event(WebSocketEvent),
375 Error(GraphError),
377 Complete,
379}
380
381pub trait LocalWebSocketDriver {
383 fn connect(
385 &self,
386 request: WebSocketRequest,
387 callback: Rc<dyn Fn(WebSocketDriverEvent)>,
388 ) -> DriverCancel;
389
390 fn send(
392 &self,
393 _request: WebSocketRequest,
394 _message: WebSocketSend,
395 _callback: Box<dyn FnOnce(Result<WebSocketSendResult, GraphError>)>,
396 ) -> Option<DriverCancel> {
397 None
398 }
399
400 fn connect_session(
402 &self,
403 _request: WebSocketRequest,
404 _callback: Rc<dyn Fn(WebSocketDriverEvent)>,
405 ) -> Option<Rc<dyn LocalWebSocketSession>> {
406 None
407 }
408}
409
410pub trait LocalWebSocketSession {
415 fn send(
417 &self,
418 message: WebSocketSend,
419 callback: Box<dyn FnOnce(Result<WebSocketSendResult, GraphError>)>,
420 ) -> DriverCancel;
421
422 fn close(&self, code: Option<u16>, reason: Option<String>);
424
425 fn cancel(&self);
427}
428
429pub enum WebhookDriverEvent {
431 Event(WebhookEvent),
433 Error(GraphError),
435 Complete,
437}
438
439pub trait LocalWebhookDriver {
441 fn register(
443 &self,
444 registration: WebhookRegistration,
445 callback: Rc<dyn Fn(WebhookDriverEvent)>,
446 ) -> DriverCancel;
447}
448
449#[cfg(feature = "tokio")]
450#[derive(Debug, Clone, Copy, Default)]
451pub struct TokioProcessDriver;
453
454#[cfg(feature = "tokio")]
455impl LocalProcessDriver for TokioProcessDriver {
456 fn run(
457 &self,
458 command: ProcessCommand,
459 callback: Box<dyn FnOnce(Result<ProcessResult, GraphError>)>,
460 ) -> DriverCancel {
461 let active = Rc::new(std::cell::Cell::new(true));
462 let active_for_task = active.clone();
463 let cancel_task = TokioLocalDriver.spawn_local(Box::pin(async move {
464 let mut cmd = tokio::process::Command::new(&command.program);
465 cmd.args(&command.args);
466 cmd.kill_on_drop(true);
467 if let Some(cwd) = command.cwd {
468 cmd.current_dir(cwd);
469 }
470 for (key, value) in command.env {
471 cmd.env(key, value);
472 }
473 match cmd.output().await {
474 Ok(output) => {
475 if active_for_task.get() {
476 callback(Ok(ProcessResult {
477 stdout: String::from_utf8_lossy(&output.stdout).into_owned(),
478 stderr: String::from_utf8_lossy(&output.stderr).into_owned(),
479 exit_code: output.status.code(),
480 signal: exit_signal(&output.status),
481 }));
482 }
483 }
484 Err(error) => {
485 if active_for_task.get() {
486 callback(Err(Box::new(error)));
487 }
488 }
489 }
490 }));
491 Box::new(move || {
492 active.set(false);
493 cancel_task();
494 })
495 }
496}
497
498#[cfg(all(feature = "tokio", unix))]
499fn exit_signal(status: &std::process::ExitStatus) -> Option<String> {
500 use std::os::unix::process::ExitStatusExt;
501 status.signal().map(|signal| signal.to_string())
502}
503
504#[cfg(all(feature = "tokio", not(unix)))]
505fn exit_signal(_status: &std::process::ExitStatus) -> Option<String> {
506 None
507}
508
509#[cfg(feature = "tokio-http")]
510#[derive(Debug, Clone, Default)]
511pub struct TokioHttpDriver {
513 client: reqwest::Client,
514}
515
516#[cfg(feature = "tokio-http")]
517impl TokioHttpDriver {
518 pub fn new() -> Self {
520 Self {
521 client: reqwest::Client::new(),
522 }
523 }
524}
525
526#[cfg(feature = "tokio-http")]
527impl LocalHttpDriver for TokioHttpDriver {
528 fn request(
529 &self,
530 request: HttpRequest,
531 callback: Box<dyn FnOnce(Result<HttpResponse, GraphError>)>,
532 ) -> DriverCancel {
533 let active = Rc::new(std::cell::Cell::new(true));
534 let active_for_task = active.clone();
535 let client = self.client.clone();
536 let cancel_task = TokioLocalDriver.spawn_local(Box::pin(async move {
537 let method = match reqwest::Method::from_bytes(request.method.as_bytes()) {
538 Ok(method) => method,
539 Err(error) => {
540 if active_for_task.get() {
541 callback(Err(Box::new(error)));
542 }
543 return;
544 }
545 };
546 let mut builder = client.request(method, request.url);
547 for (key, value) in request.headers {
548 builder = builder.header(key, value);
549 }
550 if !request.body.is_empty() {
551 builder = builder.body(request.body);
552 }
553 match builder.send().await {
554 Ok(response) => {
555 let status = response.status().as_u16();
556 let headers = response
557 .headers()
558 .iter()
559 .map(|(key, value)| {
560 (
561 key.as_str().to_owned(),
562 value.to_str().unwrap_or_default().to_owned(),
563 )
564 })
565 .collect::<Vec<_>>();
566 match response.bytes().await {
567 Ok(body) => {
568 if active_for_task.get() {
569 callback(Ok(HttpResponse {
570 status,
571 headers,
572 body: body.to_vec(),
573 }));
574 }
575 }
576 Err(error) => {
577 if active_for_task.get() {
578 callback(Err(Box::new(error)));
579 }
580 }
581 }
582 }
583 Err(error) => {
584 if active_for_task.get() {
585 callback(Err(Box::new(error)));
586 }
587 }
588 }
589 }));
590 Box::new(move || {
591 active.set(false);
592 cancel_task();
593 })
594 }
595}
596
597#[cfg(feature = "tokio-http-stream")]
598#[derive(Debug, Clone, Default)]
599pub struct TokioHttpStreamDriver {
601 client: reqwest::Client,
602}
603
604#[cfg(feature = "tokio-http-stream")]
605impl TokioHttpStreamDriver {
606 pub fn new() -> Self {
608 Self {
609 client: reqwest::Client::new(),
610 }
611 }
612}
613
614#[cfg(feature = "tokio-http-stream")]
615impl LocalHttpStreamDriver for TokioHttpStreamDriver {
616 fn stream(
617 &self,
618 request: HttpRequest,
619 callback: Rc<dyn Fn(HttpStreamDriverEvent)>,
620 ) -> DriverCancel {
621 let active = Rc::new(std::cell::Cell::new(true));
622 let active_for_task = active.clone();
623 let client = self.client.clone();
624 let cancel_task = TokioLocalDriver.spawn_local(Box::pin(async move {
625 let method = match reqwest::Method::from_bytes(request.method.as_bytes()) {
626 Ok(method) => method,
627 Err(error) => {
628 if active_for_task.replace(false) {
629 callback(HttpStreamDriverEvent::Error(Box::new(error)));
630 }
631 return;
632 }
633 };
634 let mut builder = client.request(method, request.url);
635 for (key, value) in request.headers {
636 builder = builder.header(key, value);
637 }
638 if !request.body.is_empty() {
639 builder = builder.body(request.body);
640 }
641 let mut response = match builder.send().await {
642 Ok(response) => response,
643 Err(error) => {
644 if active_for_task.replace(false) {
645 callback(HttpStreamDriverEvent::Error(Box::new(error)));
646 }
647 return;
648 }
649 };
650 let head = HttpStreamHead {
651 status: response.status().as_u16(),
652 headers: response
653 .headers()
654 .iter()
655 .map(|(key, value)| {
656 (
657 key.as_str().to_owned(),
658 value.to_str().unwrap_or_default().to_owned(),
659 )
660 })
661 .collect(),
662 };
663 if active_for_task.get() {
664 callback(HttpStreamDriverEvent::Head(head));
665 }
666 while active_for_task.get() {
667 match response.chunk().await {
668 Ok(Some(chunk)) => {
669 if !chunk.is_empty() && active_for_task.get() {
670 callback(HttpStreamDriverEvent::Chunk(chunk.to_vec()));
671 }
672 }
673 Ok(None) => {
674 if active_for_task.replace(false) {
675 callback(HttpStreamDriverEvent::Complete);
676 }
677 break;
678 }
679 Err(error) => {
680 if active_for_task.replace(false) {
681 callback(HttpStreamDriverEvent::Error(Box::new(error)));
682 }
683 break;
684 }
685 }
686 }
687 }));
688 Box::new(move || {
689 active.set(false);
690 cancel_task();
691 })
692 }
693}
694
695#[cfg(feature = "tokio-websocket")]
696#[derive(Debug, Clone, Copy, Default)]
697pub struct TokioWebSocketDriver;
699
700#[cfg(feature = "tokio-websocket")]
701impl LocalWebSocketDriver for TokioWebSocketDriver {
702 fn connect(
703 &self,
704 request: WebSocketRequest,
705 callback: Rc<dyn Fn(WebSocketDriverEvent)>,
706 ) -> DriverCancel {
707 let active = Rc::new(std::cell::Cell::new(true));
708 let active_for_task = active.clone();
709 let cancel_task = TokioLocalDriver.spawn_local(Box::pin(async move {
710 let client_request = match websocket_client_request(request) {
711 Ok(request) => request,
712 Err(error) => {
713 if active_for_task.get() {
714 callback(WebSocketDriverEvent::Error(error));
715 }
716 return;
717 }
718 };
719 match tokio_tungstenite::connect_async(client_request).await {
720 Ok((mut socket, _response)) => {
721 if active_for_task.get() {
722 callback(WebSocketDriverEvent::Event(WebSocketEvent::Open));
723 }
724 while active_for_task.get() {
725 match socket.next().await {
726 Some(Ok(message)) => {
727 if !active_for_task.get() {
728 break;
729 }
730 if let Some((event, complete)) =
731 websocket_event_from_message(message)
732 {
733 callback(WebSocketDriverEvent::Event(event));
734 if complete {
735 callback(WebSocketDriverEvent::Complete);
736 break;
737 }
738 }
739 }
740 Some(Err(error)) => {
741 if active_for_task.get() {
742 callback(WebSocketDriverEvent::Error(Box::new(error)));
743 }
744 break;
745 }
746 None => {
747 if active_for_task.get() {
748 callback(WebSocketDriverEvent::Complete);
749 }
750 break;
751 }
752 }
753 }
754 }
755 Err(error) => {
756 if active_for_task.get() {
757 callback(WebSocketDriverEvent::Error(Box::new(error)));
758 }
759 }
760 }
761 }));
762 Box::new(move || {
763 active.set(false);
764 cancel_task();
765 })
766 }
767
768 fn send(
769 &self,
770 request: WebSocketRequest,
771 message: WebSocketSend,
772 callback: Box<dyn FnOnce(Result<WebSocketSendResult, GraphError>)>,
773 ) -> Option<DriverCancel> {
774 let active = Rc::new(std::cell::Cell::new(true));
775 let active_for_task = active.clone();
776 let cancel_task = TokioLocalDriver.spawn_local(Box::pin(async move {
777 let client_request = match websocket_client_request(request) {
778 Ok(request) => request,
779 Err(error) => {
780 if active_for_task.get() {
781 callback(Err(error));
782 }
783 return;
784 }
785 };
786 match tokio_tungstenite::connect_async(client_request).await {
787 Ok((mut socket, _response)) => {
788 let result = match websocket_message_from_send(message) {
789 Ok(message) => socket
790 .send(message)
791 .await
792 .map(|()| WebSocketSendResult { sent: true })
793 .map_err(|error| Box::new(error) as GraphError),
794 Err(error) => Err(error),
795 };
796 let _ = socket.close(None).await;
797 if active_for_task.get() {
798 callback(result);
799 }
800 }
801 Err(error) => {
802 if active_for_task.get() {
803 callback(Err(Box::new(error)));
804 }
805 }
806 }
807 }));
808 Some(Box::new(move || {
809 active.set(false);
810 cancel_task();
811 }))
812 }
813
814 fn connect_session(
815 &self,
816 request: WebSocketRequest,
817 callback: Rc<dyn Fn(WebSocketDriverEvent)>,
818 ) -> Option<Rc<dyn LocalWebSocketSession>> {
819 let active = Rc::new(std::cell::Cell::new(true));
820 let opened = Rc::new(std::cell::Cell::new(false));
821 let (tx, mut rx) = tokio::sync::mpsc::channel::<TokioWebSocketSessionCommand>(1);
822 let active_for_task = active.clone();
823 let opened_for_task = opened.clone();
824 let cancel_task = TokioLocalDriver.spawn_local(Box::pin(async move {
825 let client_request = match websocket_client_request(request) {
826 Ok(request) => request,
827 Err(error) => {
828 if active_for_task.get() {
829 callback(WebSocketDriverEvent::Error(error));
830 }
831 return;
832 }
833 };
834 let (mut socket, _response) =
835 match tokio_tungstenite::connect_async(client_request).await {
836 Ok(connected) => connected,
837 Err(error) => {
838 if active_for_task.get() {
839 callback(WebSocketDriverEvent::Error(Box::new(error)));
840 }
841 return;
842 }
843 };
844 opened_for_task.set(true);
845 if active_for_task.get() {
846 callback(WebSocketDriverEvent::Event(WebSocketEvent::Open));
847 }
848 while active_for_task.get() {
849 tokio::select! {
850 command = rx.recv() => {
851 match command {
852 Some(TokioWebSocketSessionCommand::Send { message, active, callback }) => {
853 if !active.get() || !active_for_task.get() {
854 continue;
855 }
856 let result = match websocket_message_from_send(message) {
857 Ok(message) => socket
858 .send(message)
859 .await
860 .map(|()| WebSocketSendResult { sent: true })
861 .map_err(|error| Box::new(error) as GraphError),
862 Err(error) => Err(error),
863 };
864 if active.get() && active_for_task.get() {
865 callback(result);
866 }
867 }
868 Some(TokioWebSocketSessionCommand::Close { code, reason }) => {
869 let frame = websocket_close_frame(code, reason);
870 let _ = socket.close(frame).await;
871 active_for_task.set(false);
872 opened_for_task.set(false);
873 break;
874 }
875 Some(TokioWebSocketSessionCommand::Cancel) | None => {
876 let _ = socket.close(None).await;
877 active_for_task.set(false);
878 opened_for_task.set(false);
879 break;
880 }
881 }
882 }
883 message = socket.next() => {
884 match message {
885 Some(Ok(message)) => {
886 if !active_for_task.get() {
887 break;
888 }
889 if let Some((event, complete)) = websocket_event_from_message(message) {
890 callback(WebSocketDriverEvent::Event(event));
891 if complete {
892 callback(WebSocketDriverEvent::Complete);
893 active_for_task.set(false);
894 opened_for_task.set(false);
895 break;
896 }
897 }
898 }
899 Some(Err(error)) => {
900 if active_for_task.get() {
901 callback(WebSocketDriverEvent::Error(Box::new(error)));
902 }
903 active_for_task.set(false);
904 opened_for_task.set(false);
905 break;
906 }
907 None => {
908 if active_for_task.get() {
909 callback(WebSocketDriverEvent::Complete);
910 }
911 active_for_task.set(false);
912 opened_for_task.set(false);
913 break;
914 }
915 }
916 }
917 }
918 }
919 }));
920 Some(Rc::new(TokioWebSocketSession {
921 active,
922 closing: Rc::new(std::cell::Cell::new(false)),
923 opened,
924 tx,
925 cancel_task: Rc::new(std::cell::RefCell::new(Some(cancel_task))),
926 }))
927 }
928}
929
930#[cfg(feature = "tokio-websocket")]
931enum TokioWebSocketSessionCommand {
932 Send {
933 message: WebSocketSend,
934 active: Rc<std::cell::Cell<bool>>,
935 callback: Box<dyn FnOnce(Result<WebSocketSendResult, GraphError>)>,
936 },
937 Close {
938 code: Option<u16>,
939 reason: Option<String>,
940 },
941 Cancel,
942}
943
944#[cfg(feature = "tokio-websocket")]
945struct TokioWebSocketSession {
946 active: Rc<std::cell::Cell<bool>>,
947 closing: Rc<std::cell::Cell<bool>>,
948 opened: Rc<std::cell::Cell<bool>>,
949 tx: tokio::sync::mpsc::Sender<TokioWebSocketSessionCommand>,
950 cancel_task: Rc<std::cell::RefCell<Option<DriverCancel>>>,
951}
952
953#[cfg(feature = "tokio-websocket")]
954impl LocalWebSocketSession for TokioWebSocketSession {
955 fn send(
956 &self,
957 message: WebSocketSend,
958 callback: Box<dyn FnOnce(Result<WebSocketSendResult, GraphError>)>,
959 ) -> DriverCancel {
960 let send_active = Rc::new(std::cell::Cell::new(true));
961 if !self.active.get() {
962 send_active.set(false);
963 callback(Err("websocket session is closed".into()));
964 return Box::new(|| {});
965 }
966 match self.tx.try_send(TokioWebSocketSessionCommand::Send {
967 message,
968 active: send_active.clone(),
969 callback,
970 }) {
971 Ok(()) => {}
972 Err(tokio::sync::mpsc::error::TrySendError::Full(command)) => {
973 send_active.set(false);
974 let TokioWebSocketSessionCommand::Send { callback, .. } = command else {
975 return Box::new(|| {});
976 };
977 callback(Err("websocket session send queue is busy".into()));
978 return Box::new(|| {});
979 }
980 Err(tokio::sync::mpsc::error::TrySendError::Closed(command)) => {
981 send_active.set(false);
982 let TokioWebSocketSessionCommand::Send { callback, .. } = command else {
983 return Box::new(|| {});
984 };
985 callback(Err("websocket session is closed".into()));
986 return Box::new(|| {});
987 }
988 }
989 Box::new(move || {
990 send_active.set(false);
991 })
992 }
993
994 fn close(&self, code: Option<u16>, reason: Option<String>) {
995 if self.active.get() {
996 self.closing.set(true);
997 if !self.opened.get()
998 || self
999 .tx
1000 .try_send(TokioWebSocketSessionCommand::Close { code, reason })
1001 .is_err()
1002 {
1003 self.cancel();
1004 }
1005 }
1006 }
1007
1008 fn cancel(&self) {
1009 if self.active.replace(false) {
1010 let _ = self.tx.try_send(TokioWebSocketSessionCommand::Cancel);
1011 if let Some(cancel) = self.cancel_task.borrow_mut().take() {
1012 cancel();
1013 }
1014 }
1015 }
1016}
1017
1018#[cfg(feature = "tokio-websocket")]
1019impl Drop for TokioWebSocketSession {
1020 fn drop(&mut self) {
1021 if self.active.get() && !self.closing.get() {
1022 self.cancel();
1023 }
1024 }
1025}
1026
1027#[cfg(feature = "tokio-websocket")]
1028fn websocket_client_request(
1029 request: WebSocketRequest,
1030) -> Result<tokio_tungstenite::tungstenite::handshake::client::Request, GraphError> {
1031 use tokio_tungstenite::tungstenite::client::IntoClientRequest;
1032
1033 let mut client_request = request
1034 .url
1035 .as_str()
1036 .into_client_request()
1037 .map_err(|error| Box::new(error) as GraphError)?;
1038 for (key, value) in request.headers {
1039 let name = tokio_tungstenite::tungstenite::http::HeaderName::from_bytes(key.as_bytes())
1040 .map_err(|error| Box::new(error) as GraphError)?;
1041 let value = tokio_tungstenite::tungstenite::http::HeaderValue::from_str(&value)
1042 .map_err(|error| Box::new(error) as GraphError)?;
1043 client_request.headers_mut().append(name, value);
1044 }
1045 Ok(client_request)
1046}
1047
1048#[cfg(feature = "tokio-websocket")]
1049fn websocket_message_from_send(
1050 message: WebSocketSend,
1051) -> Result<tokio_tungstenite::tungstenite::Message, GraphError> {
1052 match message.frame_kind {
1053 WebSocketFrameKind::Text => String::from_utf8(message.data)
1054 .map(|text| tokio_tungstenite::tungstenite::Message::Text(text.into()))
1055 .map_err(|error| Box::new(error) as GraphError),
1056 WebSocketFrameKind::Binary => Ok(tokio_tungstenite::tungstenite::Message::Binary(
1057 message.data.into(),
1058 )),
1059 }
1060}
1061
1062#[cfg(feature = "tokio-websocket")]
1063fn websocket_event_from_message(
1064 message: tokio_tungstenite::tungstenite::Message,
1065) -> Option<(WebSocketEvent, bool)> {
1066 match message {
1067 tokio_tungstenite::tungstenite::Message::Text(text) => {
1068 Some((WebSocketEvent::Text(text.to_string()), false))
1069 }
1070 tokio_tungstenite::tungstenite::Message::Binary(bytes) => {
1071 Some((WebSocketEvent::Binary(bytes.to_vec()), false))
1072 }
1073 tokio_tungstenite::tungstenite::Message::Close(frame) => Some((
1074 WebSocketEvent::Close {
1075 code: frame.as_ref().map(|frame| u16::from(frame.code)),
1076 reason: frame.map(|frame| frame.reason.to_string()),
1077 },
1078 true,
1079 )),
1080 tokio_tungstenite::tungstenite::Message::Ping(_)
1081 | tokio_tungstenite::tungstenite::Message::Pong(_)
1082 | tokio_tungstenite::tungstenite::Message::Frame(_) => None,
1083 }
1084}
1085
1086#[cfg(feature = "tokio-websocket")]
1087fn websocket_close_frame(
1088 code: Option<u16>,
1089 reason: Option<String>,
1090) -> Option<tokio_tungstenite::tungstenite::protocol::CloseFrame> {
1091 if code.is_none() && reason.is_none() {
1092 return None;
1093 }
1094 let code = code
1095 .map(tokio_tungstenite::tungstenite::protocol::frame::coding::CloseCode::from)
1096 .unwrap_or(tokio_tungstenite::tungstenite::protocol::frame::coding::CloseCode::Normal);
1097 Some(tokio_tungstenite::tungstenite::protocol::CloseFrame {
1098 code,
1099 reason: reason.unwrap_or_default().into(),
1100 })
1101}
1102
1103#[derive(Clone, Default)]
1109pub struct EnvironmentDrivers {
1110 local_async: Option<Rc<dyn LocalAsyncDriver>>,
1111 process: Option<Rc<dyn LocalProcessDriver>>,
1112 http: Option<Rc<dyn LocalHttpDriver>>,
1113 http_stream: Option<Rc<dyn LocalHttpStreamDriver>>,
1114 sse: Option<Rc<dyn LocalSseDriver>>,
1115 websocket: Option<Rc<dyn LocalWebSocketDriver>>,
1116 webhook: Option<Rc<dyn LocalWebhookDriver>>,
1117}
1118
1119impl EnvironmentDrivers {
1120 pub fn new() -> Self {
1122 Self::default()
1123 }
1124
1125 pub fn with_local_async(mut self, driver: Rc<dyn LocalAsyncDriver>) -> Self {
1127 self.local_async = Some(driver);
1128 self
1129 }
1130
1131 pub fn with_process(mut self, driver: Rc<dyn LocalProcessDriver>) -> Self {
1133 self.process = Some(driver);
1134 self
1135 }
1136
1137 pub fn with_http(mut self, driver: Rc<dyn LocalHttpDriver>) -> Self {
1139 self.http = Some(driver);
1140 self
1141 }
1142
1143 pub fn with_http_stream(mut self, driver: Rc<dyn LocalHttpStreamDriver>) -> Self {
1145 self.http_stream = Some(driver);
1146 self
1147 }
1148
1149 pub fn with_sse(mut self, driver: Rc<dyn LocalSseDriver>) -> Self {
1151 self.sse = Some(driver);
1152 self
1153 }
1154
1155 pub fn with_websocket(mut self, driver: Rc<dyn LocalWebSocketDriver>) -> Self {
1157 self.websocket = Some(driver);
1158 self
1159 }
1160
1161 pub fn with_webhook(mut self, driver: Rc<dyn LocalWebhookDriver>) -> Self {
1163 self.webhook = Some(driver);
1164 self
1165 }
1166
1167 pub fn local_async_driver(&self) -> Option<Rc<dyn LocalAsyncDriver>> {
1169 self.local_async.clone()
1170 }
1171
1172 pub fn process_driver(&self) -> Option<Rc<dyn LocalProcessDriver>> {
1174 self.process.clone()
1175 }
1176
1177 pub fn http_driver(&self) -> Option<Rc<dyn LocalHttpDriver>> {
1179 self.http.clone()
1180 }
1181
1182 pub fn http_stream_driver(&self) -> Option<Rc<dyn LocalHttpStreamDriver>> {
1184 self.http_stream.clone()
1185 }
1186
1187 pub fn sse_driver(&self) -> Option<Rc<dyn LocalSseDriver>> {
1189 self.sse.clone()
1190 }
1191
1192 pub fn websocket_driver(&self) -> Option<Rc<dyn LocalWebSocketDriver>> {
1194 self.websocket.clone()
1195 }
1196
1197 pub fn webhook_driver(&self) -> Option<Rc<dyn LocalWebhookDriver>> {
1199 self.webhook.clone()
1200 }
1201
1202 pub(crate) fn set_local_async_driver(&mut self, driver: Option<Rc<dyn LocalAsyncDriver>>) {
1203 self.local_async = driver;
1204 }
1205}
1206
1207impl fmt::Debug for EnvironmentDrivers {
1208 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
1209 f.debug_struct("EnvironmentDrivers")
1210 .field(
1211 "local_async",
1212 &self.local_async.as_ref().map(|_| "<installed>"),
1213 )
1214 .field("process", &self.process.as_ref().map(|_| "<installed>"))
1215 .field("http", &self.http.as_ref().map(|_| "<installed>"))
1216 .field(
1217 "http_stream",
1218 &self.http_stream.as_ref().map(|_| "<installed>"),
1219 )
1220 .field("sse", &self.sse.as_ref().map(|_| "<installed>"))
1221 .field("websocket", &self.websocket.as_ref().map(|_| "<installed>"))
1222 .field("webhook", &self.webhook.as_ref().map(|_| "<installed>"))
1223 .finish()
1224 }
1225}
1226
1227#[cfg(all(test, feature = "tokio"))]
1228mod tests {
1229 use super::*;
1230 use std::cell::RefCell;
1231 use std::time::Duration;
1232
1233 #[cfg(feature = "tokio")]
1234 fn run_tokio_local<F>(future: F) -> F::Output
1235 where
1236 F: std::future::Future,
1237 {
1238 let runtime = tokio::runtime::Builder::new_current_thread()
1239 .enable_all()
1240 .build()
1241 .expect("tokio current-thread runtime");
1242 let local = tokio::task::LocalSet::new();
1243 local.block_on(&runtime, future)
1244 }
1245
1246 #[cfg(feature = "tokio")]
1247 async fn wait_until(label: &str, mut done: impl FnMut() -> bool) {
1248 tokio::time::timeout(Duration::from_secs(5), async {
1249 while !done() {
1250 tokio::task::yield_now().await;
1251 }
1252 })
1253 .await
1254 .unwrap_or_else(|_| panic!("timed out waiting for {label}"));
1255 }
1256
1257 #[cfg(feature = "tokio-http")]
1258 async fn read_http_request(stream: &mut tokio::net::TcpStream) -> String {
1259 use tokio::io::AsyncReadExt;
1260
1261 let mut buf = Vec::new();
1262 let mut chunk = [0_u8; 1024];
1263 let mut header_end = None;
1264 let mut content_length = 0_usize;
1265
1266 loop {
1267 let n = stream.read(&mut chunk).await.expect("read http request");
1268 assert_ne!(n, 0, "client closed before full http request arrived");
1269 buf.extend_from_slice(&chunk[..n]);
1270
1271 if header_end.is_none() {
1272 if let Some(pos) = buf.windows(4).position(|window| window == b"\r\n\r\n") {
1273 let headers = String::from_utf8_lossy(&buf[..pos]);
1274 content_length = headers
1275 .lines()
1276 .filter_map(|line| line.split_once(':'))
1277 .find_map(|(key, value)| {
1278 key.eq_ignore_ascii_case("content-length")
1279 .then(|| value.trim().parse::<usize>().ok())
1280 .flatten()
1281 })
1282 .unwrap_or(0);
1283 header_end = Some(pos + 4);
1284 }
1285 }
1286
1287 if let Some(end) = header_end {
1288 if buf.len() >= end + content_length {
1289 break;
1290 }
1291 }
1292 }
1293
1294 String::from_utf8_lossy(&buf).into_owned()
1295 }
1296
1297 #[cfg(feature = "tokio-websocket")]
1298 struct CaptureWebSocketHeader {
1299 key: &'static str,
1300 target: Rc<RefCell<Option<String>>>,
1301 }
1302
1303 #[cfg(feature = "tokio-websocket")]
1304 impl tokio_tungstenite::tungstenite::handshake::server::Callback for CaptureWebSocketHeader {
1305 #[allow(clippy::result_large_err)]
1306 fn on_request(
1307 self,
1308 request: &tokio_tungstenite::tungstenite::handshake::server::Request,
1309 response: tokio_tungstenite::tungstenite::handshake::server::Response,
1310 ) -> Result<
1311 tokio_tungstenite::tungstenite::handshake::server::Response,
1312 tokio_tungstenite::tungstenite::handshake::server::ErrorResponse,
1313 > {
1314 *self.target.borrow_mut() = request
1315 .headers()
1316 .get(self.key)
1317 .and_then(|value| value.to_str().ok())
1318 .map(str::to_owned);
1319 Ok(response)
1320 }
1321 }
1322
1323 #[cfg(feature = "tokio")]
1324 #[test]
1325 fn tokio_process_driver_runs_real_process_boundary() {
1326 run_tokio_local(async {
1327 let result = Rc::new(RefCell::new(None::<Result<ProcessResult, String>>));
1328 let result_for_callback = result.clone();
1329 let cancel = TokioProcessDriver.run(
1330 ProcessCommand::new("sh").args(["-c", "printf graphrefly"]),
1331 Box::new(move |value| {
1332 *result_for_callback.borrow_mut() =
1333 Some(value.map_err(|error| error.to_string()));
1334 }),
1335 );
1336
1337 wait_until("process driver callback", || result.borrow().is_some()).await;
1338 cancel();
1339
1340 let result = result
1341 .borrow_mut()
1342 .take()
1343 .expect("process driver callback fired")
1344 .expect("process completed");
1345 assert_eq!(result.stdout, "graphrefly");
1346 assert_eq!(result.stderr, "");
1347 assert_eq!(result.exit_code, Some(0));
1348 });
1349 }
1350
1351 #[cfg(feature = "tokio-http")]
1352 #[test]
1353 fn tokio_http_driver_requests_loopback_server() {
1354 run_tokio_local(async {
1355 use tokio::io::AsyncWriteExt;
1356
1357 let listener = tokio::net::TcpListener::bind("127.0.0.1:0")
1358 .await
1359 .expect("bind loopback http server");
1360 let addr = listener.local_addr().expect("loopback addr");
1361 let seen_request = Rc::new(RefCell::new(String::new()));
1362 let seen_request_for_task = seen_request.clone();
1363 tokio::task::spawn_local(async move {
1364 let (mut stream, _) = listener.accept().await.expect("accept http client");
1365 *seen_request_for_task.borrow_mut() = read_http_request(&mut stream).await;
1366 stream
1367 .write_all(
1368 b"HTTP/1.1 201 Created\r\nx-answer: 42\r\nContent-Length: 2\r\n\r\nok",
1369 )
1370 .await
1371 .expect("write http response");
1372 });
1373
1374 let result = Rc::new(RefCell::new(None::<Result<HttpResponse, String>>));
1375 let result_for_callback = result.clone();
1376 let cancel = TokioHttpDriver::new().request(
1377 HttpRequest::new("POST", format!("http://{addr}/orders"))
1378 .header("x-test", "yes")
1379 .body(b"hi".to_vec()),
1380 Box::new(move |value| {
1381 *result_for_callback.borrow_mut() =
1382 Some(value.map_err(|error| error.to_string()));
1383 }),
1384 );
1385
1386 wait_until("http driver callback", || result.borrow().is_some()).await;
1387 cancel();
1388
1389 let response = result
1390 .borrow_mut()
1391 .take()
1392 .expect("http driver callback fired")
1393 .expect("http response");
1394 assert_eq!(response.status, 201);
1395 assert_eq!(response.body, b"ok".to_vec());
1396 assert!(response
1397 .headers
1398 .iter()
1399 .any(|(key, value)| key == "x-answer" && value == "42"));
1400 let request = seen_request.borrow();
1401 assert!(request.starts_with("POST /orders HTTP/1.1"));
1402 assert!(request.contains("x-test: yes"));
1403 assert!(request.ends_with("hi"));
1404 });
1405 }
1406
1407 #[cfg(feature = "tokio-http-stream")]
1408 #[test]
1409 fn tokio_http_stream_driver_streams_loopback_response_body() {
1410 run_tokio_local(async {
1411 use tokio::io::AsyncWriteExt;
1412
1413 let listener = tokio::net::TcpListener::bind("127.0.0.1:0")
1414 .await
1415 .expect("bind loopback http stream server");
1416 let addr = listener.local_addr().expect("loopback addr");
1417 let seen_request = Rc::new(RefCell::new(String::new()));
1418 let seen_request_for_task = seen_request.clone();
1419 tokio::task::spawn_local(async move {
1420 let (mut stream, _) = listener.accept().await.expect("accept http client");
1421 *seen_request_for_task.borrow_mut() = read_http_request(&mut stream).await;
1422 stream
1423 .write_all(
1424 b"HTTP/1.1 200 OK\r\ncontent-type: text/event-stream\r\nContent-Length: 5\r\n\r\nhello",
1425 )
1426 .await
1427 .expect("write http stream response");
1428 });
1429
1430 let events = Rc::new(RefCell::new(Vec::<String>::new()));
1431 let events_for_callback = events.clone();
1432 let cancel = TokioHttpStreamDriver::new().stream(
1433 HttpRequest::get(format!("http://{addr}/events"))
1434 .header("accept", "text/event-stream"),
1435 Rc::new(move |event| match event {
1436 HttpStreamDriverEvent::Head(head) => {
1437 events_for_callback
1438 .borrow_mut()
1439 .push(format!("head:{}", head.status));
1440 }
1441 HttpStreamDriverEvent::Chunk(chunk) => {
1442 events_for_callback
1443 .borrow_mut()
1444 .push(format!("chunk:{}", String::from_utf8_lossy(&chunk)));
1445 }
1446 HttpStreamDriverEvent::Error(error) => {
1447 events_for_callback
1448 .borrow_mut()
1449 .push(format!("error:{error}"));
1450 }
1451 HttpStreamDriverEvent::Complete => {
1452 events_for_callback.borrow_mut().push("complete".to_owned());
1453 }
1454 }),
1455 );
1456
1457 wait_until("http stream complete", || {
1458 events.borrow().iter().any(|event| event == "complete")
1459 })
1460 .await;
1461 cancel();
1462
1463 assert_eq!(events.borrow().first(), Some(&"head:200".to_owned()));
1464 assert!(events.borrow().contains(&"chunk:hello".to_owned()));
1465 assert_eq!(events.borrow().last(), Some(&"complete".to_owned()));
1466 let request = seen_request.borrow();
1467 assert!(request.starts_with("GET /events HTTP/1.1"));
1468 assert!(request.contains("accept: text/event-stream"));
1469 });
1470 }
1471
1472 #[cfg(feature = "tokio-websocket")]
1473 #[test]
1474 fn tokio_websocket_driver_connects_and_streams_events() {
1475 run_tokio_local(async {
1476 let listener = tokio::net::TcpListener::bind("127.0.0.1:0")
1477 .await
1478 .expect("bind loopback websocket server");
1479 let addr = listener.local_addr().expect("loopback addr");
1480 let seen_header = Rc::new(RefCell::new(None::<String>));
1481 let seen_header_for_task = seen_header.clone();
1482 tokio::task::spawn_local(async move {
1483 let (stream, _) = listener.accept().await.expect("accept websocket client");
1484 let mut socket = tokio_tungstenite::accept_hdr_async(
1485 stream,
1486 CaptureWebSocketHeader {
1487 key: "x-graphrefly",
1488 target: seen_header_for_task,
1489 },
1490 )
1491 .await
1492 .expect("accept websocket handshake");
1493 socket
1494 .send(tokio_tungstenite::tungstenite::Message::Text(
1495 "hello".into(),
1496 ))
1497 .await
1498 .expect("send websocket text");
1499 socket.close(None).await.expect("close websocket");
1500 });
1501
1502 let events = Rc::new(RefCell::new(Vec::<String>::new()));
1503 let events_for_callback = events.clone();
1504 let cancel = TokioWebSocketDriver.connect(
1505 WebSocketRequest::new(format!("ws://{addr}")).header("x-graphrefly", "connect"),
1506 Rc::new(move |event| match event {
1507 WebSocketDriverEvent::Event(WebSocketEvent::Open) => {
1508 events_for_callback.borrow_mut().push("open".to_owned());
1509 }
1510 WebSocketDriverEvent::Event(WebSocketEvent::Text(text)) => {
1511 events_for_callback
1512 .borrow_mut()
1513 .push(format!("text:{text}"));
1514 }
1515 WebSocketDriverEvent::Event(WebSocketEvent::Binary(_))
1516 | WebSocketDriverEvent::Event(WebSocketEvent::Close { .. }) => {}
1517 WebSocketDriverEvent::Error(error) => {
1518 events_for_callback
1519 .borrow_mut()
1520 .push(format!("error:{error}"));
1521 }
1522 WebSocketDriverEvent::Complete => {
1523 events_for_callback.borrow_mut().push("complete".to_owned());
1524 }
1525 }),
1526 );
1527
1528 wait_until("websocket complete", || {
1529 events.borrow().iter().any(|event| event == "complete")
1530 })
1531 .await;
1532 cancel();
1533
1534 assert!(events.borrow().contains(&"open".to_owned()));
1535 assert!(events.borrow().contains(&"text:hello".to_owned()));
1536 assert_eq!(events.borrow().last(), Some(&"complete".to_owned()));
1537 assert_eq!(seen_header.borrow().as_deref(), Some("connect"));
1538 });
1539 }
1540
1541 #[cfg(feature = "tokio-websocket")]
1542 #[test]
1543 fn tokio_websocket_driver_sends_one_shot_message() {
1544 run_tokio_local(async {
1545 let listener = tokio::net::TcpListener::bind("127.0.0.1:0")
1546 .await
1547 .expect("bind loopback websocket server");
1548 let addr = listener.local_addr().expect("loopback addr");
1549 let received = Rc::new(RefCell::new(None::<String>));
1550 let received_for_task = received.clone();
1551 let seen_header = Rc::new(RefCell::new(None::<String>));
1552 let seen_header_for_task = seen_header.clone();
1553 tokio::task::spawn_local(async move {
1554 let (stream, _) = listener.accept().await.expect("accept websocket client");
1555 let mut socket = tokio_tungstenite::accept_hdr_async(
1556 stream,
1557 CaptureWebSocketHeader {
1558 key: "x-graphrefly",
1559 target: seen_header_for_task,
1560 },
1561 )
1562 .await
1563 .expect("accept websocket handshake");
1564 while let Some(message) = socket.next().await {
1565 match message.expect("websocket message") {
1566 tokio_tungstenite::tungstenite::Message::Binary(bytes) => {
1567 *received_for_task.borrow_mut() =
1568 Some(format!("binary:{}", String::from_utf8_lossy(&bytes)));
1569 break;
1570 }
1571 tokio_tungstenite::tungstenite::Message::Text(text) => {
1572 *received_for_task.borrow_mut() = Some(format!("text:{text}"));
1573 break;
1574 }
1575 tokio_tungstenite::tungstenite::Message::Close(_) => break,
1576 tokio_tungstenite::tungstenite::Message::Ping(_)
1577 | tokio_tungstenite::tungstenite::Message::Pong(_)
1578 | tokio_tungstenite::tungstenite::Message::Frame(_) => {}
1579 }
1580 }
1581 });
1582
1583 let result = Rc::new(RefCell::new(None::<Result<WebSocketSendResult, String>>));
1584 let result_for_callback = result.clone();
1585 let cancel = TokioWebSocketDriver
1586 .send(
1587 WebSocketRequest::new(format!("ws://{addr}")).header("x-graphrefly", "send"),
1588 WebSocketSend::text("hello"),
1589 Box::new(move |value| {
1590 *result_for_callback.borrow_mut() =
1591 Some(value.map_err(|error| error.to_string()));
1592 }),
1593 )
1594 .expect("send capability installed");
1595
1596 wait_until("websocket send callback and server receive", || {
1597 result.borrow().is_some() && received.borrow().is_some()
1598 })
1599 .await;
1600 cancel();
1601
1602 assert_eq!(
1603 result
1604 .borrow_mut()
1605 .take()
1606 .expect("websocket send callback fired")
1607 .expect("websocket send result"),
1608 WebSocketSendResult { sent: true }
1609 );
1610 assert_eq!(received.borrow().as_deref(), Some("text:hello"));
1611 assert_eq!(seen_header.borrow().as_deref(), Some("send"));
1612 });
1613 }
1614
1615 #[cfg(feature = "tokio-websocket")]
1616 #[test]
1617 fn tokio_websocket_driver_session_sends_over_same_connection_and_closes() {
1618 run_tokio_local(async {
1619 let listener = tokio::net::TcpListener::bind("127.0.0.1:0")
1620 .await
1621 .expect("bind loopback websocket server");
1622 let addr = listener.local_addr().expect("loopback addr");
1623 let received = Rc::new(RefCell::new(Vec::<String>::new()));
1624 let received_for_task = received.clone();
1625 let closed = Rc::new(std::cell::Cell::new(false));
1626 let closed_for_task = closed.clone();
1627 tokio::task::spawn_local(async move {
1628 let (stream, _) = listener.accept().await.expect("accept websocket client");
1629 let mut socket = tokio_tungstenite::accept_async(stream)
1630 .await
1631 .expect("accept websocket handshake");
1632 while let Some(message) = socket.next().await {
1633 match message.expect("websocket message") {
1634 tokio_tungstenite::tungstenite::Message::Text(text) => {
1635 received_for_task.borrow_mut().push(text.to_string());
1636 }
1637 tokio_tungstenite::tungstenite::Message::Binary(bytes) => {
1638 received_for_task
1639 .borrow_mut()
1640 .push(String::from_utf8_lossy(&bytes).into_owned());
1641 }
1642 tokio_tungstenite::tungstenite::Message::Close(_) => {
1643 closed_for_task.set(true);
1644 break;
1645 }
1646 tokio_tungstenite::tungstenite::Message::Ping(_)
1647 | tokio_tungstenite::tungstenite::Message::Pong(_)
1648 | tokio_tungstenite::tungstenite::Message::Frame(_) => {}
1649 }
1650 }
1651 });
1652
1653 let events = Rc::new(RefCell::new(Vec::<String>::new()));
1654 let events_for_callback = events.clone();
1655 let session = TokioWebSocketDriver
1656 .connect_session(
1657 WebSocketRequest::new(format!("ws://{addr}")),
1658 Rc::new(move |event| match event {
1659 WebSocketDriverEvent::Event(WebSocketEvent::Open) => {
1660 events_for_callback.borrow_mut().push("open".to_owned());
1661 }
1662 WebSocketDriverEvent::Event(WebSocketEvent::Close { .. }) => {
1663 events_for_callback.borrow_mut().push("close".to_owned());
1664 }
1665 WebSocketDriverEvent::Complete => {
1666 events_for_callback.borrow_mut().push("complete".to_owned());
1667 }
1668 WebSocketDriverEvent::Event(WebSocketEvent::Text(_))
1669 | WebSocketDriverEvent::Event(WebSocketEvent::Binary(_)) => {}
1670 WebSocketDriverEvent::Error(error) => {
1671 events_for_callback
1672 .borrow_mut()
1673 .push(format!("error:{error}"));
1674 }
1675 }),
1676 )
1677 .expect("session capability installed");
1678 wait_until("session open", || {
1679 events.borrow().iter().any(|event| event == "open")
1680 })
1681 .await;
1682
1683 let first = Rc::new(RefCell::new(None::<Result<WebSocketSendResult, String>>));
1684 let first_for_callback = first.clone();
1685 let _cancel_first = session.send(
1686 WebSocketSend::text("one"),
1687 Box::new(move |result| {
1688 *first_for_callback.borrow_mut() =
1689 Some(result.map_err(|error| error.to_string()));
1690 }),
1691 );
1692 wait_until("first session send", || {
1693 first.borrow().is_some() && received.borrow().len() == 1
1694 })
1695 .await;
1696 let second = Rc::new(RefCell::new(None::<Result<WebSocketSendResult, String>>));
1697 let second_for_callback = second.clone();
1698 let _cancel_second = session.send(
1699 WebSocketSend::text("two"),
1700 Box::new(move |result| {
1701 *second_for_callback.borrow_mut() =
1702 Some(result.map_err(|error| error.to_string()));
1703 }),
1704 );
1705
1706 wait_until("second session send", || {
1707 second.borrow().is_some() && received.borrow().len() == 2
1708 })
1709 .await;
1710 session.close(Some(1000), Some("done".to_owned()));
1711 wait_until("session close reaches server", || closed.get()).await;
1712
1713 assert_eq!(
1714 first
1715 .borrow_mut()
1716 .take()
1717 .expect("first send callback")
1718 .expect("first send result"),
1719 WebSocketSendResult { sent: true }
1720 );
1721 assert_eq!(
1722 second
1723 .borrow_mut()
1724 .take()
1725 .expect("second send callback")
1726 .expect("second send result"),
1727 WebSocketSendResult { sent: true }
1728 );
1729 assert_eq!(
1730 received.borrow().as_slice(),
1731 &["one".to_owned(), "two".to_owned()]
1732 );
1733 });
1734 }
1735}