Skip to main content

graphrefly/
environment.rs

1//! Graph-owned environment driver bag (D130/D131).
2//!
3//! Environment drivers host wall-clock, process, network, messaging, and similar
4//! boundary work outside the synchronous wave core. The bag is graph-local and
5//! is attached to node ctx for source/adapter bodies.
6
7use 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/// Process command request for graph environment process drivers.
21#[derive(Debug, Clone, PartialEq, Eq)]
22pub struct ProcessCommand {
23    /// `program` field for program.
24    pub program: String,
25    /// `args` field for args.
26    pub args: Vec<String>,
27    /// `cwd` field for cwd.
28    pub cwd: Option<PathBuf>,
29    /// `env` field for env.
30    pub env: Vec<(String, String)>,
31}
32
33impl ProcessCommand {
34    /// Creates or computes `new`.
35    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    /// Updates or reads `args`.
45    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    /// Updates or reads `cwd`.
55    pub fn cwd(mut self, cwd: impl Into<PathBuf>) -> Self {
56        self.cwd = Some(cwd.into());
57        self
58    }
59
60    /// Updates or reads `env`.
61    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/// Completed process result. Exit status is DATA, including non-zero exits.
68#[derive(Debug, Clone, PartialEq, Eq)]
69pub struct ProcessResult {
70    /// `stdout` field for stdout.
71    pub stdout: String,
72    /// `stderr` field for stderr.
73    pub stderr: String,
74    /// `exit_code` field for exit code.
75    pub exit_code: Option<i32>,
76    /// `signal` field for signal.
77    pub signal: Option<String>,
78}
79
80#[derive(Debug, Clone, PartialEq, Eq)]
81/// `HttpRequest` data container.
82pub struct HttpRequest {
83    /// `method` field for method.
84    pub method: String,
85    /// `url` field for url.
86    pub url: String,
87    /// `headers` field for headers.
88    pub headers: Vec<(String, String)>,
89    /// `body` field for body.
90    pub body: Vec<u8>,
91}
92
93impl HttpRequest {
94    /// Creates or computes `new`.
95    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    /// Creates or computes `get`.
105    pub fn get(url: impl Into<String>) -> Self {
106        Self::new("GET", url)
107    }
108
109    /// Updates or reads `header`.
110    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    /// Updates or reads `body`.
116    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)]
123/// `HttpResponse` data container.
124pub struct HttpResponse {
125    /// `status` field for status.
126    pub status: u16,
127    /// `headers` field for headers.
128    pub headers: Vec<(String, String)>,
129    /// `body` field for body.
130    pub body: Vec<u8>,
131}
132
133#[derive(Debug, Clone, PartialEq, Eq)]
134/// `HttpStreamHead` data container.
135pub struct HttpStreamHead {
136    /// `status` field for status.
137    pub status: u16,
138    /// `headers` field for headers.
139    pub headers: Vec<(String, String)>,
140}
141
142#[derive(Debug, Clone, PartialEq, Eq)]
143/// `SseRequest` data container.
144pub struct SseRequest {
145    /// `url` field for url.
146    pub url: String,
147    /// `headers` field for headers.
148    pub headers: Vec<(String, String)>,
149}
150
151impl SseRequest {
152    /// Creates or computes `new`.
153    pub fn new(url: impl Into<String>) -> Self {
154        Self {
155            url: url.into(),
156            headers: Vec::new(),
157        }
158    }
159
160    /// Updates or reads `header`.
161    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)]
168/// `SseEvent` data container.
169pub struct SseEvent {
170    /// `event` field for event.
171    pub event: Option<String>,
172    /// `data` field for data.
173    pub data: String,
174    /// `id` field for id.
175    pub id: Option<String>,
176    /// `retry_ms` field for retry ms.
177    pub retry_ms: Option<u64>,
178}
179
180#[derive(Debug, Clone, PartialEq, Eq)]
181/// `WebSocketRequest` data container.
182pub struct WebSocketRequest {
183    /// `url` field for url.
184    pub url: String,
185    /// `headers` field for headers.
186    pub headers: Vec<(String, String)>,
187}
188
189impl WebSocketRequest {
190    /// Creates or computes `new`.
191    pub fn new(url: impl Into<String>) -> Self {
192        Self {
193            url: url.into(),
194            headers: Vec::new(),
195        }
196    }
197
198    /// Updates or reads `header`.
199    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)]
206/// `WebSocketEvent` variants.
207pub enum WebSocketEvent {
208    /// `Open` variant.
209    Open,
210    /// `Text` variant.
211    Text(String),
212    /// `Binary` variant.
213    Binary(Vec<u8>),
214    /// `Close` variant.
215    Close {
216        /// `code` field for code.
217        code: Option<u16>,
218        /// `reason` field for reason.
219        reason: Option<String>,
220    },
221}
222
223#[derive(Debug, Clone, PartialEq, Eq)]
224/// `WebSocketSend` data container.
225pub struct WebSocketSend {
226    /// `data` field for data.
227    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    /// Creates or computes `text`.
239    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    /// Creates or computes `binary`.
247    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)]
256/// `WebSocketSendResult` data container.
257pub struct WebSocketSendResult {
258    /// `sent` field for sent.
259    pub sent: bool,
260}
261
262#[derive(Debug, Clone, PartialEq, Eq)]
263/// `WebhookRegistration` data container.
264pub struct WebhookRegistration {
265    /// `id` field for id.
266    pub id: String,
267    /// `method` field for method.
268    pub method: Option<String>,
269    /// `path` field for path.
270    pub path: Option<String>,
271}
272
273impl WebhookRegistration {
274    /// Creates or computes `new`.
275    pub fn new(id: impl Into<String>) -> Self {
276        Self {
277            id: id.into(),
278            method: None,
279            path: None,
280        }
281    }
282
283    /// Updates or reads `method`.
284    pub fn method(mut self, method: impl Into<String>) -> Self {
285        self.method = Some(method.into());
286        self
287    }
288
289    /// Updates or reads `path`.
290    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)]
297/// `WebhookEvent` data container.
298pub struct WebhookEvent {
299    /// `registration_id` field for registration id.
300    pub registration_id: String,
301    /// `method` field for method.
302    pub method: String,
303    /// `path` field for path.
304    pub path: String,
305    /// `headers` field for headers.
306    pub headers: Vec<(String, String)>,
307    /// `query` field for query.
308    pub query: Vec<(String, String)>,
309    /// `body` field for body.
310    pub body: Vec<u8>,
311}
312
313/// Graph-local process driver. Implementations own process/runtime details.
314pub trait LocalProcessDriver {
315    /// Updates or reads `run`.
316    fn run(
317        &self,
318        command: ProcessCommand,
319        callback: Box<dyn FnOnce(Result<ProcessResult, GraphError>)>,
320    ) -> DriverCancel;
321}
322
323/// `LocalHttpDriver` behavior contract.
324pub trait LocalHttpDriver {
325    /// Updates or reads `request`.
326    fn request(
327        &self,
328        request: HttpRequest,
329        callback: Box<dyn FnOnce(Result<HttpResponse, GraphError>)>,
330    ) -> DriverCancel;
331}
332
333/// `HttpStreamDriverEvent` variants.
334pub enum HttpStreamDriverEvent {
335    /// `Head` variant.
336    Head(HttpStreamHead),
337    /// `Chunk` variant.
338    Chunk(Vec<u8>),
339    /// `Error` variant.
340    Error(GraphError),
341    /// `Complete` variant.
342    Complete,
343}
344
345/// `LocalHttpStreamDriver` behavior contract.
346pub trait LocalHttpStreamDriver {
347    /// Updates or reads `stream`.
348    fn stream(
349        &self,
350        request: HttpRequest,
351        callback: Rc<dyn Fn(HttpStreamDriverEvent)>,
352    ) -> DriverCancel;
353}
354
355/// `SseDriverEvent` variants.
356pub enum SseDriverEvent {
357    /// `Event` variant.
358    Event(SseEvent),
359    /// `Error` variant.
360    Error(GraphError),
361    /// `Complete` variant.
362    Complete,
363}
364
365/// `LocalSseDriver` behavior contract.
366pub trait LocalSseDriver {
367    /// Updates or reads `connect`.
368    fn connect(&self, request: SseRequest, callback: Rc<dyn Fn(SseDriverEvent)>) -> DriverCancel;
369}
370
371/// `WebSocketDriverEvent` variants.
372pub enum WebSocketDriverEvent {
373    /// `Event` variant.
374    Event(WebSocketEvent),
375    /// `Error` variant.
376    Error(GraphError),
377    /// `Complete` variant.
378    Complete,
379}
380
381/// `LocalWebSocketDriver` behavior contract.
382pub trait LocalWebSocketDriver {
383    /// Updates or reads `connect`.
384    fn connect(
385        &self,
386        request: WebSocketRequest,
387        callback: Rc<dyn Fn(WebSocketDriverEvent)>,
388    ) -> DriverCancel;
389
390    /// Updates or reads `send`.
391    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    /// Updates or reads `connect_session`.
401    fn connect_session(
402        &self,
403        _request: WebSocketRequest,
404        _callback: Rc<dyn Fn(WebSocketDriverEvent)>,
405    ) -> Option<Rc<dyn LocalWebSocketSession>> {
406        None
407    }
408}
409
410/// Live same-connection WebSocket session handle for D133/D174 SessionBundles.
411///
412/// Drivers create handles; graph-visible bundles own lifecycle, retry, status,
413/// command facts, and callback fencing.
414pub trait LocalWebSocketSession {
415    /// Updates or reads `send`.
416    fn send(
417        &self,
418        message: WebSocketSend,
419        callback: Box<dyn FnOnce(Result<WebSocketSendResult, GraphError>)>,
420    ) -> DriverCancel;
421
422    /// Updates or reads `close`.
423    fn close(&self, code: Option<u16>, reason: Option<String>);
424
425    /// Updates or reads `cancel`.
426    fn cancel(&self);
427}
428
429/// `WebhookDriverEvent` variants.
430pub enum WebhookDriverEvent {
431    /// `Event` variant.
432    Event(WebhookEvent),
433    /// `Error` variant.
434    Error(GraphError),
435    /// `Complete` variant.
436    Complete,
437}
438
439/// `LocalWebhookDriver` behavior contract.
440pub trait LocalWebhookDriver {
441    /// Updates or reads `register`.
442    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)]
451/// `TokioProcessDriver` data container.
452pub 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)]
511/// `TokioHttpDriver` data container.
512pub struct TokioHttpDriver {
513    client: reqwest::Client,
514}
515
516#[cfg(feature = "tokio-http")]
517impl TokioHttpDriver {
518    /// Creates or computes `new`.
519    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)]
599/// `TokioHttpStreamDriver` data container.
600pub struct TokioHttpStreamDriver {
601    client: reqwest::Client,
602}
603
604#[cfg(feature = "tokio-http-stream")]
605impl TokioHttpStreamDriver {
606    /// Creates or computes `new`.
607    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)]
697/// `TokioWebSocketDriver` data container.
698pub 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/// Graph-local environment capabilities for source/adapter boundaries.
1104///
1105/// First slice carries the existing local async/time driver. Process, network,
1106/// messaging, and resilience driver groups grow here rather than on
1107/// [`crate::dispatcher::Dispatcher`] (D131).
1108#[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    /// Creates or computes `new`.
1121    pub fn new() -> Self {
1122        Self::default()
1123    }
1124
1125    /// Updates or reads `with_local_async`.
1126    pub fn with_local_async(mut self, driver: Rc<dyn LocalAsyncDriver>) -> Self {
1127        self.local_async = Some(driver);
1128        self
1129    }
1130
1131    /// Updates or reads `with_process`.
1132    pub fn with_process(mut self, driver: Rc<dyn LocalProcessDriver>) -> Self {
1133        self.process = Some(driver);
1134        self
1135    }
1136
1137    /// Updates or reads `with_http`.
1138    pub fn with_http(mut self, driver: Rc<dyn LocalHttpDriver>) -> Self {
1139        self.http = Some(driver);
1140        self
1141    }
1142
1143    /// Updates or reads `with_http_stream`.
1144    pub fn with_http_stream(mut self, driver: Rc<dyn LocalHttpStreamDriver>) -> Self {
1145        self.http_stream = Some(driver);
1146        self
1147    }
1148
1149    /// Updates or reads `with_sse`.
1150    pub fn with_sse(mut self, driver: Rc<dyn LocalSseDriver>) -> Self {
1151        self.sse = Some(driver);
1152        self
1153    }
1154
1155    /// Updates or reads `with_websocket`.
1156    pub fn with_websocket(mut self, driver: Rc<dyn LocalWebSocketDriver>) -> Self {
1157        self.websocket = Some(driver);
1158        self
1159    }
1160
1161    /// Updates or reads `with_webhook`.
1162    pub fn with_webhook(mut self, driver: Rc<dyn LocalWebhookDriver>) -> Self {
1163        self.webhook = Some(driver);
1164        self
1165    }
1166
1167    /// Updates or reads `local_async_driver`.
1168    pub fn local_async_driver(&self) -> Option<Rc<dyn LocalAsyncDriver>> {
1169        self.local_async.clone()
1170    }
1171
1172    /// Updates or reads `process_driver`.
1173    pub fn process_driver(&self) -> Option<Rc<dyn LocalProcessDriver>> {
1174        self.process.clone()
1175    }
1176
1177    /// Updates or reads `http_driver`.
1178    pub fn http_driver(&self) -> Option<Rc<dyn LocalHttpDriver>> {
1179        self.http.clone()
1180    }
1181
1182    /// Updates or reads `http_stream_driver`.
1183    pub fn http_stream_driver(&self) -> Option<Rc<dyn LocalHttpStreamDriver>> {
1184        self.http_stream.clone()
1185    }
1186
1187    /// Updates or reads `sse_driver`.
1188    pub fn sse_driver(&self) -> Option<Rc<dyn LocalSseDriver>> {
1189        self.sse.clone()
1190    }
1191
1192    /// Updates or reads `websocket_driver`.
1193    pub fn websocket_driver(&self) -> Option<Rc<dyn LocalWebSocketDriver>> {
1194        self.websocket.clone()
1195    }
1196
1197    /// Updates or reads `webhook_driver`.
1198    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}