1use 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
44static PG_SCRAM_MOCK_SECRET: LazyLock<[u8; PG_SCRAM_SHA256_KEY_LEN]> = LazyLock::new(rand::random);
47
48pub 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
105pub 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
132pub 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 let tmp = sha1_two(salt, hash_stage_2);
193 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 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
406pub 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 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 assert!(
657 alice
658 .verify_client_proof(b"auth-message", &[0u8; PG_SCRAM_SHA256_KEY_LEN])
659 .unwrap()
660 .is_none()
661 );
662 }
663}