1use 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
57static 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 pub static CURRENT_SESSION: Weak<dyn Any + Send + Sync>
68}
69
70pub struct PgProtocol<S, SM>
73where
74 SM: SessionManager,
75{
76 stream: PgStream<S>,
78 state: PgProtocolState,
80 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 statement_portal_dependency: HashMap<String, Vec<String>>,
96
97 tls_context: Option<SslContext>,
100
101 tls_config: Option<TlsConfig>,
103
104 ignore_util_sync: bool,
107
108 peer_addr: AddressRef,
110
111 redact_sql_option_keywords: Option<RedactSqlOptionKeywordsRef>,
112 message_memory_manager: MessageMemoryManagerRef,
113}
114
115#[derive(Debug, Clone)]
117pub struct TlsConfig {
118 pub cert: String,
120 pub key: String,
122 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 self.session_mgr.end_session(session);
170 }
171 }
172}
173
174enum 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
192pub 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
219fn 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
234fn 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 pub async fn run(&mut self) {
295 let mut notice_fut = None;
296
297 loop {
298 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(¬ice)).await {
307 tracing::error!(error = %e.as_report(), notice, "failed to send notice");
308 }
309 }
310 }));
311 }
312
313 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; }
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 pub async fn process(&mut self, msg: FeMessage) -> bool {
343 self.do_process(msg).await.is_none() || self.is_terminate
344 }
345
346 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(¤t_session.user(), &mut span);
391 }
392 span
393 }
394
395 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 let fut = Box::pin(self.do_process_inner(msg));
410
411 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 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 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 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 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 if let Err(error) = &result {
471 if cfg!(debug_assertions) && !Deployment::current().is_ci() {
472 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 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 let fut = fut.instrument(span);
498
499 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 return None;
514 }
515
516 PsqlError::StartupError(_) | PsqlError::PasswordError => {
517 self.stream
518 .write_no_flush(BeMessage::ErrorResponse {
519 error: &e,
520 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 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 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 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 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 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 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 self.stream.write(BeMessage::EncryptionResponseSsl).await?;
708 self.stream.upgrade_to_ssl(context).await?;
709 } else {
710 self.stream.write(BeMessage::EncryptionResponseNo).await?;
712 }
713
714 Ok(())
715 }
716
717 async fn process_startup_msg(&mut self, msg: FeStartupMessage) -> PsqlResult<()> {
718 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 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 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 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 let stmts =
843 Parser::parse_sql(&sql).map_err(|err| PsqlError::SimpleQueryError(err.into()))?;
844 drop(sql);
846 if stmts.is_empty() {
847 self.stream.write_no_flush(BeMessage::EmptyQueryResponse)?;
848 }
849
850 for stmt in stmts {
852 self.inner_process_query_msg_one_stmt(stmt, session.clone())
853 .await?;
854 }
855 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 let res = session.clone().run_one_query(stmt, Format::Text).await;
870
871 while let Some(notice) = session.next_notice().now_or_never() {
873 self.stream
874 .write_no_flush(BeMessage::NoticeResponse(¬ice))?;
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 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 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 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 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 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 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 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 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 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 ¶m_types.iter().map(|t| t.to_oid()).collect_vec(),
1186 ))?;
1187
1188 if row_descriptions.is_empty() {
1189 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 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 Placeholder,
1300 Unencrypted(S),
1302 Ssl(SslStream<S>),
1304}
1305
1306pub trait PgByteStream: AsyncWrite + AsyncRead + Unpin + Send + 'static {}
1308impl<S> PgByteStream for S where S: AsyncWrite + AsyncRead + Unpin + Send + 'static {}
1309
1310pub struct PgStream<S> {
1315 stream: Arc<Mutex<PgStreamInner<S>>>,
1317 write_buf: BytesMut,
1319 stream_flush_threshold_bytes: usize,
1320 read_header: Option<FeMessageHeader>,
1321}
1322
1323impl<S> PgStream<S> {
1324 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 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#[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 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 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 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; } 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
1609fn 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 assert_eq!(
1766 parse_options("-c key='value'").unwrap(),
1767 vec![("key".into(), "'value'".into())]
1768 );
1769
1770 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()); assert!(parse_options("--foo").is_err()); 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 assert_eq!(
1798 parse_options(r#"-c a=b\"#).unwrap(),
1799 vec![("a".into(), "b".into())]
1800 );
1801 }
1802}