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