Skip to main content

pgwire/
pg_protocol.rs

1// Copyright 2022 RisingWave Labs
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::any::Any;
16use std::collections::HashMap;
17use std::io::ErrorKind;
18use std::panic::AssertUnwindSafe;
19use std::pin::Pin;
20use std::str::Utf8Error;
21use std::sync::{Arc, LazyLock, Weak};
22use std::time::{Duration, Instant};
23use std::{io, str};
24
25use bytes::{Bytes, BytesMut};
26use futures::FutureExt;
27use futures::stream::StreamExt;
28use itertools::Itertools;
29use openssl::ssl::{SslAcceptor, SslContext, SslContextRef, SslMethod};
30use risingwave_common::types::DataType;
31use risingwave_common::util::deployment::Deployment;
32use risingwave_common::util::env_var::env_var_is_true;
33use risingwave_common::util::panic::FutureCatchUnwindExt;
34use risingwave_common::util::query_log::*;
35use risingwave_common::{PG_VERSION, SERVER_ENCODING, STANDARD_CONFORMING_STRINGS};
36use risingwave_sqlparser::ast::{RedactSqlOptionKeywordsRef, Statement};
37use risingwave_sqlparser::parser::Parser;
38use thiserror_ext::AsReport;
39use tokio::io::{AsyncRead, AsyncWrite, AsyncWriteExt};
40use tokio::sync::Mutex;
41use tokio_openssl::SslStream;
42use tracing::Instrument;
43
44use crate::error::{PsqlError, PsqlResult};
45use crate::error_or_notice::Severity;
46use crate::memory_manager::{MessageMemoryGuard, MessageMemoryManagerRef};
47use crate::net::AddressRef;
48use crate::pg_extended::ResultCache;
49use crate::pg_message::{
50    BeCommandCompleteMessage, BeMessage, BeParameterStatusMessage, FeBindMessage, FeCancelMessage,
51    FeCloseMessage, FeDescribeMessage, FeExecuteMessage, FeMessage, FeMessageHeader,
52    FeParseMessage, FePasswordMessage, FeStartupMessage, ServerThrottleReason, TransactionStatus,
53};
54use crate::pg_server::{Session, SessionManager, UserAuthenticator};
55use crate::types::Format;
56
57/// Truncates query log if it's longer than `RW_QUERY_LOG_TRUNCATE_LEN`, to avoid log file being too
58/// large.
59static RW_QUERY_LOG_TRUNCATE_LEN: LazyLock<usize> =
60    LazyLock::new(|| match std::env::var("RW_QUERY_LOG_TRUNCATE_LEN") {
61        Ok(len) if len.parse::<usize>().is_ok() => len.parse::<usize>().unwrap(),
62        _ => 65536,
63    });
64
65tokio::task_local! {
66    /// The current session. Concrete type is erased for different session implementations.
67    pub static CURRENT_SESSION: Weak<dyn Any + Send + Sync>
68}
69
70/// The state machine for each psql connection.
71/// Read pg messages from tcp stream and write results back.
72pub struct PgProtocol<S, SM>
73where
74    SM: SessionManager,
75{
76    /// Used for write/read pg messages.
77    stream: PgStream<S>,
78    /// Current states of pg connection.
79    state: PgProtocolState,
80    /// Whether the connection is terminated.
81    is_terminate: bool,
82
83    session_mgr: Arc<SM>,
84    session: Option<Arc<SM::Session>>,
85
86    result_cache: HashMap<String, ResultCache<<SM::Session as Session>::ValuesStream>>,
87    unnamed_prepare_statement:
88        Option<PreparedStatementData<<SM::Session as Session>::PreparedStatement>>,
89    prepare_statement_store:
90        HashMap<String, PreparedStatementData<<SM::Session as Session>::PreparedStatement>>,
91    unnamed_portal: Option<PortalData<<SM::Session as Session>::Portal>>,
92    portal_store: HashMap<String, PortalData<<SM::Session as Session>::Portal>>,
93    // Used to store the dependency of portal and prepare statement.
94    // When we close a prepare statement, we need to close all the portals that depend on it.
95    statement_portal_dependency: HashMap<String, Vec<String>>,
96
97    // Used for ssl connection.
98    // If None, not expected to build ssl connection (panic).
99    tls_context: Option<SslContext>,
100
101    // TLS configuration including SSL enforcement setting
102    tls_config: Option<TlsConfig>,
103
104    // Used in extended query protocol. When encounter error in extended query, we need to ignore
105    // the following message util sync message.
106    ignore_util_sync: bool,
107
108    // Client Address
109    peer_addr: AddressRef,
110
111    redact_sql_option_keywords: Option<RedactSqlOptionKeywordsRef>,
112    message_memory_manager: MessageMemoryManagerRef,
113}
114
115/// Configures TLS encryption for connections.
116#[derive(Debug, Clone)]
117pub struct TlsConfig {
118    /// The path to the TLS certificate.
119    pub cert: String,
120    /// The path to the TLS key.
121    pub key: String,
122    /// Whether to enforce SSL connections (reject non-SSL clients).
123    pub enforce_ssl: bool,
124}
125
126impl TlsConfig {
127    pub fn new_default() -> anyhow::Result<Option<Self>> {
128        let cert = std::env::var("RW_SSL_CERT").ok();
129        let key = std::env::var("RW_SSL_KEY").ok();
130        let enforce_ssl = env_var_is_true("RW_SSL_ENFORCE");
131
132        if cert.is_some() ^ key.is_some() {
133            return Err(anyhow::anyhow!(
134                "RW_SSL_CERT and RW_SSL_KEY must be set together"
135            ));
136        }
137
138        if enforce_ssl && cert.is_none() {
139            return Err(anyhow::anyhow!(
140                "RW_SSL_ENFORCE requires RW_SSL_CERT and RW_SSL_KEY to be set"
141            ));
142        }
143
144        let (Some(cert), Some(key)) = (cert, key) else {
145            return Ok(None);
146        };
147
148        tracing::info!(
149            "RW_SSL_CERT={}, RW_SSL_KEY={}, RW_SSL_ENFORCE={}",
150            cert,
151            key,
152            enforce_ssl
153        );
154        Ok(Some(Self {
155            cert,
156            key,
157            enforce_ssl,
158        }))
159    }
160}
161
162impl<S, SM> Drop for PgProtocol<S, SM>
163where
164    SM: SessionManager,
165{
166    fn drop(&mut self) {
167        if let Some(session) = &self.session {
168            // Clear the session in session manager.
169            self.session_mgr.end_session(session);
170        }
171    }
172}
173
174/// States flow happened from top to down.
175enum PgProtocolState {
176    Startup,
177    Regular,
178}
179
180#[derive(Clone)]
181struct PreparedStatementData<S> {
182    statement: S,
183    sql: Arc<str>,
184}
185
186#[derive(Clone)]
187struct PortalData<P> {
188    portal: P,
189    sql: Arc<str>,
190}
191
192/// Truncate 0 from C string in Bytes and stringify it (returns slice, no allocations).
193///
194/// PG protocol strings are always C strings.
195pub fn cstr_to_str(b: &Bytes) -> Result<&str, Utf8Error> {
196    let without_null = if b.last() == Some(&0) {
197        &b[..b.len() - 1]
198    } else {
199        &b[..]
200    };
201    std::str::from_utf8(without_null)
202}
203
204fn get_redacted_and_truncated_sql(
205    sql: &str,
206    redact_sql_option_keywords: Option<RedactSqlOptionKeywordsRef>,
207) -> String {
208    let redacted_sql = if let Some(keywords) = redact_sql_option_keywords
209        && !keywords.is_empty()
210    {
211        redact_sql(sql, keywords)
212    } else {
213        sql.to_owned()
214    };
215    let truncated = truncated_fmt::TruncatedFmt(&redacted_sql, *RW_QUERY_LOG_TRUNCATE_LEN);
216    truncated.to_string()
217}
218
219/// Record `sql` in the current tracing span.
220fn record_sql_in_span(
221    sql: &str,
222    redact_sql_option_keywords: Option<RedactSqlOptionKeywordsRef>,
223    span: &mut tracing::Span,
224) {
225    let redacted_and_truncated_sql =
226        get_redacted_and_truncated_sql(sql, redact_sql_option_keywords);
227    span.record("sql", tracing::field::display(&redacted_and_truncated_sql));
228}
229
230fn record_user_in_span(user: &str, span: &mut tracing::Span) {
231    span.record("user", tracing::field::display(user));
232}
233
234/// Redacts sensitive SQL fields. Data in DML is not redacted.
235fn redact_sql(sql: &str, keywords: RedactSqlOptionKeywordsRef) -> String {
236    match Parser::parse_sql(sql) {
237        Ok(sqls) => sqls
238            .into_iter()
239            .map(|sql| sql.to_redacted_string(keywords.clone()))
240            .join(";"),
241        Err(_) => sql.to_owned(),
242    }
243}
244
245#[derive(Clone)]
246pub struct ConnectionContext {
247    pub tls_config: Option<TlsConfig>,
248    pub redact_sql_option_keywords: Option<RedactSqlOptionKeywordsRef>,
249    pub message_memory_manager: MessageMemoryManagerRef,
250    pub stream_flush_threshold_bytes: usize,
251}
252
253impl<S, SM> PgProtocol<S, SM>
254where
255    S: PgByteStream,
256    SM: SessionManager,
257{
258    pub fn new(
259        stream: S,
260        session_mgr: Arc<SM>,
261        peer_addr: AddressRef,
262        context: ConnectionContext,
263    ) -> Self {
264        let ConnectionContext {
265            tls_config,
266            redact_sql_option_keywords,
267            message_memory_manager,
268            stream_flush_threshold_bytes,
269        } = context;
270        Self {
271            stream: PgStream::new(stream, stream_flush_threshold_bytes),
272            is_terminate: false,
273            state: PgProtocolState::Startup,
274            session_mgr,
275            session: None,
276            tls_context: tls_config
277                .as_ref()
278                .and_then(|e| build_ssl_ctx_from_config(e).ok()),
279            tls_config,
280            result_cache: Default::default(),
281            unnamed_prepare_statement: Default::default(),
282            prepare_statement_store: Default::default(),
283            unnamed_portal: Default::default(),
284            portal_store: Default::default(),
285            statement_portal_dependency: Default::default(),
286            ignore_util_sync: false,
287            peer_addr,
288            redact_sql_option_keywords,
289            message_memory_manager,
290        }
291    }
292
293    /// Run the protocol to serve the connection.
294    pub async fn run(&mut self) {
295        let mut notice_fut = None;
296
297        loop {
298            // Once a session is present, create a future to subscribe and send notices asynchronously.
299            if notice_fut.is_none()
300                && let Some(session) = self.session.clone()
301            {
302                let mut stream = self.stream.clone();
303                notice_fut = Some(Box::pin(async move {
304                    loop {
305                        let notice = session.next_notice().await;
306                        if let Err(e) = stream.write(BeMessage::NoticeResponse(&notice)).await {
307                            tracing::error!(error = %e.as_report(), notice, "failed to send notice");
308                        }
309                    }
310                }));
311            }
312
313            // Read and process messages.
314            let process = std::pin::pin!(async {
315                let (msg, _memory_guard) = match self.read_message().await {
316                    Ok(msg) => msg,
317                    Err(e) => {
318                        tracing::error!(error = %e.as_report(), "error when reading message");
319                        return true; // terminate the connection
320                    }
321                };
322                tracing::trace!(?msg, "received message");
323                self.process(msg).await
324            });
325
326            let terminated = if let Some(notice_fut) = notice_fut.as_mut() {
327                tokio::select! {
328                    _ = notice_fut => unreachable!(),
329                    terminated = process => terminated,
330                }
331            } else {
332                process.await
333            };
334
335            if terminated {
336                break;
337            }
338        }
339    }
340
341    /// Processes one message. Returns true if the connection is terminated.
342    pub async fn process(&mut self, msg: FeMessage) -> bool {
343        self.do_process(msg).await.is_none() || self.is_terminate
344    }
345
346    /// The root tracing span for processing a message. The target of the span is
347    /// [`PGWIRE_ROOT_SPAN_TARGET`].
348    ///
349    /// This is used to provide context for the (slow) query logs and traces.
350    ///
351    /// The span is only effective if there's a current session and the message is
352    /// query-related. Otherwise, `Span::none()` is returned.
353    fn root_span_for_msg(&self, msg: &FeMessage) -> tracing::Span {
354        let Some(session_id) = self.session.as_ref().map(|s| s.id().0) else {
355            return tracing::Span::none();
356        };
357
358        let mode = match msg {
359            FeMessage::Query(_) => "simple query",
360            FeMessage::Parse(_) => "extended query parse",
361            FeMessage::Execute(_) => "extended query execute",
362            _ => return tracing::Span::none(),
363        };
364
365        let mut span = tracing::info_span!(
366            target: PGWIRE_ROOT_SPAN_TARGET,
367            "handle_query",
368            mode,
369            session_id,
370            sql = tracing::field::Empty,
371            user = tracing::field::Empty,
372        );
373        match msg {
374            FeMessage::Execute(execute_msg) => {
375                if let Ok(portal_name) = cstr_to_str(&execute_msg.portal_name)
376                    && let Ok(sql) = self.get_portal_sql(portal_name)
377                {
378                    record_sql_in_span(&sql, self.redact_sql_option_keywords.clone(), &mut span);
379                }
380            }
381            _ => {
382                if let Ok(sql) = msg.get_sql()
383                    && let Some(sql) = sql
384                {
385                    record_sql_in_span(sql, self.redact_sql_option_keywords.clone(), &mut span);
386                }
387            }
388        }
389        if let Some(current_session) = self.session.as_ref() {
390            record_user_in_span(&current_session.user(), &mut span);
391        }
392        span
393    }
394
395    /// Return type `Option<()>` is essentially a bool, but allows `?` for early return.
396    /// - `None` means to terminate the current connection
397    /// - `Some(())` means to continue processing the next message
398    async fn do_process(&mut self, msg: FeMessage) -> Option<()> {
399        let span = self.root_span_for_msg(&msg);
400        let weak_session = self
401            .session
402            .as_ref()
403            .map(|s| Arc::downgrade(s) as Weak<dyn Any + Send + Sync>);
404
405        // Processing the message itself.
406        //
407        // Note: pin the future to avoid stack overflow as we'll wrap it multiple times
408        // in the following code.
409        let fut = Box::pin(self.do_process_inner(msg));
410
411        // Set the current session as the context when processing the message, if exists.
412        let fut = async move {
413            if let Some(session) = weak_session {
414                CURRENT_SESSION.scope(session, fut).await
415            } else {
416                fut.await
417            }
418        };
419
420        // Catch unwind.
421        let fut = async move {
422            AssertUnwindSafe(fut)
423                .rw_catch_unwind()
424                .await
425                .unwrap_or_else(|payload| {
426                    Err(PsqlError::Panic(
427                        panic_message::panic_message(&payload).to_owned(),
428                    ))
429                })
430        };
431
432        // Slow query log.
433        let fut = async move {
434            let period = *SLOW_QUERY_LOG_PERIOD;
435            let mut fut = std::pin::pin!(fut);
436            let mut elapsed = Duration::ZERO;
437
438            // Report the SQL in the log periodically if the query is slow.
439            loop {
440                match tokio::time::timeout(period, &mut fut).await {
441                    Ok(result) => break result,
442                    Err(_) => {
443                        elapsed += period;
444                        tracing::info!(
445                            target: PGWIRE_SLOW_QUERY_LOG,
446                            elapsed = %format_args!("{}ms", elapsed.as_millis()),
447                            "slow query"
448                        );
449                    }
450                }
451            }
452        };
453
454        // Query log.
455        let fut = async move {
456            if !tracing::Span::current().is_none() {
457                tracing::info!(
458                    target: PGWIRE_QUERY_LOG,
459                    status = "started",
460                );
461            }
462
463            let start = Instant::now();
464            let result = fut.await;
465            let elapsed = start.elapsed();
466
467            // Always log if an error occurs.
468            // Note: all messages will be processed through this code path, making it the
469            //       only necessary place to log errors.
470            if let Err(error) = &result {
471                if cfg!(debug_assertions) && !Deployment::current().is_ci() {
472                    // For local debugging, we print the error with backtrace.
473                    // It's useful only when:
474                    // - no additional context is added to the error
475                    // - backtrace is captured in the error
476                    // - backtrace is not printed in the middle
477                    tracing::error!(error = ?error.as_report(), "error when process message");
478                } else {
479                    tracing::error!(error = %error.as_report(), "error when process message");
480                }
481            }
482
483            // Log to optionally-enabled target `PGWIRE_QUERY_LOG`.
484            // Only log if we're currently in a tracing span set in `span_for_msg`.
485            if !tracing::Span::current().is_none() {
486                tracing::info!(
487                    target: PGWIRE_QUERY_LOG,
488                    status = if result.is_ok() { "ok" } else { "err" },
489                    time = %format_args!("{}ms", elapsed.as_millis()),
490                );
491            }
492
493            result
494        };
495
496        // Tracing span.
497        let fut = fut.instrument(span);
498
499        // Execute the future and handle the error.
500        match fut.await {
501            Ok(()) => Some(()),
502            Err(e) => {
503                match e {
504                    PsqlError::IoError(io_err) => {
505                        if io_err.kind() == std::io::ErrorKind::UnexpectedEof {
506                            return None;
507                        }
508                    }
509
510                    PsqlError::SslError(_) => {
511                        // For ssl error, because the stream has already been consumed, so there is
512                        // no way to write more message.
513                        return None;
514                    }
515
516                    PsqlError::StartupError(_) | PsqlError::PasswordError => {
517                        self.stream
518                            .write_no_flush(BeMessage::ErrorResponse {
519                                error: &e,
520                                // At this time we're not in a session, use compact error message for
521                                // better alignment with Postgres' UI.
522                                pretty: false,
523                                severity: Some(Severity::Fatal),
524                            })
525                            .ok()?;
526                        let _ = self.stream.flush().await;
527                        return None;
528                    }
529
530                    PsqlError::SimpleQueryError(_) | PsqlError::ServerThrottle(_) => {
531                        self.stream
532                            .write_no_flush(BeMessage::ErrorResponse {
533                                error: &e,
534                                pretty: true,
535                                severity: None,
536                            })
537                            .ok()?;
538                        self.ready_for_query().ok()?;
539                    }
540
541                    PsqlError::IdleInTxnTimeout | PsqlError::Panic(_) => {
542                        self.stream
543                            .write_no_flush(BeMessage::ErrorResponse {
544                                error: &e,
545                                pretty: true,
546                                severity: None,
547                            })
548                            .ok()?;
549                        let _ = self.stream.flush().await;
550
551                        // 1. Catching the panic during message processing may leave the session in an
552                        // inconsistent state. We forcefully close the connection (then end the
553                        // session) here for safety.
554                        // 2. Idle in transaction timeout should also close the connection.
555                        return None;
556                    }
557
558                    PsqlError::Uncategorized(_)
559                    | PsqlError::ExtendedPrepareError(_)
560                    | PsqlError::ExtendedExecuteError(_) => {
561                        self.stream
562                            .write_no_flush(BeMessage::ErrorResponse {
563                                error: &e,
564                                pretty: true,
565                                severity: None,
566                            })
567                            .ok()?;
568                    }
569                }
570                let _ = self.stream.flush().await;
571                Some(())
572            }
573        }
574    }
575
576    async fn do_process_inner(&mut self, msg: FeMessage) -> PsqlResult<()> {
577        // Ignore util sync message.
578        if self.ignore_util_sync {
579            if let FeMessage::Sync = msg {
580            } else {
581                tracing::trace!("ignore message {:?} until sync.", msg);
582                return Ok(());
583            }
584        }
585
586        match msg {
587            FeMessage::Gss => self.process_gss_msg().await?,
588            FeMessage::Ssl => self.process_ssl_msg().await?,
589            FeMessage::Startup(msg) => self.process_startup_msg(msg).await?,
590            FeMessage::Password(msg) => self.process_password_msg(msg).await?,
591            FeMessage::Query(query_msg) => {
592                let sql = Arc::from(query_msg.get_sql()?);
593                // The process_query_msg can be slow. Release potential large FeQueryMessage early.
594                drop(query_msg);
595                self.process_query_msg(sql).await?
596            }
597            FeMessage::CancelQuery(m) => self.process_cancel_msg(m)?,
598            FeMessage::Terminate => self.process_terminate(),
599            FeMessage::Parse(m) => {
600                if let Err(err) = self.process_parse_msg(m).await {
601                    self.ignore_util_sync = true;
602                    return Err(err);
603                }
604            }
605            FeMessage::Bind(m) => {
606                if let Err(err) = self.process_bind_msg(m) {
607                    self.ignore_util_sync = true;
608                    return Err(err);
609                }
610            }
611            FeMessage::Execute(m) => {
612                if let Err(err) = self.process_execute_msg(m).await {
613                    self.ignore_util_sync = true;
614                    return Err(err);
615                }
616            }
617            FeMessage::Describe(m) => {
618                if let Err(err) = self.process_describe_msg(m) {
619                    self.ignore_util_sync = true;
620                    return Err(err);
621                }
622            }
623            FeMessage::Sync => {
624                self.ignore_util_sync = false;
625                self.ready_for_query()?
626            }
627            FeMessage::Close(m) => {
628                if let Err(err) = self.process_close_msg(m) {
629                    self.ignore_util_sync = true;
630                    return Err(err);
631                }
632            }
633            FeMessage::Flush => {
634                if let Err(err) = self.stream.flush().await {
635                    self.ignore_util_sync = true;
636                    return Err(err.into());
637                }
638            }
639            FeMessage::HealthCheck => self.process_health_check(),
640            FeMessage::ServerThrottle(reason) => match reason {
641                ServerThrottleReason::TooLargeMessage => {
642                    return Err(PsqlError::ServerThrottle(format!(
643                        "max_single_query_size_bytes {} has been exceeded, please either reduce the query size or increase the limit",
644                        self.message_memory_manager.max_filter_bytes
645                    )));
646                }
647                ServerThrottleReason::TooManyMemoryUsage => {
648                    return Err(PsqlError::ServerThrottle(format!(
649                        "max_total_query_size_bytes {} has been exceeded, please either retry or increase the limit",
650                        self.message_memory_manager.max_running_bytes
651                    )));
652                }
653            },
654        }
655        self.stream.flush().await?;
656        Ok(())
657    }
658
659    pub async fn read_message(&mut self) -> io::Result<(FeMessage, Option<MessageMemoryGuard>)> {
660        match self.state {
661            PgProtocolState::Startup => self
662                .stream
663                .read_startup()
664                .await
665                .map(|message: FeMessage| (message, None)),
666            PgProtocolState::Regular => {
667                self.stream.read_header().await?;
668                let guard = if let Some(ref header) = self.stream.read_header {
669                    let payload_len = std::cmp::max(header.payload_len, 0) as u64;
670                    let (reason, guard) = self.message_memory_manager.add(payload_len);
671                    if let Some(reason) = reason {
672                        // Release the memory ASAP.
673                        drop(guard);
674                        self.stream.skip_body().await?;
675                        return Ok((FeMessage::ServerThrottle(reason), None));
676                    }
677                    guard
678                } else {
679                    None
680                };
681                let message = self.stream.read_body().await?;
682                Ok((message, guard))
683            }
684        }
685    }
686
687    /// Writes a `ReadyForQuery` message to the client without flushing.
688    fn ready_for_query(&mut self) -> io::Result<()> {
689        self.stream.write_no_flush(BeMessage::ReadyForQuery(
690            self.session
691                .as_ref()
692                .map(|s| s.transaction_status())
693                .unwrap_or(TransactionStatus::Idle),
694        ))
695    }
696
697    async fn process_gss_msg(&mut self) -> PsqlResult<()> {
698        // We don't support GSSAPI, so we just say no gracefully.
699        self.stream.write(BeMessage::EncryptionResponseNo).await?;
700        Ok(())
701    }
702
703    async fn process_ssl_msg(&mut self) -> PsqlResult<()> {
704        if let Some(context) = self.tls_context.as_ref() {
705            // If got and ssl context, say yes for ssl connection.
706            // Construct ssl stream and replace with current one.
707            self.stream.write(BeMessage::EncryptionResponseSsl).await?;
708            self.stream.upgrade_to_ssl(context).await?;
709        } else {
710            // If no, say no for encryption.
711            self.stream.write(BeMessage::EncryptionResponseNo).await?;
712        }
713
714        Ok(())
715    }
716
717    async fn process_startup_msg(&mut self, msg: FeStartupMessage) -> PsqlResult<()> {
718        // Check SSL enforcement: if SSL is enforced but connection is not using SSL, reject
719        if let Some(ref tls_config) = self.tls_config
720            && tls_config.enforce_ssl
721            && !self.stream.is_ssl_connection().await
722        {
723            return Err(PsqlError::StartupError(
724                "SSL connection is required but not established".into(),
725            ));
726        }
727
728        let db_name = msg
729            .config
730            .get("database")
731            .cloned()
732            .unwrap_or_else(|| "dev".to_owned());
733        let user_name = msg
734            .config
735            .get("user")
736            .cloned()
737            .unwrap_or_else(|| "root".to_owned());
738
739        let session = self
740            .session_mgr
741            .connect(&db_name, &user_name, self.peer_addr.clone())
742            .map_err(|e| PsqlError::StartupError(e.into()))?;
743
744        if let Some(options) = msg.config.get("options") {
745            for (key, value) in parse_options(options)? {
746                session
747                    .set_config(&key, value)
748                    .map_err(|e| PsqlError::StartupError(e.into()))?;
749            }
750        }
751        // dedicated `application_name` has higher priority than `options`
752        let application_name = msg.config.get("application_name");
753        if let Some(application_name) = application_name {
754            session
755                .set_config("application_name", application_name.clone())
756                .map_err(|e| PsqlError::StartupError(e.into()))?;
757        }
758
759        match session.user_authenticator() {
760            UserAuthenticator::None => {
761                self.stream.write_no_flush(BeMessage::AuthenticationOk)?;
762
763                // Cancel request need this for identify and verification. According to postgres
764                // doc, it should be written to buffer after receive AuthenticationOk.
765                self.stream
766                    .write_no_flush(BeMessage::BackendKeyData(session.id()))?;
767
768                self.stream.write_no_flush(BeMessage::ParameterStatus(
769                    BeParameterStatusMessage::TimeZone(
770                        &session
771                            .get_config("timezone")
772                            .map_err(|e| PsqlError::StartupError(e.into()))?,
773                    ),
774                ))?;
775                self.stream
776                    .write_parameter_status_msg_no_flush(&ParameterStatus {
777                        application_name: application_name.cloned(),
778                    })?;
779                self.ready_for_query()?;
780            }
781            UserAuthenticator::ClearText(_)
782            | UserAuthenticator::OAuth { .. }
783            | UserAuthenticator::Ldap(..) => {
784                self.stream
785                    .write_no_flush(BeMessage::AuthenticationCleartextPassword)?;
786            }
787            UserAuthenticator::Md5WithSalt { salt, .. } => {
788                self.stream
789                    .write_no_flush(BeMessage::AuthenticationMd5Password(salt))?;
790            }
791        }
792
793        self.session = Some(session);
794        self.state = PgProtocolState::Regular;
795        Ok(())
796    }
797
798    async fn process_password_msg(&mut self, msg: FePasswordMessage) -> PsqlResult<()> {
799        let session = self.session.as_ref().unwrap();
800        let authenticator = session.user_authenticator();
801        authenticator.authenticate(&msg.password).await?;
802        self.stream.write_no_flush(BeMessage::AuthenticationOk)?;
803        let timezone = session
804            .get_config("timezone")
805            .map_err(|e| PsqlError::StartupError(e.into()))?;
806        self.stream.write_no_flush(BeMessage::ParameterStatus(
807            BeParameterStatusMessage::TimeZone(&timezone),
808        ))?;
809        self.stream
810            .write_parameter_status_msg_no_flush(&ParameterStatus::default())?;
811        self.ready_for_query()?;
812        self.state = PgProtocolState::Regular;
813        Ok(())
814    }
815
816    fn process_cancel_msg(&mut self, m: FeCancelMessage) -> PsqlResult<()> {
817        let session_id = (m.target_process_id, m.target_secret_key);
818        tracing::trace!("cancel query in session: {:?}", session_id);
819        self.session_mgr.cancel_queries_in_session(session_id);
820        self.session_mgr.cancel_creating_jobs_in_session(session_id);
821        self.is_terminate = true;
822        Ok(())
823    }
824
825    async fn process_query_msg(&mut self, sql: Arc<str>) -> PsqlResult<()> {
826        let truncated_sql =
827            get_redacted_and_truncated_sql(&sql, self.redact_sql_option_keywords.clone());
828        let session = self.session.clone().unwrap();
829
830        session.check_idle_in_transaction_timeout()?;
831        // Store only truncated SQL in context to prevent excessive memory usage from large SQL.
832        let _exec_context_guard = session.init_exec_context(truncated_sql.into());
833        self.inner_process_query_msg(sql, session.clone()).await
834    }
835
836    async fn inner_process_query_msg(
837        &mut self,
838        sql: Arc<str>,
839        session: Arc<SM::Session>,
840    ) -> PsqlResult<()> {
841        // Parse sql.
842        let stmts =
843            Parser::parse_sql(&sql).map_err(|err| PsqlError::SimpleQueryError(err.into()))?;
844        // The following inner_process_query_msg_one_stmt can be slow. Release potential large String early.
845        drop(sql);
846        if stmts.is_empty() {
847            self.stream.write_no_flush(BeMessage::EmptyQueryResponse)?;
848        }
849
850        // Execute multiple statements in simple query. KISS later.
851        for stmt in stmts {
852            self.inner_process_query_msg_one_stmt(stmt, session.clone())
853                .await?;
854        }
855        // Put this line inside the for loop above will lead to unfinished/stuck regress test...Not
856        // sure the reason.
857        self.ready_for_query()?;
858        Ok(())
859    }
860
861    async fn inner_process_query_msg_one_stmt(
862        &mut self,
863        stmt: Statement,
864        session: Arc<SM::Session>,
865    ) -> PsqlResult<()> {
866        let session = session.clone();
867
868        // execute query
869        let res = session.clone().run_one_query(stmt, Format::Text).await;
870
871        // Take all remaining notices (if any) and send them before `CommandComplete`.
872        while let Some(notice) = session.next_notice().now_or_never() {
873            self.stream
874                .write_no_flush(BeMessage::NoticeResponse(&notice))?;
875        }
876
877        let mut res = res.map_err(|e| PsqlError::SimpleQueryError(e.into()))?;
878
879        for notice in res.notices() {
880            self.stream
881                .write_no_flush(BeMessage::NoticeResponse(notice))?;
882        }
883
884        let status = res.status();
885        if let Some(ref application_name) = status.application_name {
886            self.stream.write_no_flush(BeMessage::ParameterStatus(
887                BeParameterStatusMessage::ApplicationName(application_name),
888            ))?;
889        }
890
891        if res.is_copy_query_to_stdout() {
892            self.stream
893                .write_no_flush(BeMessage::CopyOutResponse(res.row_desc().len()))?;
894            let mut count = 0;
895            while let Some(row_set) = res.values_stream().next().await {
896                let row_set = row_set.map_err(PsqlError::SimpleQueryError)?;
897                for row in row_set {
898                    self.stream
899                        .write_streaming(BeMessage::CopyData(&row))
900                        .await?;
901                    count += 1;
902                }
903            }
904
905            self.stream.write_no_flush(BeMessage::CopyDone)?;
906
907            // Run the callback before sending the `CommandComplete` message.
908            res.run_callback().await?;
909
910            self.stream
911                .write_no_flush(BeMessage::CommandComplete(BeCommandCompleteMessage {
912                    stmt_type: res.stmt_type(),
913                    rows_cnt: count,
914                }))?;
915        } else if res.is_query() {
916            self.stream
917                .write_no_flush(BeMessage::RowDescription(res.row_desc()))?;
918
919            let mut rows_cnt = 0;
920
921            while let Some(row_set) = res.values_stream().next().await {
922                let row_set = row_set.map_err(PsqlError::SimpleQueryError)?;
923                for row in row_set {
924                    self.stream
925                        .write_streaming(BeMessage::DataRow(&row))
926                        .await?;
927                    rows_cnt += 1;
928                }
929            }
930
931            // Run the callback before sending the `CommandComplete` message.
932            res.run_callback().await?;
933
934            self.stream
935                .write_no_flush(BeMessage::CommandComplete(BeCommandCompleteMessage {
936                    stmt_type: res.stmt_type(),
937                    rows_cnt,
938                }))?;
939        } else if res.stmt_type().is_dml() && !res.stmt_type().is_returning() {
940            let first_row_set = res.values_stream().next().await;
941            let first_row_set = match first_row_set {
942                None => {
943                    return Err(PsqlError::Uncategorized(
944                        anyhow::anyhow!("no affected rows in output").into(),
945                    ));
946                }
947                Some(row) => row.map_err(PsqlError::SimpleQueryError)?,
948            };
949            let affected_rows_str = first_row_set[0].values()[0]
950                .as_ref()
951                .expect("compute node should return affected rows in output");
952
953            assert!(matches!(res.row_cnt_format(), Some(Format::Text)));
954            let affected_rows_cnt = String::from_utf8(affected_rows_str.to_vec())
955                .unwrap()
956                .parse()
957                .unwrap_or_default();
958
959            // Run the callback before sending the `CommandComplete` message.
960            res.run_callback().await?;
961
962            self.stream
963                .write_no_flush(BeMessage::CommandComplete(BeCommandCompleteMessage {
964                    stmt_type: res.stmt_type(),
965                    rows_cnt: affected_rows_cnt,
966                }))?;
967        } else {
968            // Run the callback before sending the `CommandComplete` message.
969            res.run_callback().await?;
970
971            self.stream
972                .write_no_flush(BeMessage::CommandComplete(BeCommandCompleteMessage {
973                    stmt_type: res.stmt_type(),
974                    rows_cnt: 0,
975                }))?;
976        }
977
978        Ok(())
979    }
980
981    fn process_terminate(&mut self) {
982        self.is_terminate = true;
983    }
984
985    fn process_health_check(&mut self) {
986        tracing::debug!("health check");
987        self.is_terminate = true;
988    }
989
990    async fn process_parse_msg(&mut self, mut msg: FeParseMessage) -> PsqlResult<()> {
991        let sql = Arc::from(cstr_to_str(&msg.sql_bytes).unwrap());
992        let session = self.session.clone().unwrap();
993        let statement_name = cstr_to_str(&msg.statement_name).unwrap().to_owned();
994        let type_ids = std::mem::take(&mut msg.type_ids);
995        // The inner_process_parse_msg can be slow. Release potential large FeParseMessage early.
996        drop(msg);
997        self.inner_process_parse_msg(session, sql, statement_name, type_ids)
998            .await?;
999        Ok(())
1000    }
1001
1002    async fn inner_process_parse_msg(
1003        &mut self,
1004        session: Arc<SM::Session>,
1005        sql: Arc<str>,
1006        statement_name: String,
1007        type_ids: Vec<i32>,
1008    ) -> PsqlResult<()> {
1009        if statement_name.is_empty() {
1010            // Remove the unnamed prepare statement first, in case the unsupported sql binds a wrong
1011            // prepare statement.
1012            self.unnamed_prepare_statement.take();
1013        } else if self.prepare_statement_store.contains_key(&statement_name) {
1014            return Err(PsqlError::ExtendedPrepareError(
1015                "Duplicated statement name".into(),
1016            ));
1017        }
1018
1019        let stmt = {
1020            let stmts = Parser::parse_sql(&sql)
1021                .map_err(|err| PsqlError::ExtendedPrepareError(err.into()))?;
1022            if stmts.len() > 1 {
1023                return Err(PsqlError::ExtendedPrepareError(
1024                    "Only one statement is allowed in extended query mode".into(),
1025                ));
1026            }
1027
1028            stmts.into_iter().next()
1029        };
1030
1031        let param_types: Vec<Option<DataType>> = type_ids
1032            .iter()
1033            .map(|&id| {
1034                // 0 means unspecified type
1035                // ref: https://www.postgresql.org/docs/15/protocol-message-formats.html#:~:text=Placing%20a%20zero%20here%20is%20equivalent%20to%20leaving%20the%20type%20unspecified.
1036                if id == 0 {
1037                    Ok(None)
1038                } else {
1039                    DataType::from_oid(id)
1040                        .map(Some)
1041                        .map_err(|e| PsqlError::ExtendedPrepareError(e.into()))
1042                }
1043            })
1044            .try_collect()?;
1045
1046        let prepare_statement = session
1047            .parse(stmt, param_types)
1048            .await
1049            .map_err(|e| PsqlError::ExtendedPrepareError(e.into()))?;
1050        let prepare_statement = PreparedStatementData {
1051            statement: prepare_statement,
1052            sql,
1053        };
1054
1055        if statement_name.is_empty() {
1056            self.unnamed_prepare_statement.replace(prepare_statement);
1057        } else {
1058            self.prepare_statement_store
1059                .insert(statement_name.clone(), prepare_statement);
1060        }
1061
1062        self.statement_portal_dependency
1063            .entry(statement_name)
1064            .or_default()
1065            .clear();
1066
1067        self.stream.write_no_flush(BeMessage::ParseComplete)?;
1068        Ok(())
1069    }
1070
1071    fn process_bind_msg(&mut self, msg: FeBindMessage) -> PsqlResult<()> {
1072        let statement_name = cstr_to_str(&msg.statement_name).unwrap().to_owned();
1073        let portal_name = cstr_to_str(&msg.portal_name).unwrap().to_owned();
1074        let session = self.session.clone().unwrap();
1075
1076        if self.portal_store.contains_key(&portal_name) {
1077            return Err(PsqlError::Uncategorized("Duplicated portal name".into()));
1078        }
1079
1080        let prepare_statement = self.get_statement_data(&statement_name)?.clone();
1081
1082        let result_formats = msg
1083            .result_format_codes
1084            .iter()
1085            .map(|&format_code| Format::from_i16(format_code))
1086            .try_collect()?;
1087        let param_formats = msg
1088            .param_format_codes
1089            .iter()
1090            .map(|&format_code| Format::from_i16(format_code))
1091            .try_collect()?;
1092
1093        let portal = session
1094            .bind(
1095                prepare_statement.statement,
1096                msg.params,
1097                param_formats,
1098                result_formats,
1099            )
1100            .map_err(|e| PsqlError::Uncategorized(e.into()))?;
1101        let portal = PortalData {
1102            portal,
1103            sql: prepare_statement.sql,
1104        };
1105
1106        if portal_name.is_empty() {
1107            self.result_cache.remove(&portal_name);
1108            self.unnamed_portal.replace(portal);
1109        } else {
1110            assert!(
1111                !self.result_cache.contains_key(&portal_name),
1112                "Named portal never can be overridden."
1113            );
1114            self.portal_store.insert(portal_name.clone(), portal);
1115        }
1116
1117        self.statement_portal_dependency
1118            .get_mut(&statement_name)
1119            .unwrap()
1120            .push(portal_name);
1121
1122        self.stream.write_no_flush(BeMessage::BindComplete)?;
1123        Ok(())
1124    }
1125
1126    async fn process_execute_msg(&mut self, msg: FeExecuteMessage) -> PsqlResult<()> {
1127        let portal_name = cstr_to_str(&msg.portal_name).unwrap().to_owned();
1128        let row_max = msg.max_rows as usize;
1129        drop(msg);
1130        let session = self.session.clone().unwrap();
1131
1132        match self.result_cache.remove(&portal_name) {
1133            Some(mut result_cache) => {
1134                assert!(self.portal_store.contains_key(&portal_name));
1135
1136                let is_consume_completed =
1137                    result_cache.consume::<S>(row_max, &mut self.stream).await?;
1138
1139                if !is_consume_completed {
1140                    self.result_cache.insert(portal_name, result_cache);
1141                }
1142            }
1143            _ => {
1144                let portal = self.get_portal_data(&portal_name)?.clone();
1145                let sql = format!("{}", portal.portal);
1146                let truncated_sql =
1147                    get_redacted_and_truncated_sql(&sql, self.redact_sql_option_keywords.clone());
1148                drop(sql);
1149
1150                session.check_idle_in_transaction_timeout()?;
1151                // Store only truncated SQL in context to prevent excessive memory usage from large SQL.
1152                let _exec_context_guard = session.init_exec_context(truncated_sql.into());
1153                let result = session.clone().execute(portal.portal).await;
1154
1155                let pg_response = result.map_err(|e| PsqlError::ExtendedExecuteError(e.into()))?;
1156                let mut result_cache = ResultCache::new(pg_response);
1157                let is_consume_completed =
1158                    result_cache.consume::<S>(row_max, &mut self.stream).await?;
1159                if !is_consume_completed {
1160                    self.result_cache.insert(portal_name, result_cache);
1161                }
1162            }
1163        }
1164
1165        Ok(())
1166    }
1167
1168    fn process_describe_msg(&mut self, msg: FeDescribeMessage) -> PsqlResult<()> {
1169        let name = cstr_to_str(&msg.name).unwrap().to_owned();
1170        let session = self.session.clone().unwrap();
1171        //  b'S' => Statement
1172        //  b'P' => Portal
1173
1174        assert!(msg.kind == b'S' || msg.kind == b'P');
1175        if msg.kind == b'S' {
1176            let prepare_statement = self.get_statement(&name)?;
1177
1178            let (param_types, row_descriptions) = self
1179                .session
1180                .clone()
1181                .unwrap()
1182                .describe_statement(prepare_statement)
1183                .map_err(|e| PsqlError::Uncategorized(e.into()))?;
1184            self.stream.write_no_flush(BeMessage::ParameterDescription(
1185                &param_types.iter().map(|t| t.to_oid()).collect_vec(),
1186            ))?;
1187
1188            if row_descriptions.is_empty() {
1189                // According https://www.postgresql.org/docs/current/protocol-flow.html#:~:text=The%20response%20is%20a%20RowDescri[…]0a%20query%20that%20will%20return%20rows%3B,
1190                // return NoData message if the statement is not a query.
1191                self.stream.write_no_flush(BeMessage::NoData)?;
1192            } else {
1193                self.stream
1194                    .write_no_flush(BeMessage::RowDescription(&row_descriptions))?;
1195            }
1196        } else if msg.kind == b'P' {
1197            let portal = self.get_portal(&name)?;
1198
1199            let row_descriptions = session
1200                .describe_portal(portal)
1201                .map_err(|e| PsqlError::Uncategorized(e.into()))?;
1202
1203            if row_descriptions.is_empty() {
1204                // According https://www.postgresql.org/docs/current/protocol-flow.html#:~:text=The%20response%20is%20a%20RowDescri[…]0a%20query%20that%20will%20return%20rows%3B,
1205                // return NoData message if the statement is not a query.
1206                self.stream.write_no_flush(BeMessage::NoData)?;
1207            } else {
1208                self.stream
1209                    .write_no_flush(BeMessage::RowDescription(&row_descriptions))?;
1210            }
1211        }
1212        Ok(())
1213    }
1214
1215    fn process_close_msg(&mut self, msg: FeCloseMessage) -> PsqlResult<()> {
1216        let name = cstr_to_str(&msg.name).unwrap().to_owned();
1217        assert!(msg.kind == b'S' || msg.kind == b'P');
1218        if msg.kind == b'S' {
1219            if name.is_empty() {
1220                self.unnamed_prepare_statement = None;
1221            } else {
1222                self.prepare_statement_store.remove(&name);
1223            }
1224            for portal_name in self
1225                .statement_portal_dependency
1226                .remove(&name)
1227                .unwrap_or_default()
1228            {
1229                self.remove_portal(&portal_name);
1230            }
1231        } else if msg.kind == b'P' {
1232            self.remove_portal(&name);
1233        }
1234        self.stream.write_no_flush(BeMessage::CloseComplete)?;
1235        Ok(())
1236    }
1237
1238    fn remove_portal(&mut self, portal_name: &str) {
1239        if portal_name.is_empty() {
1240            self.unnamed_portal = None;
1241        } else {
1242            self.portal_store.remove(portal_name);
1243        }
1244        self.result_cache.remove(portal_name);
1245    }
1246
1247    fn get_portal(&self, portal_name: &str) -> PsqlResult<<SM::Session as Session>::Portal> {
1248        Ok(self.get_portal_data(portal_name)?.portal.clone())
1249    }
1250
1251    fn get_portal_data(
1252        &self,
1253        portal_name: &str,
1254    ) -> PsqlResult<&PortalData<<SM::Session as Session>::Portal>> {
1255        if portal_name.is_empty() {
1256            self.unnamed_portal
1257                .as_ref()
1258                .ok_or_else(|| PsqlError::Uncategorized("unnamed portal not found".into()))
1259        } else {
1260            self.portal_store.get(portal_name).ok_or_else(|| {
1261                PsqlError::Uncategorized(format!("Portal {} not found", portal_name).into())
1262            })
1263        }
1264    }
1265
1266    fn get_statement(
1267        &self,
1268        statement_name: &str,
1269    ) -> PsqlResult<<SM::Session as Session>::PreparedStatement> {
1270        Ok(self.get_statement_data(statement_name)?.statement.clone())
1271    }
1272
1273    fn get_statement_data(
1274        &self,
1275        statement_name: &str,
1276    ) -> PsqlResult<&PreparedStatementData<<SM::Session as Session>::PreparedStatement>> {
1277        if statement_name.is_empty() {
1278            self.unnamed_prepare_statement.as_ref().ok_or_else(|| {
1279                PsqlError::Uncategorized("unnamed prepare statement not found".into())
1280            })
1281        } else {
1282            self.prepare_statement_store
1283                .get(statement_name)
1284                .ok_or_else(|| {
1285                    PsqlError::Uncategorized(
1286                        format!("Prepare statement {} not found", statement_name).into(),
1287                    )
1288                })
1289        }
1290    }
1291
1292    fn get_portal_sql(&self, portal_name: &str) -> PsqlResult<Arc<str>> {
1293        Ok(self.get_portal_data(portal_name)?.sql.clone())
1294    }
1295}
1296
1297enum PgStreamInner<S> {
1298    /// Used for the intermediate state when converting from unencrypted to ssl stream.
1299    Placeholder,
1300    /// An unencrypted stream.
1301    Unencrypted(S),
1302    /// An ssl stream.
1303    Ssl(SslStream<S>),
1304}
1305
1306/// Trait for a byte stream that can be used for pg protocol.
1307pub trait PgByteStream: AsyncWrite + AsyncRead + Unpin + Send + 'static {}
1308impl<S> PgByteStream for S where S: AsyncWrite + AsyncRead + Unpin + Send + 'static {}
1309
1310/// Wraps a byte stream and read/write pg messages.
1311///
1312/// Cloning a `PgStream` will share the same stream but a fresh & independent write buffer,
1313/// so that it can be used to write messages concurrently without interference.
1314pub struct PgStream<S> {
1315    /// The underlying stream.
1316    stream: Arc<Mutex<PgStreamInner<S>>>,
1317    /// Write into buffer before flush to stream.
1318    write_buf: BytesMut,
1319    stream_flush_threshold_bytes: usize,
1320    read_header: Option<FeMessageHeader>,
1321}
1322
1323impl<S> PgStream<S> {
1324    /// Create a new `PgStream` with the given stream and streaming flush threshold.
1325    pub fn new(stream: S, stream_flush_threshold_bytes: usize) -> Self {
1326        const DEFAULT_WRITE_BUF_CAPACITY: usize = 10 * 1024;
1327
1328        Self {
1329            stream: Arc::new(Mutex::new(PgStreamInner::Unencrypted(stream))),
1330            write_buf: BytesMut::with_capacity(DEFAULT_WRITE_BUF_CAPACITY),
1331            stream_flush_threshold_bytes,
1332            read_header: None,
1333        }
1334    }
1335
1336    /// Check if the current connection is using SSL
1337    async fn is_ssl_connection(&self) -> bool {
1338        let stream = self.stream.lock().await;
1339        matches!(*stream, PgStreamInner::Ssl(_))
1340    }
1341}
1342
1343impl<S> Clone for PgStream<S> {
1344    fn clone(&self) -> Self {
1345        Self {
1346            stream: Arc::clone(&self.stream),
1347            write_buf: BytesMut::with_capacity(self.write_buf.capacity()),
1348            stream_flush_threshold_bytes: self.stream_flush_threshold_bytes,
1349            read_header: self.read_header.clone(),
1350        }
1351    }
1352}
1353
1354/// At present there is a hard-wired set of parameters for which
1355/// ParameterStatus will be generated: they are:
1356///
1357///  * `server_version`
1358///  * `server_encoding`
1359///  * `client_encoding`
1360///  * `application_name`
1361///  * `is_superuser`
1362///  * `session_authorization`
1363///  * `DateStyle`
1364///  * `IntervalStyle`
1365///  * `TimeZone`
1366///  * `integer_datetimes`
1367///  * `standard_conforming_string`
1368///
1369/// See: <https://www.postgresql.org/docs/9.2/static/protocol-flow.html#PROTOCOL-ASYNC>.
1370#[derive(Debug, Default, Clone)]
1371pub struct ParameterStatus {
1372    pub application_name: Option<String>,
1373}
1374
1375impl<S> PgStream<S>
1376where
1377    S: PgByteStream,
1378{
1379    async fn read_startup(&mut self) -> io::Result<FeMessage> {
1380        let mut stream = self.stream.lock().await;
1381        match &mut *stream {
1382            PgStreamInner::Placeholder => unreachable!(),
1383            PgStreamInner::Unencrypted(stream) => FeStartupMessage::read(stream).await,
1384            PgStreamInner::Ssl(ssl_stream) => FeStartupMessage::read(ssl_stream).await,
1385        }
1386    }
1387
1388    async fn read_header(&mut self) -> io::Result<()> {
1389        let mut stream = self.stream.lock().await;
1390        match &mut *stream {
1391            PgStreamInner::Placeholder => unreachable!(),
1392            PgStreamInner::Unencrypted(stream) => {
1393                self.read_header = Some(FeMessage::read_header(stream).await?);
1394                Ok(())
1395            }
1396            PgStreamInner::Ssl(ssl_stream) => {
1397                self.read_header = Some(FeMessage::read_header(ssl_stream).await?);
1398                Ok(())
1399            }
1400        }
1401    }
1402
1403    async fn read_body(&mut self) -> io::Result<FeMessage> {
1404        let mut stream = self.stream.lock().await;
1405        let header = self
1406            .read_header
1407            .take()
1408            .ok_or_else(|| std::io::Error::new(ErrorKind::InvalidInput, "header not found"))?;
1409        match &mut *stream {
1410            PgStreamInner::Placeholder => unreachable!(),
1411            PgStreamInner::Unencrypted(stream) => FeMessage::read_body(stream, header).await,
1412            PgStreamInner::Ssl(ssl_stream) => FeMessage::read_body(ssl_stream, header).await,
1413        }
1414    }
1415
1416    async fn skip_body(&mut self) -> io::Result<()> {
1417        let mut stream = self.stream.lock().await;
1418        let header = self
1419            .read_header
1420            .take()
1421            .ok_or_else(|| std::io::Error::new(ErrorKind::InvalidInput, "header not found"))?;
1422        match &mut *stream {
1423            PgStreamInner::Placeholder => unreachable!(),
1424            PgStreamInner::Unencrypted(stream) => FeMessage::skip_body(stream, header).await,
1425            PgStreamInner::Ssl(ssl_stream) => FeMessage::skip_body(ssl_stream, header).await,
1426        }
1427    }
1428
1429    fn write_parameter_status_msg_no_flush(&mut self, status: &ParameterStatus) -> io::Result<()> {
1430        self.write_no_flush(BeMessage::ParameterStatus(
1431            BeParameterStatusMessage::ClientEncoding(SERVER_ENCODING),
1432        ))?;
1433        self.write_no_flush(BeMessage::ParameterStatus(
1434            BeParameterStatusMessage::StandardConformingString(STANDARD_CONFORMING_STRINGS),
1435        ))?;
1436        self.write_no_flush(BeMessage::ParameterStatus(
1437            BeParameterStatusMessage::ServerVersion(PG_VERSION),
1438        ))?;
1439        if let Some(application_name) = &status.application_name {
1440            self.write_no_flush(BeMessage::ParameterStatus(
1441                BeParameterStatusMessage::ApplicationName(application_name),
1442            ))?;
1443        }
1444        Ok(())
1445    }
1446
1447    pub fn write_no_flush(&mut self, message: BeMessage<'_>) -> io::Result<()> {
1448        BeMessage::write(&mut self.write_buf, message)
1449    }
1450
1451    /// Write a message that is part of a potentially large response, flushing periodically to
1452    /// bound the write buffer and propagate network backpressure to the result stream.
1453    pub(crate) async fn write_streaming(&mut self, message: BeMessage<'_>) -> io::Result<()> {
1454        self.write_no_flush(message)?;
1455        if self.write_buf.len() >= self.stream_flush_threshold_bytes {
1456            self.flush().await?;
1457        }
1458        Ok(())
1459    }
1460
1461    async fn write(&mut self, message: BeMessage<'_>) -> io::Result<()> {
1462        self.write_no_flush(message)?;
1463        self.flush().await?;
1464        Ok(())
1465    }
1466
1467    async fn flush(&mut self) -> io::Result<()> {
1468        let mut stream = self.stream.lock().await;
1469        match &mut *stream {
1470            PgStreamInner::Placeholder => unreachable!(),
1471            PgStreamInner::Unencrypted(stream) => {
1472                stream.write_all(&self.write_buf).await?;
1473                stream.flush().await?;
1474            }
1475            PgStreamInner::Ssl(ssl_stream) => {
1476                ssl_stream.write_all(&self.write_buf).await?;
1477                ssl_stream.flush().await?;
1478            }
1479        }
1480        self.write_buf.clear();
1481        Ok(())
1482    }
1483}
1484
1485impl<S> PgStream<S>
1486where
1487    S: PgByteStream,
1488{
1489    /// Convert the underlying stream to ssl stream based on the given context.
1490    async fn upgrade_to_ssl(&mut self, ssl_ctx: &SslContextRef) -> PsqlResult<()> {
1491        let mut stream = self.stream.lock().await;
1492
1493        match std::mem::replace(&mut *stream, PgStreamInner::Placeholder) {
1494            PgStreamInner::Unencrypted(unencrypted_stream) => {
1495                let ssl = openssl::ssl::Ssl::new(ssl_ctx).unwrap();
1496                let mut ssl_stream =
1497                    tokio_openssl::SslStream::new(ssl, unencrypted_stream).unwrap();
1498
1499                if let Err(e) = Pin::new(&mut ssl_stream).accept().await {
1500                    tracing::warn!(error = %e.as_report(), "Unable to set up an ssl connection");
1501                    let _ = ssl_stream.shutdown().await;
1502                    return Err(e.into());
1503                }
1504
1505                *stream = PgStreamInner::Ssl(ssl_stream);
1506            }
1507            PgStreamInner::Ssl(_) => panic!("the stream is already ssl"),
1508            PgStreamInner::Placeholder => unreachable!(),
1509        }
1510
1511        Ok(())
1512    }
1513}
1514
1515fn build_ssl_ctx_from_config(tls_config: &TlsConfig) -> PsqlResult<SslContext> {
1516    let mut acceptor = SslAcceptor::mozilla_intermediate_v5(SslMethod::tls()).unwrap();
1517
1518    let key_path = &tls_config.key;
1519    let cert_path = &tls_config.cert;
1520
1521    // Build ssl acceptor according to the config.
1522    // Now we set every verify to true.
1523    acceptor
1524        .set_private_key_file(key_path, openssl::ssl::SslFiletype::PEM)
1525        .map_err(|e| PsqlError::Uncategorized(e.into()))?;
1526    acceptor
1527        .set_ca_file(cert_path)
1528        .map_err(|e| PsqlError::Uncategorized(e.into()))?;
1529    acceptor
1530        .set_certificate_chain_file(cert_path)
1531        .map_err(|e| PsqlError::Uncategorized(e.into()))?;
1532    let acceptor = acceptor.build();
1533
1534    Ok(acceptor.into_context())
1535}
1536
1537pub mod truncated_fmt {
1538    use std::fmt::*;
1539
1540    struct TruncatedFormatter<'a, 'b> {
1541        remaining: usize,
1542        finished: bool,
1543        f: &'a mut Formatter<'b>,
1544    }
1545    impl Write for TruncatedFormatter<'_, '_> {
1546        fn write_str(&mut self, s: &str) -> Result {
1547            if self.finished {
1548                return Ok(());
1549            }
1550
1551            if self.remaining < s.len() {
1552                let actual = s.floor_char_boundary(self.remaining);
1553                self.f.write_str(&s[0..actual])?;
1554                self.remaining -= actual;
1555                self.f.write_str(&format!("...(truncated,{})", s.len()))?;
1556                self.finished = true; // so that ...(truncated) is printed exactly once
1557            } else {
1558                self.f.write_str(s)?;
1559                self.remaining -= s.len();
1560            }
1561            Ok(())
1562        }
1563    }
1564
1565    pub struct TruncatedFmt<'a, T>(pub &'a T, pub usize);
1566
1567    impl<T> Debug for TruncatedFmt<'_, T>
1568    where
1569        T: Debug,
1570    {
1571        fn fmt(&self, f: &mut Formatter<'_>) -> Result {
1572            TruncatedFormatter {
1573                remaining: self.1,
1574                finished: false,
1575                f,
1576            }
1577            .write_fmt(format_args!("{:?}", self.0))
1578        }
1579    }
1580
1581    impl<T> Display for TruncatedFmt<'_, T>
1582    where
1583        T: Display,
1584    {
1585        fn fmt(&self, f: &mut Formatter<'_>) -> Result {
1586            TruncatedFormatter {
1587                remaining: self.1,
1588                finished: false,
1589                f,
1590            }
1591            .write_fmt(format_args!("{}", self.0))
1592        }
1593    }
1594
1595    #[cfg(test)]
1596    mod tests {
1597        use super::*;
1598
1599        #[test]
1600        fn test_trunc_utf8() {
1601            assert_eq!(
1602                format!("{}", TruncatedFmt(&"select '🌊';", 10)),
1603                "select '...(truncated,14)",
1604            );
1605        }
1606    }
1607}
1608
1609/// Handle `options` in `StartupMessage` from client
1610///
1611/// It is like shell arguments but only respects backslash-escape and space;
1612/// quotes have no special meaning and are handled literally.
1613///
1614/// PostgreSQL allows both `-c key=value` and `--key=value`.
1615///
1616/// `key-name` is normalized as `key_name`.
1617///
1618/// * <https://github.com/postgres/postgres/blob/REL_18_1/src/backend/utils/init/postinit.c#L487>
1619/// * <https://github.com/postgres/postgres/blob/REL_18_1/src/backend/tcop/postgres.c#L3866>
1620/// * <https://github.com/postgres/postgres/blob/REL_18_1/src/backend/utils/misc/guc.c#L6361>
1621fn parse_options(options: &str) -> PsqlResult<Vec<(String, String)>> {
1622    let mut args = Vec::new();
1623    let mut current_arg = String::new();
1624    let mut chars = options.chars().peekable();
1625
1626    while let Some(c) = chars.next() {
1627        if c == '\\' {
1628            if let Some(next_c) = chars.next() {
1629                current_arg.push(next_c);
1630            }
1631        } else if c.is_ascii_whitespace() {
1632            if !current_arg.is_empty() {
1633                args.push(std::mem::take(&mut current_arg));
1634            }
1635        } else {
1636            current_arg.push(c);
1637        }
1638    }
1639    if !current_arg.is_empty() {
1640        args.push(current_arg);
1641    }
1642
1643    let mut args_iter = args.into_iter();
1644    let mut config = Vec::new();
1645
1646    while let Some(arg) = args_iter.next() {
1647        if arg == "-c" {
1648            if let Some(config_str) = args_iter.next() {
1649                if let Some((key, value)) = config_str.split_once('=') {
1650                    let key = key.replace("-", "_");
1651                    config.push((key, value.to_owned()));
1652                } else {
1653                    return Err(PsqlError::StartupError(
1654                        format!("invalid config format: {}", config_str).into(),
1655                    ));
1656                }
1657            } else {
1658                return Err(PsqlError::StartupError("missing argument for -c".into()));
1659            }
1660        } else if let Some(config_str) = arg.strip_prefix("--") {
1661            if let Some((key, value)) = config_str.split_once('=') {
1662                let key = key.replace("-", "_");
1663                config.push((key, value.to_owned()));
1664            } else {
1665                return Err(PsqlError::StartupError(
1666                    format!("invalid config format: {}", config_str).into(),
1667                ));
1668            }
1669        } else {
1670            tracing::warn!(
1671                arg,
1672                "ignoring unrecognized option for backward compatibility"
1673            );
1674        }
1675    }
1676    Ok(config)
1677}
1678
1679#[cfg(test)]
1680mod tests {
1681    use std::collections::HashSet;
1682
1683    use tokio::io::AsyncReadExt;
1684
1685    use super::*;
1686    use crate::types::Row;
1687
1688    #[tokio::test]
1689    async fn test_streaming_write_flushes_at_threshold() {
1690        const STREAM_FLUSH_THRESHOLD: usize = 64 * 1024;
1691
1692        let (server, mut client) = tokio::io::duplex(STREAM_FLUSH_THRESHOLD * 2);
1693        let mut stream = PgStream::new(server, STREAM_FLUSH_THRESHOLD);
1694
1695        let small_row = Row::new(vec![Some(Bytes::from_static(b"small"))]);
1696        stream
1697            .write_streaming(BeMessage::DataRow(&small_row))
1698            .await
1699            .unwrap();
1700        assert!(!stream.write_buf.is_empty());
1701
1702        let large_row = Row::new(vec![Some(Bytes::from(vec![0; STREAM_FLUSH_THRESHOLD]))]);
1703        stream
1704            .write_streaming(BeMessage::DataRow(&large_row))
1705            .await
1706            .unwrap();
1707        assert!(stream.write_buf.is_empty());
1708
1709        let mut message_tag = [0];
1710        client.read_exact(&mut message_tag).await.unwrap();
1711        assert_eq!(message_tag[0], b'D');
1712    }
1713
1714    #[test]
1715    fn test_redact_parsable_sql() {
1716        let keywords = Arc::new(HashSet::from(["v2".into(), "v4".into(), "b".into()]));
1717        let sql = r"
1718        create source temp (k bigint, v varchar) with (
1719            connector = 'datagen',
1720            v1 = 123,
1721            v2 = 'with',
1722            v3 = false,
1723            v4 = '',
1724        ) FORMAT plain ENCODE json (a='1',b='2')
1725        ";
1726        assert_eq!(
1727            redact_sql(sql, keywords),
1728            "CREATE SOURCE temp (k BIGINT, v CHARACTER VARYING) WITH (connector = 'datagen', v1 = 123, v2 = [REDACTED], v3 = false, v4 = [REDACTED]) FORMAT PLAIN ENCODE JSON (a = '1', b = [REDACTED])"
1729        );
1730    }
1731
1732    #[test]
1733    fn test_redact_user_password_sql() {
1734        let keywords = Arc::new(HashSet::from(["password".into()]));
1735
1736        assert_eq!(
1737            redact_sql("ALTER USER WITH PASSWORD 'rw_password_2'", keywords.clone()),
1738            "ALTER USER WITH PASSWORD [REDACTED]"
1739        );
1740        assert_eq!(
1741            redact_sql(
1742                "ALTER USER foo WITH ENCRYPTED PASSWORD 'md5827ccb0eea8a706c4c34a16891f84e7b'",
1743                keywords.clone(),
1744            ),
1745            "ALTER USER foo WITH ENCRYPTED PASSWORD [REDACTED]"
1746        );
1747        assert_eq!(
1748            redact_sql("CREATE USER foo WITH PASSWORD 'rw_password_2'", keywords),
1749            "CREATE USER foo WITH PASSWORD [REDACTED]"
1750        );
1751    }
1752
1753    #[test]
1754    fn test_parse_options() {
1755        assert_eq!(parse_options("").unwrap(), vec![]);
1756        assert_eq!(
1757            parse_options("-c a=1 -c b=2").unwrap(),
1758            vec![("a".into(), "1".into()), ("b".into(), "2".into())]
1759        );
1760        assert_eq!(
1761            parse_options("-c   key=value").unwrap(),
1762            vec![("key".into(), "value".into())]
1763        );
1764        // Custom parser treats quotes as normal characters, so they are included in value
1765        assert_eq!(
1766            parse_options("-c key='value'").unwrap(),
1767            vec![("key".into(), "'value'".into())]
1768        );
1769
1770        // Test backslash escaping for spaces (standard Postgres way)
1771        assert_eq!(
1772            parse_options(r#"-c key=value\ with\ spaces"#).unwrap(),
1773            vec![("key".into(), "value with spaces".into())]
1774        );
1775        assert_eq!(
1776            parse_options(r#"-c search_path=my\ schema"#).unwrap(),
1777            vec![("search_path".into(), "my schema".into())]
1778        );
1779
1780        assert!(parse_options("-c").is_err());
1781        assert!(parse_options("-c foo").is_err()); // missing =
1782        assert!(parse_options("--foo").is_err()); // missing = in -- option
1783
1784        assert_eq!(
1785            parse_options("--foo=bar").unwrap(),
1786            vec![("foo".into(), "bar".into())]
1787        );
1788        assert_eq!(
1789            parse_options(r#"--foo=bar\ baz"#).unwrap(),
1790            vec![("foo".into(), "bar baz".into())]
1791        );
1792        assert_eq!(
1793            parse_options("-c a=1 --b=2").unwrap(),
1794            vec![("a".into(), "1".into()), ("b".into(), "2".into())]
1795        );
1796        // Unpaired trailing backslash is silently dropped, same as PostgreSQL
1797        assert_eq!(
1798            parse_options(r#"-c a=b\"#).unwrap(),
1799            vec![("a".into(), "b".into())]
1800        );
1801    }
1802}