Skip to main content

trillium_testing/
server_connector.rs

1use crate::{Runtime, TestTransport};
2use async_channel::Receiver;
3use std::{
4    io,
5    net::{IpAddr, SocketAddr},
6    sync::Arc,
7};
8use trillium::{Handler, Transport};
9use trillium_http::HttpContext;
10use trillium_server_common::Connector;
11use url::Url;
12
13/// a bridge between trillium servers and clients
14#[derive(Debug, fieldwork::Fieldwork)]
15pub struct ServerConnector<H> {
16    /// the handler
17    #[field(get, deref = false)]
18    handler: Arc<H>,
19
20    /// the runtime
21    #[field(with, set, get, into)]
22    runtime: Runtime,
23
24    /// the server config
25    #[field(with, set, get(deref = false), into)]
26    context: Arc<HttpContext>,
27
28    pub(crate) client_peer_ips_receiver: Option<Receiver<IpAddr>>,
29    pub(crate) server_peer_ips_receiver: Option<Receiver<IpAddr>>,
30}
31
32impl<H> Clone for ServerConnector<H> {
33    fn clone(&self) -> Self {
34        Self {
35            handler: self.handler.clone(),
36            runtime: self.runtime.clone(),
37            context: self.context.clone(),
38            client_peer_ips_receiver: self.client_peer_ips_receiver.clone(),
39            server_peer_ips_receiver: self.server_peer_ips_receiver.clone(),
40        }
41    }
42}
43
44impl<H: Handler> ServerConnector<H> {
45    /// builds a new ServerConnector
46    pub fn new(handler: H) -> Self {
47        Self {
48            handler: Arc::new(handler),
49            runtime: crate::runtime().into(),
50            context: Arc::default(),
51            client_peer_ips_receiver: None,
52            server_peer_ips_receiver: None,
53        }
54    }
55
56    /// opens a new connection to this virtual server, returning the client transport
57    pub async fn connect(&self, secure: bool) -> TestTransport {
58        let (mut client_transport, mut server_transport) = TestTransport::new();
59        if let Some(server_ip) = self
60            .client_peer_ips_receiver
61            .as_ref()
62            .and_then(|channel| channel.try_recv().ok())
63        {
64            client_transport.set_peer_ip(server_ip);
65        }
66
67        if let Some(client_ip) = self
68            .server_peer_ips_receiver
69            .as_ref()
70            .and_then(|channel| channel.try_recv().ok())
71        {
72            server_transport.set_peer_ip(client_ip);
73        }
74
75        let handler = Arc::clone(&self.handler);
76        let context = Arc::clone(&self.context);
77
78        let peer_ip = server_transport
79            .peer_addr()
80            .ok()
81            .flatten()
82            .map(|addr| addr.ip());
83
84        self.runtime.spawn_detached(async move {
85            let request_handler = Arc::clone(&handler);
86            let upgrade = context
87                .run(server_transport, |mut conn| {
88                    let handler = Arc::clone(&request_handler);
89                    async move {
90                        conn.set_peer_ip(peer_ip).set_secure(secure);
91                        let conn = handler.run(conn.into()).await;
92                        let conn = handler.before_send(conn).await;
93                        let mut inner = conn.into_inner::<TestTransport>();
94                        // an upgrade conn keeps its state so has_upgrade can see it
95                        if !inner.should_upgrade() {
96                            let state = std::mem::take(inner.state_mut());
97                            *inner.transport().state().write().unwrap() = state;
98                        }
99                        inner
100                    }
101                })
102                .await
103                .unwrap();
104
105            if let Some(upgrade) = upgrade {
106                let upgrade = upgrade.into();
107                if handler.has_upgrade(&upgrade) {
108                    handler.upgrade(upgrade).await;
109                } else {
110                    log::error!("upgrade specified but no upgrade handler provided");
111                }
112            }
113        });
114
115        client_transport
116    }
117}
118
119impl<H: Handler> Connector for ServerConnector<H> {
120    type Runtime = Runtime;
121    type Transport = TestTransport;
122    type Udp = ();
123
124    async fn connect(&self, url: &Url) -> io::Result<Self::Transport> {
125        Ok(self.connect(url.scheme() == "https").await)
126    }
127
128    fn runtime(&self) -> Self::Runtime {
129        #[allow(clippy::clone_on_copy)]
130        self.runtime.clone()
131    }
132
133    async fn resolve(&self, _host: &str, _port: u16) -> io::Result<Vec<SocketAddr>> {
134        Ok(vec![SocketAddr::from(([0, 0, 0, 0], 0))])
135    }
136}
137
138/// build a connector from this handler
139pub fn connector(handler: impl Handler) -> impl Connector {
140    ServerConnector::new(handler)
141}