Skip to main content

russh/keys/agent/
server.rs

1use std::collections::HashMap;
2use std::marker::Sync;
3use std::sync::{Arc, RwLock};
4use std::time::{Duration, SystemTime};
5
6use byteorder::{BigEndian, ByteOrder};
7use bytes::Bytes;
8use futures::future::Future;
9use futures::stream::{Stream, StreamExt};
10use ssh_encoding::{Decode, Encode, Reader};
11use ssh_key::PrivateKey;
12use tokio::io::{AsyncRead, AsyncReadExt, AsyncWrite, AsyncWriteExt};
13use tokio::time::sleep;
14use {std, tokio};
15
16use super::{msg, Constraint};
17use crate::helpers::{sign_with_hash_alg, EncodedExt};
18use crate::keys::key::PrivateKeyWithHashAlg;
19use crate::keys::Error;
20use crate::CryptoVec;
21
22const MAX_AGENT_FRAME_LEN: usize = 256 * 1024;
23
24#[derive(Clone)]
25#[allow(clippy::type_complexity)]
26struct KeyStore(Arc<RwLock<HashMap<Vec<u8>, (Arc<PrivateKey>, SystemTime, Vec<Constraint>)>>>);
27
28#[derive(Clone)]
29struct Lock(Arc<RwLock<CryptoVec>>);
30
31#[allow(missing_docs)]
32#[derive(Debug)]
33pub enum ServerError<E> {
34    E(E),
35    Error(Error),
36}
37
38pub enum MessageType {
39    RequestKeys,
40    AddKeys,
41    RemoveKeys,
42    RemoveAllKeys,
43    Sign,
44    Lock,
45    Unlock,
46}
47
48#[cfg_attr(feature = "async-trait", async_trait::async_trait)]
49pub trait Agent: Clone + Send + 'static {
50    fn confirm(
51        self,
52        _pk: Arc<PrivateKey>,
53    ) -> Box<dyn Future<Output = (Self, bool)> + Unpin + Send> {
54        Box::new(futures::future::ready((self, true)))
55    }
56
57    fn confirm_request(&self, _msg: MessageType) -> impl Future<Output = bool> + Send {
58        async { true }
59    }
60}
61
62pub async fn serve<S, L, A>(mut listener: L, agent: A) -> Result<(), Error>
63where
64    S: AsyncRead + AsyncWrite + Send + Sync + Unpin + 'static,
65    L: Stream<Item = tokio::io::Result<S>> + Unpin,
66    A: Agent + Send + Sync + 'static,
67{
68    let keys = KeyStore(Arc::new(RwLock::new(HashMap::new())));
69    let lock = Lock(Arc::new(RwLock::new(CryptoVec::new())));
70    while let Some(Ok(stream)) = listener.next().await {
71        russh_util::runtime::spawn(
72            (Connection {
73                lock: lock.clone(),
74                keys: keys.clone(),
75                agent: Some(agent.clone()),
76                s: stream,
77                buf: Vec::new(),
78            })
79            .run(),
80        );
81    }
82    Ok(())
83}
84
85impl Agent for () {
86    fn confirm(self, _: Arc<PrivateKey>) -> Box<dyn Future<Output = (Self, bool)> + Unpin + Send> {
87        Box::new(futures::future::ready((self, true)))
88    }
89}
90
91struct Connection<S: AsyncRead + AsyncWrite + Send + 'static, A: Agent> {
92    lock: Lock,
93    keys: KeyStore,
94    agent: Option<A>,
95    s: S,
96    buf: Vec<u8>,
97}
98
99impl<S: AsyncRead + AsyncWrite + Send + Unpin + 'static, A: Agent + Send + Sync + 'static>
100    Connection<S, A>
101{
102    async fn read_frame(&mut self) -> Result<(), Error> {
103        self.buf.clear();
104        self.buf.resize(4, 0);
105        self.s.read_exact(&mut self.buf).await?;
106
107        let len = BigEndian::read_u32(&self.buf) as usize;
108        if len > MAX_AGENT_FRAME_LEN {
109            return Err(Error::AgentProtocolError);
110        }
111
112        self.buf.clear();
113        self.buf.resize(len, 0);
114        self.s.read_exact(&mut self.buf).await?;
115        Ok(())
116    }
117
118    async fn run(mut self) -> Result<(), Error> {
119        let mut writebuf = Vec::new();
120        loop {
121            self.read_frame().await?;
122            // respond
123            writebuf.clear();
124            self.respond(&mut writebuf).await?;
125            self.s.write_all(&writebuf).await?;
126            self.s.flush().await?
127        }
128    }
129
130    async fn respond(&mut self, writebuf: &mut Vec<u8>) -> Result<(), Error> {
131        let is_locked = {
132            if let Ok(password) = self.lock.0.read() {
133                !password.is_empty()
134            } else {
135                true
136            }
137        };
138        writebuf.extend_from_slice(&[0, 0, 0, 0]);
139        let agentref = self.agent.as_ref().ok_or(Error::AgentFailure)?;
140
141        match self.buf.split_first() {
142            Some((&11, _))
143                if !is_locked && agentref.confirm_request(MessageType::RequestKeys).await =>
144            {
145                // request identities
146                if let Ok(keys) = self.keys.0.read() {
147                    msg::IDENTITIES_ANSWER.encode(writebuf)?;
148                    (keys.len() as u32).encode(writebuf)?;
149                    for (k, _) in keys.iter() {
150                        k.encode(writebuf)?;
151                        "".encode(writebuf)?;
152                    }
153                } else {
154                    msg::FAILURE.encode(writebuf)?
155                }
156            }
157            Some((&13, mut r))
158                if !is_locked && agentref.confirm_request(MessageType::Sign).await =>
159            {
160                // sign request
161                let agent = self.agent.take().ok_or(Error::AgentFailure)?;
162                let (agent, signed) = self.try_sign(agent, &mut r, writebuf).await?;
163                self.agent = Some(agent);
164                if signed {
165                    return Ok(());
166                } else {
167                    writebuf.resize(4, 0);
168                    writebuf.push(msg::FAILURE)
169                }
170            }
171            Some((&17, mut r))
172                if !is_locked && agentref.confirm_request(MessageType::AddKeys).await =>
173            {
174                // add identity
175                if let Ok(true) = self.add_key(&mut r, false, writebuf).await {
176                } else {
177                    writebuf.push(msg::FAILURE)
178                }
179            }
180            Some((&18, mut r))
181                if !is_locked && agentref.confirm_request(MessageType::RemoveKeys).await =>
182            {
183                // remove identity
184                if let Ok(true) = self.remove_identity(&mut r) {
185                    writebuf.push(msg::SUCCESS)
186                } else {
187                    writebuf.push(msg::FAILURE)
188                }
189            }
190            Some((&19, _))
191                if !is_locked && agentref.confirm_request(MessageType::RemoveAllKeys).await =>
192            {
193                // remove all identities
194                if let Ok(mut keys) = self.keys.0.write() {
195                    keys.clear();
196                    writebuf.push(msg::SUCCESS)
197                } else {
198                    writebuf.push(msg::FAILURE)
199                }
200            }
201            Some((&22, mut r))
202                if !is_locked && agentref.confirm_request(MessageType::Lock).await =>
203            {
204                // lock
205                if let Ok(()) = self.lock(&mut r) {
206                    writebuf.push(msg::SUCCESS)
207                } else {
208                    writebuf.push(msg::FAILURE)
209                }
210            }
211            Some((&23, mut r))
212                if is_locked && agentref.confirm_request(MessageType::Unlock).await =>
213            {
214                // unlock
215                if let Ok(true) = self.unlock(&mut r) {
216                    writebuf.push(msg::SUCCESS)
217                } else {
218                    writebuf.push(msg::FAILURE)
219                }
220            }
221            Some((&25, mut r))
222                if !is_locked && agentref.confirm_request(MessageType::AddKeys).await =>
223            {
224                // add identity constrained
225                if let Ok(true) = self.add_key(&mut r, true, writebuf).await {
226                } else {
227                    writebuf.push(msg::FAILURE)
228                }
229            }
230            _ => {
231                // Message not understood
232                writebuf.push(msg::FAILURE)
233            }
234        }
235        let len = writebuf.len() - 4;
236        BigEndian::write_u32(&mut writebuf[..], len as u32);
237        Ok(())
238    }
239
240    fn lock<R: Reader>(&self, r: &mut R) -> Result<(), Error> {
241        let password = Bytes::decode(r)?;
242        let mut lock = self.lock.0.write().or(Err(Error::AgentFailure))?;
243        lock.extend(&password);
244        Ok(())
245    }
246
247    fn unlock<R: Reader>(&self, r: &mut R) -> Result<bool, Error> {
248        let password = Bytes::decode(r)?;
249        let mut lock = self.lock.0.write().or(Err(Error::AgentFailure))?;
250        if lock[..] == password {
251            lock.clear();
252            Ok(true)
253        } else {
254            Ok(false)
255        }
256    }
257
258    fn remove_identity<R: Reader>(&self, r: &mut R) -> Result<bool, Error> {
259        if let Ok(mut keys) = self.keys.0.write() {
260            if keys.remove(&Bytes::decode(r)?.to_vec()).is_some() {
261                Ok(true)
262            } else {
263                Ok(false)
264            }
265        } else {
266            Ok(false)
267        }
268    }
269
270    async fn add_key<R: Reader>(
271        &self,
272        r: &mut R,
273        constrained: bool,
274        writebuf: &mut Vec<u8>,
275    ) -> Result<bool, Error> {
276        let (blob, key_pair) = {
277            let private_key =
278                ssh_key::private::PrivateKey::new(ssh_key::private::KeypairData::decode(r)?, "")?;
279            let _comment = String::decode(r)?;
280
281            (private_key.public_key().key_data().encoded()?, private_key)
282        };
283        writebuf.push(msg::SUCCESS);
284        let mut w = self.keys.0.write().or(Err(Error::AgentFailure))?;
285        let now = SystemTime::now();
286        if constrained {
287            let mut c = Vec::new();
288            while let Ok(t) = u8::decode(r) {
289                if t == msg::CONSTRAIN_LIFETIME {
290                    let seconds = u32::decode(r)?;
291                    c.push(Constraint::KeyLifetime { seconds });
292                    let blob = blob.clone();
293                    let keys = self.keys.clone();
294                    russh_util::runtime::spawn(async move {
295                        sleep(Duration::from_secs(seconds as u64)).await;
296                        if let Ok(mut keys) = keys.0.write() {
297                            let delete = if let Some(&(_, time, _)) = keys.get(&blob) {
298                                time == now
299                            } else {
300                                false
301                            };
302                            if delete {
303                                keys.remove(&blob);
304                            }
305                        }
306                    });
307                } else if t == msg::CONSTRAIN_CONFIRM {
308                    c.push(Constraint::Confirm)
309                } else {
310                    return Ok(false);
311                }
312            }
313            w.insert(blob, (Arc::new(key_pair), now, c));
314        } else {
315            w.insert(blob, (Arc::new(key_pair), now, Vec::new()));
316        }
317        Ok(true)
318    }
319
320    async fn try_sign<R: Reader>(
321        &self,
322        agent: A,
323        r: &mut R,
324        writebuf: &mut Vec<u8>,
325    ) -> Result<(A, bool), Error> {
326        let mut needs_confirm = false;
327        let key = {
328            let blob = Bytes::decode(r)?;
329            let k = self.keys.0.read().or(Err(Error::AgentFailure))?;
330            if let Some((key, _, constraints)) = k.get(&blob.to_vec()) {
331                if constraints.contains(&Constraint::Confirm) {
332                    needs_confirm = true;
333                }
334                key.clone()
335            } else {
336                return Ok((agent, false));
337            }
338        };
339        let agent = if needs_confirm {
340            let (agent, ok) = {
341                let _pk = key.clone();
342                Box::new(futures::future::ready((agent, true)))
343            }
344            .await;
345            if !ok {
346                return Ok((agent, false));
347            }
348            agent
349        } else {
350            agent
351        };
352        writebuf.push(msg::SIGN_RESPONSE);
353        let data = Bytes::decode(r)?;
354
355        sign_with_hash_alg(&PrivateKeyWithHashAlg::new(key, None), &data)?.encode(writebuf)?;
356
357        let len = writebuf.len();
358        BigEndian::write_u32(writebuf, (len - 4) as u32);
359
360        Ok((agent, true))
361    }
362}
363
364#[cfg(test)]
365mod tests {
366    use byteorder::{BigEndian, ByteOrder};
367    use tokio::io::AsyncWriteExt;
368
369    use super::{Connection, KeyStore, Lock, MAX_AGENT_FRAME_LEN};
370    use crate::keys::Error;
371
372    #[test]
373    fn oversized_agent_request_is_rejected_before_allocation() -> std::io::Result<()> {
374        let runtime = tokio::runtime::Builder::new_current_thread()
375            .enable_all()
376            .build()?;
377
378        runtime.block_on(async {
379            let (server, mut client) = tokio::io::duplex(64);
380            let connection = Connection {
381                lock: Lock(std::sync::Arc::new(std::sync::RwLock::new(crate::CryptoVec::new()))),
382                keys: KeyStore(std::sync::Arc::new(std::sync::RwLock::new(
383                    std::collections::HashMap::new(),
384                ))),
385                agent: Some(()),
386                s: server,
387                buf: Vec::new(),
388            };
389            let server = tokio::spawn(async move { connection.run().await });
390
391            let mut frame = [0u8; 4];
392            BigEndian::write_u32(&mut frame, (MAX_AGENT_FRAME_LEN + 1) as u32);
393            client.write_all(&frame).await?;
394            drop(client);
395
396            let err = server.await.expect("server task").unwrap_err();
397            assert!(matches!(err, Error::AgentProtocolError));
398            Ok::<(), std::io::Error>(())
399        })?;
400
401        Ok(())
402    }
403}