1use base64::Engine;
4use base64::display::Base64Display;
5use base64::engine::general_purpose::STANDARD;
6use hmac::{Hmac, KeyInit, Mac};
7use rand::{self, RngExt};
8use sha2::digest::FixedOutput;
9use sha2::{Digest, Sha256};
10use std::fmt::Write;
11use std::io;
12use std::iter;
13use std::mem;
14use std::str;
15
16const NONCE_LENGTH: usize = 24;
17
18const MAX_ITERATION_COUNT: u32 = 2_000_000;
26
27pub const SCRAM_SHA_256: &str = "SCRAM-SHA-256";
29pub const SCRAM_SHA_256_PLUS: &str = "SCRAM-SHA-256-PLUS";
31
32fn normalize(pass: &[u8]) -> Vec<u8> {
36 let pass = match str::from_utf8(pass) {
37 Ok(pass) => pass,
38 Err(_) => return pass.to_vec(),
39 };
40
41 match stringprep::saslprep(pass) {
42 Ok(pass) => pass.into_owned().into_bytes(),
43 Err(_) => pass.as_bytes().to_vec(),
44 }
45}
46
47pub(crate) fn hi(str: &[u8], salt: &[u8], i: u32) -> [u8; 32] {
48 let mut hmac =
49 Hmac::<Sha256>::new_from_slice(str).expect("HMAC is able to accept all key sizes");
50 hmac.update(salt);
51 hmac.update(&[0, 0, 0, 1]);
52 let mut prev = hmac.finalize().into_bytes();
53
54 let mut hi = prev;
55
56 for _ in 1..i {
57 let mut hmac = Hmac::<Sha256>::new_from_slice(str).expect("already checked above");
58 hmac.update(&prev);
59 prev = hmac.finalize().into_bytes();
60
61 for (hi, prev) in hi.iter_mut().zip(prev) {
62 *hi ^= prev;
63 }
64 }
65
66 hi.into()
67}
68
69enum ChannelBindingInner {
70 Unrequested,
71 Unsupported,
72 TlsServerEndPoint(Vec<u8>),
73}
74
75pub struct ChannelBinding(ChannelBindingInner);
77
78impl ChannelBinding {
79 pub fn unrequested() -> ChannelBinding {
81 ChannelBinding(ChannelBindingInner::Unrequested)
82 }
83
84 pub fn unsupported() -> ChannelBinding {
86 ChannelBinding(ChannelBindingInner::Unsupported)
87 }
88
89 pub fn tls_server_end_point(signature: Vec<u8>) -> ChannelBinding {
92 ChannelBinding(ChannelBindingInner::TlsServerEndPoint(signature))
93 }
94
95 fn gs2_header(&self) -> &'static str {
96 match self.0 {
97 ChannelBindingInner::Unrequested => "y,,",
98 ChannelBindingInner::Unsupported => "n,,",
99 ChannelBindingInner::TlsServerEndPoint(_) => "p=tls-server-end-point,,",
100 }
101 }
102
103 fn cbind_data(&self) -> &[u8] {
104 match self.0 {
105 ChannelBindingInner::Unrequested | ChannelBindingInner::Unsupported => &[],
106 ChannelBindingInner::TlsServerEndPoint(ref buf) => buf,
107 }
108 }
109}
110
111enum State {
112 Update {
113 nonce: String,
114 password: Vec<u8>,
115 channel_binding: ChannelBinding,
116 },
117 Finish {
118 salted_password: [u8; 32],
119 auth_message: String,
120 },
121 Done,
122}
123
124pub struct ScramSha256 {
140 message: String,
141 state: State,
142}
143
144impl ScramSha256 {
145 pub fn new(password: &[u8], channel_binding: ChannelBinding) -> ScramSha256 {
147 let mut rng = rand::rng();
149 let nonce = (0..NONCE_LENGTH)
150 .map(|_| {
151 let mut v = rng.random_range(0x21u8..0x7e);
152 if v == 0x2c {
153 v = 0x7e
154 }
155 v as char
156 })
157 .collect::<String>();
158
159 ScramSha256::new_inner(password, channel_binding, nonce)
160 }
161
162 fn new_inner(password: &[u8], channel_binding: ChannelBinding, nonce: String) -> ScramSha256 {
163 ScramSha256 {
164 message: format!("{}n=,r={}", channel_binding.gs2_header(), nonce),
165 state: State::Update {
166 nonce,
167 password: normalize(password),
168 channel_binding,
169 },
170 }
171 }
172
173 pub fn message(&self) -> &[u8] {
175 if let State::Done = self.state {
176 panic!("invalid SCRAM state");
177 }
178 self.message.as_bytes()
179 }
180
181 pub fn update(&mut self, message: &[u8]) -> io::Result<()> {
185 let (client_nonce, password, channel_binding) =
186 match mem::replace(&mut self.state, State::Done) {
187 State::Update {
188 nonce,
189 password,
190 channel_binding,
191 } => (nonce, password, channel_binding),
192 _ => return Err(io::Error::other("invalid SCRAM state")),
193 };
194
195 let message =
196 str::from_utf8(message).map_err(|e| io::Error::new(io::ErrorKind::InvalidInput, e))?;
197
198 let parsed = Parser::new(message).server_first_message()?;
199
200 if !parsed.nonce.starts_with(&client_nonce) {
201 return Err(io::Error::new(io::ErrorKind::InvalidInput, "invalid nonce"));
202 }
203
204 if parsed.iteration_count > MAX_ITERATION_COUNT {
205 return Err(io::Error::new(
206 io::ErrorKind::InvalidInput,
207 "SCRAM iteration count exceeds the maximum allowed",
208 ));
209 }
210
211 let salt = match STANDARD.decode(parsed.salt) {
212 Ok(salt) => salt,
213 Err(e) => return Err(io::Error::new(io::ErrorKind::InvalidInput, e)),
214 };
215
216 let salted_password = hi(&password, &salt, parsed.iteration_count);
217
218 let mut hmac = Hmac::<Sha256>::new_from_slice(&salted_password)
219 .expect("HMAC is able to accept all key sizes");
220 hmac.update(b"Client Key");
221 let client_key = hmac.finalize().into_bytes();
222
223 let mut hash = Sha256::default();
224 hash.update(client_key);
225 let stored_key = hash.finalize_fixed();
226
227 let mut cbind_input = vec![];
228 cbind_input.extend(channel_binding.gs2_header().as_bytes());
229 cbind_input.extend(channel_binding.cbind_data());
230 let cbind_input = STANDARD.encode(&cbind_input);
231
232 self.message.clear();
233 write!(&mut self.message, "c={},r={}", cbind_input, parsed.nonce).unwrap();
234
235 let auth_message = format!("n=,r={},{},{}", client_nonce, message, self.message);
236
237 let mut hmac = Hmac::<Sha256>::new_from_slice(&stored_key)
238 .expect("HMAC is able to accept all key sizes");
239 hmac.update(auth_message.as_bytes());
240 let client_signature = hmac.finalize().into_bytes();
241
242 let mut client_proof = client_key;
243 for (proof, signature) in client_proof.iter_mut().zip(client_signature) {
244 *proof ^= signature;
245 }
246
247 write!(
248 &mut self.message,
249 ",p={}",
250 Base64Display::new(&client_proof, &STANDARD)
251 )
252 .unwrap();
253
254 self.state = State::Finish {
255 salted_password,
256 auth_message,
257 };
258 Ok(())
259 }
260
261 pub fn finish(&mut self, message: &[u8]) -> io::Result<()> {
266 let (salted_password, auth_message) = match mem::replace(&mut self.state, State::Done) {
267 State::Finish {
268 salted_password,
269 auth_message,
270 } => (salted_password, auth_message),
271 _ => return Err(io::Error::other("invalid SCRAM state")),
272 };
273
274 let message =
275 str::from_utf8(message).map_err(|e| io::Error::new(io::ErrorKind::InvalidInput, e))?;
276
277 let parsed = Parser::new(message).server_final_message()?;
278
279 let verifier = match parsed {
280 ServerFinalMessage::Error(e) => {
281 return Err(io::Error::other(format!("SCRAM error: {e}")));
282 }
283 ServerFinalMessage::Verifier(verifier) => verifier,
284 };
285
286 let verifier = match STANDARD.decode(verifier) {
287 Ok(verifier) => verifier,
288 Err(e) => return Err(io::Error::new(io::ErrorKind::InvalidInput, e)),
289 };
290
291 let mut hmac = Hmac::<Sha256>::new_from_slice(&salted_password)
292 .expect("HMAC is able to accept all key sizes");
293 hmac.update(b"Server Key");
294 let server_key = hmac.finalize().into_bytes();
295
296 let mut hmac = Hmac::<Sha256>::new_from_slice(&server_key)
297 .expect("HMAC is able to accept all key sizes");
298 hmac.update(auth_message.as_bytes());
299 hmac.verify_slice(&verifier)
300 .map_err(|_| io::Error::new(io::ErrorKind::InvalidInput, "SCRAM verification error"))
301 }
302}
303
304struct Parser<'a> {
305 s: &'a str,
306 it: iter::Peekable<str::CharIndices<'a>>,
307}
308
309impl<'a> Parser<'a> {
310 fn new(s: &'a str) -> Parser<'a> {
311 Parser {
312 s,
313 it: s.char_indices().peekable(),
314 }
315 }
316
317 fn eat(&mut self, target: char) -> io::Result<()> {
318 match self.it.next() {
319 Some((_, c)) if c == target => Ok(()),
320 Some((i, c)) => {
321 let m =
322 format!("unexpected character at byte {i}: expected `{target}` but got `{c}");
323 Err(io::Error::new(io::ErrorKind::InvalidInput, m))
324 }
325 None => Err(io::Error::new(
326 io::ErrorKind::UnexpectedEof,
327 "unexpected EOF",
328 )),
329 }
330 }
331
332 fn take_while<F>(&mut self, f: F) -> io::Result<&'a str>
333 where
334 F: Fn(char) -> bool,
335 {
336 let start = match self.it.peek() {
337 Some(&(i, _)) => i,
338 None => return Ok(""),
339 };
340
341 loop {
342 match self.it.peek() {
343 Some(&(_, c)) if f(c) => {
344 self.it.next();
345 }
346 Some(&(i, _)) => return Ok(&self.s[start..i]),
347 None => return Ok(&self.s[start..]),
348 }
349 }
350 }
351
352 fn printable(&mut self) -> io::Result<&'a str> {
353 self.take_while(|c| matches!(c, '\x21'..='\x2b' | '\x2d'..='\x7e'))
354 }
355
356 fn nonce(&mut self) -> io::Result<&'a str> {
357 self.eat('r')?;
358 self.eat('=')?;
359 self.printable()
360 }
361
362 fn base64(&mut self) -> io::Result<&'a str> {
363 self.take_while(|c| matches!(c, 'a'..='z' | 'A'..='Z' | '0'..='9' | '/' | '+' | '='))
364 }
365
366 fn salt(&mut self) -> io::Result<&'a str> {
367 self.eat('s')?;
368 self.eat('=')?;
369 self.base64()
370 }
371
372 fn posit_number(&mut self) -> io::Result<u32> {
373 let n = self.take_while(|c| c.is_ascii_digit())?;
374 n.parse()
375 .map_err(|e| io::Error::new(io::ErrorKind::InvalidInput, e))
376 }
377
378 fn iteration_count(&mut self) -> io::Result<u32> {
379 self.eat('i')?;
380 self.eat('=')?;
381 self.posit_number()
382 }
383
384 fn eof(&mut self) -> io::Result<()> {
385 match self.it.peek() {
386 Some(&(i, _)) => Err(io::Error::new(
387 io::ErrorKind::InvalidInput,
388 format!("unexpected trailing data at byte {i}"),
389 )),
390 None => Ok(()),
391 }
392 }
393
394 fn server_first_message(&mut self) -> io::Result<ServerFirstMessage<'a>> {
395 let nonce = self.nonce()?;
396 self.eat(',')?;
397 let salt = self.salt()?;
398 self.eat(',')?;
399 let iteration_count = self.iteration_count()?;
400 self.eof()?;
401
402 Ok(ServerFirstMessage {
403 nonce,
404 salt,
405 iteration_count,
406 })
407 }
408
409 fn value(&mut self) -> io::Result<&'a str> {
410 self.take_while(|c| !matches!(c, '\0' | '=' | ','))
411 }
412
413 fn server_error(&mut self) -> io::Result<Option<&'a str>> {
414 match self.it.peek() {
415 Some(&(_, 'e')) => {}
416 _ => return Ok(None),
417 }
418
419 self.eat('e')?;
420 self.eat('=')?;
421 self.value().map(Some)
422 }
423
424 fn verifier(&mut self) -> io::Result<&'a str> {
425 self.eat('v')?;
426 self.eat('=')?;
427 self.base64()
428 }
429
430 fn server_final_message(&mut self) -> io::Result<ServerFinalMessage<'a>> {
431 let message = match self.server_error()? {
432 Some(error) => ServerFinalMessage::Error(error),
433 None => ServerFinalMessage::Verifier(self.verifier()?),
434 };
435 self.eof()?;
436 Ok(message)
437 }
438}
439
440struct ServerFirstMessage<'a> {
441 nonce: &'a str,
442 salt: &'a str,
443 iteration_count: u32,
444}
445
446enum ServerFinalMessage<'a> {
447 Error(&'a str),
448 Verifier(&'a str),
449}
450
451#[cfg(test)]
452mod test {
453 use super::*;
454
455 #[test]
456 fn parse_server_first_message() {
457 let message = "r=fyko+d2lbbFgONRv9qkxdawL3rfcNHYJY1ZVvWVs7j,s=QSXCR+Q6sek8bf92,i=4096";
458 let message = Parser::new(message).server_first_message().unwrap();
459 assert_eq!(message.nonce, "fyko+d2lbbFgONRv9qkxdawL3rfcNHYJY1ZVvWVs7j");
460 assert_eq!(message.salt, "QSXCR+Q6sek8bf92");
461 assert_eq!(message.iteration_count, 4096);
462 }
463
464 #[test]
465 fn parse_server_error_message() {
466 let message = "e=invalid-proof";
467 match Parser::new(message).server_final_message().unwrap() {
468 ServerFinalMessage::Error(error) => assert_eq!(error, "invalid-proof"),
469 ServerFinalMessage::Verifier(_) => panic!("expected server error"),
470 }
471
472 for message in ["invalid-proof\0x", "invalid-proof=x", "invalid-proof,x"] {
474 assert_eq!(Parser::new(message).value().unwrap(), "invalid-proof");
475 }
476 }
477
478 #[test]
480 fn exchange() {
481 let password = "foobar";
482 let nonce = "9IZ2O01zb9IgiIZ1WJ/zgpJB";
483
484 let client_first = "n,,n=,r=9IZ2O01zb9IgiIZ1WJ/zgpJB";
485 let server_first = "r=9IZ2O01zb9IgiIZ1WJ/zgpJBjx/oIRLs02gGSHcw1KEty3eY,s=fs3IXBy7U7+IvVjZ,i\
486 =4096";
487 let client_final = "c=biws,r=9IZ2O01zb9IgiIZ1WJ/zgpJBjx/oIRLs02gGSHcw1KEty3eY,p=AmNKosjJzS3\
488 1NTlQYNs5BTeQjdHdk7lOflDo5re2an8=";
489 let server_final = "v=U+ppxD5XUKtradnv8e2MkeupiA8FU87Sg8CXzXHDAzw=";
490
491 let mut scram = ScramSha256::new_inner(
492 password.as_bytes(),
493 ChannelBinding::unsupported(),
494 nonce.to_string(),
495 );
496 assert_eq!(str::from_utf8(scram.message()).unwrap(), client_first);
497
498 scram.update(server_first.as_bytes()).unwrap();
499 assert_eq!(str::from_utf8(scram.message()).unwrap(), client_final);
500
501 scram.finish(server_final.as_bytes()).unwrap();
502 }
503
504 #[test]
505 fn iteration_count_limit_is_accepted() {
506 let nonce = "9IZ2O01zb9IgiIZ1WJ/zgpJB";
507 let server_first =
508 "r=9IZ2O01zb9IgiIZ1WJ/zgpJBjx/oIRLs02gGSHcw1KEty3eY,s=fs3IXBy7U7+IvVjZ,i=2000000";
509
510 let mut scram =
511 ScramSha256::new_inner(b"foobar", ChannelBinding::unsupported(), nonce.to_string());
512 scram.update(server_first.as_bytes()).unwrap();
513 }
514
515 #[test]
516 fn excessive_iteration_count_is_rejected() {
517 let nonce = "9IZ2O01zb9IgiIZ1WJ/zgpJB";
520 let server_first =
521 "r=9IZ2O01zb9IgiIZ1WJ/zgpJBjx/oIRLs02gGSHcw1KEty3eY,s=fs3IXBy7U7+IvVjZ,i=2000001";
522
523 let mut scram =
524 ScramSha256::new_inner(b"foobar", ChannelBinding::unsupported(), nonce.to_string());
525 let err = scram.update(server_first.as_bytes()).unwrap_err();
526 assert_eq!(err.kind(), std::io::ErrorKind::InvalidInput);
527 assert_eq!(
528 err.to_string(),
529 "SCRAM iteration count exceeds the maximum allowed"
530 );
531 }
532}