Skip to main content

mz_pgwire/
protocol.rs

1// Copyright Materialize, Inc. and contributors. All rights reserved.
2//
3// Use of this software is governed by the Business Source License
4// included in the LICENSE file.
5//
6// As of the Change Date specified in that file, in accordance with
7// the Business Source License, use of this software will be governed
8// by the Apache License, Version 2.0.
9
10use 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
80/// Reports whether the given stream begins with a pgwire handshake.
81///
82/// To avoid false negatives, there must be at least eight bytes in `buf`.
83pub fn match_handshake(buf: &[u8]) -> bool {
84    // The pgwire StartupMessage looks like this:
85    //
86    //     i32 - Length of entire message.
87    //     i32 - Protocol version number.
88    //     [String] - Arbitrary key-value parameters of any length.
89    //
90    // Since arbitrary parameters can be included in the StartupMessage, the
91    // first Int32 is worthless, since the message could have any length.
92    // Instead, we sniff the protocol version number.
93    if buf.len() < 8 {
94        return false;
95    }
96    let version = NetworkEndian::read_i32(&buf[4..8]);
97    VERSIONS.contains(&version)
98}
99
100/// Parameters for the [`run`] function.
101pub struct RunParams<'a, A, I>
102where
103    I: Iterator<Item = TaskMetrics> + Send,
104{
105    /// The TLS mode of the pgwire server.
106    pub tls_mode: Option<TlsMode>,
107    /// A client for the adapter.
108    pub adapter_client: mz_adapter::Client,
109    /// The connection to the client.
110    pub conn: &'a mut FramedConn<A>,
111    /// The universally unique identifier for the connection.
112    pub conn_uuid: Uuid,
113    /// The protocol version that the client provided in the startup message.
114    pub version: i32,
115    /// The parameters that the client provided in the startup message.
116    pub params: BTreeMap<String, String>,
117    /// Frontegg JWT authenticator.
118    pub frontegg: Option<FronteggAuthenticator>,
119    /// OIDC authenticator.
120    pub oidc: GenericOidcAuthenticator,
121    /// The authentication method defined by the server's listener
122    /// configuration.
123    pub authenticator_kind: listeners::AuthenticatorKind,
124    /// Global connection limit and count
125    pub active_connection_counter: ConnectionCounter,
126    /// Helm chart version
127    pub helm_chart_version: Option<String>,
128    /// Whether to allow reserved users (ie: mz_system).
129    pub allowed_roles: AllowedRoles,
130    /// Tokio metrics
131    pub tokio_metrics_intervals: I,
132}
133
134/// Runs a pgwire connection to completion.
135///
136/// This involves responding to `FrontendMessage::StartupMessage` and all future
137/// requests until the client terminates the connection or a fatal error occurs.
138///
139/// Note that this function returns successfully even upon delivering a fatal
140/// error to the client. It only returns `Err` if an unexpected I/O error occurs
141/// while communicating with the client, e.g., if the connection is severed in
142/// the middle of a request.
143#[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    // TODO move this somewhere it can be shared with HTTP
179    let is_internal_user = INTERNAL_USER_NAMES.contains(&user);
180    // this is a superset of internal users
181    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                // Create a session based on the auth session.
217                //
218                // In particular, it's important that the username come from the
219                // auth session, as Frontegg may return an email address with
220                // different casing than the user supplied via the pgwire
221                // username fN
222                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            // OIDC listener: accepts either a JWT (uses OIDC authentication) or a
253            // plain SQL password (uses SQL password authentication).
254            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                        // No invalidation of the auth session once authenticated,
280                        // so auth session lasts indefinitely.
281                        (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            // No frontegg check, so auth session lasts indefinitely.
333            (session, pending().right_future())
334        }
335        Authenticator::Sasl(adapter_client) => {
336            // Start the handshake
337            conn.send(BackendMessage::AuthenticationSASL).await?;
338            conn.flush().await?;
339            // Get the initial response indicating chosen mechanism
340            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                            // We do not support channel binding
349                            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            // No frontegg check, so auth session lasts indefinitely.
517            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            // No frontegg check, so auth session lasts indefinitely.
536            let auth_session = pending().right_future();
537            (session, auth_session)
538        }
539    };
540
541    let system_vars = adapter_client.get_system_vars().await;
542    // Startup parameters that were successfully applied. They additionally
543    // become the session's default values below, once role defaults have been
544    // applied too.
545    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            // TODO: Issuing an error here is better than what we did before
563            // (silently ignore errors on set), but erroring the connection
564            // might be the better behavior. We maybe need to support more
565            // options sent by psql and drivers before we can safely do this.
566            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    // Register session with adapter.
593    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    // Make the startup parameters the session's default values, so that RESET
599    // and DISCARD ALL restore them rather than the server defaults. This
600    // matches PostgreSQL, where client-supplied startup parameters take
601    // precedence over role defaults (which startup registration applied) both
602    // as the current value and as the reset value. Connection poolers rely on
603    // this. For example, pgbouncer's default server_reset_query is DISCARD
604    // ALL, which must not rebind a pooled connection to the default database.
605    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            // Unexpected, since the same value was accepted by set() above.
612            mz_ore::soft_panic_or_log!("failed to apply startup parameter as default: {err:?}");
613        }
614    }
615
616    // Authentication succeeded, so the connection may now carry query traffic,
617    // whose frames are far larger than any credential.
618    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            // Errors produced internally (like a malformed frame header) should send an
651            // error to the client informing them why the connection was closed. We still want to
652            // return the original error up the stack, though, so we skip error checking during conn
653            // operations.
654            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
674/// Decides if a given password is a JWT by checking
675/// if we can decode its header.
676fn is_jwt(password: &str) -> bool {
677    jsonwebtoken::decode_header(password).is_ok()
678}
679
680/// Returns (name, value) session settings pairs from an options value.
681///
682/// From Postgres, see pg_split_opts in postinit.c and process_postgres_switches
683/// in postgres.c.
684fn 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
705/// Returns the parsed key and value from option of the form `--key=value`, `-c
706/// key=value`, or `-ckey=value`. Keys replace `-` with `_`. Returns an error if
707/// there was some other prefix.
708fn 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
718/// Splits value by any number of spaces except those preceded by `\`.
719fn split_options(value: &str) -> Vec<String> {
720    let mut strs = Vec::new();
721    // Need to build a string because of the escaping, so we can't simply
722    // subslice into value, and this isn't called enough to need to make it
723    // smart so it only builds a string if needed.
724    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                    // To ignore multiple spaces in a row, only push if current
733                    // is not empty.
734                    strs.push(std::mem::take(&mut current));
735                }
736                false
737            }
738            '\\' => {
739                if was_slash {
740                    // Two slashes in a row will add a slash and not escape the
741                    // next char.
742                    current.push('\\');
743                    false
744                } else {
745                    true
746                }
747            }
748            _ => {
749                current.push(c);
750                false
751            }
752        };
753    }
754    // A `\` at the end will be ignored.
755    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
772/// Requests a cleartext password from a connection and returns it if it is valid.
773/// Sends an error response in the connection if the password
774/// is not valid.
775async 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
803/// Helper for password-based authentication using AdapterClient
804/// and returns an authenticated session.
805async 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    // Manually desugar this (don't use `async fn run`) here because a much better
880    // error message is produced if there are problems with Send or other traits
881    // somewhere within the Future.
882    #[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        // Start a new metrics interval before the `recv()` call.
903        self.tokio_metrics_intervals
904            .next()
905            .expect("infinite iterator");
906
907        // Handle timeouts first so we don't execute any statements when there's a pending timeout.
908        let message = select! {
909            biased;
910
911            // `recv_timeout()` is cancel-safe as per it's docs.
912            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                // Process the error, doing any state cleanup.
918                let error_response = err.into_response(Severity::Fatal);
919                let error_state = self.send_error_and_get_state(error_response).await;
920
921                // Terminate __after__ we do any cleanup.
922                self.adapter_client.terminate().await;
923
924                // We must wait for the client to send a request before we can send the error response.
925                // Due to the PG wire protocol, we can't send an ErrorResponse unless it is in response
926                // to a client message.
927                let _ = self.conn.recv().await?;
928                return error_state;
929            },
930            // `recv()` is cancel-safe as per it's docs.
931            message = self.conn.recv() => message?,
932        };
933
934        // Take the metrics since just before the `recv`.
935        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        // TODO(ggevay): Consider subtracting the scheduling delay from `received`. It's not obvious
942        // whether we should do this, because the result wouldn't exactly correspond to either first
943        // byte received or last byte received (for msgs that arrive in more than one network packet).
944        let received = SYSTEM_TIME();
945
946        self.adapter_client
947            .remove_idle_in_transaction_session_timeout();
948
949        // NOTE(guswynn): we could consider adding spans to all message types. Currently
950        // only a few message types seem useful.
951        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, // If `max_rows < 0`, no limit.
994                    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                // In PostgreSQL, when using the extended query protocol, some statements may
1012                // trigger an eager commit of the current implicit transaction,
1013                // see: <https://git.postgresql.org/gitweb/?p=postgresql.git&a=commitdiff&h=f92944137>.
1014                //
1015                // In Materialize we instead eagerly commit every implicit transaction that
1016                // cannot take on further statements of the same pipeline, which keeps the
1017                // single-statement optimizations available to queries issued in the extended
1018                // query protocol. The ones that can stay open, so that the pipeline commits
1019                // or rolls back as a unit. See `TransactionStatus::may_span_pipeline`.
1020                //
1021                // We don't immediately commit here to allow users to page through the portal if
1022                // necessary. Committing the transaction would destroy the portal before the next
1023                // Execute command has a chance to resume it. So we instead mark the transaction
1024                // for commit the next time that `ensure_transaction` is called.
1025                let (is_implicit, may_span_pipeline) = {
1026                    let txn = self.adapter_client.session().transaction();
1027                    (txn.is_implicit(), txn.may_span_pipeline())
1028                };
1029                // Ordered so that only a write reads the flag, keeping the catalog
1030                // snapshot off the read path.
1031                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            // Accept but ignore stray COPY subprotocol messages, mirroring
1052            // PostgreSQL. Clients stream COPY data optimistically, so when a
1053            // COPY statement fails before COPY mode is entered, its pipelined
1054            // CopyData/CopyDone/CopyFail arrive here. Draining instead would
1055            // discard unrelated messages until the next Sync, hanging simple
1056            // protocol clients that never send one.
1057            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    /// Note that `lifecycle_timestamps` belongs to the whole "Simple Query", because the whole
1100    /// Simple Query is received and parsed together. This means that if there are multiple
1101    /// statements in a Simple Query, then all of them have the same `lifecycle_timestamps`.
1102    #[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        // Bind the portal. Note that this does not set the empty string prepared
1110        // statement.
1111        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        // Maybe send row description.
1140        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        // Destroy the portal.
1177        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        // start_transaction can't error (but assert that just in case it changes in
1192        // the future.
1193        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    /// Logs an arriving frontend message at info level, when
1205    /// `enable_statement_arrival_logging` is on. Runs before the message is
1206    /// processed, so a message whose processing crashes the process still
1207    /// appears in the log. The `kind` field says which message it is, and
1208    /// thereby also whether the statement came in through the simple protocol
1209    /// (`query`) or the extended protocol (`parse`, `bind`, `execute`, ...).
1210    /// The prepared statement and portal names, together with the connection
1211    /// id, allow connecting a `bind` or `execute` back to the `parse` that
1212    /// carried the SQL text.
1213    ///
1214    /// SQL text is parsed and logged with its literals redacted, the same
1215    /// redaction the statement log applies. This means a statement that
1216    /// crashes the parser is not captured, an accepted limitation. Bind
1217    /// parameter values are data that redaction cannot reach, so only their
1218    /// count is logged. Authentication payloads are never logged. COPY data
1219    /// is logged as its length only, and only when it arrives as a stray
1220    /// message in the ready state: messages consumed by the COPY subprotocol
1221    /// or the post-error drain loop don't pass through here at all.
1222    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            // COPY payloads would flood the log. Log only their length.
1260            FrontendMessage::CopyData(data) => {
1261                info!(%conn_id, %session_uuid, kind, len = data.len(), "statement arrival");
1262            }
1263            // Authentication payloads must never be logged.
1264            FrontendMessage::Password { .. }
1265            | FrontendMessage::RawAuthentication(_)
1266            | FrontendMessage::SASLInitialResponse { .. }
1267            | FrontendMessage::SASLResponse(_) => {
1268                info!(%conn_id, %session_uuid, kind, "statement arrival");
1269            }
1270            // CopyFail carries a client-supplied free-text error message,
1271            // which we don't log.
1272            FrontendMessage::CopyFail(_) => {
1273                info!(%conn_id, %session_uuid, kind, "statement arrival");
1274            }
1275            // Log the full Debug representation for all other variants, which
1276            // carry only object names or no payload.
1277            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                // WARNING: When adding a variant here, consider whether its payload is sensitive or
1287                // bulky!
1288                //
1289                // (The field must not be named `message`, that name is
1290                // reserved for the event text in tracing.)
1291                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                // Convert our 0-based byte position to pgwire's 1-based character
1301                // position.
1302                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    /// Executes a "Simple Query", see
1316    /// <https://www.postgresql.org/docs/current/protocol-flow.html#PROTOCOL-FLOW-SIMPLE-QUERY>
1317    ///
1318    /// For implicit transaction handling, see "Multiple Statements in a Simple Query" in the above.
1319    #[instrument(level = "debug")]
1320    async fn query(&mut self, sql: String, received: EpochMillis) -> Result<State, io::Error> {
1321        // Parse first before doing any transaction checking.
1322        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        // Compare with postgres' backend/tcop/postgres.c exec_simple_query.
1333        for StatementParseResult { ast: stmt, sql } in stmts {
1334            // In an aborted transaction, reject all commands except COMMIT/ROLLBACK.
1335            if self.is_aborted_txn() && !is_txn_exit_stmt(Some(&stmt)) {
1336                self.aborted_txn_error().await?;
1337                break;
1338            }
1339
1340            // Start an implicit transaction if we aren't in any transaction and there's
1341            // more than one statement. This mirrors the `use_implicit_block` variable in
1342            // postgres.
1343            //
1344            // This needs to be done in the loop instead of once at the top because
1345            // a COMMIT/ROLLBACK statement needs to start a new transaction on next
1346            // statement.
1347            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        // Implicit transactions are closed at the end of a Query message.
1360        {
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        // Start a transaction if we aren't in one.
1381        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    /// Commits and clears the current transaction.
1447    #[instrument(level = "debug")]
1448    async fn commit_transaction(&mut self) -> Result<(), io::Error> {
1449        self.end_transaction(EndTransactionAction::Commit).await
1450    }
1451
1452    /// Rollback and clears the current transaction.
1453    #[instrument(level = "debug")]
1454    async fn rollback_transaction(&mut self) -> Result<(), io::Error> {
1455        self.end_transaction(EndTransactionAction::Rollback).await
1456    }
1457
1458    /// End a transaction and report to the user if an error occurred.
1459    ///
1460    /// The parameters this changes must be announced, exactly as an explicit
1461    /// `COMMIT`/`ROLLBACK` announces them. Otherwise a `SET LOCAL` outside an
1462    /// explicit transaction announces its new value and never its revert, and a
1463    /// client that caches parameters keeps the reverted value.
1464    #[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    /// Announces changed parameters, restricted to those the client is told
1486    /// about at startup.
1487    #[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        // Start a transaction if we aren't in one.
1520        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                        // NUL characters get the same SQLSTATE that PostgreSQL
1590                        // reports for them.
1591                        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        // Binary encodings are disabled for list, map, and aclitem types, but this doesn't
1628        // apply to COPY TO statements.
1629        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                // To be conservative, we are restricting COPY FROM to only allow list/map/aclitem types if it is not
1637                // copying from STDIN. It is likely that this works in theory, but is risky and likely to OOM anyways
1638                // as all the data will be held in a buffer in memory before being processed.
1639                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    /// `outer_ctx_extra` is Some when we are executing as part of an outer statement, e.g., a FETCH
1683    /// triggering the execution of the underlying query.
1684    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            // Check if the portal has been started and can be continued.
1698            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            // In an aborted transaction, reject all commands except COMMIT/ROLLBACK.
1724            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                    // Start a transaction if we aren't in one.
1741                    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                            // This is an error communicating with the connection.
1788                            // We consider that to be a cancelation, rather than a query error.
1789                            (Err(e), StatementEndedExecutionReason::Canceled)
1790                        }
1791                        Ok((ok, SendRowsEndedReason::Canceled)) => {
1792                            (Ok(ok), StatementEndedExecutionReason::Canceled)
1793                        }
1794                        // NOTE: For now the values for `result_size` and
1795                        // `rows_returned` in fetches are a bit confusing.
1796                        // We record `Some(n)` for the first fetch, where `n` is
1797                        // the number of bytes/rows returned by the inner
1798                        // execute (regardless of how many rows the
1799                        // fetch fetched), and `None` for subsequent fetches.
1800                        //
1801                        // This arguably makes sense since the size/rows
1802                        // returned measures how much work the compute
1803                        // layer had to do to satisfy the query, but
1804                        // we should revisit it if/when we start
1805                        // logging the inner execute separately.
1806                        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                // FETCH is an awkward command for our current architecture. In Postgres it
1831                // will extract <count> rows from the target portal, cache them, and return
1832                // them to the user as requested. Its command tag is always FETCH <num rows
1833                // extracted>. In Materialize, since we have chosen to not fully support FETCH,
1834                // we must remember the number of rows that were returned. Use this tag to
1835                // remember that information and return it.
1836                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        // Start a transaction if we aren't in one.
1879        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        // Cloning to avoid a mutable borrow issue because `send` also uses `adapter_client`
1890        let parameter_desc = BackendMessage::ParameterDescription(
1891            stmt.desc()
1892                .param_types
1893                .iter()
1894                .map(mz_pgrepr::Type::from)
1895                .collect(),
1896        );
1897        // Claim that all results will be output in text format, even
1898        // though the true result formats are not yet known. A bit
1899        // weird, but this is the behavior that PostgreSQL specifies.
1900        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        // Start a transaction if we aren't in one.
1909        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    /// Completes the portal that carried a `CLOSE`.
1947    ///
1948    /// The `CLOSE` can name that very portal, and sequencing has then already
1949    /// removed it. Marking the state is bookkeeping for a portal nobody will
1950    /// read again, so its absence is expected rather than an error.
1951    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    /// Completes the portal that carried a `DECLARE`.
1962    ///
1963    /// A successful `DECLARE` only ever adds a portal, and adding one under a
1964    /// name already in use fails with `DuplicateCursor`, so the executing portal
1965    /// is still present here.
1966    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        // Unlike Execute, no count specified in FETCH returns 1 row, and 0 means 0
1989        // instead of All.
1990        let count = count.unwrap_or(FetchDirection::ForwardCount(1));
1991
1992        // Figure out how many rows we should send back by looking at the various
1993        // combinations of the execute and fetch.
1994        //
1995        // In Postgres, Fetch will cache <count> rows from the target portal and
1996        // return those as requested (if, say, an Execute message was sent with a
1997        // max_rows < the Fetch's count). We expect that case to be incredibly rare and
1998        // so have chosen to not support it until users request it. This eases
1999        // implementation difficulty since we don't have to be able to "send" rows to
2000        // a buffer.
2001        //
2002        // TODO(mjibson): Test this somehow? Need to divide up the pgtest files in
2003        // order to have some that are not Postgres compatible.
2004        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    /// Sends a backend message to the client, after applying a severity filter.
2063    ///
2064    /// The message is only sent if its severity is above the severity set
2065    /// in the session, with the default value being NOTICE.
2066    #[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        // Flush immediately after sending an error response, as some clients
2078        // expect to be able to read the error response before sending a Sync
2079        // message. This is arguably in violation of the protocol specification,
2080        // but the specification is somewhat ambiguous, and easier to match
2081        // PostgreSQL here than to fix all the clients that have this
2082        // expectation.
2083        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        // Close the current transaction if we are in an implicit transaction.
2104        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                // This code is somewhat awkwardly structured because we
2231                // can't hold `var` across an await point.
2232                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                        // This is an error communicating with the connection.
2290                        // We consider that to be a cancelation, rather than a query error.
2291                        (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                                // This is an error communicating with the connection.
2343                                // We consider that to be a cancelation, rather than a query error.
2344                                (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                        // We don't need to finalize execution here;
2377                        // it was already done in the
2378                        // coordinator. Just extract the state and
2379                        // return that.
2380                        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                        // We don't need to finalize execution here;
2402                        // it was already done in the
2403                        // coordinator. Just extract the state and
2404                        // return that.
2405                        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    // TODO(guswynn): figure out how to get it to compile without skip_all
2500    #[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        // If this portal is being executed from a FETCH then we need to use the result
2512        // format type of the outer portal.
2513        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        // Sanity check that the various `RelationDesc`s match up.
2536        {
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        // want_rows is the maximum number of rows the client wants.
2575        let mut want_rows = match max_rows {
2576            ExecuteCount::All => usize::MAX,
2577            ExecuteCount::Count(count) => count,
2578        };
2579
2580        // Send rows while the client still wants them and there are still rows to send.
2581        loop {
2582            // Fetch next batch of rows, waiting for a possible requested
2583            // timeout or notice.
2584            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                // Biased: drain available data before checking the deadline.
2591                // This is critical for the WaitOnce case, where the deadline
2592                // is set to `Instant::now()` right after the first batch:
2593                // without `biased`, `recv()` and the already-expired deadline
2594                // race nondeterministically, so we might break the loop
2595                // before `no_more_rows` is set (or even before ready rows
2596                // are consumed). With an explicit `TIMEOUT`, missing a batch
2597                // right at the boundary is acceptable, but WaitOnce fires
2598                // immediately and the race is not.
2599                //
2600                // Trade-off: if `recv()` keeps returning Ready (unlikely in
2601                // practice—row processing + flush is slower than upstream
2602                // tick granularity), a `TIMEOUT` deadline could be delayed.
2603                // See database-issues#9470.
2604                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 is true: the first time this fn is called it blocks (same as
2642                    // deadline == None). The second time this fn is called it should behave the
2643                    // same a 0s timeout.
2644                    if wait_once && batch_rows.peek().is_some() {
2645                        deadline = Some(tokio::time::Instant::now());
2646                        wait_once = false;
2647                    }
2648
2649                    // Send a portion of the rows.
2650                    let mut sent_rows = 0;
2651                    let mut sent_bytes = 0;
2652                    let messages = (&mut batch_rows)
2653                        // TODO(parkmycar): This is a fair bit of juggling between iterator types
2654                        // to count the total number of bytes. Alternatively we could track the
2655                        // total sent bytes in this .map(...) call, but having side effects in map
2656                        // is a code smell.
2657                        .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 we have sent the number of requested rows, put the remainder of the batch
2675                    // (if any) back and stop sending.
2676                    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        // Always return rows back, even if it's empty. This prevents an unclosed
2724        // portal from re-executing after it has been emptied.
2725        *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        // Attend to metrics if there are no more rows. Only record once per stream
2737        // to avoid polluting the histogram when an exhausted cursor is FETCHed again.
2738        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                // If the result is empty, then we define time from first to last row as 0.
2756                // (Note that, currently, an empty result involves a PeekResponse with 0 rows, which
2757                // does flip `saw_rows`, so this code path is currently not exercised.)
2758                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        // Binary encoding is not implemented for some types (e.g., list, map,
2807        // and aclitem). Unlike the extended query protocol's Bind handler, COPY
2808        // does not validate this when binding the portal: the portal's result
2809        // formats describe the `CopyData` wrapper, not the COPY format itself,
2810        // so the Bind handler explicitly skips `COPY TO` statements. We must
2811        // therefore check here, before streaming any rows, otherwise
2812        // `encode_binary` would panic mid-stream (SQL-323).
2813        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        // Unlike `COPY TO <external destination>`, which is encoded in the
2836        // dataflow layer, `COPY TO STDOUT` runs in the session and so honors
2837        // the session's encoding settings, matching PostgreSQL.
2838        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        // In Postgres, binary copy has a header that is followed (in the same
2854        // CopyData) by the first row. In order to replicate their behavior, use a
2855        // common vec that we can extend one time now and then fill up with the encode
2856        // functions.
2857        let mut out = Vec::new();
2858
2859        if let CopyFormat::Binary = format {
2860            // 11-byte signature.
2861            out.extend(b"PGCOPY\n\xFF\r\n\0");
2862            // 32-bit flags field.
2863            out.extend([0, 0, 0, 0]);
2864            // 32-bit header extension length field.
2865            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        // Send required trailers.
2918        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    /// Handles the copy-in mode of the postgres protocol from transferring
2938    /// data to the server.
2939    #[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                // The connection closed gracefully without sending us a `CopyDone`,
2972                // causing us to just drop the copy request.
2973                // For the purposes of statement logging, we count this as a cancellation.
2974                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        // Set up the parallel streaming batch builders in the coordinator.
3009        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                // Drain remaining CopyData/CopyDone/CopyFail messages from the
3023                // socket. Since CopyInResponse was already sent, the client may
3024                // have pipelined copy data that we must consume before returning
3025                // the error, otherwise they'd be misinterpreted as top-level
3026                // protocol messages and cause a deadlock.
3027                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        // Batch size for splitting raw data across parallel workers (~32MB).
3051        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(&params);
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        // Receive loop: accumulate CopyData, split at row boundaries,
3069        // round-robin raw chunks to parallel batch builder workers.
3070        loop {
3071            let message = self.conn.recv().await?;
3072            match message {
3073                Some(FrontendMessage::CopyData(buf)) => {
3074                    if saw_end_marker {
3075                        // Per PostgreSQL COPY behavior, ignore all bytes after
3076                        // the end-of-copy marker until CopyDone.
3077                        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                    // Guard against pathological single rows that never terminate.
3089                    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                    // When buffer exceeds batch size, split at the last complete row
3101                    // and send the complete rows chunk to the next worker.
3102                    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, // no complete row yet
3107                        };
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                    // Worker dropped (likely errored) — stop sending,
3118                    // fall through to completion_rx for the real error.
3119                    if send_failed {
3120                        break;
3121                    }
3122                }
3123                Some(FrontendMessage::CopyDone) => {
3124                    // Send any remaining data to the next worker.
3125                    if !data.is_empty() {
3126                        let chunk = std::mem::take(&mut data);
3127                        // Ignore send failure — completion_rx will have the error.
3128                        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 the writer to signal cancellation to the background tasks.
3139                    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 we exited the receive loop before seeing `CopyDone` (e.g. because
3172        // a worker failed and dropped its channel), keep draining COPY input to
3173        // avoid desynchronizing the protocol state machine.
3174        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 all senders to signal EOF to the background batch builders.
3218        // If copy_err is set, a worker already failed — dropping the senders
3219        // will cause remaining workers to stop, and we'll get the real error
3220        // from completion_rx below.
3221        drop(writer.batch_txs);
3222
3223        // Wait for all parallel workers to finish building batches.
3224        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        // Stage all batches in the session's transaction for atomic commit.
3252        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            // Error can be called from describe and parse and so might not be in an active
3299            // transaction.
3300            TransactionStatus::Default | TransactionStatus::Failed(_) => {}
3301            // In Started (i.e., a single statement), cleanup ourselves.
3302            TransactionStatus::Started(_) => {
3303                self.rollback_transaction().await?;
3304            }
3305            // Implicit transactions also clear themselves.
3306            TransactionStatus::InTransactionImplicit(_) => {
3307                self.rollback_transaction().await?;
3308            }
3309            // Explicit transactions move to failed.
3310            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
3366// A GetResponse used by send_rows during execute messages on portals or for
3367// simple query messages.
3368fn portal_exec_message(
3369    max_rows: ExecuteCount,
3370    total_sent_rows: usize,
3371    _fetch_portal: Option<PortalRefMut>,
3372) -> BackendMessage {
3373    // If max_rows is not specified, we will always send back a CommandComplete. If
3374    // max_rows is specified, we only send CommandComplete if there were more rows
3375    // requested than were remaining. That is, if max_rows == number of rows that
3376    // were remaining before sending (not that are remaining after sending), then
3377    // we still send a PortalSuspended. The number of remaining rows after the rows
3378    // have been sent doesn't matter. This matches postgres.
3379    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
3389// A GetResponse used by send_rows during FETCH queries.
3390fn 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
3425// See postgres' backend/tcop/postgres.c IsTransactionExitStmt.
3426fn is_txn_exit_stmt(stmt: Option<&Statement<Raw>>) -> bool {
3427    match stmt {
3428        // Add PREPARE to this if we ever support it.
3429        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    // Byte offset within `data` at which the in-progress CSV record begins.
3448    // Used to verify the end-of-copy marker against the raw input bytes,
3449    // distinguishing a literal `\.` line from a quoted CSV value `"\."`
3450    // whose decoded form is also `\.`.
3451    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                                // Detect the marker against the raw input
3522                                // bytes, not the CSV-decoded record. A quoted
3523                                // data row `"\."` decodes to `\.` but must be
3524                                // imported as data; only a bare `\.` line
3525                                // terminates the COPY.
3526                                let raw = &data[self.record_start..row_end];
3527                                // csv-core ends a CRLF record after the `\r`,
3528                                // leaving the trailing `\n` as the leading byte
3529                                // of the next record's span; a CR-only record
3530                                // ends in a lone `\r`. So a `\.` marker record's
3531                                // raw span can be `\.\n` (LF), `\n\.\r` (CRLF)
3532                                // or `\.\r` (CR). Trim CR/LF from both ends
3533                                // before comparing — a trailing-only strip would
3534                                // miss the CRLF/CR forms. Quoted `"\."` data
3535                                // keeps its surrounding quotes after trimming and
3536                                // is therefore correctly rejected.
3537                                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        // `record_start` is only maintained for the CSV path; the text and
3599        // binary paths leave it at 0. For CSV, splits always occur at a
3600        // completed-row boundary, so the in-progress record (if any) starts at
3601        // the new beginning of the buffer. Assert that invariant so the
3602        // `saturating_sub` below doesn't silently paper over a bug that
3603        // bisected an in-progress record — but only when CSV is in use, since
3604        // otherwise `record_start` is meaninglessly 0.
3605        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        // The pgwire COPY row scanner must detect a bare `\.` end-of-copy
3650        // marker for every line ending, and must never mistake a quoted
3651        // `"\."` data row for it. csv-core ends a CRLF record after the `\r`
3652        // (leaving the `\n` as the next record's leading byte), so the raw
3653        // record span of a `\.` marker is `\.\n` (LF), `\n\.\r` (CRLF) or
3654        // `\.\r` (CR); a trailing-only strip would miss the CRLF/CR forms and
3655        // silently import post-marker rows.
3656        let params = CopyFormatParams::Csv(CopyCsvFormatParams::default());
3657
3658        let marker_end = |data: &[u8]| -> Option<usize> {
3659            let mut scanner = CopyRowScanner::new(&params);
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            // Bare `\.` (the marker is the second record, so record_start has
3675            // already advanced past the orphaned terminator of `first`).
3676            // csv-core reports the record after a single terminator byte, so
3677            // the marker boundary sits just past `first<eol>\.` + one byte.
3678            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            // Quoted "\." is data, not the marker.
3690            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        // Regression: `record_start` is only maintained for the CSV path; the
3698        // text and binary paths leave it at 0. `on_split` must therefore not
3699        // assert `record_start >= split_pos` for those formats — that fires on
3700        // every split of a large text/binary COPY stream (soft-assertions
3701        // panic under test). Mirrors `COPY ... FROM STDIN` (default text
3702        // format) splitting at a row boundary once the buffer fills.
3703        for params in [
3704            CopyFormatParams::Text(CopyTextFormatParams::default()),
3705            CopyFormatParams::Binary,
3706        ] {
3707            let mut scanner = CopyRowScanner::new(&params);
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            // Must not panic via the CSV-only `on_split` soft-assert.
3714            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            // Unclear what this should do.
3757            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            // Unclear what this should error, but at least test it.
3809            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        // A real JWT header decodes successfully.
3897        assert!(is_jwt("eyJhbGciOiJIUzI1NiJ9.eyJzdWIiOiIxIn0.signature"));
3898        // Not JWTs: plain strings, wrong segment count, non-JSON headers.
3899        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}