Skip to main content

agent_first_psql/db/
executor.rs

1use super::errors::{ConnectError, ExecError, map_pg_error};
2use super::params::{QueryParam, build_param_refs, build_params, validate_param_count};
3use super::rows::{fallback_columns_supported, row_to_json_fallback};
4use super::session::{
5    CancelSlot, SessionMap, connect_session, get_session, new_session_map, remove_sessions,
6    shutdown_all_sessions,
7};
8use crate::protocol::log_event;
9use crate::types::{ColumnInfo, Output, ResolvedOptions, SessionConfig, Trace};
10use agent_first_data::LogFilters;
11use async_trait::async_trait;
12use futures_util::TryStreamExt;
13use serde_json::Value;
14use std::collections::HashSet;
15use std::pin::pin;
16use tokio::sync::mpsc;
17use tokio_postgres::types::ToSql;
18
19fn row_json_size(row: &Value) -> usize {
20    serde_json::to_vec(row)
21        .map(|bytes| bytes.len())
22        .unwrap_or(0)
23}
24
25#[derive(Debug)]
26pub enum ExecOutcome {
27    Rows {
28        columns: Vec<ColumnInfo>,
29        rows: Vec<Value>,
30        /// When true, the inline row/byte limit was hit and `rows` only
31        /// contains the prefix that fit. The underlying statement still
32        /// executed in full — for `UPDATE ... RETURNING`, the writes
33        /// happened even though their RETURNING projection was capped.
34        truncated: bool,
35        /// Inline-row limit if that's what fired (otherwise None).
36        truncated_at_rows: Option<usize>,
37        /// Inline-byte limit if that's what fired (otherwise None).
38        truncated_at_bytes: Option<usize>,
39    },
40    Command {
41        affected: usize,
42    },
43}
44
45#[derive(Debug)]
46pub struct DryRunOutcome {
47    pub param_types: Vec<String>,
48    pub columns: Vec<ColumnInfo>,
49}
50
51#[derive(Debug)]
52pub enum StreamOutcome {
53    Rows {
54        row_count: usize,
55        payload_bytes: usize,
56    },
57    Command {
58        affected: usize,
59    },
60}
61
62#[async_trait]
63pub trait RowSink: Send {
64    async fn start(&mut self, columns: Vec<ColumnInfo>) -> Result<(), ExecError>;
65    async fn row(&mut self, row: Value, row_bytes: usize) -> Result<(), ExecError>;
66}
67
68pub struct ExecRequest<'a> {
69    pub session_name: &'a str,
70    pub session_cfg: &'a SessionConfig,
71    pub sql: &'a str,
72    pub params: &'a [Value],
73    pub opts: &'a ResolvedOptions,
74    pub cancel_slot: Option<CancelSlot>,
75    pub transport_log: Option<TransportLogContext>,
76}
77
78#[derive(Clone)]
79pub struct TransportLogContext {
80    pub session: String,
81    pub log: LogFilters,
82    pub writer: mpsc::Sender<Output>,
83}
84
85#[async_trait]
86pub trait DbExecutor: Send + Sync {
87    async fn execute(&self, req: ExecRequest<'_>) -> Result<ExecOutcome, ExecError>;
88
89    /// Validate `sql` and the param shape without running the statement. The
90    /// server prepares the statement inside a transaction that is rolled back,
91    /// returning the inferred parameter types and column metadata.
92    async fn prepare_only(&self, req: ExecRequest<'_>) -> Result<DryRunOutcome, ExecError>;
93
94    /// Open an explicit transaction on the named session. Subsequent
95    /// `execute`/`execute_streaming` calls on that session bypass the
96    /// implicit per-query `BEGIN..COMMIT` wrap until `tx_commit` or
97    /// `tx_rollback` is called.
98    async fn tx_begin(
99        &self,
100        _session_name: &str,
101        _session_cfg: &SessionConfig,
102        _read_only: bool,
103    ) -> Result<(), ExecError> {
104        Err(ExecError::Internal(
105            "explicit transactions not implemented for this executor".to_string(),
106        ))
107    }
108
109    async fn tx_commit(
110        &self,
111        _session_name: &str,
112        _session_cfg: &SessionConfig,
113    ) -> Result<(), ExecError> {
114        Err(ExecError::Internal(
115            "explicit transactions not implemented for this executor".to_string(),
116        ))
117    }
118
119    async fn tx_rollback(
120        &self,
121        _session_name: &str,
122        _session_cfg: &SessionConfig,
123    ) -> Result<(), ExecError> {
124        Err(ExecError::Internal(
125            "explicit transactions not implemented for this executor".to_string(),
126        ))
127    }
128
129    async fn execute_streaming(
130        &self,
131        req: ExecRequest<'_>,
132        sink: &mut (dyn RowSink + Send),
133    ) -> Result<StreamOutcome, ExecError> {
134        match self.execute(req).await? {
135            ExecOutcome::Rows { columns, rows, .. } => {
136                sink.start(columns).await?;
137                let mut row_count = 0usize;
138                let mut payload_bytes = 0usize;
139                for row in rows {
140                    let row_bytes = row_json_size(&row);
141                    payload_bytes += row_bytes;
142                    row_count += 1;
143                    sink.row(row, row_bytes).await?;
144                }
145                Ok(StreamOutcome::Rows {
146                    row_count,
147                    payload_bytes,
148                })
149            }
150            ExecOutcome::Command { affected } => Ok(StreamOutcome::Command { affected }),
151        }
152    }
153
154    async fn invalidate_sessions(&self, _session_names: &[String]) {}
155
156    async fn explicit_tx_open(&self, _session_name: &str) -> bool {
157        false
158    }
159
160    async fn shutdown(&self) {}
161}
162
163pub struct PostgresExecutor {
164    sessions: SessionMap,
165}
166
167impl Default for PostgresExecutor {
168    fn default() -> Self {
169        Self::new()
170    }
171}
172
173impl PostgresExecutor {
174    pub fn new() -> Self {
175        Self {
176            sessions: new_session_map(),
177        }
178    }
179}
180
181#[async_trait]
182impl DbExecutor for PostgresExecutor {
183    async fn explicit_tx_open(&self, session_name: &str) -> bool {
184        get_session(&self.sessions, session_name)
185            .await
186            .explicit_tx_read_only()
187            .is_some()
188    }
189
190    async fn execute(&self, req: ExecRequest<'_>) -> Result<ExecOutcome, ExecError> {
191        let session = get_session(&self.sessions, req.session_name).await;
192        let explicit_tx_read_only = session.explicit_tx_read_only();
193        if explicit_tx_read_only == Some(false) && req.opts.read_only {
194            return Err(ExecError::InvalidRequest {
195                message:
196                    "query permission is read-only, but the session has an open read-write transaction"
197                        .to_string(),
198                hint: Some(
199                    "repeat the matching write permission on every query in a read-write explicit transaction"
200                        .to_string(),
201                ),
202            });
203        }
204        let mut client_guard = session.client.lock().await;
205        let transport = ensure_connected(&mut client_guard, req.session_cfg).await?;
206        emit_transport_selected(&req, transport).await?;
207        let Some(client) = client_guard.as_mut() else {
208            return Err(ExecError::Connect(Box::new(ConnectError::new(
209                "connection unavailable",
210            ))));
211        };
212        let Some(pg_client) = client.client.as_mut() else {
213            return Err(ExecError::Connect(Box::new(ConnectError::new(
214                "connection unavailable",
215            ))));
216        };
217        install_cancel_context(
218            &req.cancel_slot,
219            pg_client.cancel_token(),
220            client.backend_pid,
221            req.session_cfg,
222        )
223        .await;
224        if cancel_requested(&req.cancel_slot) {
225            return Err(ExecError::Cancelled);
226        }
227        let result = if explicit_tx_read_only.is_some() {
228            execute_in_open_tx(pg_client, &req).await
229        } else {
230            execute_with_client(pg_client, &req).await
231        };
232        if should_drop_connection(&result) {
233            *client_guard = None;
234            // Connection dropped means the in-PG explicit tx is also gone.
235            session.set_explicit_tx(None);
236        }
237        result
238    }
239
240    async fn execute_streaming(
241        &self,
242        req: ExecRequest<'_>,
243        sink: &mut (dyn RowSink + Send),
244    ) -> Result<StreamOutcome, ExecError> {
245        let session = get_session(&self.sessions, req.session_name).await;
246        let explicit_tx_read_only = session.explicit_tx_read_only();
247        if explicit_tx_read_only == Some(false) && req.opts.read_only {
248            return Err(ExecError::InvalidRequest {
249                message:
250                    "query permission is read-only, but the session has an open read-write transaction"
251                        .to_string(),
252                hint: Some(
253                    "repeat the matching write permission on every query in a read-write explicit transaction"
254                        .to_string(),
255                ),
256            });
257        }
258        let mut client_guard = session.client.lock().await;
259        let transport = ensure_connected(&mut client_guard, req.session_cfg).await?;
260        emit_transport_selected(&req, transport).await?;
261        let Some(client) = client_guard.as_mut() else {
262            return Err(ExecError::Connect(Box::new(ConnectError::new(
263                "connection unavailable",
264            ))));
265        };
266        let Some(pg_client) = client.client.as_mut() else {
267            return Err(ExecError::Connect(Box::new(ConnectError::new(
268                "connection unavailable",
269            ))));
270        };
271        install_cancel_context(
272            &req.cancel_slot,
273            pg_client.cancel_token(),
274            client.backend_pid,
275            req.session_cfg,
276        )
277        .await;
278        if cancel_requested(&req.cancel_slot) {
279            return Err(ExecError::Cancelled);
280        }
281        let result = if explicit_tx_read_only.is_some() {
282            execute_streaming_in_open_tx(pg_client, &req, sink).await
283        } else {
284            execute_streaming_with_client(pg_client, &req, sink).await
285        };
286        if should_drop_connection(&result) {
287            *client_guard = None;
288            session.set_explicit_tx(None);
289        }
290        result
291    }
292
293    async fn tx_begin(
294        &self,
295        session_name: &str,
296        session_cfg: &SessionConfig,
297        read_only: bool,
298    ) -> Result<(), ExecError> {
299        let session = get_session(&self.sessions, session_name).await;
300        if session.explicit_tx_active() {
301            return Err(ExecError::InvalidParams(
302                "session is already in an explicit transaction; commit or rollback first"
303                    .to_string(),
304            ));
305        }
306        let mut client_guard = session.client.lock().await;
307        ensure_connected(&mut client_guard, session_cfg).await?;
308        let Some(client) = client_guard.as_mut() else {
309            return Err(ExecError::Connect(Box::new(ConnectError::new(
310                "connection unavailable",
311            ))));
312        };
313        let Some(pg_client) = client.client.as_mut() else {
314            return Err(ExecError::Connect(Box::new(ConnectError::new(
315                "connection unavailable",
316            ))));
317        };
318        let sql = if read_only {
319            "BEGIN READ ONLY"
320        } else {
321            "BEGIN"
322        };
323        pg_client.batch_execute(sql).await.map_err(map_pg_error)?;
324        session.set_explicit_tx(Some(read_only));
325        Ok(())
326    }
327
328    async fn tx_commit(
329        &self,
330        session_name: &str,
331        session_cfg: &SessionConfig,
332    ) -> Result<(), ExecError> {
333        let session = get_session(&self.sessions, session_name).await;
334        if !session.explicit_tx_active() {
335            return Err(ExecError::InvalidParams(
336                "no explicit transaction is open on this session; send `begin` first".to_string(),
337            ));
338        }
339        let mut client_guard = session.client.lock().await;
340        ensure_connected(&mut client_guard, session_cfg).await?;
341        let Some(client) = client_guard.as_mut() else {
342            return Err(ExecError::Connect(Box::new(ConnectError::new(
343                "connection unavailable",
344            ))));
345        };
346        let Some(pg_client) = client.client.as_mut() else {
347            return Err(ExecError::Connect(Box::new(ConnectError::new(
348                "connection unavailable",
349            ))));
350        };
351        let result = pg_client
352            .batch_execute("COMMIT")
353            .await
354            .map_err(map_pg_error);
355        session.set_explicit_tx(None);
356        result
357    }
358
359    async fn tx_rollback(
360        &self,
361        session_name: &str,
362        session_cfg: &SessionConfig,
363    ) -> Result<(), ExecError> {
364        let session = get_session(&self.sessions, session_name).await;
365        if !session.explicit_tx_active() {
366            return Err(ExecError::InvalidParams(
367                "no explicit transaction is open on this session; send `begin` first".to_string(),
368            ));
369        }
370        let mut client_guard = session.client.lock().await;
371        ensure_connected(&mut client_guard, session_cfg).await?;
372        let Some(client) = client_guard.as_mut() else {
373            return Err(ExecError::Connect(Box::new(ConnectError::new(
374                "connection unavailable",
375            ))));
376        };
377        let Some(pg_client) = client.client.as_mut() else {
378            return Err(ExecError::Connect(Box::new(ConnectError::new(
379                "connection unavailable",
380            ))));
381        };
382        let result = pg_client
383            .batch_execute("ROLLBACK")
384            .await
385            .map_err(map_pg_error);
386        session.set_explicit_tx(None);
387        result
388    }
389
390    async fn prepare_only(&self, req: ExecRequest<'_>) -> Result<DryRunOutcome, ExecError> {
391        let session = get_session(&self.sessions, req.session_name).await;
392        let mut client_guard = session.client.lock().await;
393        let transport = ensure_connected(&mut client_guard, req.session_cfg).await?;
394        emit_transport_selected(&req, transport).await?;
395        let Some(client) = client_guard.as_mut() else {
396            return Err(ExecError::Connect(Box::new(ConnectError::new(
397                "connection unavailable",
398            ))));
399        };
400        let Some(pg_client) = client.client.as_mut() else {
401            return Err(ExecError::Connect(Box::new(ConnectError::new(
402                "connection unavailable",
403            ))));
404        };
405        let result = prepare_only_with_client(pg_client, &req).await;
406        if should_drop_connection_dry_run(&result) {
407            *client_guard = None;
408        }
409        result
410    }
411
412    async fn invalidate_sessions(&self, session_names: &[String]) {
413        remove_sessions(&self.sessions, session_names).await;
414    }
415
416    async fn shutdown(&self) {
417        shutdown_all_sessions(&self.sessions).await;
418    }
419}
420
421async fn ensure_connected(
422    client: &mut Option<super::session::SessionClient>,
423    session_cfg: &SessionConfig,
424) -> Result<Option<super::session::TransportSelection>, ExecError> {
425    if client.as_ref().map(|c| c.is_closed()).unwrap_or(false) {
426        *client = None;
427    }
428    if client.is_none() {
429        let (connected, transport) = connect_session(session_cfg).await?;
430        *client = Some(connected);
431        return Ok(Some(transport));
432    }
433    Ok(None)
434}
435
436async fn emit_transport_selected(
437    req: &ExecRequest<'_>,
438    selected: Option<super::session::TransportSelection>,
439) -> Result<(), ExecError> {
440    let (Some(selected), Some(ctx)) = (selected, req.transport_log.as_ref()) else {
441        return Ok(());
442    };
443    emit_libpq_env_fallback(ctx, req.session_cfg).await?;
444    if !ctx.log.enabled(log_event::TRANSPORT_SELECTED) {
445        return Ok(());
446    }
447    let chain = super::session::transport_chain_summary(req.session_cfg, !ctx.log.is_empty());
448    ctx.writer
449        .send(Output::Log {
450            event: log_event::TRANSPORT_SELECTED.to_string(),
451            request_id: None,
452            session: Some(ctx.session.clone()),
453            error_code: None,
454            command_tag: None,
455            version: None,
456            config: None,
457            args: None,
458            env: None,
459            chain: Some(chain),
460            trace: Trace::only_duration(selected.duration_ms),
461        })
462        .await
463        .map_err(|_| ExecError::Internal("output channel closed".to_string()))
464}
465
466async fn emit_libpq_env_fallback(
467    ctx: &TransportLogContext,
468    cfg: &SessionConfig,
469) -> Result<(), ExecError> {
470    if !ctx.log.enabled(log_event::CONNECT_LIBPQ_ENV_FALLBACK) {
471        return Ok(());
472    }
473    let used = crate::conn::libpq_env_fallbacks_in_use(cfg);
474    if used.is_empty() {
475        return Ok(());
476    }
477    let mut config = serde_json::Map::new();
478    config.insert(
479        "env_vars".to_string(),
480        Value::Array(used.iter().map(|v| Value::from(*v)).collect()),
481    );
482    config.insert(
483        "note".to_string(),
484        Value::from(
485            "libpq PG* environment variables filled connection fields not given via flags/secrets; prefer explicit --host/--user/--password env:NAME for agent runs",
486        ),
487    );
488    ctx.writer
489        .send(Output::Log {
490            event: log_event::CONNECT_LIBPQ_ENV_FALLBACK.to_string(),
491            request_id: None,
492            session: Some(ctx.session.clone()),
493            error_code: None,
494            command_tag: None,
495            version: None,
496            config: Some(Value::Object(config)),
497            args: None,
498            env: None,
499            chain: None,
500            trace: Trace::only_duration(0),
501        })
502        .await
503        .map_err(|_| ExecError::Internal("output channel closed".to_string()))
504}
505
506async fn emit_row_encoding_degraded(
507    ctx: Option<&TransportLogContext>,
508    reason: &ExecError,
509) -> Result<(), ExecError> {
510    let Some(ctx) = ctx else {
511        return Ok(());
512    };
513    if !ctx.log.enabled(log_event::QUERY_ROW_ENCODING_DEGRADED) {
514        return Ok(());
515    }
516    // Report the wrapper's rejection in the same shape as a `sql_error`, never
517    // as a Rust `Debug` dump: an agent branches on `sqlstate`, and `position`
518    // is an offset into the wrapper's SQL rather than the caller's, which it
519    // can only interpret if it is labelled as such.
520    let mut config = serde_json::Map::new();
521    config.insert("message".to_string(), Value::String(reason.to_string()));
522    if let ExecError::Sql {
523        sqlstate,
524        detail,
525        hint,
526        position,
527        ..
528    } = reason
529    {
530        config.insert("sqlstate".to_string(), Value::String(sqlstate.clone()));
531        if let Some(detail) = detail {
532            config.insert("detail".to_string(), Value::String(detail.clone()));
533        }
534        if let Some(hint) = hint {
535            config.insert("hint".to_string(), Value::String(hint.clone()));
536        }
537        if let Some(position) = position {
538            config.insert(
539                "wrapper_position".to_string(),
540                Value::String(position.clone()),
541            );
542        }
543    }
544    ctx.writer
545        .send(Output::Log {
546            event: log_event::QUERY_ROW_ENCODING_DEGRADED.to_string(),
547            request_id: None,
548            session: Some(ctx.session.clone()),
549            error_code: None,
550            command_tag: None,
551            version: None,
552            config: Some(Value::Object(config)),
553            args: None,
554            env: None,
555            chain: None,
556            trace: Trace::only_duration(0),
557        })
558        .await
559        .map_err(|_| ExecError::Internal("output channel closed".to_string()))
560}
561
562fn should_drop_connection<T>(result: &Result<T, ExecError>) -> bool {
563    matches!(
564        result,
565        Err(ExecError::Connect(_)) | Err(ExecError::Internal(_))
566    )
567}
568
569fn should_drop_connection_dry_run(result: &Result<DryRunOutcome, ExecError>) -> bool {
570    matches!(
571        result,
572        Err(ExecError::Connect(_)) | Err(ExecError::Internal(_))
573    )
574}
575
576async fn prepare_only_with_client(
577    client: &mut tokio_postgres::Client,
578    req: &ExecRequest<'_>,
579) -> Result<DryRunOutcome, ExecError> {
580    let mut tx = start_transaction(client, true).await?;
581    let result = prepare_only_in_transaction(&mut tx, req).await;
582    // Always rollback — dry-run never commits.
583    let _ = tx.rollback().await;
584    result
585}
586
587async fn prepare_only_in_transaction(
588    tx: &mut tokio_postgres::Transaction<'_>,
589    req: &ExecRequest<'_>,
590) -> Result<DryRunOutcome, ExecError> {
591    let stmt = tx.prepare(req.sql).await.map_err(map_pg_error)?;
592    let columns = statement_columns(&stmt);
593    validate_unique_column_names(&columns)?;
594    validate_param_count(stmt.params().len(), req.params.len())?;
595    let param_types = stmt.params().iter().map(|t| t.name().to_string()).collect();
596    Ok(DryRunOutcome {
597        param_types,
598        columns,
599    })
600}
601
602fn cancel_requested(cancel_slot: &Option<CancelSlot>) -> bool {
603    cancel_slot
604        .as_ref()
605        .map(|slot| slot.is_cancelled())
606        .unwrap_or(false)
607}
608
609async fn execute_with_client(
610    client: &mut tokio_postgres::Client,
611    req: &ExecRequest<'_>,
612) -> Result<ExecOutcome, ExecError> {
613    let mut tx = start_transaction(client, req.opts.read_only).await?;
614    let result = execute_in_transaction(&mut tx, req).await;
615    finish_transaction(tx, result).await
616}
617
618/// Run a query against a client that is already inside an explicit
619/// transaction. The query is wrapped in a savepoint so a failure does not
620/// abort the user's outer transaction — the agent can retry or recover
621/// without losing prior progress.
622async fn execute_in_open_tx(
623    client: &mut tokio_postgres::Client,
624    req: &ExecRequest<'_>,
625) -> Result<ExecOutcome, ExecError> {
626    client
627        .batch_execute("SAVEPOINT afpsql_explicit")
628        .await
629        .map_err(map_pg_error)?;
630    let result = execute_in_open_tx_inner(client, req).await;
631    match &result {
632        Ok(_) => {
633            client
634                .batch_execute("RELEASE SAVEPOINT afpsql_explicit")
635                .await
636                .map_err(map_pg_error)?;
637        }
638        Err(_) => {
639            let _ = client
640                .batch_execute("ROLLBACK TO SAVEPOINT afpsql_explicit")
641                .await;
642            let _ = client
643                .batch_execute("RELEASE SAVEPOINT afpsql_explicit")
644                .await;
645        }
646    }
647    result
648}
649
650async fn execute_in_open_tx_inner(
651    client: &mut tokio_postgres::Client,
652    req: &ExecRequest<'_>,
653) -> Result<ExecOutcome, ExecError> {
654    apply_query_settings_client(client, req.opts).await?;
655    let stmt = client.prepare(req.sql).await.map_err(map_pg_error)?;
656    let columns = statement_columns(&stmt);
657    validate_unique_column_names(&columns)?;
658    validate_param_count(stmt.params().len(), req.params.len())?;
659    if columns.is_empty() {
660        let query_params = build_params(req.params, stmt.params())?;
661        let bind_refs = build_param_refs(&query_params);
662        let affected = client
663            .execute(&stmt, &bind_refs)
664            .await
665            .map_err(map_pg_error)? as usize;
666        return Ok(ExecOutcome::Command { affected });
667    }
668
669    let mut collector =
670        InlineRowCollector::new(columns, req.opts.inline_max_rows, req.opts.inline_max_bytes);
671    match prepare_wrapped_on_client(client, req).await {
672        Ok(wrapped_stmt) => {
673            let wrapped_params = build_params(req.params, wrapped_stmt.params())?;
674            let wrapped_refs = build_param_refs(&wrapped_params);
675            let stream = client
676                .query_raw(&wrapped_stmt, wrapped_refs)
677                .await
678                .map_err(map_pg_error)?;
679            let mut rows = pin!(stream);
680            while let Some(row) = rows.try_next().await.map_err(map_pg_error)? {
681                let value = wrapped_row_to_json(&row)?;
682                let row_bytes = row_json_size(&value);
683                let _ = collector.push(value, row_bytes)?;
684                if collector.is_truncated() {
685                    break;
686                }
687            }
688        }
689        Err(ExecError::InvalidParams(message)) => return Err(ExecError::InvalidParams(message)),
690        Err(error) => {
691            if !fallback_columns_supported(&stmt) {
692                return Err(error);
693            }
694            emit_row_encoding_degraded(req.transport_log.as_ref(), &error).await?;
695            let query_params = build_params(req.params, stmt.params())?;
696            let bind_refs = build_param_refs(&query_params);
697            let stream = client
698                .query_raw(&stmt, bind_refs)
699                .await
700                .map_err(map_pg_error)?;
701            let mut rows = pin!(stream);
702            while let Some(row) = rows.try_next().await.map_err(map_pg_error)? {
703                let value = row_to_json_fallback(&row)?;
704                let row_bytes = row_json_size(&value);
705                let _ = collector.push(value, row_bytes)?;
706                if collector.is_truncated() {
707                    break;
708                }
709            }
710        }
711    }
712    Ok(ExecOutcome::Rows {
713        truncated: collector.is_truncated(),
714        truncated_at_rows: collector.truncated_at_rows,
715        truncated_at_bytes: collector.truncated_at_bytes,
716        columns: collector.columns,
717        rows: collector.rows,
718    })
719}
720
721async fn execute_streaming_in_open_tx(
722    client: &mut tokio_postgres::Client,
723    req: &ExecRequest<'_>,
724    sink: &mut (dyn RowSink + Send),
725) -> Result<StreamOutcome, ExecError> {
726    client
727        .batch_execute("SAVEPOINT afpsql_explicit")
728        .await
729        .map_err(map_pg_error)?;
730    let result = execute_streaming_in_open_tx_inner(client, req, sink).await;
731    match &result {
732        Ok(_) => {
733            client
734                .batch_execute("RELEASE SAVEPOINT afpsql_explicit")
735                .await
736                .map_err(map_pg_error)?;
737        }
738        Err(_) => {
739            let _ = client
740                .batch_execute("ROLLBACK TO SAVEPOINT afpsql_explicit")
741                .await;
742            let _ = client
743                .batch_execute("RELEASE SAVEPOINT afpsql_explicit")
744                .await;
745        }
746    }
747    result
748}
749
750async fn execute_streaming_in_open_tx_inner(
751    client: &mut tokio_postgres::Client,
752    req: &ExecRequest<'_>,
753    sink: &mut (dyn RowSink + Send),
754) -> Result<StreamOutcome, ExecError> {
755    apply_query_settings_client(client, req.opts).await?;
756    let stmt = client.prepare(req.sql).await.map_err(map_pg_error)?;
757    let columns = statement_columns(&stmt);
758    validate_unique_column_names(&columns)?;
759    validate_param_count(stmt.params().len(), req.params.len())?;
760    if columns.is_empty() {
761        let query_params = build_params(req.params, stmt.params())?;
762        let bind_refs = build_param_refs(&query_params);
763        let affected = client
764            .execute(&stmt, &bind_refs)
765            .await
766            .map_err(map_pg_error)? as usize;
767        return Ok(StreamOutcome::Command { affected });
768    }
769
770    // Start the sink only once the wrapper is known to be usable, so a wrap
771    // failure cannot emit `result_start` without a matching `result_end`.
772    let mut row_count = 0usize;
773    let mut payload_bytes = 0usize;
774    match prepare_wrapped_on_client(client, req).await {
775        Ok(wrapped_stmt) => {
776            let wrapped_params = build_params(req.params, wrapped_stmt.params())?;
777            let wrapped_refs = build_param_refs(&wrapped_params);
778            let stream = client
779                .query_raw(&wrapped_stmt, wrapped_refs)
780                .await
781                .map_err(map_pg_error)?;
782            sink.start(columns).await?;
783            let mut rows = pin!(stream);
784            while let Some(row) = rows.try_next().await.map_err(map_pg_error)? {
785                let value = wrapped_row_to_json(&row)?;
786                let row_bytes = row_json_size(&value);
787                payload_bytes += row_bytes;
788                row_count += 1;
789                sink.row(value, row_bytes).await?;
790            }
791        }
792        Err(ExecError::InvalidParams(message)) => return Err(ExecError::InvalidParams(message)),
793        Err(error) => {
794            if !fallback_columns_supported(&stmt) {
795                return Err(error);
796            }
797            emit_row_encoding_degraded(req.transport_log.as_ref(), &error).await?;
798            let query_params = build_params(req.params, stmt.params())?;
799            let bind_refs = build_param_refs(&query_params);
800            let stream = client
801                .query_raw(&stmt, bind_refs)
802                .await
803                .map_err(map_pg_error)?;
804            sink.start(columns).await?;
805            let mut rows = pin!(stream);
806            while let Some(row) = rows.try_next().await.map_err(map_pg_error)? {
807                let value = row_to_json_fallback(&row)?;
808                let row_bytes = row_json_size(&value);
809                payload_bytes += row_bytes;
810                row_count += 1;
811                sink.row(value, row_bytes).await?;
812            }
813        }
814    }
815    Ok(StreamOutcome::Rows {
816        row_count,
817        payload_bytes,
818    })
819}
820
821async fn apply_query_settings_client(
822    client: &mut tokio_postgres::Client,
823    opts: &ResolvedOptions,
824) -> Result<(), ExecError> {
825    let statement_timeout = format!("{}ms", opts.statement_timeout_ms);
826    client
827        .execute(
828            "select set_config('statement_timeout', $1, true)",
829            &[&statement_timeout],
830        )
831        .await
832        .map_err(map_pg_error)?;
833
834    let lock_timeout = format!("{}ms", opts.lock_timeout_ms);
835    client
836        .execute(
837            "select set_config('lock_timeout', $1, true)",
838            &[&lock_timeout],
839        )
840        .await
841        .map_err(map_pg_error)?;
842    Ok(())
843}
844
845async fn execute_in_transaction(
846    tx: &mut tokio_postgres::Transaction<'_>,
847    req: &ExecRequest<'_>,
848) -> Result<ExecOutcome, ExecError> {
849    apply_query_settings(tx, req.opts).await?;
850    let prepared = prepare_bound_statement(tx, req.sql, req.params).await?;
851    let bind_refs = build_param_refs(&prepared.query_params);
852
853    if !prepared.columns.is_empty() {
854        let mut collector = InlineRowCollector::new(
855            prepared.columns,
856            req.opts.inline_max_rows,
857            req.opts.inline_max_bytes,
858        );
859        collect_rows_wrapped_or_direct(tx, req, &prepared.stmt, bind_refs, &mut collector).await?;
860
861        return Ok(ExecOutcome::Rows {
862            truncated: collector.is_truncated(),
863            truncated_at_rows: collector.truncated_at_rows,
864            truncated_at_bytes: collector.truncated_at_bytes,
865            columns: collector.columns,
866            rows: collector.rows,
867        });
868    }
869
870    let affected = tx
871        .execute(&prepared.stmt, &bind_refs)
872        .await
873        .map_err(map_pg_error)? as usize;
874
875    Ok(ExecOutcome::Command { affected })
876}
877
878async fn execute_streaming_with_client(
879    client: &mut tokio_postgres::Client,
880    req: &ExecRequest<'_>,
881    sink: &mut (dyn RowSink + Send),
882) -> Result<StreamOutcome, ExecError> {
883    let mut tx = start_transaction(client, req.opts.read_only).await?;
884    let result = execute_streaming_in_transaction(&mut tx, req, sink).await;
885    finish_transaction(tx, result).await
886}
887
888async fn execute_streaming_in_transaction(
889    tx: &mut tokio_postgres::Transaction<'_>,
890    req: &ExecRequest<'_>,
891    sink: &mut (dyn RowSink + Send),
892) -> Result<StreamOutcome, ExecError> {
893    apply_query_settings(tx, req.opts).await?;
894    let prepared = prepare_bound_statement(tx, req.sql, req.params).await?;
895    let bind_refs = build_param_refs(&prepared.query_params);
896
897    if prepared.columns.is_empty() {
898        let affected = tx
899            .execute(&prepared.stmt, &bind_refs)
900            .await
901            .map_err(map_pg_error)? as usize;
902        return Ok(StreamOutcome::Command { affected });
903    }
904
905    let stats =
906        stream_rows_wrapped_or_direct(tx, req, &prepared.stmt, bind_refs, prepared.columns, sink)
907            .await?;
908
909    Ok(StreamOutcome::Rows {
910        row_count: stats.row_count,
911        payload_bytes: stats.payload_bytes,
912    })
913}
914
915async fn finish_transaction<T>(
916    tx: tokio_postgres::Transaction<'_>,
917    result: Result<T, ExecError>,
918) -> Result<T, ExecError> {
919    match result {
920        Ok(outcome) => {
921            tx.commit().await.map_err(map_pg_error)?;
922            Ok(outcome)
923        }
924        Err(err) => {
925            tx.rollback().await.map_err(map_pg_error)?;
926            Err(err)
927        }
928    }
929}
930
931async fn start_transaction(
932    client: &mut tokio_postgres::Client,
933    read_only: bool,
934) -> Result<tokio_postgres::Transaction<'_>, ExecError> {
935    client
936        .build_transaction()
937        .read_only(read_only)
938        .start()
939        .await
940        .map_err(map_pg_error)
941}
942
943async fn install_cancel_context(
944    slot: &Option<CancelSlot>,
945    token: tokio_postgres::CancelToken,
946    backend_pid: i32,
947    session_cfg: &SessionConfig,
948) {
949    if let Some(slot) = slot {
950        slot.set_context(token, backend_pid, session_cfg).await;
951    }
952}
953
954struct PreparedStatement {
955    stmt: tokio_postgres::Statement,
956    columns: Vec<ColumnInfo>,
957    query_params: Vec<QueryParam>,
958}
959
960async fn prepare_bound_statement(
961    tx: &mut tokio_postgres::Transaction<'_>,
962    sql: &str,
963    params: &[Value],
964) -> Result<PreparedStatement, ExecError> {
965    let stmt = tx.prepare(sql).await.map_err(map_pg_error)?;
966    let columns = statement_columns(&stmt);
967    validate_unique_column_names(&columns)?;
968    validate_param_count(stmt.params().len(), params.len())?;
969    let query_params = build_params(params, stmt.params())?;
970    Ok(PreparedStatement {
971        stmt,
972        columns,
973        query_params,
974    })
975}
976
977fn statement_columns(stmt: &tokio_postgres::Statement) -> Vec<ColumnInfo> {
978    stmt.columns()
979        .iter()
980        .map(|col| ColumnInfo {
981            name: col.name().to_string(),
982            type_name: col.type_().name().to_string(),
983        })
984        .collect()
985}
986
987fn validate_unique_column_names(columns: &[ColumnInfo]) -> Result<(), ExecError> {
988    let mut seen = HashSet::new();
989    let mut duplicate_seen = HashSet::new();
990    let mut duplicates = Vec::new();
991
992    for column in columns {
993        let name = column.name.as_str();
994        if !seen.insert(name) && duplicate_seen.insert(name) {
995            duplicates.push(column.name.clone());
996        }
997    }
998
999    if duplicates.is_empty() {
1000        return Ok(());
1001    }
1002
1003    Err(ExecError::InvalidParams(format!(
1004        "query result has duplicate column name(s): {}. JSON object rows cannot safely represent duplicate keys; use AS aliases such as `a.id AS a_id` and `b.id AS b_id` to make output column names unique",
1005        format_column_names(&duplicates)
1006    )))
1007}
1008
1009fn format_column_names(names: &[String]) -> String {
1010    names
1011        .iter()
1012        .map(|name| format!("`{name}`"))
1013        .collect::<Vec<_>>()
1014        .join(", ")
1015}
1016
1017#[cfg(test)]
1018#[allow(clippy::items_after_test_module)]
1019mod tests {
1020    use super::*;
1021    use crate::types::{ContainerConfig, SshConfig};
1022    use tokio::sync::mpsc;
1023
1024    fn column(name: &str) -> ColumnInfo {
1025        ColumnInfo {
1026            name: name.to_string(),
1027            type_name: "int4".to_string(),
1028        }
1029    }
1030
1031    fn test_request<'a>(
1032        cfg: &'a SessionConfig,
1033        opts: &'a ResolvedOptions,
1034        log: LogFilters,
1035        writer: mpsc::Sender<Output>,
1036    ) -> ExecRequest<'a> {
1037        ExecRequest {
1038            session_name: "default",
1039            session_cfg: cfg,
1040            sql: "select 1",
1041            params: &[],
1042            opts,
1043            cancel_slot: None,
1044            transport_log: Some(TransportLogContext {
1045                session: "default".to_string(),
1046                log,
1047                writer,
1048            }),
1049        }
1050    }
1051
1052    fn default_opts() -> ResolvedOptions {
1053        ResolvedOptions {
1054            stream_rows: false,
1055            batch_rows: 1024,
1056            batch_bytes: 1 << 20,
1057            statement_timeout_ms: 0,
1058            lock_timeout_ms: 0,
1059            read_only: true,
1060            inline_max_rows: 100,
1061            inline_max_bytes: 1 << 20,
1062        }
1063    }
1064
1065    #[test]
1066    fn validate_unique_column_names_accepts_aliases() {
1067        let columns = vec![column("a_id"), column("b_id")];
1068        assert!(validate_unique_column_names(&columns).is_ok());
1069    }
1070
1071    #[test]
1072    fn validate_unique_column_names_rejects_duplicates_once() {
1073        let columns = vec![
1074            column("id"),
1075            column("name"),
1076            column("id"),
1077            column("name"),
1078            column("id"),
1079        ];
1080
1081        assert!(matches!(
1082            validate_unique_column_names(&columns),
1083            Err(ExecError::InvalidParams(message))
1084                if message.contains("`id`, `name`")
1085                    && message.contains("JSON object rows")
1086                    && message.contains("AS aliases")
1087        ));
1088    }
1089
1090    #[test]
1091    fn wrapped_rows_sql_removes_trailing_semicolons_and_protects_line_comments() {
1092        assert_eq!(
1093            wrapped_rows_sql("select now();  "),
1094            "with __afpsql_rows as (select now()\n) select to_jsonb(__afpsql_rows) as row_json from __afpsql_rows"
1095        );
1096        assert_eq!(
1097            wrapped_rows_sql("select 1 -- comment"),
1098            "with __afpsql_rows as (select 1 -- comment\n) select to_jsonb(__afpsql_rows) as row_json from __afpsql_rows"
1099        );
1100    }
1101
1102    #[tokio::test]
1103    async fn emit_transport_selected_skips_when_log_filter_empty() {
1104        let cfg = SessionConfig {
1105            host: Some("127.0.0.1".to_string()),
1106            port: Some(5432),
1107            ..Default::default()
1108        };
1109        let opts = default_opts();
1110        let (tx, mut rx) = mpsc::channel::<Output>(4);
1111        let req = test_request(&cfg, &opts, LogFilters::default(), tx);
1112        let selected = super::super::session::TransportSelection { duration_ms: 7 };
1113        assert!(emit_transport_selected(&req, Some(selected)).await.is_ok());
1114        assert!(
1115            rx.try_recv().is_err(),
1116            "log filter empty must suppress emission"
1117        );
1118    }
1119
1120    #[tokio::test]
1121    async fn emit_transport_selected_skips_when_selection_none() {
1122        let cfg = SessionConfig::default();
1123        let opts = default_opts();
1124        let (tx, mut rx) = mpsc::channel::<Output>(4);
1125        let req = test_request(&cfg, &opts, LogFilters::new(["transport"]), tx);
1126        assert!(emit_transport_selected(&req, None).await.is_ok());
1127        assert!(
1128            rx.try_recv().is_err(),
1129            "no selection must suppress emission"
1130        );
1131    }
1132
1133    async fn assert_transport_event(cfg: SessionConfig, chain_substring: &str, duration_ms: u64) {
1134        let opts = default_opts();
1135        let (tx, mut rx) = mpsc::channel::<Output>(4);
1136        let req = test_request(&cfg, &opts, LogFilters::new(["transport"]), tx);
1137        let selected = super::super::session::TransportSelection { duration_ms };
1138        assert!(emit_transport_selected(&req, Some(selected)).await.is_ok());
1139        let received = rx.try_recv().ok();
1140        assert!(
1141            matches!(received, Some(Output::Log { .. })),
1142            "expected Output::Log, got {received:?}"
1143        );
1144        let Some(Output::Log {
1145            event,
1146            session,
1147            chain,
1148            trace,
1149            ..
1150        }) = received
1151        else {
1152            return;
1153        };
1154        assert_eq!(event, "transport.selected");
1155        assert_eq!(session.as_deref(), Some("default"));
1156        let chain = chain.unwrap_or_default();
1157        assert!(
1158            chain.contains(chain_substring),
1159            "chain {chain:?} missing {chain_substring:?}"
1160        );
1161        assert_eq!(trace.duration_ms, duration_ms);
1162        assert!(trace.row_count.is_none());
1163        assert!(trace.payload_bytes.is_none());
1164    }
1165
1166    #[tokio::test]
1167    async fn emit_transport_selected_direct_chain_includes_postgres_endpoint() {
1168        let cfg = SessionConfig {
1169            host: Some("127.0.0.1".to_string()),
1170            port: Some(5432),
1171            ..Default::default()
1172        };
1173        assert_transport_event(cfg, "127.0.0.1:5432", 11).await;
1174    }
1175
1176    #[tokio::test]
1177    async fn emit_transport_selected_ssh_chain_includes_ssh_segment() {
1178        let cfg = SessionConfig {
1179            ssh: SshConfig {
1180                destination: Some("root@example.com".to_string()),
1181                ..Default::default()
1182            },
1183            host: Some("127.0.0.1".to_string()),
1184            port: Some(5432),
1185            ..Default::default()
1186        };
1187        assert_transport_event(cfg, "ssh:root@example.com ->", 22).await;
1188    }
1189
1190    #[tokio::test]
1191    async fn emit_transport_selected_container_chain_includes_exec_segment() {
1192        let cfg = SessionConfig {
1193            container: ContainerConfig {
1194                kubectl_pod: Some("app-pod".to_string()),
1195                kubectl_container: Some("postgres".to_string()),
1196                ..Default::default()
1197            },
1198            host: Some("127.0.0.1".to_string()),
1199            port: Some(5432),
1200            ..Default::default()
1201        };
1202        assert_transport_event(cfg, "kubectl exec app-pod -c postgres", 33).await;
1203    }
1204}
1205
1206const INLINE_PORTAL_BATCH_ROWS: usize = 1024;
1207
1208struct InlineRowCollector {
1209    columns: Vec<ColumnInfo>,
1210    rows: Vec<Value>,
1211    row_count: usize,
1212    payload_bytes: usize,
1213    max_rows: usize,
1214    max_bytes: usize,
1215    truncated_at_rows: Option<usize>,
1216    truncated_at_bytes: Option<usize>,
1217}
1218
1219impl InlineRowCollector {
1220    fn new(columns: Vec<ColumnInfo>, max_rows: usize, max_bytes: usize) -> Self {
1221        Self {
1222            columns,
1223            rows: vec![],
1224            row_count: 0,
1225            payload_bytes: 0,
1226            max_rows,
1227            max_bytes,
1228            truncated_at_rows: None,
1229            truncated_at_bytes: None,
1230        }
1231    }
1232
1233    fn is_truncated(&self) -> bool {
1234        self.truncated_at_rows.is_some() || self.truncated_at_bytes.is_some()
1235    }
1236
1237    /// Try to append a row. Returns `Ok(true)` if the row was accepted;
1238    /// `Ok(false)` if the inline limit fired and the collector now refuses
1239    /// further rows. Never errors — callers should treat `Ok(false)` as a
1240    /// signal to stop fetching from the portal.
1241    fn push(&mut self, row: Value, row_bytes: usize) -> Result<bool, ExecError> {
1242        if self.is_truncated() {
1243            return Ok(false);
1244        }
1245        let next_row_count = self.row_count.saturating_add(1);
1246        let next_payload_bytes = self.payload_bytes.saturating_add(row_bytes);
1247        if next_row_count > self.max_rows {
1248            self.truncated_at_rows = Some(self.max_rows);
1249            return Ok(false);
1250        }
1251        if next_payload_bytes > self.max_bytes {
1252            self.truncated_at_bytes = Some(self.max_bytes);
1253            return Ok(false);
1254        }
1255
1256        self.row_count = next_row_count;
1257        self.payload_bytes = next_payload_bytes;
1258        self.rows.push(row);
1259        Ok(true)
1260    }
1261}
1262
1263async fn collect_rows_wrapped_or_direct(
1264    tx: &mut tokio_postgres::Transaction<'_>,
1265    req: &ExecRequest<'_>,
1266    stmt: &tokio_postgres::Statement,
1267    bind_refs: Vec<&(dyn ToSql + Sync)>,
1268    collector: &mut InlineRowCollector,
1269) -> Result<(), ExecError> {
1270    let wrapped = wrapped_rows_sql(req.sql);
1271    tx.execute("savepoint afpsql_wrap", &[])
1272        .await
1273        .map_err(map_pg_error)?;
1274
1275    let batch_rows = req.opts.batch_rows.clamp(1, INLINE_PORTAL_BATCH_ROWS);
1276    match bind_wrapped_rows(tx, &wrapped, req.params).await {
1277        Ok(portal) => {
1278            let result = collect_portal_rows(tx, &portal, collector, batch_rows).await;
1279            if let Err(error) = result {
1280                rollback_wrap_savepoint(tx).await?;
1281                return Err(error);
1282            }
1283            release_wrap_savepoint(tx).await?;
1284            Ok(())
1285        }
1286        Err(ExecError::InvalidParams(message)) => {
1287            rollback_wrap_savepoint(tx).await?;
1288            Err(ExecError::InvalidParams(message))
1289        }
1290        Err(error) => {
1291            rollback_wrap_savepoint(tx).await?;
1292            if !fallback_columns_supported(stmt) {
1293                return Err(error);
1294            }
1295            emit_row_encoding_degraded(req.transport_log.as_ref(), &error).await?;
1296            let portal = tx.bind(stmt, &bind_refs).await.map_err(map_pg_error)?;
1297            collect_portal_rows_fallback(tx, &portal, collector, batch_rows).await
1298        }
1299    }
1300}
1301
1302/// Prepare the `to_jsonb` wrapper against a client that already sits inside the
1303/// caller's explicit transaction. Parse failures (a utility statement such as
1304/// `EXPLAIN`/`SHOW`, which cannot appear in a CTE) abort the surrounding
1305/// transaction, so the attempt runs under its own savepoint and the caller can
1306/// still fall back to the unwrapped statement. Parse never executes the
1307/// statement, so falling back cannot run it twice.
1308async fn prepare_wrapped_on_client(
1309    client: &tokio_postgres::Client,
1310    req: &ExecRequest<'_>,
1311) -> Result<tokio_postgres::Statement, ExecError> {
1312    client
1313        .batch_execute("SAVEPOINT afpsql_wrap")
1314        .await
1315        .map_err(map_pg_error)?;
1316    match client.prepare(&wrapped_rows_sql(req.sql)).await {
1317        Ok(stmt) => {
1318            client
1319                .batch_execute("RELEASE SAVEPOINT afpsql_wrap")
1320                .await
1321                .map_err(map_pg_error)?;
1322            validate_param_count(stmt.params().len(), req.params.len())?;
1323            Ok(stmt)
1324        }
1325        Err(error) => {
1326            client
1327                .batch_execute("ROLLBACK TO SAVEPOINT afpsql_wrap; RELEASE SAVEPOINT afpsql_wrap")
1328                .await
1329                .map_err(map_pg_error)?;
1330            Err(map_pg_error(error))
1331        }
1332    }
1333}
1334
1335async fn bind_wrapped_rows(
1336    tx: &mut tokio_postgres::Transaction<'_>,
1337    wrapped_sql: &str,
1338    params: &[Value],
1339) -> Result<tokio_postgres::Portal, ExecError> {
1340    let wrapped_stmt = tx.prepare(wrapped_sql).await.map_err(map_pg_error)?;
1341    validate_param_count(wrapped_stmt.params().len(), params.len())?;
1342    let wrapped_params = build_params(params, wrapped_stmt.params())?;
1343    let wrapped_refs = build_param_refs(&wrapped_params);
1344    tx.bind(&wrapped_stmt, &wrapped_refs)
1345        .await
1346        .map_err(map_pg_error)
1347}
1348
1349async fn collect_portal_rows(
1350    tx: &mut tokio_postgres::Transaction<'_>,
1351    portal: &tokio_postgres::Portal,
1352    collector: &mut InlineRowCollector,
1353    batch_rows: usize,
1354) -> Result<(), ExecError> {
1355    loop {
1356        let fetch_rows = inline_fetch_rows(collector, batch_rows);
1357        if drain_portal_batch(tx, portal, collector, fetch_rows).await? {
1358            return Ok(());
1359        }
1360    }
1361}
1362
1363async fn collect_portal_rows_fallback(
1364    tx: &mut tokio_postgres::Transaction<'_>,
1365    portal: &tokio_postgres::Portal,
1366    collector: &mut InlineRowCollector,
1367    batch_rows: usize,
1368) -> Result<(), ExecError> {
1369    loop {
1370        let fetch_rows = inline_fetch_rows(collector, batch_rows);
1371        let stream = tx
1372            .query_portal_raw(portal, fetch_rows)
1373            .await
1374            .map_err(map_pg_error)?;
1375        let mut rows = pin!(stream);
1376        while let Some(row) = rows.try_next().await.map_err(map_pg_error)? {
1377            let value = row_to_json_fallback(&row)?;
1378            let row_bytes = row_json_size(&value);
1379            let _ = collector.push(value, row_bytes)?;
1380        }
1381        if rows.rows_affected().is_some() || collector.is_truncated() {
1382            return Ok(());
1383        }
1384    }
1385}
1386
1387fn inline_fetch_rows(collector: &InlineRowCollector, batch_rows: usize) -> i32 {
1388    let remaining = collector.max_rows.saturating_sub(collector.row_count);
1389    let fetch_rows = remaining.saturating_add(1).min(batch_rows).max(1);
1390    fetch_rows.min(i32::MAX as usize) as i32
1391}
1392
1393async fn drain_portal_batch(
1394    tx: &mut tokio_postgres::Transaction<'_>,
1395    portal: &tokio_postgres::Portal,
1396    collector: &mut InlineRowCollector,
1397    fetch_rows: i32,
1398) -> Result<bool, ExecError> {
1399    let stream = tx
1400        .query_portal_raw(portal, fetch_rows)
1401        .await
1402        .map_err(map_pg_error)?;
1403    let mut rows = pin!(stream);
1404
1405    while let Some(row) = rows.try_next().await.map_err(map_pg_error)? {
1406        let value = wrapped_row_to_json(&row)?;
1407        let row_bytes = row_json_size(&value);
1408        // collector.push returns Ok(false) once the inline cap is hit; we
1409        // keep draining the current portal batch so PG's protocol stays in
1410        // a clean state, but stop accepting new rows.
1411        let _ = collector.push(value, row_bytes)?;
1412    }
1413
1414    let portal_exhausted = rows.rows_affected().is_some();
1415    Ok(portal_exhausted || collector.is_truncated())
1416}
1417
1418struct StreamStats {
1419    row_count: usize,
1420    payload_bytes: usize,
1421}
1422
1423async fn stream_rows_wrapped_or_direct(
1424    tx: &mut tokio_postgres::Transaction<'_>,
1425    req: &ExecRequest<'_>,
1426    stmt: &tokio_postgres::Statement,
1427    bind_refs: Vec<&(dyn ToSql + Sync)>,
1428    columns: Vec<ColumnInfo>,
1429    sink: &mut (dyn RowSink + Send),
1430) -> Result<StreamStats, ExecError> {
1431    let wrapped = wrapped_rows_sql(req.sql);
1432    tx.execute("savepoint afpsql_wrap", &[])
1433        .await
1434        .map_err(map_pg_error)?;
1435
1436    match bind_wrapped_rows(tx, &wrapped, req.params).await {
1437        Ok(portal) => {
1438            sink.start(columns).await?;
1439            let stats = drain_portal_stream(tx, &portal, sink).await;
1440            if let Err(error) = stats {
1441                rollback_wrap_savepoint(tx).await?;
1442                return Err(error);
1443            }
1444            release_wrap_savepoint(tx).await?;
1445            stats
1446        }
1447        Err(error) => {
1448            rollback_wrap_savepoint(tx).await?;
1449            if !fallback_columns_supported(stmt) {
1450                return Err(error);
1451            }
1452            emit_row_encoding_degraded(req.transport_log.as_ref(), &error).await?;
1453            let stream = tx.query_raw(stmt, bind_refs).await.map_err(map_pg_error)?;
1454            sink.start(columns).await?;
1455            drain_row_stream_fallback(stream, sink).await
1456        }
1457    }
1458}
1459
1460async fn drain_row_stream_fallback(
1461    stream: tokio_postgres::RowStream,
1462    sink: &mut (dyn RowSink + Send),
1463) -> Result<StreamStats, ExecError> {
1464    let mut rows = pin!(stream);
1465    let mut row_count = 0usize;
1466    let mut payload_bytes = 0usize;
1467    while let Some(row) = rows.try_next().await.map_err(map_pg_error)? {
1468        let value = row_to_json_fallback(&row)?;
1469        let row_bytes = row_json_size(&value);
1470        payload_bytes += row_bytes;
1471        row_count += 1;
1472        sink.row(value, row_bytes).await?;
1473    }
1474    Ok(StreamStats {
1475        row_count,
1476        payload_bytes,
1477    })
1478}
1479
1480async fn drain_portal_stream(
1481    tx: &mut tokio_postgres::Transaction<'_>,
1482    portal: &tokio_postgres::Portal,
1483    sink: &mut (dyn RowSink + Send),
1484) -> Result<StreamStats, ExecError> {
1485    let mut row_count = 0usize;
1486    let mut payload_bytes = 0usize;
1487    loop {
1488        let stream = tx
1489            .query_portal_raw(portal, INLINE_PORTAL_BATCH_ROWS as i32)
1490            .await
1491            .map_err(map_pg_error)?;
1492        let mut rows = pin!(stream);
1493        while let Some(row) = rows.try_next().await.map_err(map_pg_error)? {
1494            let value = wrapped_row_to_json(&row)?;
1495            let row_bytes = row_json_size(&value);
1496            payload_bytes += row_bytes;
1497            row_count += 1;
1498            sink.row(value, row_bytes).await?;
1499        }
1500        if rows.rows_affected().is_some() {
1501            return Ok(StreamStats {
1502                row_count,
1503                payload_bytes,
1504            });
1505        }
1506    }
1507}
1508
1509fn wrapped_row_to_json(row: &tokio_postgres::Row) -> Result<Value, ExecError> {
1510    row.try_get::<_, Value>("row_json").map_err(|error| {
1511        ExecError::Internal(format!(
1512            "wrapped PostgreSQL row did not contain valid row_json: {error}"
1513        ))
1514    })
1515}
1516
1517fn wrapped_rows_sql(sql: &str) -> String {
1518    // Preserve PostgreSQL's own type serialization for SELECT and RETURNING-style rows.
1519    let sql = super::trim_trailing_statement_terminators(sql);
1520    format!(
1521        "with __afpsql_rows as ({sql}\n) select to_jsonb(__afpsql_rows) as row_json from __afpsql_rows"
1522    )
1523}
1524
1525async fn rollback_wrap_savepoint(
1526    tx: &mut tokio_postgres::Transaction<'_>,
1527) -> Result<(), ExecError> {
1528    tx.execute("rollback to savepoint afpsql_wrap", &[])
1529        .await
1530        .map_err(map_pg_error)?;
1531    release_wrap_savepoint(tx).await
1532}
1533
1534async fn release_wrap_savepoint(tx: &mut tokio_postgres::Transaction<'_>) -> Result<(), ExecError> {
1535    tx.execute("release savepoint afpsql_wrap", &[])
1536        .await
1537        .map_err(map_pg_error)?;
1538    Ok(())
1539}
1540
1541async fn apply_query_settings(
1542    tx: &mut tokio_postgres::Transaction<'_>,
1543    opts: &ResolvedOptions,
1544) -> Result<(), ExecError> {
1545    let statement_timeout = format!("{}ms", opts.statement_timeout_ms);
1546    tx.execute(
1547        "select set_config('statement_timeout', $1, true)",
1548        &[&statement_timeout],
1549    )
1550    .await
1551    .map_err(map_pg_error)?;
1552
1553    let lock_timeout = format!("{}ms", opts.lock_timeout_ms);
1554    tx.execute(
1555        "select set_config('lock_timeout', $1, true)",
1556        &[&lock_timeout],
1557    )
1558    .await
1559    .map_err(map_pg_error)?;
1560
1561    Ok(())
1562}