Skip to main content

servers/postgres/
auth_handler.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::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    // pass generated process id and secret key to client, this information will
221    // be sent to postgres client for query cancellation.
222    // use all 0 before we actually supported query cancellation
223    client.set_pid_and_secret_key(0, SecretKey::I32(0));
224    // set userinfo outside
225}
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                // check ssl requirement
242                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                // check if db is valid
254                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                // try to set TimeZone
267                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    /// Records a rejected SCRAM attempt in [`METRIC_AUTH_FAILURE`]. The label is
352    /// intentionally uniform (never `UserNotFound`), so the counter cannot be
353    /// used to distinguish a wrong password from an unknown user.
354    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                // A labeled block funnels every rejection through a single
388                // `record_scram_failure`, so failed SCRAM attempts still show up
389                // in `METRIC_AUTH_FAILURE` without sprinkling the metric call
390                // across each early return.
391                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                        // `authorize` already recorded this failure; return early
479                        // to bypass `record_scram_failure` and avoid double-counting.
480                        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
595/// A function extracted to resolve lifetime and readability issues:
596async 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}