1use async_trait::async_trait;
11use base64::Engine as _;
12use faucet_core::{AuthProvider, Credential, FaucetError};
13use hmac::{Hmac, Mac};
14use percent_encoding::{AsciiSet, NON_ALPHANUMERIC, utf8_percent_encode};
15use sha2::Sha256;
16use std::collections::BTreeMap;
17use std::sync::atomic::{AtomicU64, Ordering};
18
19type HmacSha256 = Hmac<Sha256>;
20
21const OAUTH_ENCODE: &AsciiSet = &NON_ALPHANUMERIC
25 .remove(b'-')
26 .remove(b'.')
27 .remove(b'_')
28 .remove(b'~');
29
30fn enc(s: &str) -> String {
31 utf8_percent_encode(s, OAUTH_ENCODE).to_string()
32}
33
34pub struct OAuth1Provider {
36 consumer_key: String,
37 consumer_secret: String,
38 token: String,
39 token_secret: String,
40 realm: Option<String>,
41 nonce_counter: AtomicU64,
44}
45
46impl std::fmt::Debug for OAuth1Provider {
49 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
50 f.debug_struct("OAuth1Provider")
51 .field("consumer_key", &self.consumer_key)
52 .field("consumer_secret", &"***")
53 .field("token", &self.token)
54 .field("token_secret", &"***")
55 .field("realm", &self.realm)
56 .finish()
57 }
58}
59
60impl OAuth1Provider {
61 pub fn from_config(config: &serde_json::Value) -> Result<Self, FaucetError> {
65 let req = |k: &str| -> Result<String, FaucetError> {
66 config
67 .get(k)
68 .and_then(|v| v.as_str())
69 .filter(|s| !s.is_empty())
70 .map(str::to_string)
71 .ok_or_else(|| FaucetError::Config(format!("oauth1 auth provider: missing `{k}`")))
72 };
73 if let Some(m) = config.get("signature_method").and_then(|v| v.as_str())
74 && !m.eq_ignore_ascii_case("HMAC-SHA256")
75 {
76 return Err(FaucetError::Config(format!(
77 "oauth1: unsupported signature_method {m:?} (only HMAC-SHA256 is supported)"
78 )));
79 }
80 Ok(Self {
81 consumer_key: req("consumer_key")?,
82 consumer_secret: req("consumer_secret")?,
83 token: req("token")?,
84 token_secret: req("token_secret")?,
85 realm: config
86 .get("realm")
87 .and_then(|v| v.as_str())
88 .map(str::to_string),
89 nonce_counter: AtomicU64::new(0),
90 })
91 }
92}
93
94#[async_trait]
95impl AuthProvider for OAuth1Provider {
96 async fn credential(&self) -> Result<Credential, FaucetError> {
97 Err(FaucetError::Auth(
99 "oauth1 provider signs each request — it has no standalone credential; the connector \
100 must call sign_request()"
101 .into(),
102 ))
103 }
104
105 async fn sign_request(
106 &self,
107 method: &str,
108 url: &str,
109 query: &BTreeMap<String, String>,
110 ) -> Result<Option<Credential>, FaucetError> {
111 let timestamp = std::time::SystemTime::now()
112 .duration_since(std::time::UNIX_EPOCH)
113 .map(|d| d.as_secs())
114 .unwrap_or(0);
115 let counter = self.nonce_counter.fetch_add(1, Ordering::Relaxed);
116 let nonce = format!("{timestamp}{counter}");
117 let header = self.authorization_header(method, url, query, &nonce, timestamp);
118 Ok(Some(Credential::Header {
119 name: "Authorization".to_string(),
120 value: header,
121 }))
122 }
123
124 fn provider_name(&self) -> &'static str {
125 "oauth1"
126 }
127}
128
129impl OAuth1Provider {
130 fn authorization_header(
137 &self,
138 method: &str,
139 url: &str,
140 query: &BTreeMap<String, String>,
141 nonce: &str,
142 timestamp: u64,
143 ) -> String {
144 let ts = timestamp.to_string();
145 let oauth_params: [(&str, &str); 6] = [
147 ("oauth_consumer_key", &self.consumer_key),
148 ("oauth_nonce", nonce),
149 ("oauth_signature_method", "HMAC-SHA256"),
150 ("oauth_timestamp", &ts),
151 ("oauth_token", &self.token),
152 ("oauth_version", "1.0"),
153 ];
154
155 let mut encoded: Vec<(String, String)> = oauth_params
159 .iter()
160 .map(|(k, v)| (enc(k), enc(v)))
161 .chain(query.iter().map(|(k, v)| (enc(k), enc(v))))
162 .collect();
163 encoded.sort();
164 let param_string = encoded
165 .iter()
166 .map(|(k, v)| format!("{k}={v}"))
167 .collect::<Vec<_>>()
168 .join("&");
169
170 let base_url = url.split('?').next().unwrap_or(url);
171 let base_string = format!(
172 "{}&{}&{}",
173 method.to_uppercase(),
174 enc(base_url),
175 enc(¶m_string)
176 );
177
178 let signing_key = format!("{}&{}", enc(&self.consumer_secret), enc(&self.token_secret));
179 let mut mac = HmacSha256::new_from_slice(signing_key.as_bytes())
180 .expect("HMAC accepts any key length");
181 mac.update(base_string.as_bytes());
182 let signature =
183 base64::engine::general_purpose::STANDARD.encode(mac.finalize().into_bytes());
184
185 let mut parts: Vec<String> = Vec::new();
188 if let Some(realm) = &self.realm {
189 parts.push(format!("realm=\"{}\"", enc(realm)));
190 }
191 for (k, v) in oauth_params {
192 parts.push(format!("{}=\"{}\"", enc(k), enc(v)));
193 }
194 parts.push(format!("oauth_signature=\"{}\"", enc(&signature)));
195 format!("OAuth {}", parts.join(", "))
196 }
197}
198
199#[cfg(test)]
200mod tests {
201 use super::*;
202
203 fn provider() -> OAuth1Provider {
204 OAuth1Provider::from_config(&serde_json::json!({
205 "consumer_key": "ck",
206 "consumer_secret": "cs",
207 "token": "tk",
208 "token_secret": "ts",
209 "realm": "ACCT123",
210 }))
211 .unwrap()
212 }
213
214 #[test]
215 fn header_is_stable_for_fixed_nonce_and_timestamp() {
216 let p = provider();
217 let mut q = BTreeMap::new();
218 q.insert("limit".to_string(), "10".to_string());
219 let h1 = p.authorization_header(
220 "GET",
221 "https://api.example.com/records",
222 &q,
223 "nonce1",
224 1700000000,
225 );
226 let h2 = p.authorization_header(
227 "GET",
228 "https://api.example.com/records",
229 &q,
230 "nonce1",
231 1700000000,
232 );
233 assert_eq!(h1, h2, "same inputs → same signature");
234 assert!(h1.starts_with("OAuth "));
235 assert!(h1.contains("realm=\"ACCT123\""));
236 assert!(h1.contains("oauth_signature_method=\"HMAC-SHA256\""));
237 assert!(h1.contains("oauth_consumer_key=\"ck\""));
238 assert!(h1.contains("oauth_signature=\""));
239 }
240
241 #[test]
242 fn signature_changes_with_method_and_query() {
243 let p = provider();
244 let q = BTreeMap::new();
245 let get = p.authorization_header("GET", "https://api.example.com/x", &q, "n", 1);
246 let post = p.authorization_header("POST", "https://api.example.com/x", &q, "n", 1);
247 assert_ne!(get, post, "method is part of the base string");
248 let mut q2 = BTreeMap::new();
249 q2.insert("a".to_string(), "b".to_string());
250 let with_q = p.authorization_header("GET", "https://api.example.com/x", &q2, "n", 1);
251 assert_ne!(get, with_q, "query params are part of the base string");
252 }
253
254 #[test]
255 fn query_string_on_url_is_stripped_from_base_url() {
256 let p = provider();
259 let mut q = BTreeMap::new();
260 q.insert("a".to_string(), "b".to_string());
261 let with_qs = p.authorization_header("GET", "https://api.example.com/x?a=b", &q, "n", 1);
262 let clean = p.authorization_header("GET", "https://api.example.com/x", &q, "n", 1);
263 assert_eq!(with_qs, clean);
264 }
265
266 #[tokio::test]
267 async fn sign_request_returns_authorization_header() {
268 let p = provider();
269 let cred = p
270 .sign_request("GET", "https://api.example.com/x", &BTreeMap::new())
271 .await
272 .unwrap()
273 .expect("oauth1 signs the request");
274 match cred {
275 Credential::Header { name, value } => {
276 assert_eq!(name, "Authorization");
277 assert!(value.starts_with("OAuth "));
278 }
279 other => panic!("expected a Header credential, got {other:?}"),
280 }
281 }
282
283 #[tokio::test]
284 async fn credential_errors_directing_to_sign_request() {
285 assert!(provider().credential().await.is_err());
286 }
287
288 #[test]
289 fn rejects_missing_fields_and_bad_signature_method() {
290 assert!(OAuth1Provider::from_config(&serde_json::json!({"consumer_key": "x"})).is_err());
291 assert!(
292 OAuth1Provider::from_config(&serde_json::json!({
293 "consumer_key": "ck", "consumer_secret": "cs", "token": "tk", "token_secret": "ts",
294 "signature_method": "HMAC-SHA1"
295 }))
296 .is_err()
297 );
298 }
299
300 #[test]
301 fn debug_does_not_leak_secrets() {
302 let s = format!("{:?}", provider());
303 assert!(!s.contains("cs") || !s.contains("\"cs\""));
304 assert!(s.contains("***"));
305 assert!(!s.contains("token_secret\": \"ts"));
306 }
307}