1use std::io::{ErrorKind, Read, Write};
6use std::net::{TcpStream, ToSocketAddrs};
7use std::sync::Arc;
8use std::time::{Duration, Instant};
9
10use serde_json::{Value, json};
11
12use super::wire::Assertion;
13use crate::error::{Error, Result};
14use crate::org::OrgId;
15use crate::server::ssh::{self, SshRequest};
16use crate::server::terminal::{Pty, PtyOutput, TermRequest};
17
18const CONNECT_TIMEOUT: Duration = Duration::from_secs(10);
19const MAX_RESPONSE: usize = 64 * 1024 * 1024;
20
21type Tls = rustls::StreamOwned<rustls::ClientConnection, TcpStream>;
22
23#[derive(Clone)]
25pub struct AgentClient {
26 pub name: String,
27 pub address: String,
28 pub port: u16,
29 tls: Arc<rustls::ClientConfig>,
30}
31
32impl std::fmt::Debug for AgentClient {
33 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
34 write!(
35 f,
36 "AgentClient({} at {}:{})",
37 self.name, self.address, self.port
38 )
39 }
40}
41
42#[derive(Debug, Clone)]
44pub struct Answer {
45 pub status: u16,
46 pub headers: Vec<(String, String)>,
47 pub body: Vec<u8>,
48}
49
50fn unreachable(name: &str, step: &str, e: impl std::fmt::Display) -> Error {
51 Error::OperationFailed {
52 step: format!("reach server {name} ({step})"),
53 message: e.to_string(),
54 }
55}
56
57impl AgentClient {
58 pub fn new(name: &str, address: &str, port: u16, tls: Arc<rustls::ClientConfig>) -> Self {
59 AgentClient {
60 name: name.into(),
61 address: address.into(),
62 port,
63 tls,
64 }
65 }
66
67 fn server_name(&self) -> Result<rustls::pki_types::ServerName<'static>> {
68 rustls::pki_types::ServerName::try_from(self.address.clone())
69 .map_err(|e| Error::invalid(format!("server address {}: {e}", self.address)))
70 }
71
72 fn authority(&self) -> String {
73 if self.address.contains(':') {
74 format!("[{}]:{}", self.address, self.port)
75 } else {
76 format!("{}:{}", self.address, self.port)
77 }
78 }
79
80 fn connect(&self, timeout: Duration) -> Result<Tls> {
83 let addr = (self.address.as_str(), self.port)
84 .to_socket_addrs()
85 .map_err(|e| unreachable(&self.name, "resolve", e))?
86 .next()
87 .ok_or_else(|| unreachable(&self.name, "resolve", "no address"))?;
88 let sock = TcpStream::connect_timeout(&addr, CONNECT_TIMEOUT.min(timeout))
89 .map_err(|e| unreachable(&self.name, "connect", e))?;
90 sock.set_read_timeout(Some(CONNECT_TIMEOUT.min(timeout)))?;
91 sock.set_write_timeout(Some(CONNECT_TIMEOUT.min(timeout)))?;
92 let _ = sock.set_nodelay(true);
93 let conn = rustls::ClientConnection::new(self.tls.clone(), self.server_name()?)
94 .map_err(|e| unreachable(&self.name, "TLS", e))?;
95 let mut s = rustls::StreamOwned::new(conn, sock);
96 while s.conn.is_handshaking() {
97 s.conn
98 .complete_io(&mut s.sock)
99 .map_err(|e| unreachable(&self.name, "TLS handshake", e))?;
100 }
101 Ok(s)
102 }
103
104 pub fn peer_fingerprint(&self) -> Result<String> {
106 let s = self.connect(CONNECT_TIMEOUT)?;
107 let c = s
108 .conn
109 .peer_certificates()
110 .and_then(|c| c.first())
111 .ok_or_else(|| unreachable(&self.name, "TLS", "no server certificate"))?;
112 Ok(super::pki::hex(
113 ring::digest::digest(&ring::digest::SHA256, c.as_ref()).as_ref(),
114 ))
115 }
116
117 pub fn request(
119 &self,
120 method: &str,
121 path: &str,
122 headers: &[(String, String)],
123 body: &[u8],
124 timeout: Duration,
125 ) -> Result<Answer> {
126 let started = Instant::now();
127 let mut s = self.connect(timeout)?;
128 let mut head = format!(
129 "{method} {path} HTTP/1.1\r\nHost: {}\r\nUser-Agent: isb/{}\r\nContent-Length: {}\r\nConnection: close\r\n",
130 self.authority(),
131 env!("CARGO_PKG_VERSION"),
132 body.len()
133 );
134 for (k, v) in headers {
135 if k.contains(['\r', '\n', ':']) || v.contains(['\r', '\n']) {
136 continue;
137 }
138 head.push_str(&format!("{k}: {v}\r\n"));
139 }
140 head.push_str("\r\n");
141 let io = |e: std::io::Error| unreachable(&self.name, "send", e);
142 s.write_all(head.as_bytes()).map_err(io)?;
143 s.write_all(body).map_err(io)?;
144 s.flush().map_err(io)?;
145 let mut buf = Vec::with_capacity(8192);
146 let mut chunk = [0u8; 16384];
147 loop {
148 let left = timeout.saturating_sub(started.elapsed());
149 if left.is_zero() {
150 return Err(Error::Io(std::io::Error::new(
151 ErrorKind::TimedOut,
152 format!(
153 "server {} did not answer {method} {path} within {timeout:?}",
154 self.name
155 ),
156 )));
157 }
158 s.sock.set_read_timeout(Some(left))?;
159 match s.read(&mut chunk) {
160 Ok(0) => break,
161 Ok(n) => {
162 buf.extend_from_slice(&chunk[..n]);
163 if buf.len() > MAX_RESPONSE {
164 return Err(Error::Protocol(format!(
165 "server {}: response over {MAX_RESPONSE} bytes",
166 self.name
167 )));
168 }
169 }
170 Err(e) if e.kind() == ErrorKind::Interrupted => {}
171 Err(e) if e.kind() == ErrorKind::UnexpectedEof => break,
173 Err(e) if matches!(e.kind(), ErrorKind::WouldBlock | ErrorKind::TimedOut) => {}
174 Err(e) => return Err(unreachable(&self.name, "read", e)),
175 }
176 }
177 parse_answer(&buf).ok_or_else(|| {
178 Error::Protocol(format!(
179 "server {}: {method} {path}: an unreadable response",
180 self.name
181 ))
182 })
183 }
184
185 pub fn call(
188 &self,
189 tool: &str,
190 args: &Value,
191 who: &Assertion,
192 org: Option<&OrgId>,
193 request_id: Option<&str>,
194 timeout: Duration,
195 ) -> Result<Value> {
196 let path = match org {
197 Some(o) => format!("/orgs/{o}/api/v1/tools/{tool}"),
198 None => format!("/api/v1/tools/{tool}"),
199 };
200 let mut h = vec![
201 ("Authorization".to_string(), who.header()),
202 ("Content-Type".to_string(), "application/json".to_string()),
203 ];
204 if let Some(r) = request_id {
205 h.push(("X-Request-Id".into(), r.into()));
206 }
207 let a = self.request("POST", &path, &h, &serde_json::to_vec(args)?, timeout)?;
208 tool_answer(&self.name, &a)
209 }
210
211 pub fn internal(
213 &self,
214 method: &str,
215 path: &str,
216 body: Option<&Value>,
217 timeout: Duration,
218 ) -> Result<Value> {
219 let h = vec![
220 (
221 "Authorization".to_string(),
222 Assertion::control_plane().header(),
223 ),
224 ("Content-Type".to_string(), "application/json".to_string()),
225 ];
226 let body = match body {
227 Some(b) => serde_json::to_vec(b)?,
228 None => Vec::new(),
229 };
230 let a = self.request(method, path, &h, &body, timeout)?;
231 internal_answer(&self.name, path, &a)
232 }
233
234 pub fn terminal(&self, who: &Assertion, org: &OrgId, t: &TermRequest) -> Result<Box<dyn Pty>> {
236 let path = format!("/orgs/{org}/api/v1/terminal?{}", t.query());
237 let target = format!("{}:{}", self.name, t.target());
238 Ok(Box::new(self.websocket(&path, who, &[], target)?))
239 }
240
241 pub fn ssh(
244 &self,
245 who: &Assertion,
246 org: &OrgId,
247 s: &SshRequest,
248 keys: &[String],
249 ) -> Result<Box<dyn Pty>> {
250 let req = SshRequest {
251 instance: s.instance.clone(),
252 keys_of: None,
253 forwarded_keys: None,
254 };
255 let path = format!("/orgs/{org}/api/v1/ssh?{}", req.query());
256 let header = (ssh::KEYS_HEADER, ssh::keys_header(keys));
257 let target = format!("{}:{}", self.name, s.instance);
258 Ok(Box::new(self.websocket(&path, who, &[header], target)?))
259 }
260
261 fn websocket(
263 &self,
264 path: &str,
265 who: &Assertion,
266 headers: &[(&str, String)],
267 target: String,
268 ) -> Result<RemotePty> {
269 use tungstenite::client::IntoClientRequest;
270 let s = self.connect(CONNECT_TIMEOUT)?;
271 let url = format!("wss://{}{path}", self.authority());
272 let mut req = url
273 .into_client_request()
274 .map_err(|e| Error::WebSocket(e.to_string()))?;
275 let mut put = |k: &str, v: &str| -> Result<()> {
276 let name = tungstenite::http::HeaderName::from_bytes(k.as_bytes())
277 .map_err(|_| Error::invalid(format!("header {k}")))?;
278 let value = v
279 .parse()
280 .map_err(|_| Error::invalid(format!("header {k}")))?;
281 req.headers_mut().insert(name, value);
282 Ok(())
283 };
284 put("authorization", &who.header())?;
285 for (k, v) in headers {
286 put(k, v)?;
287 }
288 let (ws, _) = tungstenite::client(req, s).map_err(|e| match e {
289 tungstenite::HandshakeError::Failure(tungstenite::Error::Http(r)) => {
290 let body = r
291 .body()
292 .as_ref()
293 .map(|b| String::from_utf8_lossy(b).into_owned())
294 .unwrap_or_default();
295 let status = r.status().as_u16();
296 let v: Value = serde_json::from_str(&body).unwrap_or(Value::Null);
297 let message = match v["message"].as_str().unwrap_or(&body) {
298 "" => format!("server {} refused it (HTTP {status})", self.name),
300 m => m.to_string(),
301 };
302 match v["error"].as_str() {
303 Some("forbidden") => Error::Forbidden(message),
304 None if matches!(status, 401 | 403) => Error::Forbidden(message),
305 code => Error::Remote {
306 code: code.unwrap_or("server_error").into(),
307 message,
308 data: Value::Null,
309 },
310 }
311 }
312 e => Error::WebSocket(format!("server {}: {e}", self.name)),
313 })?;
314 Ok(RemotePty {
315 ws,
316 done: false,
317 target,
318 })
319 }
320
321 pub fn internal_bytes(&self, path: &str, body: &[u8], timeout: Duration) -> Result<Value> {
323 let h = vec![
324 (
325 "Authorization".to_string(),
326 Assertion::control_plane().header(),
327 ),
328 (
329 "Content-Type".to_string(),
330 "application/octet-stream".to_string(),
331 ),
332 ];
333 let a = self.request("POST", path, &h, body, timeout)?;
334 internal_answer(&self.name, path, &a)
335 }
336}
337
338pub fn parse_answer(buf: &[u8]) -> Option<Answer> {
340 let mut hs = [httparse::EMPTY_HEADER; 64];
341 let mut r = httparse::Response::new(&mut hs);
342 let n = match r.parse(buf).ok()? {
343 httparse::Status::Complete(n) => n,
344 httparse::Status::Partial => return None,
345 };
346 let headers: Vec<(String, String)> = r
347 .headers
348 .iter()
349 .map(|h| {
350 (
351 h.name.to_string(),
352 String::from_utf8_lossy(h.value).trim().to_string(),
353 )
354 })
355 .collect();
356 let mut body = buf[n..].to_vec();
357 if let Some(l) = headers
358 .iter()
359 .find(|(k, _)| k.eq_ignore_ascii_case("content-length"))
360 .and_then(|(_, v)| v.parse::<usize>().ok())
361 {
362 if body.len() < l {
363 return None;
364 }
365 body.truncate(l);
366 }
367 Some(Answer {
368 status: r.code?,
369 headers,
370 body,
371 })
372}
373
374fn internal_answer(name: &str, path: &str, a: &Answer) -> Result<Value> {
376 let v: Value = serde_json::from_slice(&a.body).unwrap_or(Value::Null);
377 if a.status != 200 {
378 return Err(Error::Remote {
379 code: v["error"].as_str().unwrap_or("server_error").into(),
380 message: format!(
381 "server {name}: {path}: HTTP {}: {}",
382 a.status,
383 v["message"]
384 .as_str()
385 .unwrap_or_else(|| std::str::from_utf8(&a.body).unwrap_or(""))
386 ),
387 data: Value::Null,
388 });
389 }
390 Ok(v)
391}
392
393pub fn tool_answer(server: &str, a: &Answer) -> Result<Value> {
395 let v: Value = serde_json::from_slice(&a.body).map_err(|e| {
396 Error::Protocol(format!(
397 "server {server}: HTTP {}, undecodable body ({e}): {}",
398 a.status,
399 String::from_utf8_lossy(&a.body[..a.body.len().min(200)])
400 ))
401 })?;
402 if a.status == 200 {
403 return Ok(v.get("result").cloned().unwrap_or(Value::Null));
404 }
405 let code = v["error"].as_str().unwrap_or("server_error").to_string();
406 let message = v["message"]
407 .as_str()
408 .map(String::from)
409 .unwrap_or_else(|| format!("HTTP {}", a.status));
410 let message = match code.as_str() {
412 "forbidden" => message
413 .strip_prefix("forbidden: ")
414 .unwrap_or(&message)
415 .to_string(),
416 _ => message,
417 };
418 if code == "forbidden" {
419 return Err(Error::Forbidden(message));
420 }
421 Err(Error::Remote {
422 code,
423 message,
424 data: v.get("data").cloned().unwrap_or(Value::Null),
425 })
426}
427
428struct RemotePty {
430 ws: tungstenite::WebSocket<Tls>,
431 done: bool,
432 target: String,
433}
434
435fn would_block(e: &tungstenite::Error) -> bool {
436 matches!(e, tungstenite::Error::Io(i) if matches!(i.kind(), ErrorKind::WouldBlock | ErrorKind::TimedOut))
437}
438
439impl RemotePty {
440 fn blocking(&mut self) {
441 let s = &self.ws.get_ref().sock;
442 let _ = s.set_nonblocking(false);
443 let _ = s.set_read_timeout(Some(Duration::from_secs(10)));
444 }
445
446 fn send(&mut self, m: tungstenite::Message) -> Result<()> {
447 self.blocking();
448 self.ws.send(m).map_err(|e| Error::WebSocket(e.to_string()))
449 }
450}
451
452impl Pty for RemotePty {
453 fn input(&mut self, data: &[u8]) -> Result<()> {
454 self.send(tungstenite::Message::binary(data.to_vec()))
455 }
456
457 fn resize(&mut self, cols: u16, rows: u16) {
458 let _ = self.send(tungstenite::Message::text(
459 json!({"type": "resize", "cols": cols, "rows": rows}).to_string(),
460 ));
461 }
462
463 fn output(&mut self, wait: Duration) -> PtyOutput {
464 if self.done {
465 return PtyOutput::Exit(None);
466 }
467 {
468 let s = &self.ws.get_ref().sock;
469 if wait.is_zero() {
470 let _ = s.set_nonblocking(true);
471 } else {
472 let _ = s.set_nonblocking(false);
473 let _ = s.set_read_timeout(Some(wait));
474 }
475 }
476 match self.ws.read() {
477 Ok(tungstenite::Message::Binary(b)) => PtyOutput::Data(b.to_vec()),
478 Ok(tungstenite::Message::Text(t)) => {
479 let v: Value = serde_json::from_str(t.as_str()).unwrap_or(Value::Null);
480 match v["type"].as_str() {
481 Some("exit") => {
482 self.done = true;
483 PtyOutput::Exit(v["code"].as_i64().map(|c| c as i32))
484 }
485 Some("error") => {
486 self.done = true;
487 PtyOutput::Failed(v["message"].as_str().unwrap_or("error").to_string())
488 }
489 Some("accepted") => PtyOutput::Note(v),
492 _ => PtyOutput::Idle,
493 }
494 }
495 Ok(tungstenite::Message::Close(_)) => {
496 self.done = true;
497 PtyOutput::Exit(None)
498 }
499 Ok(_) => PtyOutput::Idle,
500 Err(e) if would_block(&e) => PtyOutput::Idle,
501 Err(e) => {
502 self.done = true;
503 PtyOutput::Failed(format!("the server's terminal broke: {e}"))
504 }
505 }
506 }
507
508 fn close(&mut self) {
509 if !self.done {
510 self.done = true;
511 self.blocking();
512 let _ = self.ws.close(None);
513 let _ = self.ws.flush();
514 }
515 }
516
517 fn target(&self) -> Option<String> {
518 Some(self.target.clone())
519 }
520}
521
522#[cfg(test)]
523mod tests {
524 use super::*;
525
526 #[test]
527 fn answers_parse_and_map_errors() {
528 let a = parse_answer(b"HTTP/1.1 200 OK\r\nContent-Length: 12\r\n\r\n{\"result\":1}xyz")
529 .unwrap();
530 assert_eq!(a.status, 200);
531 assert_eq!(tool_answer("s", &a).unwrap(), json!(1));
532 let a = parse_answer(
533 b"HTTP/1.1 403 Forbidden\r\n\r\n{\"error\":\"forbidden\",\"message\":\"forbidden: no access to org b\"}",
534 )
535 .unwrap();
536 match tool_answer("s", &a).unwrap_err() {
537 Error::Forbidden(m) => assert_eq!(m, "no access to org b"),
538 e => panic!("{e:?}"),
539 }
540 let a = parse_answer(b"HTTP/1.1 404 Not Found\r\n\r\n{\"error\":\"not_found\",\"message\":\"app x not found\"}").unwrap();
541 assert!(tool_answer("s", &a).unwrap_err().is_not_found());
542 assert!(parse_answer(b"HTTP/1.1 200 OK\r\nContent-Length: 99\r\n\r\n{}").is_none());
543 }
544}