1use std::time::{Duration, SystemTime, UNIX_EPOCH};
7
8use base64::engine::general_purpose::URL_SAFE_NO_PAD;
9use base64::Engine;
10use recall_wire::signature::{self, normalize_authority, SigningKey, Target};
11use recall_wire::ErrorResponse;
12use reqwest::StatusCode;
13use serde::de::DeserializeOwned;
14use serde::Serialize;
15
16#[derive(Debug, thiserror::Error)]
18pub enum ApiError {
19 #[error("{0}")]
21 Transport(String),
22 #[error("{status}: {message}")]
24 Status {
25 status: StatusCode,
27 message: String,
29 retry_after: Option<u64>,
31 },
32 #[error("unexpected answer: {0}")]
34 Body(String),
35}
36
37impl ApiError {
38 pub fn status(&self) -> Option<StatusCode> {
40 match self {
41 ApiError::Status { status, .. } => Some(*status),
42 _ => None,
43 }
44 }
45
46 pub fn message(&self) -> &str {
48 match self {
49 ApiError::Status { message, .. } => message,
50 _ => "",
51 }
52 }
53}
54
55#[derive(Debug, Clone)]
57pub struct Api {
58 http: reqwest::Client,
59 base: String,
62 authority: String,
64 prefix: String,
66}
67
68impl Api {
69 pub fn new(url: &str) -> Result<Self, ApiError> {
71 let parsed = reqwest::Url::parse(url).map_err(|e| ApiError::Transport(e.to_string()))?;
72 let host = parsed
73 .host_str()
74 .ok_or_else(|| ApiError::Transport(format!("{url} has no host")))?;
75 let authority = normalize_authority(&match parsed.port() {
79 Some(port) => format!("{host}:{port}"),
80 None => host.to_string(),
81 });
82 let http = reqwest::Client::builder()
83 .user_agent(crate::user_agent())
84 .connect_timeout(Duration::from_secs(10))
85 .redirect(reqwest::redirect::Policy::none())
90 .build()
91 .map_err(|e| ApiError::Transport(e.to_string()))?;
92 Ok(Self {
93 http,
94 base: url.trim_end_matches('/').to_string(),
95 authority,
96 prefix: parsed.path().trim_end_matches('/').to_string(),
97 })
98 }
99
100 pub async fn get<T: DeserializeOwned>(
102 &self,
103 path: &str,
104 timeout: Duration,
105 ) -> Result<T, ApiError> {
106 let req = self
107 .http
108 .get(format!("{}{path}", self.base))
109 .timeout(timeout)
110 .header(
111 recall_wire::PROTOCOL_HEADER,
112 recall_wire::PROTOCOL.to_string(),
113 );
114 answer(req).await
115 }
116
117 pub async fn post<B: Serialize, T: DeserializeOwned>(
119 &self,
120 path: &str,
121 body: &B,
122 timeout: Duration,
123 ) -> Result<T, ApiError> {
124 self.send(path, body, None, timeout).await
125 }
126
127 pub async fn post_signed<B: Serialize, T: DeserializeOwned>(
129 &self,
130 path: &str,
131 body: &B,
132 key: &SigningKey,
133 keyid: &str,
134 timeout: Duration,
135 ) -> Result<T, ApiError> {
136 self.send(path, body, Some((key, keyid)), timeout).await
137 }
138
139 async fn send<B: Serialize, T: DeserializeOwned>(
140 &self,
141 path: &str,
142 body: &B,
143 signer: Option<(&SigningKey, &str)>,
144 timeout: Duration,
145 ) -> Result<T, ApiError> {
146 let bytes = serde_json::to_vec(body).map_err(|e| ApiError::Body(e.to_string()))?;
147 let mut req = self
148 .http
149 .post(format!("{}{path}", self.base))
150 .timeout(timeout)
151 .header("content-type", "application/json")
152 .header(
153 recall_wire::PROTOCOL_HEADER,
154 recall_wire::PROTOCOL.to_string(),
155 );
156 if let Some((key, keyid)) = signer {
157 let full_path = format!("{}{path}", self.prefix);
158 let signed = signature::sign_request(
159 key,
160 keyid,
161 &Target {
162 method: "POST",
163 authority: &self.authority,
164 path: &full_path,
165 query: None,
166 },
167 &recall_wire::PROTOCOL.to_string(),
168 &bytes,
169 unix_now(),
170 &nonce()?,
171 )
172 .map_err(|e| ApiError::Body(e.to_string()))?;
173 req = req
174 .header(signature::CONTENT_DIGEST_HEADER, signed.content_digest)
175 .header(signature::SIGNATURE_INPUT_HEADER, signed.signature_input)
176 .header(signature::SIGNATURE_HEADER, signed.signature);
177 }
178 answer(req.body(bytes)).await
179 }
180}
181
182async fn answer<T: DeserializeOwned>(req: reqwest::RequestBuilder) -> Result<T, ApiError> {
184 let resp = req
185 .send()
186 .await
187 .map_err(|e| ApiError::Transport(describe(&e)))?;
188 let status = resp.status();
189 let retry_after = resp
190 .headers()
191 .get("retry-after")
192 .and_then(|v| v.to_str().ok())
193 .and_then(|v| v.trim().parse().ok());
194 let text = resp
195 .bytes()
196 .await
197 .map_err(|e| ApiError::Transport(describe(&e)))?;
198 if !status.is_success() {
199 let message = serde_json::from_slice::<ErrorResponse>(&text)
200 .map(|e| e.error)
201 .unwrap_or_else(|_| String::from_utf8_lossy(&text).trim().to_string());
202 return Err(ApiError::Status {
203 status,
204 message,
205 retry_after,
206 });
207 }
208 serde_json::from_slice(&text).map_err(|e| ApiError::Body(e.to_string()))
209}
210
211fn describe(e: &reqwest::Error) -> String {
213 let mut out = e.to_string();
214 let mut source = std::error::Error::source(e);
215 while let Some(s) = source {
216 out.push_str(": ");
217 out.push_str(&s.to_string());
218 source = s.source();
219 }
220 out
221}
222
223fn unix_now() -> i64 {
224 SystemTime::now()
225 .duration_since(UNIX_EPOCH)
226 .map(|d| d.as_secs() as i64)
227 .unwrap_or(0)
228}
229
230fn nonce() -> Result<String, ApiError> {
232 let mut bytes = [0u8; 16];
233 getrandom::fill(&mut bytes).map_err(|e| ApiError::Transport(format!("no randomness: {e}")))?;
234 Ok(URL_SAFE_NO_PAD.encode(bytes))
235}
236
237#[cfg(test)]
238mod tests {
239 use super::*;
240
241 #[test]
242 fn the_authority_is_what_the_server_will_read() {
243 let api = Api::new("http://recall-server:8787").unwrap();
244 assert_eq!(
245 (api.authority.as_str(), api.prefix.as_str()),
246 ("recall-server:8787", "")
247 );
248 let api = Api::new("https://Recall.Example.com/").unwrap();
249 assert_eq!(api.authority, "recall.example.com");
250 assert_eq!(api.base, "https://Recall.Example.com");
251 let api = Api::new("https://example.com:8443/recall/").unwrap();
252 assert_eq!(
253 (api.authority.as_str(), api.prefix.as_str()),
254 ("example.com:8443", "/recall")
255 );
256 }
257
258 #[test]
259 fn nonces_differ() {
260 assert_ne!(nonce().unwrap(), nonce().unwrap());
261 }
262}