Skip to main content

postgres_protocol/authentication/
sasl.rs

1//! SASL-based authentication support.
2
3use 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
18/// The maximum SCRAM iteration count the client will accept from the server.
19///
20/// The iteration count is sent by the server and drives a PBKDF2 loop, so an
21/// unbounded value lets a malicious or impersonating server force the client to
22/// perform an arbitrary number of HMAC operations before authentication even
23/// completes (a denial of service). Allow servers configured with up to
24/// 2_000_000 iterations while keeping authentication work bounded.
25const MAX_ITERATION_COUNT: u32 = 2_000_000;
26
27/// The identifier of the SCRAM-SHA-256 SASL authentication mechanism.
28pub const SCRAM_SHA_256: &str = "SCRAM-SHA-256";
29/// The identifier of the SCRAM-SHA-256-PLUS SASL authentication mechanism.
30pub const SCRAM_SHA_256_PLUS: &str = "SCRAM-SHA-256-PLUS";
31
32// since postgres passwords are not required to exclude saslprep-prohibited
33// characters or even be valid UTF8, we run saslprep if possible and otherwise
34// return the raw password.
35fn 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
75/// The channel binding configuration for a SCRAM authentication exchange.
76pub struct ChannelBinding(ChannelBindingInner);
77
78impl ChannelBinding {
79    /// The server did not request channel binding.
80    pub fn unrequested() -> ChannelBinding {
81        ChannelBinding(ChannelBindingInner::Unrequested)
82    }
83
84    /// The server requested channel binding but the client is unable to provide it.
85    pub fn unsupported() -> ChannelBinding {
86        ChannelBinding(ChannelBindingInner::Unsupported)
87    }
88
89    /// The server requested channel binding and the client will use the `tls-server-end-point`
90    /// method.
91    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
124/// A type which handles the client side of the SCRAM-SHA-256/SCRAM-SHA-256-PLUS authentication
125/// process.
126///
127/// During the authentication process, if the backend sends an `AuthenticationSASL` message which
128/// includes `SCRAM-SHA-256` as an authentication mechanism, this type can be used.
129///
130/// After a `ScramSha256` is constructed, the buffer returned by the `message()` method should be
131/// sent to the backend in a `SASLInitialResponse` message along with the mechanism name.
132///
133/// The server will reply with an `AuthenticationSASLContinue` message. Its contents should be
134/// passed to the `update()` method, after which the buffer returned by the `message()` method
135/// should be sent to the backend in a `SASLResponse` message.
136///
137/// The server will reply with an `AuthenticationSASLFinal` message. Its contents should be passed
138/// to the `finish()` method, after which the authentication process is complete.
139pub struct ScramSha256 {
140    message: String,
141    state: State,
142}
143
144impl ScramSha256 {
145    /// Constructs a new instance which will use the provided password for authentication.
146    pub fn new(password: &[u8], channel_binding: ChannelBinding) -> ScramSha256 {
147        // rand 0.5's ThreadRng is cryptographically secure
148        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    /// Returns the message which should be sent to the backend in an `SASLResponse` message.
174    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    /// Updates the state machine with the response from the backend.
182    ///
183    /// This should be called when an `AuthenticationSASLContinue` message is received.
184    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    /// Finalizes the authentication process.
262    ///
263    /// This should be called when the backend sends an `AuthenticationSASLFinal` message.
264    /// Authentication has only succeeded if this method returns `Ok(())`.
265    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        // the error value ends at the first '\0', '=' or ','
473        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    // recorded auth exchange from psql
479    #[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        // a malicious server cannot force an unbounded PBKDF2 loop; the iteration
518        // count is rejected before `hi()` runs.
519        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}