1use std::fmt::Debug;
16
17use ::auth::{
18 BEARER_TOKEN_USER, Identity, Password, PgAuthInfo, PgScramSha256Verifier, UserInfoRef,
19 UserProviderRef, userinfo_by_name,
20};
21use async_trait::async_trait;
22use base64::Engine;
23use base64::engine::general_purpose::STANDARD as BASE64;
24use common_catalog::parse_catalog_and_schema_from_db_string;
25use common_error::ext::ErrorExt;
26use common_error::status_code::StatusCode;
27use common_time::Timezone;
28use futures::{Sink, SinkExt};
29use pgwire::api::auth::StartupHandler;
30use pgwire::api::auth::sasl::SCRAM_SHA_256_METHOD;
31use pgwire::api::{ClientInfo, PgWireConnectionState, auth};
32use pgwire::error::{ErrorInfo, PgWireError, PgWireResult};
33use pgwire::messages::response::ErrorResponse;
34use pgwire::messages::startup::{Authentication, PasswordMessageFamily, SecretKey};
35use pgwire::messages::{PgWireBackendMessage, PgWireFrontendMessage};
36use session::Session;
37use snafu::IntoError;
38use tokio::sync::Mutex;
39
40use crate::error::{AuthSnafu, Result};
41use crate::metrics::METRIC_AUTH_FAILURE;
42use crate::postgres::PostgresServerHandlerInner;
43use crate::postgres::types::PgErrorCode;
44use crate::postgres::utils::convert_err;
45use crate::query_handler::sql::ServerSqlQueryHandlerRef;
46
47pub(crate) struct PgLoginVerifier {
48 user_provider: Option<UserProviderRef>,
49 state: Mutex<PgAuthenticationState>,
50}
51
52impl PgLoginVerifier {
53 pub(crate) fn new(user_provider: Option<UserProviderRef>) -> Self {
54 Self {
55 user_provider,
56 state: Mutex::new(PgAuthenticationState::Initial),
57 }
58 }
59}
60
61enum PgAuthenticationState {
62 Initial,
63 Cleartext,
64 SaslInitial {
65 auth_info: PgAuthInfo,
66 },
67 SaslFinal {
68 verifier: PgScramSha256Verifier,
69 user_info: Option<UserInfoRef>,
70 channel_binding: String,
71 nonce: String,
72 client_first_bare: String,
73 server_first: String,
74 },
75}
76
77#[allow(dead_code)]
78struct LoginInfo {
79 user: Option<String>,
80 catalog: Option<String>,
81 schema: Option<String>,
82 host: String,
83}
84
85impl LoginInfo {
86 pub fn from_client_info<C>(client: &C) -> LoginInfo
87 where
88 C: ClientInfo,
89 {
90 LoginInfo {
91 user: client.metadata().get(super::METADATA_USER).map(Into::into),
92 catalog: client
93 .metadata()
94 .get(super::METADATA_CATALOG)
95 .map(Into::into),
96 schema: client
97 .metadata()
98 .get(super::METADATA_SCHEMA)
99 .map(Into::into),
100 host: client.socket_addr().ip().to_string(),
101 }
102 }
103}
104
105impl PgLoginVerifier {
106 async fn auth(&self, login: &LoginInfo, password: &str) -> Result<Option<UserInfoRef>> {
107 let user_provider = match &self.user_provider {
108 Some(provider) => provider,
109 None => return Ok(None),
110 };
111
112 let user_name = match &login.user {
113 Some(name) => name,
114 None => return Ok(None),
115 };
116 let catalog = match &login.catalog {
117 Some(name) => name,
118 None => return Ok(None),
119 };
120 let schema = match &login.schema {
121 Some(name) => name,
122 None => return Ok(None),
123 };
124
125 let result = if user_name == BEARER_TOKEN_USER {
126 user_provider
127 .auth_bearer_token(password, catalog, schema)
128 .await
129 } else {
130 user_provider
131 .auth(
132 Identity::UserId(user_name, None),
133 Password::PlainText(password.to_string().into()),
134 catalog,
135 schema,
136 )
137 .await
138 };
139 match result {
140 Err(e) => {
141 METRIC_AUTH_FAILURE
142 .with_label_values(&[e.status_code().as_ref()])
143 .inc();
144 Err(AuthSnafu.into_error(e))
145 }
146 Ok(user_info) => Ok(Some(user_info)),
147 }
148 }
149
150 async fn postgres_auth_info(&self, login: &LoginInfo) -> Result<PgAuthInfo> {
151 let user_provider = match &self.user_provider {
152 Some(provider) => provider,
153 None => return Ok(PgAuthInfo::Cleartext),
154 };
155
156 let user_name = match &login.user {
157 Some(name) => name,
158 None => return Ok(PgAuthInfo::Cleartext),
159 };
160 if user_name == BEARER_TOKEN_USER {
161 return Ok(PgAuthInfo::Cleartext);
162 }
163 let catalog = match &login.catalog {
164 Some(name) => name,
165 None => return Ok(PgAuthInfo::Cleartext),
166 };
167
168 match user_provider
169 .postgres_auth_info(Identity::UserId(user_name, None), catalog)
170 .await
171 {
172 Err(e) => {
173 METRIC_AUTH_FAILURE
174 .with_label_values(&[e.status_code().as_ref()])
175 .inc();
176 Err(AuthSnafu.into_error(e))
177 }
178 Ok(auth_info) => Ok(auth_info),
179 }
180 }
181
182 async fn authorize(&self, login: &LoginInfo, user_info: &UserInfoRef) -> Result<()> {
183 let user_provider = match &self.user_provider {
184 Some(provider) => provider,
185 None => return Ok(()),
186 };
187
188 let catalog = match &login.catalog {
189 Some(name) => name,
190 None => return Ok(()),
191 };
192 let schema = match &login.schema {
193 Some(name) => name,
194 None => return Ok(()),
195 };
196
197 match user_provider.authorize(catalog, schema, user_info).await {
198 Err(e) => {
199 METRIC_AUTH_FAILURE
200 .with_label_values(&[e.status_code().as_ref()])
201 .inc();
202 Err(AuthSnafu.into_error(e))
203 }
204 Ok(()) => Ok(()),
205 }
206 }
207}
208
209fn set_client_info<C>(client: &mut C, session: &Session)
210where
211 C: ClientInfo,
212{
213 if let Some(current_catalog) = client.metadata().get(super::METADATA_CATALOG) {
214 session.set_catalog(current_catalog.clone());
215 }
216 if let Some(current_schema) = client.metadata().get(super::METADATA_SCHEMA) {
217 session.set_schema(current_schema.clone());
218 }
219
220 client.set_pid_and_secret_key(0, SecretKey::I32(0));
224 }
226
227#[async_trait]
228impl StartupHandler for PostgresServerHandlerInner {
229 async fn on_startup<C>(
230 &self,
231 client: &mut C,
232 message: PgWireFrontendMessage,
233 ) -> PgWireResult<()>
234 where
235 C: ClientInfo + Sink<PgWireBackendMessage> + Unpin + Send,
236 C::Error: Debug,
237 PgWireError: From<<C as Sink<PgWireBackendMessage>>::Error>,
238 {
239 match message {
240 PgWireFrontendMessage::Startup(ref startup) => {
241 if !client.is_secure() && self.force_tls {
243 send_error(
244 client,
245 PgErrorCode::Ec28000.to_err_info("No encryption".to_string()),
246 )
247 .await?;
248 return Ok(());
249 }
250
251 auth::save_startup_parameters_to_metadata(client, startup);
252
253 match resolve_db_info(client, self.query_handler.clone()).await? {
255 DbResolution::Resolved(catalog, schema) => {
256 let metadata = client.metadata_mut();
257 let _ = metadata.insert(super::METADATA_CATALOG.to_owned(), catalog);
258 let _ = metadata.insert(super::METADATA_SCHEMA.to_owned(), schema);
259 }
260 DbResolution::NotFound(msg) => {
261 send_error(client, PgErrorCode::Ec3D000.to_err_info(msg)).await?;
262 return Ok(());
263 }
264 }
265
266 if let Some(tz) = client.metadata().get("TimeZone") {
268 match Timezone::from_tz_string(tz) {
269 Ok(tz) => self.session.set_timezone(tz),
270 Err(_) => {
271 send_error(
272 client,
273 PgErrorCode::Ec22023
274 .to_err_info(format!("Invalid TimeZone: {}", tz)),
275 )
276 .await?;
277
278 return Ok(());
279 }
280 }
281 }
282
283 if self.login_verifier.user_provider.is_some() {
284 let login_info = LoginInfo::from_client_info(client);
285 let auth_info = match self.login_verifier.postgres_auth_info(&login_info).await
286 {
287 Ok(auth_info) => auth_info,
288 Err(_) => {
289 return send_password_authentication_failed(client).await;
290 }
291 };
292 client.set_state(PgWireConnectionState::AuthenticationInProgress);
293 match auth_info {
294 PgAuthInfo::ScramSha256 { .. } => {
295 *self.login_verifier.state.lock().await =
296 PgAuthenticationState::SaslInitial { auth_info };
297 client
298 .send(PgWireBackendMessage::Authentication(Authentication::SASL(
299 vec![SCRAM_SHA_256_METHOD.to_string()],
300 )))
301 .await?;
302 }
303 PgAuthInfo::Cleartext => {
304 *self.login_verifier.state.lock().await =
305 PgAuthenticationState::Cleartext;
306 client
307 .send(PgWireBackendMessage::Authentication(
308 Authentication::CleartextPassword,
309 ))
310 .await?;
311 }
312 }
313 } else {
314 self.session.set_user_info(userinfo_by_name(
315 client.metadata().get(super::METADATA_USER).cloned(),
316 ));
317 set_client_info(client, &self.session);
318 auth::finish_authentication(client, self.param_provider.as_ref()).await?;
319 }
320 }
321 PgWireFrontendMessage::PasswordMessageFamily(pwd) => {
322 let login_info = LoginInfo::from_client_info(client);
323 match self
324 .authenticate_password_message(client, &login_info, pwd)
325 .await?
326 {
327 PgAuthenticationResult::Continue => {}
328 PgAuthenticationResult::Success(user_info) => {
329 self.session.set_user_info(user_info);
330 set_client_info(client, &self.session);
331 auth::finish_authentication(client, self.param_provider.as_ref()).await?;
332 }
333 PgAuthenticationResult::Failed => {
334 return send_password_authentication_failed(client).await;
335 }
336 }
337 }
338 _ => {}
339 }
340 Ok(())
341 }
342}
343
344enum PgAuthenticationResult {
345 Continue,
346 Success(UserInfoRef),
347 Failed,
348}
349
350impl PostgresServerHandlerInner {
351 fn record_scram_failure(result: PgAuthenticationResult) -> PgAuthenticationResult {
355 if matches!(result, PgAuthenticationResult::Failed) {
356 METRIC_AUTH_FAILURE
357 .with_label_values(&[StatusCode::UserPasswordMismatch.as_ref()])
358 .inc();
359 }
360 result
361 }
362
363 async fn authenticate_password_message<C>(
364 &self,
365 client: &mut C,
366 login_info: &LoginInfo,
367 pwd: PasswordMessageFamily,
368 ) -> PgWireResult<PgAuthenticationResult>
369 where
370 C: ClientInfo + Sink<PgWireBackendMessage> + Unpin + Send,
371 C::Error: Debug,
372 PgWireError: From<<C as Sink<PgWireBackendMessage>>::Error>,
373 {
374 let mut state = self.login_verifier.state.lock().await;
375 let current_state = std::mem::replace(&mut *state, PgAuthenticationState::Initial);
376
377 match current_state {
378 PgAuthenticationState::Cleartext => {
379 let pwd = pwd.into_password()?;
380 drop(state);
381 match self.login_verifier.auth(login_info, &pwd.password).await {
382 Ok(Some(user_info)) => Ok(PgAuthenticationResult::Success(user_info)),
383 Ok(None) | Err(_) => Ok(PgAuthenticationResult::Failed),
384 }
385 }
386 PgAuthenticationState::SaslInitial { auth_info } => {
387 let result = 'sasl: {
392 let sasl_initial = pwd.into_sasl_initial_response()?;
393 if sasl_initial.auth_method != SCRAM_SHA_256_METHOD {
394 break 'sasl PgAuthenticationResult::Failed;
395 }
396 let Some(data) = sasl_initial.data else {
397 break 'sasl PgAuthenticationResult::Failed;
398 };
399 let Some(client_first) = ScramClientFirst::parse(&data) else {
400 break 'sasl PgAuthenticationResult::Failed;
401 };
402 let PgAuthInfo::ScramSha256 {
403 verifier,
404 user_info,
405 } = auth_info
406 else {
407 break 'sasl PgAuthenticationResult::Failed;
408 };
409
410 let server_nonce = BASE64.encode(rand::random::<[u8; 18]>());
411 let nonce = format!("{}{}", client_first.nonce, server_nonce);
412 let server_first = format!(
413 "r={},s={},i={}",
414 nonce,
415 BASE64.encode(verifier.salt()),
416 verifier.iterations()
417 );
418 client
419 .send(PgWireBackendMessage::Authentication(
420 Authentication::SASLContinue(server_first.clone().into()),
421 ))
422 .await?;
423 *state = PgAuthenticationState::SaslFinal {
424 verifier,
425 user_info,
426 channel_binding: client_first.channel_binding,
427 nonce,
428 client_first_bare: client_first.bare,
429 server_first,
430 };
431 PgAuthenticationResult::Continue
432 };
433 Ok(Self::record_scram_failure(result))
434 }
435 PgAuthenticationState::SaslFinal {
436 verifier,
437 user_info,
438 channel_binding,
439 nonce,
440 client_first_bare,
441 server_first,
442 } => {
443 let result = 'sasl: {
444 let sasl_response = pwd.into_sasl_response()?;
445 let Some(client_final) = ScramClientFinal::parse(&sasl_response.data) else {
446 break 'sasl PgAuthenticationResult::Failed;
447 };
448 if client_final.channel_binding != channel_binding {
449 break 'sasl PgAuthenticationResult::Failed;
450 }
451 if client_final.nonce != nonce {
452 break 'sasl PgAuthenticationResult::Failed;
453 }
454
455 let auth_message = format!(
456 "{},{},{}",
457 client_first_bare, server_first, client_final.without_proof
458 );
459 let Ok(client_proof) = BASE64.decode(client_final.proof.as_bytes()) else {
460 break 'sasl PgAuthenticationResult::Failed;
461 };
462 let Ok(Some(server_signature)) =
463 verifier.verify_client_proof(auth_message.as_bytes(), &client_proof)
464 else {
465 break 'sasl PgAuthenticationResult::Failed;
466 };
467 let Some(user_info) = user_info else {
468 break 'sasl PgAuthenticationResult::Failed;
469 };
470
471 drop(state);
472 if self
473 .login_verifier
474 .authorize(login_info, &user_info)
475 .await
476 .is_err()
477 {
478 return Ok(PgAuthenticationResult::Failed);
481 }
482
483 client
484 .send(PgWireBackendMessage::Authentication(
485 Authentication::SASLFinal(
486 format!("v={}", BASE64.encode(server_signature)).into(),
487 ),
488 ))
489 .await?;
490 PgAuthenticationResult::Success(user_info)
491 };
492 Ok(Self::record_scram_failure(result))
493 }
494 PgAuthenticationState::Initial => Ok(PgAuthenticationResult::Failed),
495 }
496 }
497}
498
499struct ScramClientFirst {
500 channel_binding: String,
501 bare: String,
502 nonce: String,
503}
504
505impl ScramClientFirst {
506 fn parse(data: &[u8]) -> Option<Self> {
507 let message = std::str::from_utf8(data).ok()?;
508 let mut parts = message.splitn(3, ',');
509 let cbind = parts.next()?;
510 if !matches!(cbind, "n" | "y") {
511 return None;
512 }
513 let authzid = parts.next()?;
514 if !authzid.is_empty() {
515 return None;
516 }
517 let channel_binding = BASE64.encode(format!("{cbind},,"));
518 let bare = parts.next()?.to_string();
519 let nonce = bare
520 .split(',')
521 .find_map(|chunk| chunk.strip_prefix("r="))?
522 .to_string();
523 Some(Self {
524 channel_binding,
525 bare,
526 nonce,
527 })
528 }
529}
530
531struct ScramClientFinal {
532 channel_binding: String,
533 nonce: String,
534 without_proof: String,
535 proof: String,
536}
537
538impl ScramClientFinal {
539 fn parse(data: &[u8]) -> Option<Self> {
540 let message = std::str::from_utf8(data).ok()?;
541 let proof_pos = message.rfind(",p=")?;
542 let without_proof = message[..proof_pos].to_string();
543 let proof = message[proof_pos + 3..].to_string();
544 let channel_binding = without_proof
545 .split(',')
546 .find_map(|chunk| chunk.strip_prefix("c="))?;
547 let nonce = without_proof
548 .split(',')
549 .find_map(|chunk| chunk.strip_prefix("r="))?;
550 if nonce.is_empty() || proof.is_empty() {
551 return None;
552 }
553
554 Some(Self {
555 channel_binding: channel_binding.to_string(),
556 nonce: nonce.to_string(),
557 without_proof,
558 proof,
559 })
560 }
561}
562
563async fn send_error<C>(client: &mut C, err_info: ErrorInfo) -> PgWireResult<()>
564where
565 C: ClientInfo + Sink<PgWireBackendMessage> + Unpin + Send,
566 C::Error: Debug,
567 PgWireError: From<<C as Sink<PgWireBackendMessage>>::Error>,
568{
569 let error = ErrorResponse::from(err_info);
570 client
571 .feed(PgWireBackendMessage::ErrorResponse(error))
572 .await?;
573 client.close().await?;
574 Ok(())
575}
576
577async fn send_password_authentication_failed<C>(client: &mut C) -> PgWireResult<()>
578where
579 C: ClientInfo + Sink<PgWireBackendMessage> + Unpin + Send,
580 C::Error: Debug,
581 PgWireError: From<<C as Sink<PgWireBackendMessage>>::Error>,
582{
583 send_error(
584 client,
585 PgErrorCode::Ec28P01.to_err_info("password authentication failed".to_string()),
586 )
587 .await
588}
589
590enum DbResolution {
591 Resolved(String, String),
592 NotFound(String),
593}
594
595async fn resolve_db_info<C>(
597 client: &mut C,
598 query_handler: ServerSqlQueryHandlerRef,
599) -> PgWireResult<DbResolution>
600where
601 C: ClientInfo + Unpin + Send,
602{
603 let db_ref = client.metadata().get(super::METADATA_DATABASE);
604 if let Some(db) = db_ref {
605 let (catalog, schema) = parse_catalog_and_schema_from_db_string(db);
606 if query_handler
607 .is_valid_schema(&catalog, &schema)
608 .await
609 .map_err(convert_err)?
610 {
611 Ok(DbResolution::Resolved(catalog, schema))
612 } else {
613 Ok(DbResolution::NotFound(format!("Database not found: {db}")))
614 }
615 } else {
616 Ok(DbResolution::NotFound("Database not specified".to_owned()))
617 }
618}
619
620#[cfg(test)]
621mod tests {
622 use super::*;
623
624 #[test]
625 fn test_scram_client_first_parse() {
626 let message = b"n,,n=greptime,r=clientnonce";
627 let client_first = ScramClientFirst::parse(message).unwrap();
628 assert_eq!("biws", client_first.channel_binding);
629 assert_eq!("n=greptime,r=clientnonce", client_first.bare);
630 assert_eq!("clientnonce", client_first.nonce);
631
632 let message = b"y,,n=greptime,r=clientnonce";
633 let client_first = ScramClientFirst::parse(message).unwrap();
634 assert_eq!("eSws", client_first.channel_binding);
635 assert_eq!("n=greptime,r=clientnonce", client_first.bare);
636 assert_eq!("clientnonce", client_first.nonce);
637
638 assert!(ScramClientFirst::parse(b"p=tls-server-end-point,,n=greptime,r=nonce").is_none());
639 assert!(ScramClientFirst::parse(b"n,a=authzid,n=greptime,r=nonce").is_none());
640 }
641
642 #[test]
643 fn test_scram_client_final_parse() {
644 let message = b"c=biws,r=clientnonceservernonce,p=dGVzdA==";
645 let client_final = ScramClientFinal::parse(message).unwrap();
646 assert_eq!("biws", client_final.channel_binding);
647 assert_eq!("clientnonceservernonce", client_final.nonce);
648 assert_eq!(
649 "c=biws,r=clientnonceservernonce",
650 client_final.without_proof
651 );
652 assert_eq!("dGVzdA==", client_final.proof);
653
654 assert!(ScramClientFinal::parse(b"r=nonce,p=dGVzdA==").is_none());
655 assert!(ScramClientFinal::parse(b"c=biws,r=nonce").is_none());
656 assert!(ScramClientFinal::parse(b"c=biws,r=,p=dGVzdA==").is_none());
657 }
658}