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 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 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 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 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 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 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 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 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 if let Ok(true) = self.add_key(&mut r, true, writebuf).await {
226 } else {
227 writebuf.push(msg::FAILURE)
228 }
229 }
230 _ => {
231 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}