Skip to main content

auc_tool/application/
server.rs

1use std::future::Future;
2use std::sync::Arc;
3use std::time::Duration;
4
5use capulus::managed::PeerCredentials;
6use rustix::net::sockopt::socket_peercred;
7use tokio::io::{AsyncReadExt, AsyncWriteExt};
8use tokio::net::{UnixListener, UnixStream};
9use tokio::sync::Semaphore;
10
11use super::protocol::{
12    ApplicationError, ApplicationRequest, ApplicationResponse, ErrorCode, PROTOCOL_MAJOR,
13    ProtocolError, RequestEnvelope, ResponseBody, ResponseEnvelope, decode, encode,
14};
15
16pub trait ApplicationHandler: Send + Sync + 'static {
17    fn handle(
18        &self,
19        peer: PeerCredentials,
20        request: ApplicationRequest,
21    ) -> impl Future<Output = Result<ApplicationResponse, ProtocolError>> + Send;
22}
23
24pub struct ApplicationServer<H> {
25    listener: UnixListener,
26    handler: Arc<H>,
27}
28
29impl<H: ApplicationHandler> ApplicationServer<H> {
30    pub fn new(listener: UnixListener, handler: Arc<H>) -> Self {
31        Self { listener, handler }
32    }
33
34    pub async fn run(self) -> Result<(), ApplicationError> {
35        let permits = Arc::new(Semaphore::new(32));
36        loop {
37            let permit = Arc::clone(&permits)
38                .acquire_owned()
39                .await
40                .expect("application connection semaphore remains open");
41            let (stream, _) = self.listener.accept().await?;
42            let handler = Arc::clone(&self.handler);
43            tokio::spawn(async move {
44                let _permit = permit;
45                if let Err(error) =
46                    tokio::time::timeout(Duration::from_secs(30), serve_connection(stream, handler))
47                        .await
48                        .map_err(|_| {
49                            ApplicationError::Io(std::io::Error::from(std::io::ErrorKind::TimedOut))
50                        })
51                        .and_then(|result| result)
52                {
53                    eprintln!("auc application connection failed: {error}");
54                }
55            });
56        }
57    }
58}
59
60async fn serve_connection<H: ApplicationHandler>(
61    mut stream: UnixStream,
62    handler: Arc<H>,
63) -> Result<(), ApplicationError> {
64    let credentials = socket_peercred(&stream).map_err(|error| {
65        ApplicationError::Io(std::io::Error::from_raw_os_error(error.raw_os_error()))
66    })?;
67    let peer = PeerCredentials {
68        pid: credentials.pid.as_raw_nonzero().get() as u32,
69        uid: credentials.uid.as_raw(),
70        gid: credentials.gid.as_raw(),
71    };
72    let request: RequestEnvelope = decode(&read_frame(&mut stream).await?)?;
73    let body = if request.minimum_protocol_major <= PROTOCOL_MAJOR
74        && request.maximum_protocol_major >= PROTOCOL_MAJOR
75    {
76        match handler.handle(peer, request.request).await {
77            Ok(response) => ResponseBody::Ok(response),
78            Err(error) => ResponseBody::Error(error),
79        }
80    } else {
81        ResponseBody::Error(ProtocolError::new(
82            ErrorCode::UnsupportedProtocol,
83            format!("auc-agent supports application protocol v{PROTOCOL_MAJOR}"),
84        ))
85    };
86    let payload = encode(&ResponseEnvelope {
87        request_id: request.request_id,
88        protocol_major: PROTOCOL_MAJOR,
89        body,
90    })?;
91    stream.write_u32(payload.len() as u32).await?;
92    stream.write_all(&payload).await?;
93    stream.flush().await?;
94    Ok(())
95}
96
97async fn read_frame(stream: &mut UnixStream) -> Result<Vec<u8>, ApplicationError> {
98    let length = match stream.read_u32().await {
99        Ok(length) => length as usize,
100        Err(error) if error.kind() == std::io::ErrorKind::UnexpectedEof => {
101            return Err(ApplicationError::EarlyEof);
102        }
103        Err(error) => return Err(error.into()),
104    };
105    if length > super::protocol::MAX_FRAME_BYTES {
106        return Err(ApplicationError::FrameTooLarge);
107    }
108    let mut payload = vec![0_u8; length];
109    stream.read_exact(&mut payload).await.map_err(|error| {
110        if error.kind() == std::io::ErrorKind::UnexpectedEof {
111            ApplicationError::EarlyEof
112        } else {
113            error.into()
114        }
115    })?;
116    Ok(payload)
117}
118
119#[cfg(test)]
120mod tests {
121    use std::sync::Mutex;
122
123    use tokio::sync::oneshot;
124
125    use super::*;
126    use crate::application::{ApplicationClient, Status};
127
128    struct RecordingHandler {
129        credentials: Mutex<Option<oneshot::Sender<PeerCredentials>>>,
130    }
131
132    impl ApplicationHandler for RecordingHandler {
133        async fn handle(
134            &self,
135            peer: PeerCredentials,
136            _request: ApplicationRequest,
137        ) -> Result<ApplicationResponse, ProtocolError> {
138            if let Some(sender) = self.credentials.lock().expect("lock").take() {
139                let _ = sender.send(peer);
140            }
141            Ok(ApplicationResponse::Status(Status {
142                product: "auc".to_string(),
143                package: "auc-tool".to_string(),
144                version: env!("CARGO_PKG_VERSION").to_string(),
145                protocol_major: PROTOCOL_MAJOR,
146                device_present: false,
147                pending_touch: false,
148                credential_count: 0,
149            }))
150        }
151    }
152
153    #[tokio::test]
154    async fn server_uses_kernel_peer_credentials() {
155        let directory = tempfile::tempdir().unwrap();
156        let path = directory.path().join("agent.sock");
157        let listener = UnixListener::bind(&path).unwrap();
158        let (credentials_tx, credentials_rx) = oneshot::channel();
159        let server = tokio::spawn(
160            ApplicationServer::new(
161                listener,
162                Arc::new(RecordingHandler {
163                    credentials: Mutex::new(Some(credentials_tx)),
164                }),
165            )
166            .run(),
167        );
168        let response = tokio::task::spawn_blocking(move || {
169            ApplicationClient::new(path).request(ApplicationRequest::Status)
170        })
171        .await
172        .unwrap()
173        .unwrap();
174        let credentials = tokio::time::timeout(Duration::from_secs(1), credentials_rx)
175            .await
176            .unwrap()
177            .unwrap();
178        server.abort();
179
180        assert!(matches!(response, ApplicationResponse::Status(_)));
181        assert_eq!(credentials.pid, std::process::id());
182        assert_eq!(credentials.uid, rustix::process::geteuid().as_raw());
183        assert_eq!(credentials.gid, rustix::process::getegid().as_raw());
184    }
185}