1use std::collections::BTreeMap;
11use std::convert::TryFrom;
12use std::future::Future;
13use std::ops::Deref;
14use std::sync::Arc;
15use std::time::{Duration, Instant};
16use std::{iter, mem};
17
18use base64::prelude::*;
19use byteorder::{ByteOrder, NetworkEndian};
20use csv_core::ReadRecordResult;
21use futures::future::{BoxFuture, FutureExt, pending};
22use itertools::Itertools;
23use mz_adapter::client::{RecordFirstRowStream, redact_sql_for_logging};
24use mz_adapter::session::{
25 EndTransactionAction, InProgressRows, LifecycleTimestamps, PortalRefMut, PortalState, Session,
26 SessionConfig, TransactionStatus,
27};
28use mz_adapter::statement_logging::{StatementEndedExecutionReason, StatementExecutionStrategy};
29use mz_adapter::{
30 AdapterError, AdapterNotice, ExecuteContextGuard, ExecuteResponse, PeekResponseUnary, metrics,
31 verify_datum_desc,
32};
33use mz_adapter_types::dyncfgs::OIDC_GROUP_CLAIM;
34use mz_auth::Authenticated;
35use mz_auth::password::Password;
36use mz_authenticator::{Authenticator, GenericOidcAuthenticator};
37use mz_frontegg_auth::Authenticator as FronteggAuthenticator;
38use mz_ore::cast::CastFrom;
39use mz_ore::netio::AsyncReady;
40use mz_ore::now::{EpochMillis, SYSTEM_TIME};
41use mz_ore::str::StrExt;
42use mz_ore::{assert_none, assert_ok, instrument, soft_assert_eq_or_log, soft_assert_or_log};
43use mz_pgcopy::{CopyCsvFormatParams, CopyFormatParams, CopyTextFormatParams};
44use mz_pgwire_common::{
45 ConnectionCounter, Cursor, ErrorResponse, Format, FrontendMessage, Severity, VERSION_3,
46 VERSIONS,
47};
48use mz_repr::{
49 CatalogItemId, ColumnIndex, Datum, RelationDesc, RowArena, RowIterator, RowRef,
50 SqlRelationType, SqlScalarType,
51};
52use mz_server_core::TlsMode;
53use mz_server_core::listeners;
54use mz_server_core::listeners::AllowedRoles;
55use mz_sql::ast::display::AstDisplay;
56use mz_sql::ast::{
57 CopyDirection, CopyStatement, CopyTarget, FetchDirection, Ident, Raw, Statement,
58};
59use mz_sql::parse::StatementParseResult;
60use mz_sql::plan::{CopyFormat, ExecuteTimeout, StatementDesc};
61use mz_sql::session::metadata::SessionMetadata;
62use mz_sql::session::user::INTERNAL_USER_NAMES;
63use mz_sql::session::vars::VarInput;
64use postgres::error::SqlState;
65use tokio::io::{self, AsyncRead, AsyncWrite};
66use tokio::select;
67use tokio::time::{self};
68use tokio_metrics::TaskMetrics;
69use tracing::{Instrument, debug, debug_span, info, warn};
70use uuid::Uuid;
71
72use crate::codec::{
73 FramedConn, decode_password, decode_sasl_initial_response, decode_sasl_response,
74};
75use crate::message::{
76 self, BackendMessage, SASLServerFinalMessage, SASLServerFinalMessageKinds,
77 SASLServerFirstMessage,
78};
79
80pub fn match_handshake(buf: &[u8]) -> bool {
84 if buf.len() < 8 {
94 return false;
95 }
96 let version = NetworkEndian::read_i32(&buf[4..8]);
97 VERSIONS.contains(&version)
98}
99
100pub struct RunParams<'a, A, I>
102where
103 I: Iterator<Item = TaskMetrics> + Send,
104{
105 pub tls_mode: Option<TlsMode>,
107 pub adapter_client: mz_adapter::Client,
109 pub conn: &'a mut FramedConn<A>,
111 pub conn_uuid: Uuid,
113 pub version: i32,
115 pub params: BTreeMap<String, String>,
117 pub frontegg: Option<FronteggAuthenticator>,
119 pub oidc: GenericOidcAuthenticator,
121 pub authenticator_kind: listeners::AuthenticatorKind,
124 pub active_connection_counter: ConnectionCounter,
126 pub helm_chart_version: Option<String>,
128 pub allowed_roles: AllowedRoles,
130 pub tokio_metrics_intervals: I,
132}
133
134#[mz_ore::instrument(level = "debug")]
144pub async fn run<'a, A, I>(
145 RunParams {
146 tls_mode,
147 adapter_client,
148 conn,
149 conn_uuid,
150 version,
151 mut params,
152 frontegg,
153 oidc,
154 authenticator_kind,
155 active_connection_counter,
156 helm_chart_version,
157 allowed_roles,
158 tokio_metrics_intervals,
159 }: RunParams<'a, A, I>,
160) -> Result<(), io::Error>
161where
162 A: AsyncRead + AsyncWrite + AsyncReady + Send + Sync + Unpin,
163 I: Iterator<Item = TaskMetrics> + Send,
164{
165 if version != VERSION_3 {
166 return conn
167 .send(ErrorResponse::fatal(
168 SqlState::SQLSERVER_REJECTED_ESTABLISHMENT_OF_SQLCONNECTION,
169 "server does not support the client's requested protocol version",
170 ))
171 .await;
172 }
173
174 let user = params.remove("user").unwrap_or_else(String::new);
175 let options = parse_options(params.get("options").unwrap_or(&String::new()));
176 let authenticator =
177 get_authenticator(authenticator_kind, frontegg, oidc, adapter_client.clone());
178 let is_internal_user = INTERNAL_USER_NAMES.contains(&user);
180 let is_reserved_user = mz_adapter::catalog::is_reserved_role_name(user.as_str());
182 let role_allowed = match allowed_roles {
183 AllowedRoles::Normal => !is_reserved_user,
184 AllowedRoles::Internal => is_internal_user,
185 AllowedRoles::NormalAndInternal => !is_reserved_user || is_internal_user,
186 };
187 if !role_allowed {
188 let msg = format!("unauthorized login to user '{user}'");
189 return conn
190 .send(ErrorResponse::fatal(SqlState::INSUFFICIENT_PRIVILEGE, msg))
191 .await;
192 }
193
194 if let Err(err) = conn.inner().ensure_tls_compatibility(&tls_mode) {
195 return conn.send(err).await;
196 }
197
198 let authenticator_kind = authenticator.kind();
199
200 let (mut session, expired) = match authenticator {
201 Authenticator::Frontegg(frontegg) => {
202 let password = match request_cleartext_password(conn).await {
203 Ok(password) => password,
204 Err(PasswordRequestError::IoError(e)) => return Err(e),
205 Err(PasswordRequestError::InvalidPasswordError(e)) => {
206 return conn.send(e).await;
207 }
208 };
209
210 let group_claim =
211 OIDC_GROUP_CLAIM.get(adapter_client.get_system_vars().await.dyncfgs());
212 let auth_response = frontegg
213 .authenticate(&user, &password, Some(&group_claim))
214 .await;
215 match auth_response {
216 Ok((mut auth_session, authenticated)) => {
223 let groups = auth_session.groups();
224 let session = adapter_client.new_session(
225 SessionConfig {
226 conn_id: conn.conn_id().clone(),
227 uuid: conn_uuid,
228 user: auth_session.user().into(),
229 client_ip: conn.peer_addr().clone(),
230 external_metadata_rx: Some(auth_session.external_metadata_rx()),
231 helm_chart_version,
232 authenticator_kind,
233 groups,
234 },
235 authenticated,
236 );
237 let expired = async move { auth_session.expired().await };
238 (session, expired.left_future())
239 }
240 Err(err) => {
241 warn!(?err, "pgwire connection failed authentication");
242 return conn
243 .send(ErrorResponse::fatal(
244 SqlState::INVALID_PASSWORD,
245 "invalid password",
246 ))
247 .await;
248 }
249 }
250 }
251 Authenticator::Oidc(oidc) => {
252 let password = match request_cleartext_password(conn).await {
255 Ok(password) => password,
256 Err(PasswordRequestError::IoError(e)) => return Err(e),
257 Err(PasswordRequestError::InvalidPasswordError(e)) => {
258 return conn.send(e).await;
259 }
260 };
261 if is_jwt(&password) {
262 let auth_response = oidc.authenticate(&password, Some(&user)).await;
263 match auth_response {
264 Ok((mut claims, authenticated)) => {
265 let groups = claims.groups.take();
266 let session = adapter_client.new_session(
267 SessionConfig {
268 conn_id: conn.conn_id().clone(),
269 uuid: conn_uuid,
270 user: std::mem::take(&mut claims.user),
271 client_ip: conn.peer_addr().clone(),
272 external_metadata_rx: None,
273 helm_chart_version,
274 authenticator_kind,
275 groups,
276 },
277 authenticated,
278 );
279 (session, pending().right_future())
282 }
283 Err(err) => {
284 warn!(?err, "pgwire connection failed authentication");
285 return conn.send(err.into_response()).await;
286 }
287 }
288 } else {
289 let session = match authenticate_with_password(
290 conn,
291 &adapter_client,
292 user,
293 Password(password),
294 conn_uuid,
295 helm_chart_version,
296 )
297 .await
298 {
299 Ok(session) => session,
300 Err(PasswordRequestError::IoError(e)) => return Err(e),
301 Err(PasswordRequestError::InvalidPasswordError(e)) => {
302 return conn.send(e).await;
303 }
304 };
305 (session, pending().right_future())
306 }
307 }
308 Authenticator::Password(adapter_client) => {
309 let password = match request_cleartext_password(conn).await {
310 Ok(password) => password,
311 Err(PasswordRequestError::IoError(e)) => return Err(e),
312 Err(PasswordRequestError::InvalidPasswordError(e)) => {
313 return conn.send(e).await;
314 }
315 };
316 let session = match authenticate_with_password(
317 conn,
318 &adapter_client,
319 user,
320 Password(password),
321 conn_uuid,
322 helm_chart_version,
323 )
324 .await
325 {
326 Ok(session) => session,
327 Err(PasswordRequestError::IoError(e)) => return Err(e),
328 Err(PasswordRequestError::InvalidPasswordError(e)) => {
329 return conn.send(e).await;
330 }
331 };
332 (session, pending().right_future())
334 }
335 Authenticator::Sasl(adapter_client) => {
336 conn.send(BackendMessage::AuthenticationSASL).await?;
338 conn.flush().await?;
339 let (mechanism, initial_response) = match conn.recv().await? {
341 Some(FrontendMessage::RawAuthentication(data)) => {
342 match decode_sasl_initial_response(Cursor::new(&data)).ok() {
343 Some(FrontendMessage::SASLInitialResponse {
344 gs2_header,
345 mechanism,
346 initial_response,
347 }) => {
348 if gs2_header.channel_binding_enabled() {
350 return conn
351 .send(ErrorResponse::fatal(
352 SqlState::PROTOCOL_VIOLATION,
353 "channel binding not supported",
354 ))
355 .await;
356 }
357 (mechanism, initial_response)
358 }
359 _ => {
360 return conn
361 .send(ErrorResponse::fatal(
362 SqlState::INVALID_AUTHORIZATION_SPECIFICATION,
363 "expected SASLInitialResponse message",
364 ))
365 .await;
366 }
367 }
368 }
369 _ => {
370 return conn
371 .send(ErrorResponse::fatal(
372 SqlState::INVALID_AUTHORIZATION_SPECIFICATION,
373 "expected SASLInitialResponse message",
374 ))
375 .await;
376 }
377 };
378
379 if mechanism != "SCRAM-SHA-256" {
380 return conn
381 .send(ErrorResponse::fatal(
382 SqlState::INVALID_AUTHORIZATION_SPECIFICATION,
383 "unsupported SASL mechanism",
384 ))
385 .await;
386 }
387
388 if initial_response.nonce.len() > 256 {
389 return conn
390 .send(ErrorResponse::fatal(
391 SqlState::INVALID_AUTHORIZATION_SPECIFICATION,
392 "nonce too long",
393 ))
394 .await;
395 }
396
397 let (server_first_message_raw, mock_hash) = match adapter_client
398 .generate_sasl_challenge(&user, &initial_response.nonce)
399 .await
400 {
401 Ok(response) => {
402 let server_first_message_raw = format!(
403 "r={},s={},i={}",
404 response.nonce, response.salt, response.iteration_count
405 );
406
407 let client_key = [0u8; 32];
408 let server_key = [1u8; 32];
409 let mock_hash = format!(
410 "SCRAM-SHA-256${}:{}${}:{}",
411 response.iteration_count,
412 response.salt,
413 BASE64_STANDARD.encode(client_key),
414 BASE64_STANDARD.encode(server_key)
415 );
416
417 conn.send(BackendMessage::AuthenticationSASLContinue(
418 SASLServerFirstMessage {
419 iteration_count: response.iteration_count,
420 nonce: response.nonce,
421 salt: response.salt,
422 },
423 ))
424 .await?;
425 conn.flush().await?;
426 (server_first_message_raw, mock_hash)
427 }
428 Err(e) => {
429 return conn.send(e.into_response(Severity::Fatal)).await;
430 }
431 };
432
433 let authenticated = match conn.recv().await? {
434 Some(FrontendMessage::RawAuthentication(data)) => {
435 match decode_sasl_response(Cursor::new(&data)).ok() {
436 Some(FrontendMessage::SASLResponse(response)) => {
437 let auth_message = format!(
438 "{},{},{}",
439 initial_response.client_first_message_bare_raw,
440 server_first_message_raw,
441 response.client_final_message_bare_raw
442 );
443 if response.proof.len() > 1024 {
444 return conn
445 .send(ErrorResponse::fatal(
446 SqlState::INVALID_AUTHORIZATION_SPECIFICATION,
447 "proof too long",
448 ))
449 .await;
450 }
451 match adapter_client
452 .verify_sasl_proof(
453 &user,
454 &response.proof,
455 &auth_message,
456 &mock_hash,
457 )
458 .await
459 {
460 Ok((proof_response, authenticated)) => {
461 conn.send(BackendMessage::AuthenticationSASLFinal(
462 SASLServerFinalMessage {
463 kind: SASLServerFinalMessageKinds::Verifier(
464 proof_response.verifier,
465 ),
466 extensions: vec![],
467 },
468 ))
469 .await?;
470 conn.flush().await?;
471 authenticated
472 }
473 Err(_) => {
474 return conn
475 .send(ErrorResponse::fatal(
476 SqlState::INVALID_PASSWORD,
477 "invalid password",
478 ))
479 .await;
480 }
481 }
482 }
483 _ => {
484 return conn
485 .send(ErrorResponse::fatal(
486 SqlState::INVALID_AUTHORIZATION_SPECIFICATION,
487 "expected SASLResponse message",
488 ))
489 .await;
490 }
491 }
492 }
493 _ => {
494 return conn
495 .send(ErrorResponse::fatal(
496 SqlState::INVALID_AUTHORIZATION_SPECIFICATION,
497 "expected SASLResponse message",
498 ))
499 .await;
500 }
501 };
502
503 let session = adapter_client.new_session(
504 SessionConfig {
505 conn_id: conn.conn_id().clone(),
506 uuid: conn_uuid,
507 user,
508 client_ip: conn.peer_addr().clone(),
509 external_metadata_rx: None,
510 helm_chart_version,
511 authenticator_kind,
512 groups: None,
513 },
514 authenticated,
515 );
516 let auth_session = pending().right_future();
518 (session, auth_session)
519 }
520
521 Authenticator::None => {
522 let session = adapter_client.new_session(
523 SessionConfig {
524 conn_id: conn.conn_id().clone(),
525 uuid: conn_uuid,
526 user,
527 client_ip: conn.peer_addr().clone(),
528 external_metadata_rx: None,
529 helm_chart_version,
530 authenticator_kind,
531 groups: None,
532 },
533 Authenticated,
534 );
535 let auth_session = pending().right_future();
537 (session, auth_session)
538 }
539 };
540
541 let system_vars = adapter_client.get_system_vars().await;
542 let mut applied_params = vec![];
546 for (name, value) in params {
547 let settings = match name.as_str() {
548 "options" => match &options {
549 Ok(opts) => opts,
550 Err(()) => {
551 session.add_notice(AdapterNotice::BadStartupSetting {
552 name,
553 reason: "could not parse".into(),
554 });
555 continue;
556 }
557 },
558 _ => &vec![(name, value)],
559 };
560 for (key, val) in settings {
561 const LOCAL: bool = false;
562 match session
567 .vars_mut()
568 .set(&system_vars, key, VarInput::Flat(val), LOCAL)
569 {
570 Ok(()) => applied_params.push((key.clone(), val.clone())),
571 Err(err) => {
572 session.add_notice(AdapterNotice::BadStartupSetting {
573 name: key.clone(),
574 reason: err.to_string(),
575 });
576 }
577 }
578 }
579 }
580 session
581 .vars_mut()
582 .end_transaction(EndTransactionAction::Commit);
583
584 let _guard = match active_connection_counter.allocate_connection(session.user()) {
585 Ok(drop_connection) => drop_connection,
586 Err(e) => {
587 let e: AdapterError = e.into();
588 return conn.send(e.into_response(Severity::Fatal)).await;
589 }
590 };
591
592 let mut adapter_client = match adapter_client.startup(session).await {
594 Ok(adapter_client) => adapter_client,
595 Err(e) => return conn.send(e.into_response(Severity::Fatal)).await,
596 };
597
598 for (key, val) in applied_params {
606 if let Err(err) = adapter_client
607 .session()
608 .vars_mut()
609 .set_default(&key, VarInput::Flat(&val))
610 {
611 mz_ore::soft_panic_or_log!("failed to apply startup parameter as default: {err:?}");
613 }
614 }
615
616 conn.allow_post_auth_frames();
619
620 let mut buf = vec![BackendMessage::AuthenticationOk];
621 for var in adapter_client.session().vars().notify_set() {
622 buf.push(BackendMessage::ParameterStatus(var.name(), var.value()));
623 }
624 buf.push(BackendMessage::BackendKeyData {
625 conn_id: adapter_client.session().conn_id().unhandled(),
626 secret_key: adapter_client.session().secret_key(),
627 });
628 buf.extend(
629 adapter_client
630 .session()
631 .drain_notices()
632 .into_iter()
633 .map(|notice| BackendMessage::ErrorResponse(notice.into_response())),
634 );
635 buf.push(BackendMessage::ReadyForQuery(
636 adapter_client.session().transaction().into(),
637 ));
638 conn.send_all(buf).await?;
639 conn.flush().await?;
640
641 let machine = StateMachine {
642 conn,
643 adapter_client,
644 txn_needs_commit: false,
645 tokio_metrics_intervals,
646 };
647
648 select! {
649 r = machine.run() => {
650 if let Err(err) = &r {
655 let _ = conn
656 .send(ErrorResponse::fatal(
657 SqlState::CONNECTION_FAILURE,
658 err.to_string(),
659 ))
660 .await;
661 let _ = conn.flush().await;
662 }
663 r
664 },
665 _ = expired => {
666 conn
667 .send(ErrorResponse::fatal(SqlState::INVALID_AUTHORIZATION_SPECIFICATION, "authentication expired"))
668 .await?;
669 conn.flush().await
670 }
671 }
672}
673
674fn is_jwt(password: &str) -> bool {
677 jsonwebtoken::decode_header(password).is_ok()
678}
679
680fn parse_options(value: &str) -> Result<Vec<(String, String)>, ()> {
685 let opts = split_options(value);
686 let mut pairs = Vec::with_capacity(opts.len());
687 let mut seen_prefix = false;
688 for opt in opts {
689 if !seen_prefix {
690 if opt == "-c" {
691 seen_prefix = true;
692 } else {
693 let (key, val) = parse_option(&opt)?;
694 pairs.push((key.to_owned(), val.to_owned()));
695 }
696 } else {
697 let (key, val) = opt.split_once('=').ok_or(())?;
698 pairs.push((key.to_owned(), val.to_owned()));
699 seen_prefix = false;
700 }
701 }
702 Ok(pairs)
703}
704
705fn parse_option(option: &str) -> Result<(&str, &str), ()> {
709 let (key, value) = option.split_once('=').ok_or(())?;
710 for prefix in &["-c", "--"] {
711 if let Some(key) = key.strip_prefix(prefix) {
712 return Ok((key, value));
713 }
714 }
715 Err(())
716}
717
718fn split_options(value: &str) -> Vec<String> {
720 let mut strs = Vec::new();
721 let mut current = String::new();
725 let mut was_slash = false;
726 for c in value.chars() {
727 was_slash = match c {
728 ' ' => {
729 if was_slash {
730 current.push(' ');
731 } else if !current.is_empty() {
732 strs.push(std::mem::take(&mut current));
735 }
736 false
737 }
738 '\\' => {
739 if was_slash {
740 current.push('\\');
743 false
744 } else {
745 true
746 }
747 }
748 _ => {
749 current.push(c);
750 false
751 }
752 };
753 }
754 if !current.is_empty() {
756 strs.push(current);
757 }
758 strs
759}
760
761enum PasswordRequestError {
762 InvalidPasswordError(ErrorResponse),
763 IoError(io::Error),
764}
765
766impl From<io::Error> for PasswordRequestError {
767 fn from(e: io::Error) -> Self {
768 PasswordRequestError::IoError(e)
769 }
770}
771
772async fn request_cleartext_password<A>(
776 conn: &mut FramedConn<A>,
777) -> Result<String, PasswordRequestError>
778where
779 A: AsyncRead + AsyncWrite + AsyncReady + Send + Sync + Unpin,
780{
781 conn.send(BackendMessage::AuthenticationCleartextPassword)
782 .await?;
783 conn.flush().await?;
784
785 if let Some(message) = conn.recv().await? {
786 if let FrontendMessage::RawAuthentication(data) = message {
787 if let Some(FrontendMessage::Password { password }) =
788 decode_password(Cursor::new(&data)).ok()
789 {
790 return Ok(password);
791 }
792 }
793 }
794
795 Err(PasswordRequestError::InvalidPasswordError(
796 ErrorResponse::fatal(
797 SqlState::INVALID_AUTHORIZATION_SPECIFICATION,
798 "expected Password message",
799 ),
800 ))
801}
802
803async fn authenticate_with_password<A>(
806 conn: &FramedConn<A>,
807 adapter_client: &mz_adapter::Client,
808 user: String,
809 password: Password,
810 conn_uuid: Uuid,
811 helm_chart_version: Option<String>,
812) -> Result<Session, PasswordRequestError>
813where
814 A: AsyncRead + AsyncWrite + AsyncReady + Send + Sync + Unpin,
815{
816 let authenticated = match adapter_client.authenticate(&user, &password).await {
817 Ok(authenticated) => authenticated,
818 Err(err) => {
819 warn!(?err, "pgwire connection failed authentication");
820 return Err(PasswordRequestError::InvalidPasswordError(
821 ErrorResponse::fatal(SqlState::INVALID_PASSWORD, "invalid password"),
822 ));
823 }
824 };
825
826 let session = adapter_client.new_session(
827 SessionConfig {
828 conn_id: conn.conn_id().clone(),
829 uuid: conn_uuid,
830 user,
831 client_ip: conn.peer_addr().clone(),
832 external_metadata_rx: None,
833 helm_chart_version,
834 authenticator_kind: mz_auth::AuthenticatorKind::Password,
835 groups: None,
836 },
837 authenticated,
838 );
839
840 Ok(session)
841}
842
843#[derive(Debug)]
844enum State {
845 Ready,
846 Drain,
847 Done,
848}
849
850struct StateMachine<'a, A, I>
851where
852 I: Iterator<Item = TaskMetrics> + Send + 'a,
853{
854 conn: &'a mut FramedConn<A>,
855 adapter_client: mz_adapter::SessionClient,
856 txn_needs_commit: bool,
857 tokio_metrics_intervals: I,
858}
859
860enum SendRowsEndedReason {
861 Success {
862 result_size: u64,
863 rows_returned: u64,
864 },
865 Errored {
866 error: String,
867 },
868 Canceled,
869}
870
871const ABORTED_TXN_MSG: &str =
872 "current transaction is aborted, commands ignored until end of transaction block";
873
874impl<'a, A, I> StateMachine<'a, A, I>
875where
876 A: AsyncRead + AsyncWrite + AsyncReady + Send + Sync + Unpin + 'a,
877 I: Iterator<Item = TaskMetrics> + Send + 'a,
878{
879 #[allow(clippy::manual_async_fn)]
883 #[mz_ore::instrument(level = "debug")]
884 fn run(mut self) -> impl Future<Output = Result<(), io::Error>> + Send + 'a {
885 async move {
886 let mut state = State::Ready;
887 loop {
888 self.send_pending_notices().await?;
889 state = match state {
890 State::Ready => self.advance_ready().await?,
891 State::Drain => self.advance_drain().await?,
892 State::Done => return Ok(()),
893 };
894 self.adapter_client
895 .add_idle_in_transaction_session_timeout();
896 }
897 }
898 }
899
900 #[instrument(level = "debug")]
901 async fn advance_ready(&mut self) -> Result<State, io::Error> {
902 self.tokio_metrics_intervals
904 .next()
905 .expect("infinite iterator");
906
907 let message = select! {
909 biased;
910
911 Some(timeout) = self.adapter_client.recv_timeout() => {
913 let err: AdapterError = timeout.into();
914 let conn_id = self.adapter_client.session().conn_id();
915 tracing::warn!("session timed out, conn_id {}", conn_id);
916
917 let error_response = err.into_response(Severity::Fatal);
919 let error_state = self.send_error_and_get_state(error_response).await;
920
921 self.adapter_client.terminate().await;
923
924 let _ = self.conn.recv().await?;
928 return error_state;
929 },
930 message = self.conn.recv() => message?,
932 };
933
934 let interval = self
936 .tokio_metrics_intervals
937 .next()
938 .expect("infinite iterator");
939 let recv_scheduling_delay_ms = interval.total_scheduled_duration.as_secs_f64() * 1000.0;
940
941 let received = SYSTEM_TIME();
945
946 self.adapter_client
947 .remove_idle_in_transaction_session_timeout();
948
949 let message_name = message.as_ref().map(|m| m.name()).unwrap_or_default();
952
953 if let Some(message) = &message {
954 self.maybe_log_message_arrival(message).await;
955 }
956
957 let start = message.as_ref().map(|_| Instant::now());
958 let next_state = match message {
959 Some(FrontendMessage::Query { sql }) => {
960 let query_root_span =
961 tracing::info_span!(parent: None, "advance_ready", otel.name = message_name);
962 query_root_span.follows_from(tracing::Span::current());
963 self.query(sql, received)
964 .instrument(query_root_span)
965 .await?
966 }
967 Some(FrontendMessage::Parse {
968 name,
969 sql,
970 param_types,
971 }) => self.parse(name, sql, param_types).await?,
972 Some(FrontendMessage::Bind {
973 portal_name,
974 statement_name,
975 param_formats,
976 raw_params,
977 result_formats,
978 }) => {
979 self.bind(
980 portal_name,
981 statement_name,
982 param_formats,
983 raw_params,
984 result_formats,
985 )
986 .await?
987 }
988 Some(FrontendMessage::Execute {
989 portal_name,
990 max_rows,
991 }) => {
992 let max_rows = match usize::try_from(max_rows) {
993 Ok(0) | Err(_) => ExecuteCount::All, Ok(n) => ExecuteCount::Count(n),
995 };
996 let execute_root_span =
997 tracing::info_span!(parent: None, "advance_ready", otel.name = message_name);
998 execute_root_span.follows_from(tracing::Span::current());
999 let state = self
1000 .execute(
1001 portal_name,
1002 max_rows,
1003 portal_exec_message,
1004 None,
1005 ExecuteTimeout::None,
1006 None,
1007 Some(received),
1008 )
1009 .instrument(execute_root_span)
1010 .await?;
1011 let (is_implicit, may_span_pipeline) = {
1026 let txn = self.adapter_client.session().transaction();
1027 (txn.is_implicit(), txn.may_span_pipeline())
1028 };
1029 let spans_pipeline = may_span_pipeline
1032 && self
1033 .adapter_client
1034 .extended_protocol_implicit_transaction_enabled()
1035 .await;
1036 if is_implicit && !spans_pipeline {
1037 self.txn_needs_commit = true;
1038 }
1039 state
1040 }
1041 Some(FrontendMessage::DescribeStatement { name }) => {
1042 self.describe_statement(&name).await?
1043 }
1044 Some(FrontendMessage::DescribePortal { name }) => self.describe_portal(&name).await?,
1045 Some(FrontendMessage::CloseStatement { name }) => self.close_statement(name).await?,
1046 Some(FrontendMessage::ClosePortal { name }) => self.close_portal(name).await?,
1047 Some(FrontendMessage::Flush) => self.flush().await?,
1048 Some(FrontendMessage::Sync) => self.sync().await?,
1049 Some(FrontendMessage::Terminate) => State::Done,
1050
1051 Some(FrontendMessage::CopyData(_))
1058 | Some(FrontendMessage::CopyDone)
1059 | Some(FrontendMessage::CopyFail(_)) => State::Ready,
1060
1061 Some(FrontendMessage::Password { .. })
1062 | Some(FrontendMessage::RawAuthentication(_))
1063 | Some(FrontendMessage::SASLInitialResponse { .. })
1064 | Some(FrontendMessage::SASLResponse(_)) => State::Drain,
1065 None => State::Done,
1066 };
1067
1068 if let Some(start) = start {
1069 self.adapter_client
1070 .inner()
1071 .metrics()
1072 .pgwire_message_processing_seconds
1073 .with_label_values(&[message_name])
1074 .observe(start.elapsed().as_secs_f64());
1075 }
1076 self.adapter_client
1077 .inner()
1078 .metrics()
1079 .pgwire_recv_scheduling_delay_ms
1080 .with_label_values(&[message_name])
1081 .observe(recv_scheduling_delay_ms);
1082
1083 Ok(next_state)
1084 }
1085
1086 async fn advance_drain(&mut self) -> Result<State, io::Error> {
1087 let message = self.conn.recv().await?;
1088 if message.is_some() {
1089 self.adapter_client
1090 .remove_idle_in_transaction_session_timeout();
1091 }
1092 match message {
1093 Some(FrontendMessage::Sync) => self.sync().await,
1094 None => Ok(State::Done),
1095 _ => Ok(State::Drain),
1096 }
1097 }
1098
1099 #[instrument(level = "debug")]
1103 async fn one_query(
1104 &mut self,
1105 stmt: Statement<Raw>,
1106 sql: String,
1107 lifecycle_timestamps: LifecycleTimestamps,
1108 ) -> Result<State, io::Error> {
1109 const EMPTY_PORTAL: &str = "";
1112 if let Err(e) = self
1113 .adapter_client
1114 .declare(EMPTY_PORTAL.to_string(), stmt, sql)
1115 .await
1116 {
1117 return self
1118 .send_error_and_get_state(e.into_response(Severity::Error))
1119 .await;
1120 }
1121 let portal = self
1122 .adapter_client
1123 .session()
1124 .get_portal_unverified_mut(EMPTY_PORTAL)
1125 .expect("unnamed portal should be present");
1126
1127 *portal.lifecycle_timestamps = Some(lifecycle_timestamps);
1128
1129 let stmt_desc = portal.desc.clone();
1130 if !stmt_desc.param_types.is_empty() {
1131 return self
1132 .send_error_and_get_state(ErrorResponse::error(
1133 SqlState::UNDEFINED_PARAMETER,
1134 "there is no parameter $1",
1135 ))
1136 .await;
1137 }
1138
1139 if let Some(relation_desc) = &stmt_desc.relation_desc {
1141 if !stmt_desc.is_copy {
1142 let formats = vec![Format::Text; stmt_desc.arity()];
1143 self.send(BackendMessage::RowDescription(
1144 message::encode_row_description(relation_desc, &formats),
1145 ))
1146 .await?;
1147 }
1148 }
1149
1150 let result = match self
1151 .adapter_client
1152 .execute(EMPTY_PORTAL.to_string(), self.conn.wait_closed(), None)
1153 .await
1154 {
1155 Ok((response, execute_started)) => {
1156 self.send_pending_notices().await?;
1157 self.send_execute_response(
1158 response,
1159 stmt_desc.relation_desc,
1160 EMPTY_PORTAL.to_string(),
1161 ExecuteCount::All,
1162 portal_exec_message,
1163 None,
1164 ExecuteTimeout::None,
1165 execute_started,
1166 )
1167 .await
1168 }
1169 Err(e) => {
1170 self.send_pending_notices().await?;
1171 self.send_error_and_get_state(e.into_response(Severity::Error))
1172 .await
1173 }
1174 };
1175
1176 self.adapter_client.session().remove_portal(EMPTY_PORTAL);
1178
1179 result
1180 }
1181
1182 async fn ensure_transaction(
1183 &mut self,
1184 num_stmts: usize,
1185 message_type: &str,
1186 ) -> Result<(), io::Error> {
1187 let start = Instant::now();
1188 if self.txn_needs_commit {
1189 self.commit_transaction().await?;
1190 }
1191 let res = self.adapter_client.start_transaction(Some(num_stmts));
1194 assert_ok!(res);
1195 self.adapter_client
1196 .inner()
1197 .metrics()
1198 .pgwire_ensure_transaction_seconds
1199 .with_label_values(&[message_type])
1200 .observe(start.elapsed().as_secs_f64());
1201 Ok(())
1202 }
1203
1204 async fn maybe_log_message_arrival(&mut self, message: &FrontendMessage) {
1223 if !self
1224 .adapter_client
1225 .statement_arrival_logging_enabled()
1226 .await
1227 {
1228 return;
1229 }
1230 let session = self.adapter_client.session();
1231 let conn_id = session.conn_id();
1232 let session_uuid = session.uuid();
1233 let kind = message.name();
1234 match message {
1235 FrontendMessage::Query { sql } => {
1236 info!(
1237 %conn_id, %session_uuid, kind, sql = %redact_sql_for_logging(sql),
1238 "statement arrival"
1239 );
1240 }
1241 FrontendMessage::Parse { name, sql, .. } => {
1242 info!(
1243 %conn_id, %session_uuid, kind, name, sql = %redact_sql_for_logging(sql),
1244 "statement arrival"
1245 );
1246 }
1247 FrontendMessage::Bind {
1248 portal_name,
1249 statement_name,
1250 raw_params,
1251 ..
1252 } => {
1253 info!(
1254 %conn_id, %session_uuid, kind, portal_name, statement_name,
1255 num_params = raw_params.len(),
1256 "statement arrival"
1257 );
1258 }
1259 FrontendMessage::CopyData(data) => {
1261 info!(%conn_id, %session_uuid, kind, len = data.len(), "statement arrival");
1262 }
1263 FrontendMessage::Password { .. }
1265 | FrontendMessage::RawAuthentication(_)
1266 | FrontendMessage::SASLInitialResponse { .. }
1267 | FrontendMessage::SASLResponse(_) => {
1268 info!(%conn_id, %session_uuid, kind, "statement arrival");
1269 }
1270 FrontendMessage::CopyFail(_) => {
1273 info!(%conn_id, %session_uuid, kind, "statement arrival");
1274 }
1275 FrontendMessage::DescribeStatement { .. }
1278 | FrontendMessage::DescribePortal { .. }
1279 | FrontendMessage::Execute { .. }
1280 | FrontendMessage::Flush
1281 | FrontendMessage::Sync
1282 | FrontendMessage::CloseStatement { .. }
1283 | FrontendMessage::ClosePortal { .. }
1284 | FrontendMessage::Terminate
1285 | FrontendMessage::CopyDone => {
1286 info!(%conn_id, %session_uuid, kind, contents = ?message, "statement arrival");
1292 }
1293 }
1294 }
1295
1296 fn parse_sql<'b>(&self, sql: &'b str) -> Result<Vec<StatementParseResult<'b>>, ErrorResponse> {
1297 let parse_start = Instant::now();
1298 let result = match self.adapter_client.parse(sql) {
1299 Ok(result) => result.map_err(|e| {
1300 let pos = sql[..e.error.pos].chars().count() + 1;
1303 ErrorResponse::error(SqlState::SYNTAX_ERROR, e.error.message).with_position(pos)
1304 }),
1305 Err(msg) => Err(ErrorResponse::error(SqlState::PROGRAM_LIMIT_EXCEEDED, msg)),
1306 };
1307 self.adapter_client
1308 .inner()
1309 .metrics()
1310 .parse_seconds
1311 .observe(parse_start.elapsed().as_secs_f64());
1312 result
1313 }
1314
1315 #[instrument(level = "debug")]
1320 async fn query(&mut self, sql: String, received: EpochMillis) -> Result<State, io::Error> {
1321 let stmts = match self.parse_sql(&sql) {
1323 Ok(stmts) => stmts,
1324 Err(err) => {
1325 self.send_error_and_get_state(err).await?;
1326 return self.ready().await;
1327 }
1328 };
1329
1330 let num_stmts = stmts.len();
1331
1332 for StatementParseResult { ast: stmt, sql } in stmts {
1334 if self.is_aborted_txn() && !is_txn_exit_stmt(Some(&stmt)) {
1336 self.aborted_txn_error().await?;
1337 break;
1338 }
1339
1340 self.ensure_transaction(num_stmts, "query").await?;
1348
1349 match self
1350 .one_query(stmt, sql.to_string(), LifecycleTimestamps { received })
1351 .await?
1352 {
1353 State::Ready => (),
1354 State::Drain => break,
1355 State::Done => return Ok(State::Done),
1356 }
1357 }
1358
1359 {
1361 if self.adapter_client.session().transaction().is_implicit() {
1362 self.commit_transaction().await?;
1363 }
1364 }
1365
1366 if num_stmts == 0 {
1367 self.send(BackendMessage::EmptyQueryResponse).await?;
1368 }
1369
1370 self.ready().await
1371 }
1372
1373 #[instrument(level = "debug")]
1374 async fn parse(
1375 &mut self,
1376 name: String,
1377 sql: String,
1378 param_oids: Vec<u32>,
1379 ) -> Result<State, io::Error> {
1380 self.ensure_transaction(1, "parse").await?;
1382
1383 let mut param_types = vec![];
1384 for oid in param_oids {
1385 match mz_pgrepr::Type::from_oid(oid) {
1386 Ok(ty) => match SqlScalarType::try_from(&ty) {
1387 Ok(ty) => param_types.push(Some(ty)),
1388 Err(err) => {
1389 return self
1390 .send_error_and_get_state(ErrorResponse::error(
1391 SqlState::INVALID_PARAMETER_VALUE,
1392 err.to_string(),
1393 ))
1394 .await;
1395 }
1396 },
1397 Err(_) if oid == 0 => param_types.push(None),
1398 Err(e) => {
1399 return self
1400 .send_error_and_get_state(ErrorResponse::error(
1401 SqlState::PROTOCOL_VIOLATION,
1402 e.to_string(),
1403 ))
1404 .await;
1405 }
1406 }
1407 }
1408
1409 let stmts = match self.parse_sql(&sql) {
1410 Ok(stmts) => stmts,
1411 Err(err) => {
1412 return self.send_error_and_get_state(err).await;
1413 }
1414 };
1415 if stmts.len() > 1 {
1416 return self
1417 .send_error_and_get_state(ErrorResponse::error(
1418 SqlState::INTERNAL_ERROR,
1419 "cannot insert multiple commands into a prepared statement",
1420 ))
1421 .await;
1422 }
1423 let (maybe_stmt, sql) = match stmts.into_iter().next() {
1424 None => (None, ""),
1425 Some(StatementParseResult { ast, sql }) => (Some(ast), sql),
1426 };
1427 if self.is_aborted_txn() && !is_txn_exit_stmt(maybe_stmt.as_ref()) {
1428 return self.aborted_txn_error().await;
1429 }
1430 match self
1431 .adapter_client
1432 .prepare(name, maybe_stmt, sql.to_string(), param_types)
1433 .await
1434 {
1435 Ok(()) => {
1436 self.send(BackendMessage::ParseComplete).await?;
1437 Ok(State::Ready)
1438 }
1439 Err(e) => {
1440 self.send_error_and_get_state(e.into_response(Severity::Error))
1441 .await
1442 }
1443 }
1444 }
1445
1446 #[instrument(level = "debug")]
1448 async fn commit_transaction(&mut self) -> Result<(), io::Error> {
1449 self.end_transaction(EndTransactionAction::Commit).await
1450 }
1451
1452 #[instrument(level = "debug")]
1454 async fn rollback_transaction(&mut self) -> Result<(), io::Error> {
1455 self.end_transaction(EndTransactionAction::Rollback).await
1456 }
1457
1458 #[instrument(level = "debug")]
1465 async fn end_transaction(&mut self, action: EndTransactionAction) -> Result<(), io::Error> {
1466 self.txn_needs_commit = false;
1467 match self.adapter_client.end_transaction(action).await {
1468 Ok(
1469 ExecuteResponse::TransactionCommitted { params }
1470 | ExecuteResponse::TransactionRolledBack { params },
1471 ) => {
1472 self.send_parameter_statuses(params).await?;
1473 }
1474 Ok(_) => {}
1475 Err(err) => {
1476 self.send(BackendMessage::ErrorResponse(
1477 err.into_response(Severity::Error),
1478 ))
1479 .await?;
1480 }
1481 }
1482 Ok(())
1483 }
1484
1485 #[instrument(level = "debug")]
1488 async fn send_parameter_statuses(
1489 &mut self,
1490 params: BTreeMap<&'static str, String>,
1491 ) -> Result<(), io::Error> {
1492 let notify_set: mz_ore::collections::HashSet<String> = self
1493 .adapter_client
1494 .session()
1495 .vars()
1496 .notify_set()
1497 .map(|v| v.name().to_string())
1498 .collect();
1499
1500 for (name, value) in params
1501 .into_iter()
1502 .filter(|(name, _value)| notify_set.contains(*name))
1503 {
1504 self.send(BackendMessage::ParameterStatus(name, value))
1505 .await?;
1506 }
1507 Ok(())
1508 }
1509
1510 #[instrument(level = "debug")]
1511 async fn bind(
1512 &mut self,
1513 portal_name: String,
1514 statement_name: String,
1515 param_formats: Vec<Format>,
1516 raw_params: Vec<Option<Vec<u8>>>,
1517 result_formats: Vec<Format>,
1518 ) -> Result<State, io::Error> {
1519 self.ensure_transaction(1, "bind").await?;
1521
1522 let aborted_txn = self.is_aborted_txn();
1523 let stmt = match self
1524 .adapter_client
1525 .get_prepared_statement(&statement_name)
1526 .await
1527 {
1528 Ok(stmt) => stmt,
1529 Err(err) => {
1530 return self
1531 .send_error_and_get_state(err.into_response(Severity::Error))
1532 .await;
1533 }
1534 };
1535
1536 let param_types = &stmt.desc().param_types;
1537 if param_types.len() != raw_params.len() {
1538 let message = format!(
1539 "bind message supplies {actual} parameters, \
1540 but prepared statement \"{name}\" requires {expected}",
1541 name = statement_name,
1542 actual = raw_params.len(),
1543 expected = param_types.len()
1544 );
1545 return self
1546 .send_error_and_get_state(ErrorResponse::error(
1547 SqlState::PROTOCOL_VIOLATION,
1548 message,
1549 ))
1550 .await;
1551 }
1552 let param_formats = match pad_formats(param_formats, raw_params.len()) {
1553 Ok(param_formats) => param_formats,
1554 Err(msg) => {
1555 return self
1556 .send_error_and_get_state(ErrorResponse::error(
1557 SqlState::PROTOCOL_VIOLATION,
1558 msg,
1559 ))
1560 .await;
1561 }
1562 };
1563 if aborted_txn && !is_txn_exit_stmt(stmt.stmt()) {
1564 return self.aborted_txn_error().await;
1565 }
1566 let buf = RowArena::new();
1567 let mut params = vec![];
1568 for ((raw_param, mz_typ), format) in raw_params
1569 .into_iter()
1570 .zip_eq(param_types)
1571 .zip_eq(param_formats)
1572 {
1573 let pg_typ = mz_pgrepr::Type::from(mz_typ);
1574 let datum = match raw_param {
1575 None => Datum::Null,
1576 Some(bytes) => match mz_pgrepr::Value::decode(format, &pg_typ, &bytes) {
1577 Ok(param) => match param.into_datum_decode_error(&buf, &pg_typ, "parameter") {
1578 Ok(datum) => datum,
1579 Err(msg) => {
1580 return self
1581 .send_error_and_get_state(ErrorResponse::error(
1582 SqlState::INVALID_PARAMETER_VALUE,
1583 msg,
1584 ))
1585 .await;
1586 }
1587 },
1588 Err(err) => {
1589 let (code, msg) = if err.is::<mz_pgrepr::NulCharacterError>() {
1592 (SqlState::CHARACTER_NOT_IN_REPERTOIRE, err.to_string())
1593 } else {
1594 (
1595 SqlState::INVALID_PARAMETER_VALUE,
1596 format!("unable to decode parameter: {}", err),
1597 )
1598 };
1599 return self
1600 .send_error_and_get_state(ErrorResponse::error(code, msg))
1601 .await;
1602 }
1603 },
1604 };
1605 params.push((datum, mz_typ.clone()))
1606 }
1607
1608 let result_formats = match pad_formats(
1609 result_formats,
1610 stmt.desc()
1611 .relation_desc
1612 .clone()
1613 .map(|desc| desc.typ().column_types.len())
1614 .unwrap_or(0),
1615 ) {
1616 Ok(result_formats) => result_formats,
1617 Err(msg) => {
1618 return self
1619 .send_error_and_get_state(ErrorResponse::error(
1620 SqlState::PROTOCOL_VIOLATION,
1621 msg,
1622 ))
1623 .await;
1624 }
1625 };
1626
1627 if !stmt.stmt().map_or(false, |stmt| match stmt {
1630 Statement::Copy(CopyStatement {
1631 direction: CopyDirection::To,
1632 ..
1633 }) => true,
1634 Statement::Copy(CopyStatement {
1635 direction: CopyDirection::From,
1636 target: CopyTarget::Expr(_),
1640 ..
1641 }) => true,
1642 _ => false,
1643 }) {
1644 if let Some(desc) = stmt.desc().relation_desc.clone() {
1645 for (format, ty) in result_formats.iter().zip_eq(desc.iter_types()) {
1646 if let Format::Binary = format {
1647 if let Err(msg) = mz_pgrepr::Value::binary_encoding_error(&ty.scalar_type) {
1648 return self
1649 .send_error_and_get_state(ErrorResponse::error(
1650 SqlState::UNDEFINED_FUNCTION,
1651 msg,
1652 ))
1653 .await;
1654 }
1655 }
1656 }
1657 }
1658 }
1659
1660 let desc = stmt.desc().clone();
1661 let logging = Arc::clone(stmt.logging());
1662 let stmt_ast = stmt.stmt().cloned();
1663 let state_revision = stmt.state_revision;
1664 if let Err(err) = self.adapter_client.session().set_portal(
1665 portal_name,
1666 desc,
1667 stmt_ast,
1668 logging,
1669 params,
1670 result_formats,
1671 state_revision,
1672 ) {
1673 return self
1674 .send_error_and_get_state(err.into_response(Severity::Error))
1675 .await;
1676 }
1677
1678 self.send(BackendMessage::BindComplete).await?;
1679 Ok(State::Ready)
1680 }
1681
1682 fn execute(
1685 &mut self,
1686 portal_name: String,
1687 max_rows: ExecuteCount,
1688 get_response: GetResponse,
1689 fetch_portal_name: Option<String>,
1690 timeout: ExecuteTimeout,
1691 outer_ctx_extra: Option<ExecuteContextGuard>,
1692 received: Option<EpochMillis>,
1693 ) -> BoxFuture<'_, Result<State, io::Error>> {
1694 async move {
1695 let aborted_txn = self.is_aborted_txn();
1696
1697 let portal = match self
1699 .adapter_client
1700 .session()
1701 .get_portal_unverified_mut(&portal_name)
1702 {
1703 Some(portal) => portal,
1704 None => {
1705 let msg = format!("portal {} does not exist", portal_name.quoted());
1706 if let Some(outer_ctx_extra) = outer_ctx_extra {
1707 self.adapter_client.retire_execute(
1708 outer_ctx_extra,
1709 StatementEndedExecutionReason::Errored { error: msg.clone() },
1710 );
1711 }
1712 return self
1713 .send_error_and_get_state(ErrorResponse::error(
1714 SqlState::INVALID_CURSOR_NAME,
1715 msg,
1716 ))
1717 .await;
1718 }
1719 };
1720
1721 *portal.lifecycle_timestamps = received.map(LifecycleTimestamps::new);
1722
1723 let txn_exit_stmt = is_txn_exit_stmt(portal.stmt.as_deref());
1725 if aborted_txn && !txn_exit_stmt {
1726 if let Some(outer_ctx_extra) = outer_ctx_extra {
1727 self.adapter_client.retire_execute(
1728 outer_ctx_extra,
1729 StatementEndedExecutionReason::Errored {
1730 error: ABORTED_TXN_MSG.to_string(),
1731 },
1732 );
1733 }
1734 return self.aborted_txn_error().await;
1735 }
1736
1737 let row_desc = portal.desc.relation_desc.clone();
1738 match portal.state {
1739 PortalState::NotStarted => {
1740 self.ensure_transaction(1, "execute").await?;
1742 match self
1743 .adapter_client
1744 .execute(
1745 portal_name.clone(),
1746 self.conn.wait_closed(),
1747 outer_ctx_extra,
1748 )
1749 .await
1750 {
1751 Ok((response, execute_started)) => {
1752 self.send_pending_notices().await?;
1753 self.send_execute_response(
1754 response,
1755 row_desc,
1756 portal_name,
1757 max_rows,
1758 get_response,
1759 fetch_portal_name,
1760 timeout,
1761 execute_started,
1762 )
1763 .await
1764 }
1765 Err(e) => {
1766 self.send_pending_notices().await?;
1767 self.send_error_and_get_state(e.into_response(Severity::Error))
1768 .await
1769 }
1770 }
1771 }
1772 PortalState::InProgress(rows) => {
1773 let rows = rows.take().expect("InProgress rows must be populated");
1774 let (result, statement_ended_execution_reason) = match self
1775 .send_rows(
1776 row_desc.expect("portal missing row desc on resumption"),
1777 portal_name,
1778 rows,
1779 max_rows,
1780 get_response,
1781 fetch_portal_name,
1782 timeout,
1783 )
1784 .await
1785 {
1786 Err(e) => {
1787 (Err(e), StatementEndedExecutionReason::Canceled)
1790 }
1791 Ok((ok, SendRowsEndedReason::Canceled)) => {
1792 (Ok(ok), StatementEndedExecutionReason::Canceled)
1793 }
1794 Ok((
1807 ok,
1808 SendRowsEndedReason::Success {
1809 result_size: _,
1810 rows_returned: _,
1811 },
1812 )) => (
1813 Ok(ok),
1814 StatementEndedExecutionReason::Success {
1815 result_size: None,
1816 rows_returned: None,
1817 execution_strategy: None,
1818 },
1819 ),
1820 Ok((ok, SendRowsEndedReason::Errored { error })) => {
1821 (Ok(ok), StatementEndedExecutionReason::Errored { error })
1822 }
1823 };
1824 if let Some(outer_ctx_extra) = outer_ctx_extra {
1825 self.adapter_client
1826 .retire_execute(outer_ctx_extra, statement_ended_execution_reason);
1827 }
1828 result
1829 }
1830 PortalState::Completed(Some(tag)) => {
1837 let tag = tag.to_string();
1838 if let Some(outer_ctx_extra) = outer_ctx_extra {
1839 self.adapter_client.retire_execute(
1840 outer_ctx_extra,
1841 StatementEndedExecutionReason::Success {
1842 result_size: None,
1843 rows_returned: None,
1844 execution_strategy: None,
1845 },
1846 );
1847 }
1848 self.send(BackendMessage::CommandComplete { tag }).await?;
1849 Ok(State::Ready)
1850 }
1851 PortalState::Completed(None) => {
1852 let error = format!(
1853 "portal {} cannot be run",
1854 Ident::new_unchecked(portal_name).to_ast_string_stable()
1855 );
1856 if let Some(outer_ctx_extra) = outer_ctx_extra {
1857 self.adapter_client.retire_execute(
1858 outer_ctx_extra,
1859 StatementEndedExecutionReason::Errored {
1860 error: error.clone(),
1861 },
1862 );
1863 }
1864 self.send_error_and_get_state(ErrorResponse::error(
1865 SqlState::OBJECT_NOT_IN_PREREQUISITE_STATE,
1866 error,
1867 ))
1868 .await
1869 }
1870 }
1871 }
1872 .instrument(debug_span!("execute"))
1873 .boxed()
1874 }
1875
1876 #[instrument(level = "debug")]
1877 async fn describe_statement(&mut self, name: &str) -> Result<State, io::Error> {
1878 self.ensure_transaction(1, "describe_statement").await?;
1880
1881 let stmt = match self.adapter_client.get_prepared_statement(name).await {
1882 Ok(stmt) => stmt,
1883 Err(err) => {
1884 return self
1885 .send_error_and_get_state(err.into_response(Severity::Error))
1886 .await;
1887 }
1888 };
1889 let parameter_desc = BackendMessage::ParameterDescription(
1891 stmt.desc()
1892 .param_types
1893 .iter()
1894 .map(mz_pgrepr::Type::from)
1895 .collect(),
1896 );
1897 let formats = vec![Format::Text; stmt.desc().arity()];
1901 let row_desc = describe_rows(stmt.desc(), &formats);
1902 self.send_all([parameter_desc, row_desc]).await?;
1903 Ok(State::Ready)
1904 }
1905
1906 #[instrument(level = "debug")]
1907 async fn describe_portal(&mut self, name: &str) -> Result<State, io::Error> {
1908 self.ensure_transaction(1, "describe_portal").await?;
1910
1911 let session = self.adapter_client.session();
1912 let row_desc = session
1913 .get_portal_unverified(name)
1914 .map(|portal| describe_rows(&portal.desc, &portal.result_formats));
1915 match row_desc {
1916 Some(row_desc) => {
1917 self.send(row_desc).await?;
1918 Ok(State::Ready)
1919 }
1920 None => {
1921 self.send_error_and_get_state(ErrorResponse::error(
1922 SqlState::INVALID_CURSOR_NAME,
1923 format!("portal {} does not exist", name.quoted()),
1924 ))
1925 .await
1926 }
1927 }
1928 }
1929
1930 #[instrument(level = "debug")]
1931 async fn close_statement(&mut self, name: String) -> Result<State, io::Error> {
1932 self.adapter_client
1933 .session()
1934 .remove_prepared_statement(&name);
1935 self.send(BackendMessage::CloseComplete).await?;
1936 Ok(State::Ready)
1937 }
1938
1939 #[instrument(level = "debug")]
1940 async fn close_portal(&mut self, name: String) -> Result<State, io::Error> {
1941 self.adapter_client.session().remove_portal(&name);
1942 self.send(BackendMessage::CloseComplete).await?;
1943 Ok(State::Ready)
1944 }
1945
1946 fn complete_portal_after_close(&mut self, name: &str) {
1952 if let Some(portal) = self
1953 .adapter_client
1954 .session()
1955 .get_portal_unverified_mut(name)
1956 {
1957 *portal.state = PortalState::Completed(None);
1958 }
1959 }
1960
1961 fn complete_portal_after_declare(&mut self, name: &str) {
1967 match self
1968 .adapter_client
1969 .session()
1970 .get_portal_unverified_mut(name)
1971 {
1972 Some(portal) => *portal.state = PortalState::Completed(None),
1973 None => {
1974 mz_ore::soft_panic_or_log!("portal {} vanished across a DECLARE", name.quoted())
1975 }
1976 }
1977 }
1978
1979 async fn fetch(
1980 &mut self,
1981 name: String,
1982 count: Option<FetchDirection>,
1983 max_rows: ExecuteCount,
1984 fetch_portal_name: Option<String>,
1985 timeout: ExecuteTimeout,
1986 ctx_extra: ExecuteContextGuard,
1987 ) -> Result<State, io::Error> {
1988 let count = count.unwrap_or(FetchDirection::ForwardCount(1));
1991
1992 let count = match (max_rows, count) {
2005 (ExecuteCount::Count(max_rows), FetchDirection::ForwardCount(count)) => {
2006 let count = usize::cast_from(count);
2007 if max_rows < count {
2008 let msg = "Execute with max_rows < a FETCH's count is not supported";
2009 self.adapter_client.retire_execute(
2010 ctx_extra,
2011 StatementEndedExecutionReason::Errored {
2012 error: msg.to_string(),
2013 },
2014 );
2015 return self
2016 .send_error_and_get_state(ErrorResponse::error(
2017 SqlState::FEATURE_NOT_SUPPORTED,
2018 msg,
2019 ))
2020 .await;
2021 }
2022 ExecuteCount::Count(count)
2023 }
2024 (ExecuteCount::Count(_), FetchDirection::ForwardAll) => {
2025 let msg = "Execute with max_rows of a FETCH ALL is not supported";
2026 self.adapter_client.retire_execute(
2027 ctx_extra,
2028 StatementEndedExecutionReason::Errored {
2029 error: msg.to_string(),
2030 },
2031 );
2032 return self
2033 .send_error_and_get_state(ErrorResponse::error(
2034 SqlState::FEATURE_NOT_SUPPORTED,
2035 msg,
2036 ))
2037 .await;
2038 }
2039 (ExecuteCount::All, FetchDirection::ForwardAll) => ExecuteCount::All,
2040 (ExecuteCount::All, FetchDirection::ForwardCount(count)) => {
2041 ExecuteCount::Count(usize::cast_from(count))
2042 }
2043 };
2044 let cursor_name = name.to_string();
2045 self.execute(
2046 cursor_name,
2047 count,
2048 fetch_message,
2049 fetch_portal_name,
2050 timeout,
2051 Some(ctx_extra),
2052 None,
2053 )
2054 .await
2055 }
2056
2057 async fn flush(&mut self) -> Result<State, io::Error> {
2058 self.conn.flush().await?;
2059 Ok(State::Ready)
2060 }
2061
2062 #[instrument(level = "debug")]
2067 async fn send<M>(&mut self, message: M) -> Result<(), io::Error>
2068 where
2069 M: Into<BackendMessage>,
2070 {
2071 let message: BackendMessage = message.into();
2072 let is_error =
2073 matches!(&message, BackendMessage::ErrorResponse(e) if e.severity.is_error());
2074
2075 self.conn.send(message).await?;
2076
2077 if is_error {
2084 self.conn.flush().await?;
2085 }
2086
2087 Ok(())
2088 }
2089
2090 #[instrument(level = "debug")]
2091 pub async fn send_all(
2092 &mut self,
2093 messages: impl IntoIterator<Item = BackendMessage>,
2094 ) -> Result<(), io::Error> {
2095 for m in messages {
2096 self.send(m).await?;
2097 }
2098 Ok(())
2099 }
2100
2101 #[instrument(level = "debug")]
2102 async fn sync(&mut self) -> Result<State, io::Error> {
2103 if self.adapter_client.session().transaction().is_implicit() {
2105 self.commit_transaction().await?;
2106 }
2107 self.ready().await
2108 }
2109
2110 #[instrument(level = "debug")]
2111 async fn ready(&mut self) -> Result<State, io::Error> {
2112 let txn_state = self.adapter_client.session().transaction().into();
2113 self.send(BackendMessage::ReadyForQuery(txn_state)).await?;
2114 self.flush().await
2115 }
2116
2117 #[allow(clippy::too_many_arguments)]
2118 #[instrument(level = "debug")]
2119 async fn send_execute_response(
2120 &mut self,
2121 response: ExecuteResponse,
2122 row_desc: Option<RelationDesc>,
2123 portal_name: String,
2124 max_rows: ExecuteCount,
2125 get_response: GetResponse,
2126 fetch_portal_name: Option<String>,
2127 timeout: ExecuteTimeout,
2128 execute_started: Instant,
2129 ) -> Result<State, io::Error> {
2130 let mut tag = response.tag();
2131
2132 macro_rules! command_complete {
2133 () => {{
2134 self.send(BackendMessage::CommandComplete {
2135 tag: tag
2136 .take()
2137 .expect("command_complete only called on tag-generating results"),
2138 })
2139 .await?;
2140 Ok(State::Ready)
2141 }};
2142 }
2143
2144 let r = match response {
2145 ExecuteResponse::ClosedCursor => {
2146 self.complete_portal_after_close(&portal_name);
2147 command_complete!()
2148 }
2149 ExecuteResponse::DeclaredCursor => {
2150 self.complete_portal_after_declare(&portal_name);
2151 command_complete!()
2152 }
2153 ExecuteResponse::EmptyQuery => {
2154 self.send(BackendMessage::EmptyQueryResponse).await?;
2155 Ok(State::Ready)
2156 }
2157 ExecuteResponse::Fetch {
2158 name,
2159 count,
2160 timeout,
2161 ctx_extra,
2162 } => {
2163 self.fetch(
2164 name,
2165 count,
2166 max_rows,
2167 Some(portal_name.to_string()),
2168 timeout,
2169 ctx_extra,
2170 )
2171 .await
2172 }
2173 ExecuteResponse::SendingRowsStreaming {
2174 rows,
2175 instance_id,
2176 strategy,
2177 } => {
2178 let row_desc = row_desc
2179 .expect("missing row description for ExecuteResponse::SendingRowsStreaming");
2180
2181 let span = tracing::debug_span!("sending_rows_streaming");
2182
2183 self.send_rows(
2184 row_desc,
2185 portal_name,
2186 InProgressRows::new(RecordFirstRowStream::new(
2187 Box::new(rows),
2188 execute_started,
2189 &self.adapter_client,
2190 Some(instance_id),
2191 Some(strategy),
2192 )),
2193 max_rows,
2194 get_response,
2195 fetch_portal_name,
2196 timeout,
2197 )
2198 .instrument(span)
2199 .await
2200 .map(|(state, _)| state)
2201 }
2202 ExecuteResponse::SendingRowsImmediate { rows } => {
2203 let row_desc = row_desc
2204 .expect("missing row description for ExecuteResponse::SendingRowsImmediate");
2205
2206 let span = tracing::debug_span!("sending_rows_immediate");
2207
2208 let stream =
2209 futures::stream::once(futures::future::ready(PeekResponseUnary::Rows(rows)));
2210 self.send_rows(
2211 row_desc,
2212 portal_name,
2213 InProgressRows::new(RecordFirstRowStream::new(
2214 Box::new(stream),
2215 execute_started,
2216 &self.adapter_client,
2217 None,
2218 Some(StatementExecutionStrategy::Constant),
2219 )),
2220 max_rows,
2221 get_response,
2222 fetch_portal_name,
2223 timeout,
2224 )
2225 .instrument(span)
2226 .await
2227 .map(|(state, _)| state)
2228 }
2229 ExecuteResponse::SetVariable { name, .. } => {
2230 let qn = name.to_string();
2233 let msg = if let Some(var) = self
2234 .adapter_client
2235 .session()
2236 .vars_mut()
2237 .notify_set()
2238 .find(|v| v.name() == qn)
2239 {
2240 Some(BackendMessage::ParameterStatus(var.name(), var.value()))
2241 } else {
2242 None
2243 };
2244 if let Some(msg) = msg {
2245 self.send(msg).await?;
2246 }
2247 command_complete!()
2248 }
2249 ExecuteResponse::Subscribing {
2250 rx,
2251 ctx_extra,
2252 instance_id,
2253 } => {
2254 if fetch_portal_name.is_none() {
2255 let mut msg = ErrorResponse::notice(
2256 SqlState::WARNING,
2257 "streaming SUBSCRIBE rows directly requires a client that does not buffer output",
2258 );
2259 if self.adapter_client.session().vars().application_name() == "psql" {
2260 msg.hint = Some(
2261 "Wrap your SUBSCRIBE statement in `COPY (SUBSCRIBE ...) TO STDOUT`."
2262 .into(),
2263 )
2264 }
2265 self.send(msg).await?;
2266 self.conn.flush().await?;
2267 }
2268 let row_desc =
2269 row_desc.expect("missing row description for ExecuteResponse::Subscribing");
2270 let (result, statement_ended_execution_reason) = match self
2271 .send_rows(
2272 row_desc,
2273 portal_name,
2274 InProgressRows::new(RecordFirstRowStream::new(
2275 rx,
2276 execute_started,
2277 &self.adapter_client,
2278 Some(instance_id),
2279 None,
2280 )),
2281 max_rows,
2282 get_response,
2283 fetch_portal_name,
2284 timeout,
2285 )
2286 .await
2287 {
2288 Err(e) => {
2289 (Err(e), StatementEndedExecutionReason::Canceled)
2292 }
2293 Ok((ok, SendRowsEndedReason::Canceled)) => {
2294 (Ok(ok), StatementEndedExecutionReason::Canceled)
2295 }
2296 Ok((
2297 ok,
2298 SendRowsEndedReason::Success {
2299 result_size,
2300 rows_returned,
2301 },
2302 )) => (
2303 Ok(ok),
2304 StatementEndedExecutionReason::Success {
2305 result_size: Some(result_size),
2306 rows_returned: Some(rows_returned),
2307 execution_strategy: None,
2308 },
2309 ),
2310 Ok((ok, SendRowsEndedReason::Errored { error })) => {
2311 (Ok(ok), StatementEndedExecutionReason::Errored { error })
2312 }
2313 };
2314 self.adapter_client
2315 .retire_execute(ctx_extra, statement_ended_execution_reason);
2316 return result;
2317 }
2318 ExecuteResponse::CopyTo { format, resp } => {
2319 let row_desc =
2320 row_desc.expect("missing row description for ExecuteResponse::CopyTo");
2321 match *resp {
2322 ExecuteResponse::Subscribing {
2323 rx,
2324 ctx_extra,
2325 instance_id,
2326 } => {
2327 let (result, statement_ended_execution_reason) = match self
2328 .copy_rows(
2329 format,
2330 row_desc,
2331 RecordFirstRowStream::new(
2332 rx,
2333 execute_started,
2334 &self.adapter_client,
2335 Some(instance_id),
2336 None,
2337 ),
2338 )
2339 .await
2340 {
2341 Err(e) => {
2342 (Err(e), StatementEndedExecutionReason::Canceled)
2345 }
2346 Ok((
2347 state,
2348 SendRowsEndedReason::Success {
2349 result_size,
2350 rows_returned,
2351 },
2352 )) => (
2353 Ok(state),
2354 StatementEndedExecutionReason::Success {
2355 result_size: Some(result_size),
2356 rows_returned: Some(rows_returned),
2357 execution_strategy: None,
2358 },
2359 ),
2360 Ok((state, SendRowsEndedReason::Errored { error })) => {
2361 (Ok(state), StatementEndedExecutionReason::Errored { error })
2362 }
2363 Ok((state, SendRowsEndedReason::Canceled)) => {
2364 (Ok(state), StatementEndedExecutionReason::Canceled)
2365 }
2366 };
2367 self.adapter_client
2368 .retire_execute(ctx_extra, statement_ended_execution_reason);
2369 return result;
2370 }
2371 ExecuteResponse::SendingRowsStreaming {
2372 rows,
2373 instance_id,
2374 strategy,
2375 } => {
2376 return self
2381 .copy_rows(
2382 format,
2383 row_desc,
2384 RecordFirstRowStream::new(
2385 Box::new(rows),
2386 execute_started,
2387 &self.adapter_client,
2388 Some(instance_id),
2389 Some(strategy),
2390 ),
2391 )
2392 .await
2393 .map(|(state, _)| state);
2394 }
2395 ExecuteResponse::SendingRowsImmediate { rows } => {
2396 let span = tracing::debug_span!("sending_rows_immediate");
2397
2398 let rows = futures::stream::once(futures::future::ready(
2399 PeekResponseUnary::Rows(rows),
2400 ));
2401 return self
2406 .copy_rows(
2407 format,
2408 row_desc,
2409 RecordFirstRowStream::new(
2410 Box::new(rows),
2411 execute_started,
2412 &self.adapter_client,
2413 None,
2414 Some(StatementExecutionStrategy::Constant),
2415 ),
2416 )
2417 .instrument(span)
2418 .await
2419 .map(|(state, _)| state);
2420 }
2421 _ => {
2422 return self
2423 .send_error_and_get_state(ErrorResponse::error(
2424 SqlState::INTERNAL_ERROR,
2425 "unsupported COPY response type".to_string(),
2426 ))
2427 .await;
2428 }
2429 };
2430 }
2431 ExecuteResponse::CopyFrom {
2432 target_id,
2433 target_name,
2434 columns,
2435 params,
2436 ctx_extra,
2437 } => {
2438 let row_desc =
2439 row_desc.expect("missing row description for ExecuteResponse::CopyFrom");
2440 self.copy_from(target_id, target_name, columns, params, row_desc, ctx_extra)
2441 .await
2442 }
2443 ExecuteResponse::TransactionCommitted { params }
2444 | ExecuteResponse::TransactionRolledBack { params }
2445 | ExecuteResponse::DiscardedAll { params } => {
2446 self.send_parameter_statuses(params).await?;
2447 command_complete!()
2448 }
2449
2450 ExecuteResponse::AlteredDefaultPrivileges
2451 | ExecuteResponse::AlteredObject(..)
2452 | ExecuteResponse::AlteredRole
2453 | ExecuteResponse::AlteredSystemConfiguration
2454 | ExecuteResponse::CreatedCluster { .. }
2455 | ExecuteResponse::CreatedClusterReplica { .. }
2456 | ExecuteResponse::CreatedConnection { .. }
2457 | ExecuteResponse::CreatedDatabase { .. }
2458 | ExecuteResponse::CreatedIndex { .. }
2459 | ExecuteResponse::CreatedMetricSink { .. }
2460 | ExecuteResponse::CreatedIntrospectionSubscribe
2461 | ExecuteResponse::CreatedMaterializedView { .. }
2462 | ExecuteResponse::CreatedRole
2463 | ExecuteResponse::CreatedSchema { .. }
2464 | ExecuteResponse::CreatedSecret { .. }
2465 | ExecuteResponse::CreatedSink { .. }
2466 | ExecuteResponse::CreatedSource { .. }
2467 | ExecuteResponse::CreatedTable { .. }
2468 | ExecuteResponse::CreatedType
2469 | ExecuteResponse::CreatedView { .. }
2470 | ExecuteResponse::CreatedViews { .. }
2471 | ExecuteResponse::CreatedNetworkPolicy
2472 | ExecuteResponse::Comment
2473 | ExecuteResponse::Deallocate { .. }
2474 | ExecuteResponse::Deleted(..)
2475 | ExecuteResponse::DiscardedTemp
2476 | ExecuteResponse::DroppedObject(_)
2477 | ExecuteResponse::DroppedOwned
2478 | ExecuteResponse::GrantedPrivilege
2479 | ExecuteResponse::GrantedRole
2480 | ExecuteResponse::Inserted(..)
2481 | ExecuteResponse::Copied(..)
2482 | ExecuteResponse::Prepare
2483 | ExecuteResponse::Raised
2484 | ExecuteResponse::ReassignOwned
2485 | ExecuteResponse::RevokedPrivilege
2486 | ExecuteResponse::RevokedRole
2487 | ExecuteResponse::StartedTransaction { .. }
2488 | ExecuteResponse::Updated(..)
2489 | ExecuteResponse::ValidatedConnection => {
2490 command_complete!()
2491 }
2492 };
2493
2494 assert_none!(tag, "tag created but not consumed: {:?}", tag);
2495 r
2496 }
2497
2498 #[allow(clippy::too_many_arguments)]
2499 #[mz_ore::instrument(level = "debug")]
2501 async fn send_rows(
2502 &mut self,
2503 row_desc: RelationDesc,
2504 portal_name: String,
2505 mut rows: InProgressRows,
2506 max_rows: ExecuteCount,
2507 get_response: GetResponse,
2508 fetch_portal_name: Option<String>,
2509 timeout: ExecuteTimeout,
2510 ) -> Result<(State, SendRowsEndedReason), io::Error> {
2511 let result_format_portal_name: &str = if let Some(ref name) = fetch_portal_name {
2514 name
2515 } else {
2516 &portal_name
2517 };
2518 let result_formats = self
2519 .adapter_client
2520 .session()
2521 .get_portal_unverified(result_format_portal_name)
2522 .expect("valid fetch portal name for send rows")
2523 .result_formats
2524 .clone();
2525
2526 let (mut wait_once, mut deadline) = match timeout {
2527 ExecuteTimeout::None => (false, None),
2528 ExecuteTimeout::Seconds(t) => (
2529 false,
2530 Some(tokio::time::Instant::now() + tokio::time::Duration::from_secs_f64(t)),
2531 ),
2532 ExecuteTimeout::WaitOnce => (true, None),
2533 };
2534
2535 {
2537 let portal_name_desc = &self
2538 .adapter_client
2539 .session()
2540 .get_portal_unverified(portal_name.as_str())
2541 .expect("portal should exist")
2542 .desc
2543 .relation_desc;
2544 if let Some(portal_name_desc) = portal_name_desc {
2545 soft_assert_eq_or_log!(portal_name_desc, &row_desc);
2546 }
2547 if let Some(fetch_portal_name) = &fetch_portal_name {
2548 let fetch_portal_desc = &self
2549 .adapter_client
2550 .session()
2551 .get_portal_unverified(fetch_portal_name)
2552 .expect("portal should exist")
2553 .desc
2554 .relation_desc;
2555 if let Some(fetch_portal_desc) = fetch_portal_desc {
2556 soft_assert_eq_or_log!(fetch_portal_desc, &row_desc);
2557 }
2558 }
2559 }
2560
2561 self.conn.set_encode_state(
2562 row_desc
2563 .typ()
2564 .column_types
2565 .iter()
2566 .map(|ty| mz_pgrepr::Type::from(&ty.scalar_type))
2567 .zip_eq(result_formats)
2568 .collect(),
2569 self.adapter_client.session().vars().text_encode_settings(),
2570 );
2571
2572 let mut total_sent_rows = 0;
2573 let mut total_sent_bytes = 0;
2574 let mut want_rows = match max_rows {
2576 ExecuteCount::All => usize::MAX,
2577 ExecuteCount::Count(count) => count,
2578 };
2579
2580 loop {
2582 let batch = if rows.current.is_some() {
2585 FetchResult::Rows(rows.current.take())
2586 } else if want_rows == 0 {
2587 FetchResult::Rows(None)
2588 } else {
2589 let notice_fut = self.adapter_client.session().recv_notice();
2590 tokio::select! {
2605 biased;
2606 err = self.conn.wait_closed() => return Err(err),
2607 batch = rows.remaining.recv() => match batch {
2608 None => FetchResult::Rows(None),
2609 Some(PeekResponseUnary::Rows(rows)) => FetchResult::Rows(Some(rows)),
2610 Some(PeekResponseUnary::Error(err)) => {
2611 FetchResult::Error(err.into_response(Severity::Error))
2612 }
2613 Some(PeekResponseUnary::DependencyDropped(dep)) => {
2614 FetchResult::Error(
2615 dep.to_concurrent_dependency_drop()
2616 .into_response(Severity::Error),
2617 )
2618 }
2619 Some(PeekResponseUnary::Canceled) => FetchResult::Canceled,
2620 },
2621 notice = notice_fut => {
2622 FetchResult::Notice(notice)
2623 }
2624 _ = time::sleep_until(
2625 deadline.unwrap_or_else(tokio::time::Instant::now),
2626 ), if deadline.is_some() => FetchResult::Rows(None),
2627 }
2628 };
2629
2630 match batch {
2631 FetchResult::Rows(None) => break,
2632 FetchResult::Rows(Some(mut batch_rows)) => {
2633 if let Err(err) = verify_datum_desc(&row_desc, &mut batch_rows) {
2634 let msg = err.to_string();
2635 return self
2636 .send_error_and_get_state(err.into_response(Severity::Error))
2637 .await
2638 .map(|state| (state, SendRowsEndedReason::Errored { error: msg }));
2639 }
2640
2641 if wait_once && batch_rows.peek().is_some() {
2645 deadline = Some(tokio::time::Instant::now());
2646 wait_once = false;
2647 }
2648
2649 let mut sent_rows = 0;
2651 let mut sent_bytes = 0;
2652 let messages = (&mut batch_rows)
2653 .map(|row| {
2658 let row_len = row.byte_len();
2659 let values = mz_pgrepr::values_from_row(row, row_desc.typ());
2660 (row_len, BackendMessage::DataRow(values))
2661 })
2662 .inspect(|(row_len, _)| {
2663 sent_bytes += row_len;
2664 sent_rows += 1
2665 })
2666 .map(|(_row_len, row)| row)
2667 .take(want_rows);
2668 self.send_all(messages).await?;
2669
2670 total_sent_rows += sent_rows;
2671 total_sent_bytes += sent_bytes;
2672 want_rows -= sent_rows;
2673
2674 if want_rows == 0 {
2677 if batch_rows.peek().is_some() {
2678 rows.current = Some(batch_rows);
2679 }
2680 break;
2681 }
2682
2683 self.conn.flush().await?;
2684 }
2685 FetchResult::Notice(notice) => {
2686 self.send(notice.into_response()).await?;
2687 self.conn.flush().await?;
2688 }
2689 FetchResult::Error(err) => {
2690 let text = err.message.clone();
2691 return self
2692 .send_error_and_get_state(err)
2693 .await
2694 .map(|state| (state, SendRowsEndedReason::Errored { error: text }));
2695 }
2696 FetchResult::Canceled => {
2697 return self
2698 .send_error_and_get_state(ErrorResponse::error(
2699 SqlState::QUERY_CANCELED,
2700 "canceling statement due to user request",
2701 ))
2702 .await
2703 .map(|state| (state, SendRowsEndedReason::Canceled));
2704 }
2705 }
2706 }
2707
2708 let portal = self
2709 .adapter_client
2710 .session()
2711 .get_portal_unverified_mut(&portal_name)
2712 .expect("valid portal name for send rows");
2713
2714 let saw_rows = rows.remaining.saw_rows;
2715 let no_more_rows = rows.no_more_rows();
2716 let metric_recorded = rows.remaining.metric_recorded;
2717 let recorded_first_row_instant = rows.remaining.recorded_first_row_instant;
2718
2719 if no_more_rows && !metric_recorded {
2720 rows.remaining.metric_recorded = true;
2721 }
2722
2723 *portal.state = PortalState::InProgress(Some(rows));
2726
2727 let fetch_portal = fetch_portal_name.map(|name| {
2728 self.adapter_client
2729 .session()
2730 .get_portal_unverified_mut(&name)
2731 .expect("valid fetch portal")
2732 });
2733 let response_message = get_response(max_rows, total_sent_rows, fetch_portal);
2734 self.send(response_message).await?;
2735
2736 if no_more_rows && !metric_recorded {
2739 let statement_type = if let Some(stmt) = &self
2740 .adapter_client
2741 .session()
2742 .get_portal_unverified(&portal_name)
2743 .expect("valid portal name for send_rows")
2744 .stmt
2745 {
2746 metrics::statement_type_label_value(stmt.deref())
2747 } else {
2748 "no-statement"
2749 };
2750 let duration = if saw_rows {
2751 recorded_first_row_instant
2752 .expect("recorded_first_row_instant because saw_rows")
2753 .elapsed()
2754 } else {
2755 Duration::ZERO
2759 };
2760 self.adapter_client
2761 .inner()
2762 .metrics()
2763 .result_rows_first_to_last_byte_seconds
2764 .with_label_values(&[statement_type])
2765 .observe(duration.as_secs_f64());
2766 }
2767
2768 Ok((
2769 State::Ready,
2770 SendRowsEndedReason::Success {
2771 result_size: u64::cast_from(total_sent_bytes),
2772 rows_returned: u64::cast_from(total_sent_rows),
2773 },
2774 ))
2775 }
2776
2777 #[mz_ore::instrument(level = "debug")]
2778 async fn copy_rows(
2779 &mut self,
2780 format: CopyFormat,
2781 row_desc: RelationDesc,
2782 mut stream: RecordFirstRowStream,
2783 ) -> Result<(State, SendRowsEndedReason), io::Error> {
2784 let (row_format, encode_format) = match format {
2785 CopyFormat::Text => (
2786 CopyFormatParams::Text(CopyTextFormatParams::default()),
2787 Format::Text,
2788 ),
2789 CopyFormat::Binary => (CopyFormatParams::Binary, Format::Binary),
2790 CopyFormat::Csv => (
2791 CopyFormatParams::Csv(CopyCsvFormatParams::default()),
2792 Format::Text,
2793 ),
2794 CopyFormat::Parquet => {
2795 let text = "Parquet format is not supported".to_string();
2796 return self
2797 .send_error_and_get_state(ErrorResponse::error(
2798 SqlState::INTERNAL_ERROR,
2799 text.clone(),
2800 ))
2801 .await
2802 .map(|state| (state, SendRowsEndedReason::Errored { error: text }));
2803 }
2804 };
2805
2806 if let CopyFormat::Binary = format {
2814 if let Some(msg) = row_desc
2815 .iter_types()
2816 .find_map(|ty| mz_pgrepr::Value::binary_encoding_error(&ty.scalar_type).err())
2817 {
2818 return self
2819 .send_error_and_get_state(ErrorResponse::error(
2820 SqlState::UNDEFINED_FUNCTION,
2821 msg,
2822 ))
2823 .await
2824 .map(|state| {
2825 (
2826 state,
2827 SendRowsEndedReason::Errored {
2828 error: msg.to_string(),
2829 },
2830 )
2831 });
2832 }
2833 }
2834
2835 let text_settings = self.adapter_client.session().vars().text_encode_settings();
2839 let encode_fn = |row: &RowRef, typ: &SqlRelationType, out: &mut Vec<u8>| {
2840 mz_pgcopy::encode_copy_format(&row_format, row, typ, out, text_settings)
2841 };
2842
2843 let typ = row_desc.typ();
2844 let column_formats = iter::repeat(encode_format)
2845 .take(typ.column_types.len())
2846 .collect();
2847 self.send(BackendMessage::CopyOutResponse {
2848 overall_format: encode_format,
2849 column_formats,
2850 })
2851 .await?;
2852
2853 let mut out = Vec::new();
2858
2859 if let CopyFormat::Binary = format {
2860 out.extend(b"PGCOPY\n\xFF\r\n\0");
2862 out.extend([0, 0, 0, 0]);
2864 out.extend([0, 0, 0, 0]);
2866 }
2867
2868 let mut count = 0;
2869 let mut total_sent_bytes = 0;
2870 loop {
2871 tokio::select! {
2872 e = self.conn.wait_closed() => return Err(e),
2873 batch = stream.recv() => match batch {
2874 None => break,
2875 Some(PeekResponseUnary::Error(err)) => {
2876 let text = err.to_string();
2877 return self
2878 .send_error_and_get_state(err.into_response(Severity::Error))
2879 .await
2880 .map(|state| (state, SendRowsEndedReason::Errored { error: text }));
2881 }
2882 Some(PeekResponseUnary::DependencyDropped(dep)) => {
2883 let err = dep.to_concurrent_dependency_drop();
2884 let text = err.to_string();
2885 let resp = err.into_response(Severity::Error);
2886 return self
2887 .send_error_and_get_state(resp)
2888 .await
2889 .map(|state| (state, SendRowsEndedReason::Errored { error: text }));
2890 }
2891 Some(PeekResponseUnary::Canceled) => {
2892 return self.send_error_and_get_state(ErrorResponse::error(
2893 SqlState::QUERY_CANCELED,
2894 "canceling statement due to user request",
2895 ))
2896 .await.map(|state| (state, SendRowsEndedReason::Canceled));
2897 }
2898 Some(PeekResponseUnary::Rows(mut rows)) => {
2899 count += rows.count();
2900 while let Some(row) = rows.next() {
2901 total_sent_bytes += row.byte_len();
2902 encode_fn(row, typ, &mut out)?;
2903 self.send(BackendMessage::CopyData(mem::take(&mut out)))
2904 .await?;
2905 }
2906 }
2907 },
2908 notice = self.adapter_client.session().recv_notice() => {
2909 self.send(notice.into_response())
2910 .await?;
2911 self.conn.flush().await?;
2912 }
2913 }
2914
2915 self.conn.flush().await?;
2916 }
2917 if let CopyFormat::Binary = format {
2919 let trailer: i16 = -1;
2920 out.extend(trailer.to_be_bytes());
2921 self.send(BackendMessage::CopyData(mem::take(&mut out)))
2922 .await?;
2923 }
2924
2925 let tag = format!("COPY {}", count);
2926 self.send(BackendMessage::CopyDone).await?;
2927 self.send(BackendMessage::CommandComplete { tag }).await?;
2928 Ok((
2929 State::Ready,
2930 SendRowsEndedReason::Success {
2931 result_size: u64::cast_from(total_sent_bytes),
2932 rows_returned: u64::cast_from(count),
2933 },
2934 ))
2935 }
2936
2937 #[instrument(level = "debug")]
2940 async fn copy_from(
2941 &mut self,
2942 target_id: CatalogItemId,
2943 target_name: String,
2944 columns: Vec<ColumnIndex>,
2945 params: CopyFormatParams<'static>,
2946 row_desc: RelationDesc,
2947 mut ctx_extra: ExecuteContextGuard,
2948 ) -> Result<State, io::Error> {
2949 let res = self
2950 .copy_from_inner(
2951 target_id,
2952 target_name,
2953 columns,
2954 params,
2955 row_desc,
2956 &mut ctx_extra,
2957 )
2958 .await;
2959 match &res {
2960 Ok(State::Ready) => {
2961 self.adapter_client.retire_execute(
2962 ctx_extra,
2963 StatementEndedExecutionReason::Success {
2964 result_size: None,
2965 rows_returned: None,
2966 execution_strategy: None,
2967 },
2968 );
2969 }
2970 Ok(State::Done) => {
2971 self.adapter_client
2975 .retire_execute(ctx_extra, StatementEndedExecutionReason::Canceled);
2976 }
2977 Err(e) => {
2978 self.adapter_client.retire_execute(
2979 ctx_extra,
2980 StatementEndedExecutionReason::Errored {
2981 error: format!("{e}"),
2982 },
2983 );
2984 }
2985 Ok(State::Drain) => {}
2986 }
2987 res
2988 }
2989
2990 async fn copy_from_inner(
2991 &mut self,
2992 target_id: CatalogItemId,
2993 target_name: String,
2994 columns: Vec<ColumnIndex>,
2995 params: CopyFormatParams<'static>,
2996 row_desc: RelationDesc,
2997 ctx_extra: &mut ExecuteContextGuard,
2998 ) -> Result<State, io::Error> {
2999 let typ = row_desc.typ();
3000 let column_formats = vec![Format::Text; typ.column_types.len()];
3001 self.send(BackendMessage::CopyInResponse {
3002 overall_format: Format::Text,
3003 column_formats,
3004 })
3005 .await?;
3006 self.conn.flush().await?;
3007
3008 let writer = match self
3010 .adapter_client
3011 .start_copy_from_stdin(
3012 target_id,
3013 target_name.clone(),
3014 columns.clone(),
3015 row_desc.clone(),
3016 params.clone(),
3017 )
3018 .await
3019 {
3020 Ok(writer) => writer,
3021 Err(e) => {
3022 loop {
3028 match self.conn.recv().await? {
3029 Some(FrontendMessage::CopyData(_)) => {}
3030 Some(FrontendMessage::CopyDone) | Some(FrontendMessage::CopyFail(_)) => {
3031 break;
3032 }
3033 Some(FrontendMessage::Flush) | Some(FrontendMessage::Sync) => {}
3034 Some(_) => break,
3035 None => return Ok(State::Done),
3036 }
3037 }
3038 self.adapter_client.retire_execute(
3039 std::mem::take(ctx_extra),
3040 StatementEndedExecutionReason::Errored {
3041 error: e.to_string(),
3042 },
3043 );
3044 return self
3045 .send_error_and_get_state(e.into_response(Severity::Error))
3046 .await;
3047 }
3048 };
3049
3050 const BATCH_SIZE: usize = 32 * 1024 * 1024;
3052 let max_copy_from_row_size = self
3053 .adapter_client
3054 .get_system_vars()
3055 .await
3056 .max_copy_from_row_size()
3057 .try_into()
3058 .unwrap_or(usize::MAX);
3059
3060 let mut data = Vec::new();
3061 let mut row_scanner = CopyRowScanner::new(¶ms);
3062 let num_workers = writer.batch_txs.len();
3063 let mut next_worker: usize = 0;
3064 let mut saw_copy_done = false;
3065 let mut saw_end_marker = false;
3066 let mut copy_from_error: Option<(SqlState, String)> = None;
3067
3068 loop {
3071 let message = self.conn.recv().await?;
3072 match message {
3073 Some(FrontendMessage::CopyData(buf)) => {
3074 if saw_end_marker {
3075 continue;
3078 }
3079 data.extend(buf);
3080 row_scanner.scan_new_bytes(&data);
3081
3082 if let Some(end_pos) = row_scanner.end_marker_end() {
3083 data.truncate(end_pos);
3084 row_scanner.on_truncate(end_pos);
3085 saw_end_marker = true;
3086 }
3087
3088 if row_scanner.current_row_size(data.len()) > max_copy_from_row_size {
3090 copy_from_error = Some((
3091 SqlState::INSUFFICIENT_RESOURCES,
3092 format!(
3093 "COPY FROM STDIN row exceeded max_copy_from_row_size \
3094 ({max_copy_from_row_size} bytes)"
3095 ),
3096 ));
3097 break;
3098 }
3099
3100 let mut send_failed = false;
3103 while data.len() >= BATCH_SIZE {
3104 let split_pos = match row_scanner.last_row_end() {
3105 Some(pos) => pos,
3106 None => break, };
3108 let remainder = data.split_off(split_pos);
3109 let chunk = std::mem::replace(&mut data, remainder);
3110 row_scanner.on_split(split_pos);
3111 if writer.batch_txs[next_worker].send(chunk).await.is_err() {
3112 send_failed = true;
3113 break;
3114 }
3115 next_worker = (next_worker + 1) % num_workers;
3116 }
3117 if send_failed {
3120 break;
3121 }
3122 }
3123 Some(FrontendMessage::CopyDone) => {
3124 if !data.is_empty() {
3126 let chunk = std::mem::take(&mut data);
3127 let _ = writer.batch_txs[next_worker].send(chunk).await;
3129 }
3130 saw_copy_done = true;
3131 break;
3132 }
3133 Some(FrontendMessage::CopyFail(err)) => {
3134 self.adapter_client.retire_execute(
3135 std::mem::take(ctx_extra),
3136 StatementEndedExecutionReason::Canceled,
3137 );
3138 drop(writer);
3140 return self
3141 .send_error_and_get_state(ErrorResponse::error(
3142 SqlState::QUERY_CANCELED,
3143 format!("COPY from stdin failed: {}", err),
3144 ))
3145 .await;
3146 }
3147 Some(FrontendMessage::Flush) | Some(FrontendMessage::Sync) => {}
3148 Some(_) => {
3149 let msg = "unexpected message type during COPY from stdin";
3150 self.adapter_client.retire_execute(
3151 std::mem::take(ctx_extra),
3152 StatementEndedExecutionReason::Errored {
3153 error: msg.to_string(),
3154 },
3155 );
3156 drop(writer);
3157 return self
3158 .send_error_and_get_state(ErrorResponse::error(
3159 SqlState::PROTOCOL_VIOLATION,
3160 msg,
3161 ))
3162 .await;
3163 }
3164 None => {
3165 drop(writer);
3166 return Ok(State::Done);
3167 }
3168 }
3169 }
3170
3171 if !saw_copy_done {
3175 loop {
3176 match self.conn.recv().await? {
3177 Some(FrontendMessage::CopyData(_)) => {}
3178 Some(FrontendMessage::CopyDone) | Some(FrontendMessage::CopyFail(_)) => {
3179 break;
3180 }
3181 Some(FrontendMessage::Flush) | Some(FrontendMessage::Sync) => {}
3182 Some(_) => {
3183 let msg = "unexpected message type during COPY from stdin";
3184 self.adapter_client.retire_execute(
3185 std::mem::take(ctx_extra),
3186 StatementEndedExecutionReason::Errored {
3187 error: msg.to_string(),
3188 },
3189 );
3190 drop(writer);
3191 return self
3192 .send_error_and_get_state(ErrorResponse::error(
3193 SqlState::PROTOCOL_VIOLATION,
3194 msg,
3195 ))
3196 .await;
3197 }
3198 None => {
3199 drop(writer);
3200 return Ok(State::Done);
3201 }
3202 }
3203 }
3204 }
3205
3206 if let Some((code, msg)) = copy_from_error {
3207 self.adapter_client.retire_execute(
3208 std::mem::take(ctx_extra),
3209 StatementEndedExecutionReason::Errored { error: msg.clone() },
3210 );
3211 drop(writer);
3212 return self
3213 .send_error_and_get_state(ErrorResponse::error(code, msg))
3214 .await;
3215 }
3216
3217 drop(writer.batch_txs);
3222
3223 let (proto_batches, row_count) = match writer.completion_rx.await {
3225 Ok(Ok(result)) => result,
3226 Ok(Err(e)) => {
3227 self.adapter_client.retire_execute(
3228 std::mem::take(ctx_extra),
3229 StatementEndedExecutionReason::Errored {
3230 error: e.to_string(),
3231 },
3232 );
3233 return self
3234 .send_error_and_get_state(e.into_response(Severity::Error))
3235 .await;
3236 }
3237 Err(_) => {
3238 let msg = "COPY FROM STDIN: background batch builder tasks dropped";
3239 self.adapter_client.retire_execute(
3240 std::mem::take(ctx_extra),
3241 StatementEndedExecutionReason::Errored {
3242 error: msg.to_string(),
3243 },
3244 );
3245 return self
3246 .send_error_and_get_state(ErrorResponse::error(SqlState::INTERNAL_ERROR, msg))
3247 .await;
3248 }
3249 };
3250
3251 if let Err(e) = self
3253 .adapter_client
3254 .stage_copy_from_stdin_batches(target_id, proto_batches)
3255 {
3256 self.adapter_client.retire_execute(
3257 std::mem::take(ctx_extra),
3258 StatementEndedExecutionReason::Errored {
3259 error: e.to_string(),
3260 },
3261 );
3262 return self
3263 .send_error_and_get_state(e.into_response(Severity::Error))
3264 .await;
3265 }
3266
3267 let tag = format!("COPY {}", row_count);
3268 self.send(BackendMessage::CommandComplete { tag }).await?;
3269
3270 Ok(State::Ready)
3271 }
3272
3273 #[instrument(level = "debug")]
3274 async fn send_pending_notices(&mut self) -> Result<(), io::Error> {
3275 let notices = self
3276 .adapter_client
3277 .session()
3278 .drain_notices()
3279 .into_iter()
3280 .map(|notice| BackendMessage::ErrorResponse(notice.into_response()));
3281 self.send_all(notices).await?;
3282 Ok(())
3283 }
3284
3285 #[instrument(level = "debug")]
3286 async fn send_error_and_get_state(&mut self, err: ErrorResponse) -> Result<State, io::Error> {
3287 assert!(err.severity.is_error());
3288 debug!(
3289 "cid={} error code={}",
3290 self.adapter_client.session().conn_id(),
3291 err.code.code()
3292 );
3293 let is_fatal = err.severity.is_fatal();
3294 self.send(BackendMessage::ErrorResponse(err)).await?;
3295
3296 let txn = self.adapter_client.session().transaction();
3297 match txn {
3298 TransactionStatus::Default | TransactionStatus::Failed(_) => {}
3301 TransactionStatus::Started(_) => {
3303 self.rollback_transaction().await?;
3304 }
3305 TransactionStatus::InTransactionImplicit(_) => {
3307 self.rollback_transaction().await?;
3308 }
3309 TransactionStatus::InTransaction(_) => {
3311 self.adapter_client.fail_transaction();
3312 }
3313 };
3314 if is_fatal {
3315 Ok(State::Done)
3316 } else {
3317 Ok(State::Drain)
3318 }
3319 }
3320
3321 #[instrument(level = "debug")]
3322 async fn aborted_txn_error(&mut self) -> Result<State, io::Error> {
3323 self.send(BackendMessage::ErrorResponse(ErrorResponse::error(
3324 SqlState::IN_FAILED_SQL_TRANSACTION,
3325 ABORTED_TXN_MSG,
3326 )))
3327 .await?;
3328 Ok(State::Drain)
3329 }
3330
3331 fn is_aborted_txn(&mut self) -> bool {
3332 matches!(
3333 self.adapter_client.session().transaction(),
3334 TransactionStatus::Failed(_)
3335 )
3336 }
3337}
3338
3339fn pad_formats(formats: Vec<Format>, n: usize) -> Result<Vec<Format>, String> {
3340 match (formats.len(), n) {
3341 (0, e) => Ok(vec![Format::Text; e]),
3342 (1, e) => Ok(iter::repeat(formats[0]).take(e).collect()),
3343 (a, e) if a == e => Ok(formats),
3344 (a, e) => Err(format!(
3345 "expected {} field format specifiers, but got {}",
3346 e, a
3347 )),
3348 }
3349}
3350
3351fn describe_rows(stmt_desc: &StatementDesc, formats: &[Format]) -> BackendMessage {
3352 match &stmt_desc.relation_desc {
3353 Some(desc) if !stmt_desc.is_copy => {
3354 BackendMessage::RowDescription(message::encode_row_description(desc, formats))
3355 }
3356 _ => BackendMessage::NoData,
3357 }
3358}
3359
3360type GetResponse = fn(
3361 max_rows: ExecuteCount,
3362 total_sent_rows: usize,
3363 fetch_portal: Option<PortalRefMut>,
3364) -> BackendMessage;
3365
3366fn portal_exec_message(
3369 max_rows: ExecuteCount,
3370 total_sent_rows: usize,
3371 _fetch_portal: Option<PortalRefMut>,
3372) -> BackendMessage {
3373 match max_rows {
3380 ExecuteCount::Count(max_rows) if max_rows <= total_sent_rows => {
3381 BackendMessage::PortalSuspended
3382 }
3383 _ => BackendMessage::CommandComplete {
3384 tag: format!("SELECT {}", total_sent_rows),
3385 },
3386 }
3387}
3388
3389fn fetch_message(
3391 _max_rows: ExecuteCount,
3392 total_sent_rows: usize,
3393 fetch_portal: Option<PortalRefMut>,
3394) -> BackendMessage {
3395 let tag = format!("FETCH {}", total_sent_rows);
3396 if let Some(portal) = fetch_portal {
3397 *portal.state = PortalState::Completed(Some(tag.clone()));
3398 }
3399 BackendMessage::CommandComplete { tag }
3400}
3401
3402fn get_authenticator(
3403 authenticator_kind: listeners::AuthenticatorKind,
3404 frontegg: Option<FronteggAuthenticator>,
3405 oidc: GenericOidcAuthenticator,
3406 adapter_client: mz_adapter::Client,
3407) -> Authenticator {
3408 match authenticator_kind {
3409 listeners::AuthenticatorKind::Frontegg => Authenticator::Frontegg(frontegg.expect(
3410 "Frontegg authenticator should exist with listeners::AuthenticatorKind::Frontegg",
3411 )),
3412 listeners::AuthenticatorKind::Password => Authenticator::Password(adapter_client),
3413 listeners::AuthenticatorKind::Sasl => Authenticator::Sasl(adapter_client),
3414 listeners::AuthenticatorKind::Oidc => Authenticator::Oidc(oidc),
3415 listeners::AuthenticatorKind::None => Authenticator::None,
3416 }
3417}
3418
3419#[derive(Debug, Copy, Clone)]
3420enum ExecuteCount {
3421 All,
3422 Count(usize),
3423}
3424
3425fn is_txn_exit_stmt(stmt: Option<&Statement<Raw>>) -> bool {
3427 match stmt {
3428 Some(stmt) => matches!(stmt, Statement::Commit(_) | Statement::Rollback(_)),
3430 None => false,
3431 }
3432}
3433
3434#[derive(Debug)]
3435enum FetchResult {
3436 Rows(Option<Box<dyn RowIterator + Send + Sync>>),
3437 Canceled,
3438 Error(ErrorResponse),
3439 Notice(AdapterNotice),
3440}
3441
3442#[derive(Debug)]
3443struct CopyRowScanner {
3444 scan_pos: usize,
3445 last_row_end: Option<usize>,
3446 end_marker_end: Option<usize>,
3447 record_start: usize,
3452 csv: Option<CsvScanState>,
3453}
3454
3455#[derive(Debug)]
3456struct CsvScanState {
3457 reader: csv_core::Reader,
3458 output: Vec<u8>,
3459 ends: Vec<usize>,
3460 skip_first_record: bool,
3461}
3462
3463impl CopyRowScanner {
3464 fn new(params: &CopyFormatParams<'_>) -> Self {
3465 let csv = match params {
3466 CopyFormatParams::Csv(CopyCsvFormatParams {
3467 delimiter,
3468 quote,
3469 escape,
3470 header,
3471 ..
3472 }) => Some(CsvScanState::new(*delimiter, *quote, *escape, *header)),
3473 _ => None,
3474 };
3475
3476 CopyRowScanner {
3477 scan_pos: 0,
3478 last_row_end: None,
3479 end_marker_end: None,
3480 record_start: 0,
3481 csv,
3482 }
3483 }
3484
3485 fn scan_new_bytes(&mut self, data: &[u8]) {
3486 if self.scan_pos >= data.len() {
3487 return;
3488 }
3489
3490 if let Some(csv) = self.csv.as_mut() {
3491 let mut input = &data[self.scan_pos..];
3492 let mut consumed = 0usize;
3493 while !input.is_empty() {
3494 let (result, n_input, _n_output, _n_ends) =
3495 csv.reader
3496 .read_record(input, &mut csv.output, &mut csv.ends);
3497 consumed += n_input;
3498 input = &input[n_input..];
3499
3500 match result {
3501 ReadRecordResult::InputEmpty => break,
3502 ReadRecordResult::OutputFull => {
3503 if n_input == 0 {
3504 csv.output
3505 .resize(csv.output.len().saturating_mul(2).max(1), 0);
3506 }
3507 }
3508 ReadRecordResult::OutputEndsFull => {
3509 if n_input == 0 {
3510 csv.ends.resize(csv.ends.len().saturating_mul(2).max(1), 0);
3511 }
3512 }
3513 ReadRecordResult::Record | ReadRecordResult::End => {
3514 let row_end = self.scan_pos + consumed;
3515 self.last_row_end = Some(row_end);
3516 if self.end_marker_end.is_none() {
3517 let is_marker = if csv.skip_first_record {
3518 csv.skip_first_record = false;
3519 false
3520 } else {
3521 let raw = &data[self.record_start..row_end];
3527 let start = raw
3538 .iter()
3539 .take_while(|&&b| b == b'\r' || b == b'\n')
3540 .count();
3541 let trailing = raw[start..]
3542 .iter()
3543 .rev()
3544 .take_while(|&&b| b == b'\r' || b == b'\n')
3545 .count();
3546 let trimmed = &raw[start..raw.len() - trailing];
3547 trimmed == b"\\."
3548 };
3549 if is_marker {
3550 self.end_marker_end = Some(row_end);
3551 self.record_start = row_end;
3552 break;
3553 }
3554 }
3555 self.record_start = row_end;
3556 }
3557 }
3558 }
3559 } else {
3560 let mut row_start = self.last_row_end.unwrap_or(0);
3561 for (offset, b) in data[self.scan_pos..].iter().enumerate() {
3562 if *b == b'\n' {
3563 let row_end = self.scan_pos + offset + 1;
3564 self.last_row_end = Some(row_end);
3565 if self.end_marker_end.is_none() {
3566 let row = &data[row_start..row_end];
3567 if row.get(0..2) == Some(b"\\.") {
3568 self.end_marker_end = Some(row_end);
3569 break;
3570 }
3571 }
3572 row_start = row_end;
3573 }
3574 }
3575 }
3576
3577 self.scan_pos = data.len();
3578 }
3579
3580 fn last_row_end(&self) -> Option<usize> {
3581 self.last_row_end
3582 }
3583
3584 fn end_marker_end(&self) -> Option<usize> {
3585 self.end_marker_end
3586 }
3587
3588 fn current_row_size(&self, data_len: usize) -> usize {
3589 data_len.saturating_sub(self.last_row_end.unwrap_or(0))
3590 }
3591
3592 fn on_split(&mut self, split_pos: usize) {
3593 self.scan_pos = self.scan_pos.saturating_sub(split_pos);
3594 self.last_row_end = None;
3595 self.end_marker_end = self
3596 .end_marker_end
3597 .and_then(|end| end.checked_sub(split_pos));
3598 soft_assert_or_log!(
3606 self.csv.is_none() || self.record_start >= split_pos,
3607 "split bisected an in-progress CSV record: record_start={} < split_pos={}",
3608 self.record_start,
3609 split_pos,
3610 );
3611 self.record_start = self.record_start.saturating_sub(split_pos);
3612 }
3613
3614 fn on_truncate(&mut self, new_len: usize) {
3615 self.scan_pos = self.scan_pos.min(new_len);
3616 self.last_row_end = self.last_row_end.filter(|&end| end <= new_len);
3617 self.end_marker_end = self.end_marker_end.filter(|&end| end <= new_len);
3618 self.record_start = self.record_start.min(new_len);
3619 }
3620}
3621
3622impl CsvScanState {
3623 fn new(delimiter: u8, quote: u8, escape: u8, header: bool) -> Self {
3624 let (double_quote, escape) = if quote == escape {
3625 (true, None)
3626 } else {
3627 (false, Some(escape))
3628 };
3629 CsvScanState {
3630 reader: csv_core::ReaderBuilder::new()
3631 .delimiter(delimiter)
3632 .quote(quote)
3633 .double_quote(double_quote)
3634 .escape(escape)
3635 .build(),
3636 output: vec![0; 1],
3637 ends: vec![0; 1],
3638 skip_first_record: header,
3639 }
3640 }
3641}
3642
3643#[cfg(test)]
3644mod test {
3645 use super::*;
3646
3647 #[mz_ore::test]
3648 fn test_copy_row_scanner_end_marker_line_endings() {
3649 let params = CopyFormatParams::Csv(CopyCsvFormatParams::default());
3657
3658 let marker_end = |data: &[u8]| -> Option<usize> {
3659 let mut scanner = CopyRowScanner::new(¶ms);
3660 scanner.scan_new_bytes(data);
3661 scanner.end_marker_end()
3662 };
3663
3664 for eol in [&b"\n"[..], b"\r\n", b"\r"] {
3665 let join = |lines: &[&str]| -> Vec<u8> {
3666 let mut out = Vec::new();
3667 for line in lines {
3668 out.extend_from_slice(line.as_bytes());
3669 out.extend_from_slice(eol);
3670 }
3671 out
3672 };
3673
3674 let data = join(&["first", "\\.", "after"]);
3679 let mut prefix = Vec::new();
3680 prefix.extend_from_slice(b"first");
3681 prefix.extend_from_slice(eol);
3682 prefix.extend_from_slice(b"\\.");
3683 assert_eq!(
3684 marker_end(&data),
3685 Some(prefix.len() + 1),
3686 "bare marker, eol={eol:?}"
3687 );
3688
3689 let data = join(&["before", "\"\\.\"", "after"]);
3691 assert_eq!(marker_end(&data), None, "quoted marker, eol={eol:?}");
3692 }
3693 }
3694
3695 #[mz_ore::test]
3696 fn test_copy_row_scanner_non_csv_split() {
3697 for params in [
3704 CopyFormatParams::Text(CopyTextFormatParams::default()),
3705 CopyFormatParams::Binary,
3706 ] {
3707 let mut scanner = CopyRowScanner::new(¶ms);
3708 let data = b"1\thello world\t2\tsome text value here\n\
3709 3\thello world\t6\tsome text value here\n";
3710 scanner.scan_new_bytes(data);
3711 let split_pos = scanner.last_row_end().expect("a complete row");
3712 assert!(split_pos > 0, "params={params:?}");
3713 scanner.on_split(split_pos);
3715 assert_eq!(scanner.record_start, 0, "params={params:?}");
3716 }
3717 }
3718
3719 #[mz_ore::test]
3720 fn test_parse_options() {
3721 struct TestCase {
3722 input: &'static str,
3723 expect: Result<Vec<(&'static str, &'static str)>, ()>,
3724 }
3725 let tests = vec![
3726 TestCase {
3727 input: "",
3728 expect: Ok(vec![]),
3729 },
3730 TestCase {
3731 input: "--key",
3732 expect: Err(()),
3733 },
3734 TestCase {
3735 input: "--key=val",
3736 expect: Ok(vec![("key", "val")]),
3737 },
3738 TestCase {
3739 input: r#"--key=val -ckey2=val2 -c key3=val3 -c key4=val4 -ckey5=val5"#,
3740 expect: Ok(vec![
3741 ("key", "val"),
3742 ("key2", "val2"),
3743 ("key3", "val3"),
3744 ("key4", "val4"),
3745 ("key5", "val5"),
3746 ]),
3747 },
3748 TestCase {
3749 input: r#"-c\ key=val"#,
3750 expect: Ok(vec![(" key", "val")]),
3751 },
3752 TestCase {
3753 input: "--key=val -ckey2 val2",
3754 expect: Err(()),
3755 },
3756 TestCase {
3758 input: "--key=",
3759 expect: Ok(vec![("key", "")]),
3760 },
3761 ];
3762 for test in tests {
3763 let got = parse_options(test.input);
3764 let expect = test.expect.map(|r| {
3765 r.into_iter()
3766 .map(|(k, v)| (k.to_owned(), v.to_owned()))
3767 .collect()
3768 });
3769 assert_eq!(got, expect, "input: {}", test.input);
3770 }
3771 }
3772
3773 #[mz_ore::test]
3774 fn test_parse_option() {
3775 struct TestCase {
3776 input: &'static str,
3777 expect: Result<(&'static str, &'static str), ()>,
3778 }
3779 let tests = vec![
3780 TestCase {
3781 input: "",
3782 expect: Err(()),
3783 },
3784 TestCase {
3785 input: "--",
3786 expect: Err(()),
3787 },
3788 TestCase {
3789 input: "--c",
3790 expect: Err(()),
3791 },
3792 TestCase {
3793 input: "a=b",
3794 expect: Err(()),
3795 },
3796 TestCase {
3797 input: "--a=b",
3798 expect: Ok(("a", "b")),
3799 },
3800 TestCase {
3801 input: "--ca=b",
3802 expect: Ok(("ca", "b")),
3803 },
3804 TestCase {
3805 input: "-ca=b",
3806 expect: Ok(("a", "b")),
3807 },
3808 TestCase {
3810 input: "--=",
3811 expect: Ok(("", "")),
3812 },
3813 ];
3814 for test in tests {
3815 let got = parse_option(test.input);
3816 assert_eq!(got, test.expect, "input: {}", test.input);
3817 }
3818 }
3819
3820 #[mz_ore::test]
3821 fn test_split_options() {
3822 struct TestCase {
3823 input: &'static str,
3824 expect: Vec<&'static str>,
3825 }
3826 let tests = vec![
3827 TestCase {
3828 input: "",
3829 expect: vec![],
3830 },
3831 TestCase {
3832 input: " ",
3833 expect: vec![],
3834 },
3835 TestCase {
3836 input: " a ",
3837 expect: vec!["a"],
3838 },
3839 TestCase {
3840 input: " ab cd ",
3841 expect: vec!["ab", "cd"],
3842 },
3843 TestCase {
3844 input: r#" ab\ cd "#,
3845 expect: vec!["ab ", "cd"],
3846 },
3847 TestCase {
3848 input: r#" ab\\ cd "#,
3849 expect: vec![r#"ab\"#, "cd"],
3850 },
3851 TestCase {
3852 input: r#" ab\\\ cd "#,
3853 expect: vec![r#"ab\ "#, "cd"],
3854 },
3855 TestCase {
3856 input: r#" ab\\\ cd "#,
3857 expect: vec![r#"ab\ cd"#],
3858 },
3859 TestCase {
3860 input: r#" ab\\\cd "#,
3861 expect: vec![r#"ab\cd"#],
3862 },
3863 TestCase {
3864 input: r#"a\"#,
3865 expect: vec!["a"],
3866 },
3867 TestCase {
3868 input: r#"a\ "#,
3869 expect: vec!["a "],
3870 },
3871 TestCase {
3872 input: r#"\"#,
3873 expect: vec![],
3874 },
3875 TestCase {
3876 input: r#"\ "#,
3877 expect: vec![r#" "#],
3878 },
3879 TestCase {
3880 input: r#" \ "#,
3881 expect: vec![r#" "#],
3882 },
3883 TestCase {
3884 input: r#"\ "#,
3885 expect: vec![r#" "#],
3886 },
3887 ];
3888 for test in tests {
3889 let got = split_options(test.input);
3890 assert_eq!(got, test.expect, "input: {}", test.input);
3891 }
3892 }
3893
3894 #[mz_ore::test]
3895 fn test_is_jwt() {
3896 assert!(is_jwt("eyJhbGciOiJIUzI1NiJ9.eyJzdWIiOiIxIn0.signature"));
3898 for s in [
3900 "",
3901 "secure_password",
3902 "p4ss.w0rd",
3903 "aaa.bbb.ccc",
3904 "eyJhbGciOiJIUzI1NiJ9.eyJzdWIiOiIxIn0",
3905 "eyJhbGciOiJIUzI1NiJ9.eyJzdWIiOiIxIn0.sig.extra",
3906 ] {
3907 assert!(!is_jwt(s), "is_jwt({s:?})");
3908 }
3909 }
3910}