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 truncated: bool,
29 truncated_at_rows: Option<usize>,
31 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 async fn prepare_only(&self, req: ExecRequest<'_>) -> Result<DryRunOutcome, ExecError>;
87
88 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 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 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
523async 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 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 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 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}