Skip to main content

auth/
common.rs

1// Copyright 2023 Greptime Team
2//
3// Licensed under the Apache License, Version 2.0 (the "License");
4// you may not use this file except in compliance with the License.
5// You may obtain a copy of the License at
6//
7//     http://www.apache.org/licenses/LICENSE-2.0
8//
9// Unless required by applicable law or agreed to in writing, software
10// distributed under the License is distributed on an "AS IS" BASIS,
11// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12// See the License for the specific language governing permissions and
13// limitations under the License.
14
15use std::sync::{Arc, LazyLock};
16
17use common_base::secrets::SecretString;
18use digest::Digest;
19use hmac::{Hmac, Mac};
20use pbkdf2::pbkdf2_hmac;
21use sha1::Sha1;
22use sha2::Sha256;
23use snafu::{OptionExt, ensure};
24use subtle::ConstantTimeEq;
25
26use crate::error::{IllegalParamSnafu, InvalidConfigSnafu, Result, UserPasswordMismatchSnafu};
27use crate::user_info::DefaultUserInfo;
28use crate::user_provider::static_user_provider::{STATIC_USER_PROVIDER, StaticUserProvider};
29use crate::user_provider::watch_file_user_provider::{
30    WATCH_FILE_USER_PROVIDER, WatchFileUserProvider,
31};
32use crate::{UserInfoRef, UserProviderRef};
33
34pub(crate) const DEFAULT_USERNAME: &str = "greptime";
35pub const DEFAULT_PBKDF2_SHA256_ITERATIONS: u32 = 4096;
36pub const DEFAULT_PBKDF2_SHA256_SALT_LEN: usize = 16;
37pub const PBKDF2_SHA256_HASH_LEN: usize = 32;
38pub const MAX_PBKDF2_SHA256_ITERATIONS: u32 = 1_000_000;
39pub const MAX_PBKDF2_SHA256_SALT_LEN: usize = 1024;
40pub const PG_SCRAM_SHA256_KEY_LEN: usize = 32;
41
42type HmacSha256 = Hmac<Sha256>;
43
44/// Process-wide secret used to derive a stable mock salt for unknown users, so
45/// the SCRAM `server-first-message` can't be used to enumerate usernames.
46static PG_SCRAM_MOCK_SECRET: LazyLock<[u8; PG_SCRAM_SHA256_KEY_LEN]> = LazyLock::new(rand::random);
47
48/// construct a [`UserInfo`](crate::user_info::UserInfo) impl with name
49/// use default username `greptime` if None is provided
50pub fn userinfo_by_name(username: Option<String>) -> UserInfoRef {
51    DefaultUserInfo::with_name(username.unwrap_or_else(|| DEFAULT_USERNAME.to_string()))
52}
53
54pub fn user_provider_from_option(opt: &str) -> Result<UserProviderRef> {
55    let (name, content) = opt.split_once(':').with_context(|| InvalidConfigSnafu {
56        value: opt.to_string(),
57        msg: "UserProviderOption must be in format `<option>:<value>`",
58    })?;
59    match name {
60        STATIC_USER_PROVIDER => {
61            let provider =
62                StaticUserProvider::new(content).map(|p| Arc::new(p) as UserProviderRef)?;
63            Ok(provider)
64        }
65        WATCH_FILE_USER_PROVIDER => {
66            WatchFileUserProvider::new(content).map(|p| Arc::new(p) as UserProviderRef)
67        }
68        _ => InvalidConfigSnafu {
69            value: name.to_string(),
70            msg: "Invalid UserProviderOption",
71        }
72        .fail(),
73    }
74}
75
76pub fn static_user_provider_from_option(opt: &str) -> Result<StaticUserProvider> {
77    let (name, content) = opt.split_once(':').with_context(|| InvalidConfigSnafu {
78        value: opt.to_string(),
79        msg: "UserProviderOption must be in format `<option>:<value>`",
80    })?;
81    match name {
82        STATIC_USER_PROVIDER => {
83            let provider = StaticUserProvider::new(content)?;
84            Ok(provider)
85        }
86        _ => InvalidConfigSnafu {
87            value: name.to_string(),
88            msg: format!("Invalid UserProviderOption, expect only {STATIC_USER_PROVIDER}"),
89        }
90        .fail(),
91    }
92}
93
94type Username<'a> = &'a str;
95type HostOrIp<'a> = &'a str;
96
97#[derive(Debug, Clone)]
98pub enum Identity<'a> {
99    UserId(Username<'a>, Option<HostOrIp<'a>>),
100}
101
102pub type HashedPassword<'a> = &'a [u8];
103pub type Salt<'a> = &'a [u8];
104
105/// Authentication information sent by the client.
106pub enum Password<'a> {
107    PlainText(SecretString),
108    MysqlNativePassword(HashedPassword<'a>, Salt<'a>),
109    PgMD5(HashedPassword<'a>, Salt<'a>),
110}
111
112impl Password<'_> {
113    pub fn r#type(&self) -> &str {
114        match self {
115            Password::PlainText(_) => "plain_text",
116            Password::MysqlNativePassword(_, _) => "mysql_native_password",
117            Password::PgMD5(_, _) => "pg_md5",
118        }
119    }
120}
121
122pub fn auth_mysql(
123    auth_data: HashedPassword,
124    salt: Salt,
125    username: &str,
126    save_pwd: &[u8],
127) -> Result<()> {
128    let hash_stage_2 = mysql_native_password_hash(save_pwd);
129    auth_mysql_with_hash_stage_2(auth_data, salt, username, &hash_stage_2)
130}
131
132/// Authenticates a MySQL native password response against an encoded verifier.
133pub fn auth_mysql_with_verifier(
134    auth_data: HashedPassword,
135    salt: Salt,
136    username: &str,
137    verifier: &str,
138) -> Result<()> {
139    let hash_stage_2 = parse_mysql_native_password_verifier(verifier)?;
140    auth_mysql_with_hash_stage_2(auth_data, salt, username, &hash_stage_2)
141}
142
143pub fn validate_mysql_native_password_verifier(verifier: &str) -> Result<()> {
144    parse_mysql_native_password_verifier(verifier).map(|_| ())
145}
146
147pub(crate) fn parse_mysql_native_password_verifier(verifier: &str) -> Result<Vec<u8>> {
148    let Some(verifier) = verifier.strip_prefix("mysql_native_password:") else {
149        return InvalidConfigSnafu {
150            value: "mysql_native_password".to_string(),
151            msg: "Invalid mysql native password verifier format",
152        }
153        .fail();
154    };
155    let Ok(hash_stage_2) = hex::decode(verifier) else {
156        return InvalidConfigSnafu {
157            value: "mysql_native_password".to_string(),
158            msg: "Invalid mysql native password verifier encoding",
159        }
160        .fail();
161    };
162    ensure!(
163        hash_stage_2.len() == 20,
164        InvalidConfigSnafu {
165            value: "mysql_native_password".to_string(),
166            msg: "Illegal mysql native password verifier length",
167        }
168    );
169    Ok(hash_stage_2)
170}
171
172pub(crate) fn auth_mysql_with_hash_stage_2(
173    auth_data: HashedPassword,
174    salt: Salt,
175    username: &str,
176    hash_stage_2: &[u8],
177) -> Result<()> {
178    ensure!(
179        auth_data.len() == 20,
180        IllegalParamSnafu {
181            msg: "Illegal mysql password length"
182        }
183    );
184    ensure!(
185        hash_stage_2.len() == 20,
186        InvalidConfigSnafu {
187            value: hash_stage_2.len().to_string(),
188            msg: "Illegal mysql native password verifier length",
189        }
190    );
191    // ref: https://github.com/mysql/mysql-server/blob/a246bad76b9271cb4333634e954040a970222e0a/sql/auth/password.cc#L62
192    let tmp = sha1_two(salt, hash_stage_2);
193    // xor auth_data and tmp
194    let mut xor_result = [0u8; 20];
195    for i in 0..20 {
196        xor_result[i] = auth_data[i] ^ tmp[i];
197    }
198    let candidate_stage_2 = sha1_one(&xor_result);
199    if candidate_stage_2 == hash_stage_2 {
200        Ok(())
201    } else {
202        UserPasswordMismatchSnafu {
203            username: username.to_string(),
204        }
205        .fail()
206    }
207}
208
209pub fn mysql_native_password_hash(save_pwd: &[u8]) -> Vec<u8> {
210    double_sha1(save_pwd)
211}
212
213pub fn format_mysql_native_password_verifier(password: &[u8]) -> String {
214    format!(
215        "mysql_native_password:{}",
216        hex::encode(mysql_native_password_hash(password))
217    )
218}
219
220pub fn format_pbkdf2_sha256_password_verifier(
221    password: &[u8],
222    salt: &[u8],
223    iterations: u32,
224) -> Result<String> {
225    ensure!(
226        iterations > 0 && iterations <= MAX_PBKDF2_SHA256_ITERATIONS,
227        IllegalParamSnafu {
228            msg: format!(
229                "pbkdf2_sha256 iterations must be in 1..={}",
230                MAX_PBKDF2_SHA256_ITERATIONS
231            )
232        }
233    );
234    ensure!(
235        !salt.is_empty() && salt.len() <= MAX_PBKDF2_SHA256_SALT_LEN,
236        IllegalParamSnafu {
237            msg: format!(
238                "pbkdf2_sha256 salt length must be in 1..={}",
239                MAX_PBKDF2_SHA256_SALT_LEN
240            )
241        }
242    );
243
244    let mut hash = [0u8; PBKDF2_SHA256_HASH_LEN];
245    pbkdf2_hmac::<Sha256>(password, salt, iterations, &mut hash);
246    Ok(format!(
247        "pbkdf2_sha256:{iterations}:{}:{}",
248        hex::encode(salt),
249        hex::encode(hash)
250    ))
251}
252
253#[derive(Clone, PartialEq, Eq)]
254pub struct PgScramSha256Verifier {
255    iterations: u32,
256    salt: Vec<u8>,
257    stored_key: Vec<u8>,
258    server_key: Vec<u8>,
259}
260
261impl std::fmt::Debug for PgScramSha256Verifier {
262    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
263        f.debug_struct("PgScramSha256Verifier")
264            .field("iterations", &self.iterations)
265            .field("salt", &"<REDACTED>")
266            .field("stored_key", &"<REDACTED>")
267            .field("server_key", &"<REDACTED>")
268            .finish()
269    }
270}
271
272impl PgScramSha256Verifier {
273    pub fn new(
274        iterations: u32,
275        salt: Vec<u8>,
276        stored_key: Vec<u8>,
277        server_key: Vec<u8>,
278    ) -> Result<Self> {
279        ensure!(
280            iterations > 0 && iterations <= MAX_PBKDF2_SHA256_ITERATIONS,
281            IllegalParamSnafu {
282                msg: format!(
283                    "pg_scram_sha256 iterations must be in 1..={}",
284                    MAX_PBKDF2_SHA256_ITERATIONS
285                )
286            }
287        );
288        ensure!(
289            !salt.is_empty() && salt.len() <= MAX_PBKDF2_SHA256_SALT_LEN,
290            IllegalParamSnafu {
291                msg: format!(
292                    "pg_scram_sha256 salt length must be in 1..={}",
293                    MAX_PBKDF2_SHA256_SALT_LEN
294                )
295            }
296        );
297        ensure!(
298            stored_key.len() == PG_SCRAM_SHA256_KEY_LEN,
299            IllegalParamSnafu {
300                msg: "pg_scram_sha256 stored key must be 32 bytes"
301            }
302        );
303        ensure!(
304            server_key.len() == PG_SCRAM_SHA256_KEY_LEN,
305            IllegalParamSnafu {
306                msg: "pg_scram_sha256 server key must be 32 bytes"
307            }
308        );
309
310        Ok(Self {
311            iterations,
312            salt,
313            stored_key,
314            server_key,
315        })
316    }
317
318    pub fn from_password(password: &[u8], salt: &[u8], iterations: u32) -> Result<Self> {
319        let salted_password = pg_scram_sha256_salted_password(password, salt, iterations)?;
320        Self::from_salted_password(salted_password, salt.to_vec(), iterations)
321    }
322
323    pub fn from_salted_password(
324        salted_password: Vec<u8>,
325        salt: Vec<u8>,
326        iterations: u32,
327    ) -> Result<Self> {
328        ensure!(
329            salted_password.len() == PG_SCRAM_SHA256_KEY_LEN,
330            IllegalParamSnafu {
331                msg: "pg_scram_sha256 salted password must be 32 bytes"
332            }
333        );
334        let client_key = hmac_sha256(&salted_password, b"Client Key");
335        let stored_key = sha256(&client_key);
336        let server_key = hmac_sha256(&salted_password, b"Server Key");
337        Self::new(iterations, salt, stored_key, server_key)
338    }
339
340    /// Builds a throwaway verifier for an unknown user. The salt is derived
341    /// deterministically from the username and a process-wide secret, so the
342    /// SCRAM `server-first-message` (salt + iterations) stays stable per
343    /// username and indistinguishable from a real user across reconnects. No
344    /// PBKDF2 is run (avoids a CPU-exhaustion DoS keyed on unknown usernames),
345    /// and the random keys guarantee the client proof never matches.
346    pub fn mock_for_unknown_user(username: &[u8]) -> Self {
347        let mock_salt = hmac_sha256(PG_SCRAM_MOCK_SECRET.as_slice(), username);
348        Self {
349            iterations: DEFAULT_PBKDF2_SHA256_ITERATIONS,
350            salt: mock_salt[..DEFAULT_PBKDF2_SHA256_SALT_LEN].to_vec(),
351            stored_key: rand::random::<[u8; PG_SCRAM_SHA256_KEY_LEN]>().to_vec(),
352            server_key: rand::random::<[u8; PG_SCRAM_SHA256_KEY_LEN]>().to_vec(),
353        }
354    }
355
356    pub fn iterations(&self) -> u32 {
357        self.iterations
358    }
359
360    pub fn salt(&self) -> &[u8] {
361        &self.salt
362    }
363
364    pub fn verify_plain_password(&self, password: &[u8]) -> Result<bool> {
365        let salted_password =
366            pg_scram_sha256_salted_password(password, &self.salt, self.iterations)?;
367        let client_key = hmac_sha256(&salted_password, b"Client Key");
368        let stored_key = sha256(&client_key);
369        Ok(self.stored_key.ct_eq(&stored_key).into())
370    }
371
372    pub fn verify_client_proof(
373        &self,
374        auth_message: &[u8],
375        client_proof: &[u8],
376    ) -> Result<Option<Vec<u8>>> {
377        if client_proof.len() != PG_SCRAM_SHA256_KEY_LEN {
378            return Ok(None);
379        }
380
381        let client_signature = hmac_sha256(&self.stored_key, auth_message);
382        let client_key = xor(client_proof, &client_signature);
383        let stored_key = sha256(&client_key);
384        if self.stored_key.ct_eq(&stored_key).into() {
385            Ok(Some(hmac_sha256(&self.server_key, auth_message)))
386        } else {
387            Ok(None)
388        }
389    }
390}
391
392pub fn format_pg_scram_sha256_password_verifier(
393    password: &[u8],
394    salt: &[u8],
395    iterations: u32,
396) -> Result<String> {
397    let verifier = PgScramSha256Verifier::from_password(password, salt, iterations)?;
398    Ok(format!(
399        "pg_scram_sha256:{iterations}:{}:{}:{}",
400        hex::encode(verifier.salt),
401        hex::encode(verifier.stored_key),
402        hex::encode(verifier.server_key)
403    ))
404}
405
406/// Parses and validates an encoded PostgreSQL SCRAM-SHA-256 verifier.
407pub fn parse_pg_scram_sha256_password_verifier(verifier: &str) -> Result<PgScramSha256Verifier> {
408    let Some(verifier) = verifier.strip_prefix("pg_scram_sha256:") else {
409        return InvalidConfigSnafu {
410            value: "pg_scram_sha256".to_string(),
411            msg: "Invalid pg scram sha256 verifier format",
412        }
413        .fail();
414    };
415    let mut parts = verifier.split(':');
416    let (Some(iterations), Some(salt), Some(stored_key), Some(server_key), None) = (
417        parts.next(),
418        parts.next(),
419        parts.next(),
420        parts.next(),
421        parts.next(),
422    ) else {
423        return InvalidConfigSnafu {
424            value: "pg_scram_sha256".to_string(),
425            msg: "Invalid pg scram sha256 verifier format",
426        }
427        .fail();
428    };
429    let (Ok(iterations), Ok(salt), Ok(stored_key), Ok(server_key)) = (
430        iterations.parse::<u32>(),
431        hex::decode(salt),
432        hex::decode(stored_key),
433        hex::decode(server_key),
434    ) else {
435        return InvalidConfigSnafu {
436            value: "pg_scram_sha256".to_string(),
437            msg: "Invalid pg scram sha256 verifier encoding",
438        }
439        .fail();
440    };
441
442    PgScramSha256Verifier::new(iterations, salt, stored_key, server_key)
443}
444
445fn pg_scram_sha256_salted_password(
446    password: &[u8],
447    salt: &[u8],
448    iterations: u32,
449) -> Result<Vec<u8>> {
450    ensure!(
451        iterations > 0 && iterations <= MAX_PBKDF2_SHA256_ITERATIONS,
452        IllegalParamSnafu {
453            msg: format!(
454                "pg_scram_sha256 iterations must be in 1..={}",
455                MAX_PBKDF2_SHA256_ITERATIONS
456            )
457        }
458    );
459    ensure!(
460        !salt.is_empty() && salt.len() <= MAX_PBKDF2_SHA256_SALT_LEN,
461        IllegalParamSnafu {
462            msg: format!(
463                "pg_scram_sha256 salt length must be in 1..={}",
464                MAX_PBKDF2_SHA256_SALT_LEN
465            )
466        }
467    );
468
469    let prepared_password = std::str::from_utf8(password)
470        .ok()
471        .and_then(|password| stringprep::saslprep(password).ok());
472    let password = prepared_password
473        .as_deref()
474        .map(str::as_bytes)
475        .unwrap_or(password);
476
477    let mut salted_password = [0u8; PG_SCRAM_SHA256_KEY_LEN];
478    pbkdf2_hmac::<Sha256>(password, salt, iterations, &mut salted_password);
479    Ok(salted_password.to_vec())
480}
481
482fn hmac_sha256(key: &[u8], msg: &[u8]) -> Vec<u8> {
483    let mut mac = HmacSha256::new_from_slice(key).expect("HMAC accepts any key length");
484    mac.update(msg);
485    mac.finalize().into_bytes().to_vec()
486}
487
488fn sha256(data: &[u8]) -> Vec<u8> {
489    let mut hasher = Sha256::new();
490    hasher.update(data);
491    hasher.finalize().to_vec()
492}
493
494fn xor(lhs: &[u8], rhs: &[u8]) -> Vec<u8> {
495    lhs.iter().zip(rhs).map(|(l, r)| l ^ r).collect()
496}
497
498fn sha1_two(input_1: &[u8], input_2: &[u8]) -> Vec<u8> {
499    let mut hasher = Sha1::new();
500    hasher.update(input_1);
501    hasher.update(input_2);
502    hasher.finalize().to_vec()
503}
504
505fn sha1_one(data: &[u8]) -> Vec<u8> {
506    let mut hasher = Sha1::new();
507    hasher.update(data);
508    hasher.finalize().to_vec()
509}
510
511fn double_sha1(data: &[u8]) -> Vec<u8> {
512    sha1_one(&sha1_one(data))
513}
514
515#[cfg(test)]
516mod tests {
517    use super::*;
518
519    fn mysql_native_password_auth_data(password: &[u8], salt: &[u8]) -> Vec<u8> {
520        let hash_stage_1 = sha1_one(password);
521        let hash_stage_2 = mysql_native_password_hash(password);
522        let scramble = sha1_two(salt, &hash_stage_2);
523        hash_stage_1
524            .iter()
525            .zip(scramble)
526            .map(|(lhs, rhs)| lhs ^ rhs)
527            .collect()
528    }
529
530    #[test]
531    fn test_sha() {
532        let sha_1_answer: Vec<u8> = vec![
533            124, 74, 141, 9, 202, 55, 98, 175, 97, 229, 149, 32, 148, 61, 194, 100, 148, 248, 148,
534            27,
535        ];
536        let sha_1 = sha1_one("123456".as_bytes());
537        assert_eq!(sha_1, sha_1_answer);
538
539        let double_sha1_answer: Vec<u8> = vec![
540            107, 180, 131, 126, 183, 67, 41, 16, 94, 228, 86, 141, 218, 125, 198, 126, 210, 202,
541            42, 217,
542        ];
543        let double_sha1 = double_sha1("123456".as_bytes());
544        assert_eq!(double_sha1, double_sha1_answer);
545
546        let sha1_2_answer: Vec<u8> = vec![
547            132, 115, 215, 211, 99, 186, 164, 206, 168, 152, 217, 192, 117, 47, 240, 252, 142, 244,
548            37, 204,
549        ];
550        let sha1_2 = sha1_two("123456".as_bytes(), "654321".as_bytes());
551        assert_eq!(sha1_2, sha1_2_answer);
552    }
553
554    #[test]
555    fn test_format_mysql_native_password_verifier() {
556        let verifier = format_mysql_native_password_verifier("123456".as_bytes());
557        assert_eq!(
558            "mysql_native_password:6bb4837eb74329105ee4568dda7dc67ed2ca2ad9",
559            verifier
560        );
561        assert_eq!(
562            mysql_native_password_hash(b"123456"),
563            parse_mysql_native_password_verifier(&verifier).unwrap()
564        );
565        assert!(parse_mysql_native_password_verifier("mysql_native_password:00").is_err());
566
567        let salt = b"01234567890123456789";
568        let auth_data = mysql_native_password_auth_data(b"123456", salt);
569        auth_mysql_with_verifier(&auth_data, salt, "greptime", &verifier).unwrap();
570        assert!(
571            auth_mysql_with_verifier(
572                &auth_data,
573                salt,
574                "greptime",
575                &format_mysql_native_password_verifier(b"wrong")
576            )
577            .is_err()
578        );
579    }
580
581    #[test]
582    fn test_format_pbkdf2_sha256_password_verifier() {
583        let verifier =
584            format_pbkdf2_sha256_password_verifier("password".as_bytes(), b"salt", 4096).unwrap();
585        assert_eq!(
586            "pbkdf2_sha256:4096:73616c74:c5e478d59288c841aa530db6845c4c8d962893a001ce4e11a4963873aa98134a",
587            verifier
588        );
589
590        assert!(format_pbkdf2_sha256_password_verifier(b"password", b"", 4096).is_err());
591        assert!(format_pbkdf2_sha256_password_verifier(b"password", b"salt", 0).is_err());
592        assert!(
593            format_pbkdf2_sha256_password_verifier(
594                b"password",
595                b"salt",
596                MAX_PBKDF2_SHA256_ITERATIONS + 1,
597            )
598            .is_err()
599        );
600    }
601
602    #[test]
603    fn test_format_pg_scram_sha256_password_verifier() {
604        let verifier =
605            format_pg_scram_sha256_password_verifier("password".as_bytes(), b"salt", 4096).unwrap();
606        assert_eq!(
607            "pg_scram_sha256:4096:73616c74:945e1c466fc9932efadc23781edc5d1e78d5e10f005933652af1a6105154f084:b9bf0e811b1fb6793671c0cc3adedf7c75cd72291191092ad65878c5a02aad2c",
608            verifier
609        );
610        let parsed = parse_pg_scram_sha256_password_verifier(&verifier).unwrap();
611        assert_eq!(4096, parsed.iterations());
612        assert_eq!(b"salt", parsed.salt());
613        assert!(parse_pg_scram_sha256_password_verifier("pg_scram_sha256:bad").is_err());
614    }
615
616    #[test]
617    fn test_pg_scram_sha256_applies_saslprep() {
618        let normalized = PgScramSha256Verifier::from_password(b"pass word", b"salt", 4096).unwrap();
619        let non_breaking_space =
620            PgScramSha256Verifier::from_password("pass\u{00a0}word".as_bytes(), b"salt", 4096)
621                .unwrap();
622
623        assert_eq!(normalized, non_breaking_space);
624        assert!(
625            non_breaking_space
626                .verify_plain_password(b"pass word")
627                .unwrap()
628        );
629    }
630
631    #[test]
632    fn test_pg_scram_sha256_uses_original_bytes_when_saslprep_is_not_possible() {
633        let invalid_utf8 = b"password\xff";
634        let prohibited = b"password\x07";
635
636        for password in [invalid_utf8.as_slice(), prohibited.as_slice()] {
637            let verifier = PgScramSha256Verifier::from_password(password, b"salt", 4096).unwrap();
638            assert!(verifier.verify_plain_password(password).unwrap());
639        }
640    }
641
642    #[test]
643    fn test_mock_verifier_is_stable_per_username() {
644        let alice = PgScramSha256Verifier::mock_for_unknown_user(b"alice");
645        let alice_again = PgScramSha256Verifier::mock_for_unknown_user(b"alice");
646        let bob = PgScramSha256Verifier::mock_for_unknown_user(b"bob");
647
648        // Same username must yield the same salt and iterations across reconnects,
649        // otherwise the SCRAM server-first message leaks user (non-)existence.
650        assert_eq!(alice.salt(), alice_again.salt());
651        assert_ne!(alice.salt(), bob.salt());
652        assert_eq!(alice.iterations(), DEFAULT_PBKDF2_SHA256_ITERATIONS);
653        assert_eq!(alice.salt().len(), DEFAULT_PBKDF2_SHA256_SALT_LEN);
654
655        // The mock verifier must never accept a client proof.
656        assert!(
657            alice
658                .verify_client_proof(b"auth-message", &[0u8; PG_SCRAM_SHA256_KEY_LEN])
659                .unwrap()
660                .is_none()
661        );
662    }
663}