Skip to main content

agent_first_psql/
cli.rs

1use std::io::{Read, Write};
2
3use crate::limits::{MAX_PARAMS, MAX_SQL_BYTES};
4use crate::types::{Permission, QueryOptions, SessionConfig};
5use agent_first_data::{cli_parse_log_filters, cli_parse_output, OutputFormat};
6use clap::{Args, CommandFactory, Parser, Subcommand, ValueEnum};
7use serde_json::{json, Value};
8use std::collections::{btree_map::Entry, BTreeMap};
9
10const STARTUP_ENV_KEYS: &[&str] = &[
11    "AFPSQL_DSN_SECRET",
12    "AFPSQL_CONNINFO_SECRET",
13    "AFPSQL_HOST",
14    "AFPSQL_PORT",
15    "AFPSQL_USER",
16    "AFPSQL_DBNAME",
17    "AFPSQL_PASSWORD_SECRET",
18    "AFPSQL_SSH",
19    "AFPSQL_SSH_LOCAL_HOST",
20    "AFPSQL_SSH_LOCAL_PORT",
21    "AFPSQL_SSH_REMOTE_SOCKET",
22    "AFPSQL_SSH_SUDO_USER",
23    "PGHOST",
24    "PGPORT",
25    "PGUSER",
26    "PGDATABASE",
27    "PGPASSWORD",
28    "PGSSLMODE",
29];
30
31pub enum Mode {
32    Cli(CliRequest),
33    Pipe(PipeInit),
34    PsqlAdmin(PsqlAdminRequest),
35    PsqlUnsupported(PsqlUnsupportedRequest),
36}
37
38pub struct PipeInit {
39    pub output: OutputFormat,
40    pub session: SessionConfig,
41    pub log: Vec<String>,
42    pub startup_args: Value,
43    pub startup_env: Value,
44    pub startup_requested: bool,
45}
46
47#[derive(Debug, Clone)]
48pub struct PsqlAdminRequest {
49    pub action: PsqlAdminAction,
50    pub output: OutputFormat,
51}
52
53#[derive(Debug, Clone)]
54pub enum PsqlAdminAction {
55    Status { bin_dir: Option<String> },
56    Install { bin_dir: Option<String> },
57    Uninstall { bin_dir: Option<String> },
58}
59
60pub struct CliRequest {
61    pub sql: String,
62    pub params: Vec<Value>,
63    pub options: QueryOptions,
64    pub session: SessionConfig,
65    pub output: OutputFormat,
66    pub output_file: Option<String>,
67    pub log_file: Option<String>,
68    pub log: Vec<String>,
69    pub startup_args: Value,
70    pub startup_env: Value,
71    pub startup_requested: bool,
72    pub dry_run: bool,
73}
74
75pub struct PsqlUnsupportedRequest {
76    pub reason: String,
77}
78
79#[derive(Clone, Copy, Debug, PartialEq, Eq, ValueEnum)]
80enum RuntimeMode {
81    Cli,
82    Pipe,
83    #[value(name = "psql")]
84    Psql,
85}
86
87#[derive(Subcommand)]
88enum AfdCommand {
89    /// Manage the local psql wrapper for afpsql --mode psql.
90    Psql(PsqlCommand),
91}
92
93#[derive(Args)]
94struct PsqlCommand {
95    #[command(subcommand)]
96    action: PsqlCliAction,
97}
98
99#[derive(Subcommand)]
100enum PsqlCliAction {
101    /// Show whether the afpsql-managed psql wrapper is installed and active.
102    Status(PsqlPathArgs),
103    /// Install an afpsql-managed psql wrapper.
104    Install(PsqlPathArgs),
105    /// Remove an afpsql-managed psql wrapper.
106    Uninstall(PsqlPathArgs),
107}
108
109#[derive(Args)]
110struct PsqlPathArgs {
111    /// Directory that contains the psql wrapper. Defaults to the afpsql executable directory.
112    #[arg(long = "bin-dir")]
113    bin_dir: Option<String>,
114}
115
116#[doc = r#"Agent-First PostgreSQL client.
117
118`afpsql` gives agents a reliable PostgreSQL contract: structured stdout
119events, explicit write permissions, stable pipe sessions, and machine-readable
120failures.
121
122### Interface Policy
123
124- default mode is canonical agent-first CLI
125- `--mode psql` is argument translation only; runtime output stays JSONL
126- stdout carries protocol events; stderr is not a protocol channel
127- native CLI and pipe mode default to read-only transactions; writes require permission
128
129### Query Sources and Parameters
130
131- use `--sql` for inline SQL or `--sql-file` for a file
132- use repeatable `--param N=value` for positional binds
133- placeholder count is validated from prepared-statement metadata, not by SQL text scanning
134
135### Connection Sources
136
137- `--dsn-secret` for a PostgreSQL URI
138- `--conninfo-secret` for libpq-style conninfo
139- or discrete `--host`, `--port`, `--user`, `--dbname`, `--password-secret`
140- add `--ssh user@server` when PostgreSQL is reachable only from the server
141- agent-first environment fallbacks: `AFPSQL_*`
142- PostgreSQL environment fallbacks: `PGHOST`, `PGPORT`, `PGUSER`, `PGDATABASE`, `PGPASSWORD`, `PGSSLMODE`
143
144### Result Shaping
145
146- default mode buffers a bounded inline result
147- use `--stream-rows` for large result sets, with `--batch-rows` and `--batch-bytes` to tune chunk size
148- `--output json|yaml|plain` changes rendering only, not the runtime schema
149
150### Examples
151
152```text
153afpsql --sql "select now() as now_rfc3339"
154afpsql --sql-file ./query.sql
155afpsql --sql "select * from users where id = $1" --param 1=123
156afpsql --dsn-secret-env DATABASE_URL --sql "select 1"
157afpsql --ssh user@server --host 127.0.0.1 --port 5432 --user app --dbname appdb --sql "select 1"
158afpsql --mode psql -h 127.0.0.1 -p 5432 -U app -d appdb -c "select 1"
159afpsql --sql "select * from big_table" --stream-rows --batch-rows 1000
160afpsql --mode pipe
161afpsql psql status
162afpsql psql install
163```
164
165### Exit Codes
166
167- `0`: query completed successfully
168- `1`: SQL error or runtime error
169- `2`: invalid CLI arguments
170"#]
171#[derive(Parser)]
172#[command(name = "afpsql", version, verbatim_doc_comment)]
173pub struct AfdCli {
174    /// Inline SQL string to execute.
175    #[arg(long, allow_hyphen_values = true, help_heading = "Query")]
176    sql: Option<String>,
177    /// Read SQL from a file.
178    #[arg(long = "sql-file", allow_hyphen_values = true, help_heading = "Query")]
179    sql_file: Option<String>,
180    /// Positional bind parameter in `N=value` form. Repeat for additional parameters.
181    #[arg(long = "param", help_heading = "Query")]
182    param: Vec<String>,
183    /// Stream large result sets as `result_rows` batches instead of a single inline result.
184    #[arg(long = "stream-rows", help_heading = "Query")]
185    stream_rows: bool,
186    /// Maximum rows per streamed batch.
187    #[arg(long = "batch-rows", help_heading = "Query")]
188    batch_rows: Option<usize>,
189    /// Soft byte target per streamed batch.
190    #[arg(long = "batch-bytes", help_heading = "Query")]
191    batch_bytes: Option<usize>,
192    /// Per-query statement timeout in milliseconds.
193    #[arg(long = "statement-timeout-ms", help_heading = "Query")]
194    statement_timeout_ms: Option<u64>,
195    /// Per-query lock timeout in milliseconds.
196    #[arg(long = "lock-timeout-ms", help_heading = "Query")]
197    lock_timeout_ms: Option<u64>,
198    /// Maximum inline rows before returning `result_too_large`.
199    #[arg(long = "inline-max-rows", help_heading = "Query")]
200    inline_max_rows: Option<usize>,
201    /// Maximum inline payload bytes before returning `result_too_large`.
202    #[arg(long = "inline-max-bytes", help_heading = "Query")]
203    inline_max_bytes: Option<usize>,
204    /// Query permission: read, write, ssh-read, or ssh-write. Defaults to read, or ssh-read with --ssh.
205    #[arg(long = "permission", value_parser = parse_permission_arg, help_heading = "Query")]
206    permission: Option<Permission>,
207    /// Preview the query without executing it
208    #[arg(long, help_heading = "Query")]
209    dry_run: bool,
210
211    /// PostgreSQL DSN URI. Redacted in structured output.
212    #[arg(long = "dsn-secret", help_heading = "Connection")]
213    dsn_secret: Option<String>,
214    /// Read PostgreSQL DSN URI from an environment variable.
215    #[arg(long = "dsn-secret-env", help_heading = "Connection")]
216    dsn_secret_env: Option<String>,
217    /// libpq-style conninfo string. Redacted in structured output.
218    #[arg(long = "conninfo-secret", help_heading = "Connection")]
219    conninfo_secret: Option<String>,
220    /// PostgreSQL host.
221    #[arg(long, help_heading = "Connection")]
222    host: Option<String>,
223    /// PostgreSQL port.
224    #[arg(long, help_heading = "Connection")]
225    port: Option<u16>,
226    /// PostgreSQL user name.
227    #[arg(long, help_heading = "Connection")]
228    user: Option<String>,
229    /// PostgreSQL database name.
230    #[arg(long, help_heading = "Connection")]
231    dbname: Option<String>,
232    /// PostgreSQL password. Redacted in structured output.
233    #[arg(long = "password-secret", help_heading = "Connection")]
234    password_secret: Option<String>,
235    /// Read PostgreSQL password from an environment variable.
236    #[arg(long = "password-secret-env", help_heading = "Connection")]
237    password_secret_env: Option<String>,
238    /// Open an SSH transport to USER@HOST before connecting to PostgreSQL.
239    #[arg(long = "ssh", help_heading = "SSH Transport")]
240    ssh: Option<String>,
241    /// Additional OpenSSH -o option. Repeat for multiple options.
242    #[arg(long = "ssh-option", help_heading = "SSH Transport")]
243    ssh_options: Vec<String>,
244    /// Local bind host for the SSH tunnel.
245    #[arg(long = "ssh-local-host", help_heading = "SSH Transport")]
246    ssh_local_host: Option<String>,
247    /// Local bind port for the SSH tunnel. Defaults to an ephemeral port.
248    #[arg(long = "ssh-local-port", help_heading = "SSH Transport")]
249    ssh_local_port: Option<u16>,
250    /// Explicit remote PostgreSQL Unix socket path for SSH forwarding.
251    #[arg(long = "ssh-remote-socket", help_heading = "SSH Transport")]
252    ssh_remote_socket: Option<String>,
253    /// Remote OS user for sudo -n Unix-socket bridge mode; requires an explicit socket.
254    #[arg(long = "ssh-sudo-user", help_heading = "SSH Transport")]
255    ssh_sudo_user: Option<String>,
256
257    /// Output format: json (default), yaml, or plain.
258    #[arg(long, default_value = "json", help_heading = "Runtime")]
259    output: String,
260    /// Diagnostic log categories.
261    #[arg(long = "log", value_delimiter = ',', help_heading = "Runtime")]
262    log: Vec<String>,
263    /// Runtime mode: canonical cli, pipe, or `psql` translation mode.
264    #[arg(long, value_enum, default_value_t = RuntimeMode::Cli, help_heading = "Runtime")]
265    mode: RuntimeMode,
266
267    #[command(subcommand)]
268    command: Option<AfdCommand>,
269}
270
271pub fn parse_args() -> Result<Mode, String> {
272    let raw: Vec<String> = std::env::args().collect();
273    if is_psql_mode_requested(&raw) {
274        return parse_psql_mode(&raw);
275    }
276    let startup_requested = startup_requested_from_raw(&raw);
277
278    // --help: recursive plain-text help (all subcommands expanded)
279    if top_level_help_requested(&raw) {
280        let _ = writeln!(
281            std::io::stdout(),
282            "{}",
283            agent_first_data::cli_render_help(&AfdCli::command(), &[])
284        );
285        std::process::exit(0);
286    }
287    // --help-markdown: Markdown for doc generation
288    if top_level_help_markdown_requested(&raw) {
289        let _ = writeln!(
290            std::io::stdout(),
291            "{}",
292            agent_first_data::cli_render_help_markdown(&AfdCli::command(), &[])
293        );
294        std::process::exit(0);
295    }
296
297    let cli = match AfdCli::try_parse_from(&raw) {
298        Ok(c) => c,
299        Err(e) => {
300            use clap::error::ErrorKind;
301            if matches!(e.kind(), ErrorKind::DisplayVersion | ErrorKind::DisplayHelp) {
302                let _ = writeln!(std::io::stdout(), "{e}");
303                std::process::exit(0);
304            }
305            return Err(e.to_string());
306        }
307    };
308    let output = parse_output(&cli.output)?;
309    let log = parse_log_categories(&cli.log);
310    let dsn_secret = resolve_secret_value(
311        "--dsn-secret",
312        cli.dsn_secret,
313        cli.dsn_secret_env.as_deref(),
314    )?;
315    let password_secret = resolve_secret_value(
316        "--password-secret",
317        cli.password_secret,
318        cli.password_secret_env.as_deref(),
319    )?;
320    let session = SessionConfig {
321        dsn_secret,
322        conninfo_secret: cli.conninfo_secret,
323        host: cli.host,
324        port: cli.port,
325        user: cli.user,
326        dbname: cli.dbname,
327        password_secret,
328        ssh: cli.ssh.or_else(|| std::env::var("AFPSQL_SSH").ok()),
329        ssh_options: cli.ssh_options,
330        ssh_local_host: cli
331            .ssh_local_host
332            .or_else(|| std::env::var("AFPSQL_SSH_LOCAL_HOST").ok()),
333        ssh_local_port: cli.ssh_local_port.or_else(|| {
334            std::env::var("AFPSQL_SSH_LOCAL_PORT")
335                .ok()
336                .and_then(|v| v.parse().ok())
337        }),
338        ssh_remote_socket: cli
339            .ssh_remote_socket
340            .or_else(|| std::env::var("AFPSQL_SSH_REMOTE_SOCKET").ok()),
341        ssh_sudo_user: cli
342            .ssh_sudo_user
343            .or_else(|| std::env::var("AFPSQL_SSH_SUDO_USER").ok()),
344    };
345    let mode_name = match cli.mode {
346        RuntimeMode::Cli => "cli",
347        RuntimeMode::Pipe => "pipe",
348        RuntimeMode::Psql => "psql",
349    };
350    let startup_env = startup_env_snapshot();
351
352    if let Some(command) = cli.command {
353        return Ok(match command {
354            AfdCommand::Psql(psql) => Mode::PsqlAdmin(PsqlAdminRequest {
355                action: psql_admin_action(psql.action),
356                output,
357            }),
358        });
359    }
360
361    match cli.mode {
362        RuntimeMode::Pipe => {
363            return Ok(Mode::Pipe(PipeInit {
364                output,
365                session,
366                log: log.clone(),
367                startup_args: startup_args(mode_name, None, None, 0),
368                startup_env,
369                startup_requested,
370            }));
371        }
372        RuntimeMode::Cli | RuntimeMode::Psql => {}
373    }
374
375    let startup_sql_file = cli.sql_file.clone();
376    let sql = load_sql(cli.sql, cli.sql_file)?;
377    let params = parse_params(&cli.param)?;
378    let startup_args = startup_args(
379        mode_name,
380        Some(&sql),
381        startup_sql_file.as_deref(),
382        params.len(),
383    );
384
385    let options = QueryOptions {
386        stream_rows: cli.stream_rows,
387        batch_rows: cli.batch_rows,
388        batch_bytes: cli.batch_bytes,
389        statement_timeout_ms: cli.statement_timeout_ms,
390        lock_timeout_ms: cli.lock_timeout_ms,
391        permission: cli.permission,
392        inline_max_rows: cli.inline_max_rows,
393        inline_max_bytes: cli.inline_max_bytes,
394    };
395
396    Ok(Mode::Cli(CliRequest {
397        sql,
398        params,
399        options,
400        session,
401        output,
402        output_file: None,
403        log_file: None,
404        log,
405        startup_args,
406        startup_env,
407        startup_requested,
408        dry_run: cli.dry_run,
409    }))
410}
411
412fn parse_psql_mode(raw: &[String]) -> Result<Mode, String> {
413    let startup_requested = startup_requested_from_raw(raw);
414    let mut state = PsqlModeState::default();
415
416    let mut i = 1usize;
417    while i < raw.len() {
418        let arg = raw[i].as_str();
419        if arg == "--" {
420            i += 1;
421            while i < raw.len() {
422                state.positionals.push(raw[i].clone());
423                i += 1;
424            }
425            break;
426        }
427        if arg.starts_with("--") {
428            parse_psql_long_arg(raw, &mut i, &mut state)?;
429            continue;
430        }
431        if arg.starts_with('-') && arg.len() > 1 {
432            parse_psql_short_arg(raw, &mut i, &mut state)?;
433            continue;
434        }
435        state.positionals.push(raw[i].clone());
436        i += 1;
437    }
438
439    if let Some(reason) = state.interactive_reason {
440        return Ok(Mode::PsqlUnsupported(PsqlUnsupportedRequest { reason }));
441    }
442
443    apply_psql_positionals(&mut state)?;
444    if state.list_databases {
445        state.sql = Some(psql_list_databases_sql());
446        state.sql_file = None;
447    }
448    if state.sql.is_none() && state.sql_file.is_none() {
449        return Ok(Mode::PsqlUnsupported(PsqlUnsupportedRequest {
450            reason: "no -c/--command, -f/--file, or -l/--list was provided".to_string(),
451        }));
452    }
453
454    let dsn_secret = resolve_secret_value(
455        "--dsn-secret",
456        state.dsn_secret,
457        state.dsn_secret_env.as_deref(),
458    )?;
459    let password_secret = resolve_secret_value(
460        "--password-secret",
461        state.password_secret,
462        state.password_secret_env.as_deref(),
463    )?;
464    let session = SessionConfig {
465        dsn_secret,
466        conninfo_secret: state.conninfo_secret,
467        host: state.host,
468        port: state.port,
469        user: state.user,
470        dbname: state.dbname,
471        password_secret,
472        ssh: None,
473        ssh_options: vec![],
474        ssh_local_host: None,
475        ssh_local_port: None,
476        ssh_remote_socket: None,
477        ssh_sudo_user: None,
478    };
479
480    let startup_sql_file = state.sql_file.clone();
481    let sql = load_sql(state.sql, state.sql_file)?;
482    let params = parse_params(&state.params_kv)?;
483    let startup_args = psql_startup_args(PsqlStartupArgs {
484        mode: "psql",
485        sql: Some(&sql),
486        sql_file: startup_sql_file,
487        param_count: params.len(),
488    });
489    Ok(Mode::Cli(CliRequest {
490        sql,
491        params,
492        options: QueryOptions {
493            permission: Some(Permission::Write),
494            ..Default::default()
495        },
496        session,
497        output: state.output,
498        output_file: state.output_file,
499        log_file: state.log_file,
500        log: parse_log_categories(&state.log_entries),
501        startup_args,
502        startup_env: startup_env_snapshot(),
503        startup_requested,
504        dry_run: false,
505    }))
506}
507
508struct PsqlModeState {
509    sql: Option<String>,
510    sql_file: Option<String>,
511    host: Option<String>,
512    port: Option<u16>,
513    user: Option<String>,
514    dbname: Option<String>,
515    dsn_secret: Option<String>,
516    dsn_secret_env: Option<String>,
517    conninfo_secret: Option<String>,
518    password_secret: Option<String>,
519    password_secret_env: Option<String>,
520    params_kv: Vec<String>,
521    output: OutputFormat,
522    log_entries: Vec<String>,
523    output_file: Option<String>,
524    log_file: Option<String>,
525    list_databases: bool,
526    positionals: Vec<String>,
527    interactive_reason: Option<String>,
528}
529
530impl Default for PsqlModeState {
531    fn default() -> Self {
532        Self {
533            sql: None,
534            sql_file: None,
535            host: None,
536            port: None,
537            user: None,
538            dbname: None,
539            dsn_secret: None,
540            dsn_secret_env: None,
541            conninfo_secret: None,
542            password_secret: None,
543            password_secret_env: None,
544            params_kv: vec![],
545            output: OutputFormat::Json,
546            log_entries: vec![],
547            output_file: None,
548            log_file: None,
549            list_databases: false,
550            positionals: vec![],
551            interactive_reason: None,
552        }
553    }
554}
555
556impl PsqlModeState {
557    fn set_sql(&mut self, sql: String, flag: &str) -> Result<(), String> {
558        if self.sql.is_some() || self.sql_file.is_some() {
559            return Err(format!(
560                "psql mode currently supports only one -c/--command or -f/--file source; repeated source at {flag}"
561            ));
562        }
563        self.sql = Some(sql);
564        Ok(())
565    }
566
567    fn set_sql_file(&mut self, path: String, flag: &str) -> Result<(), String> {
568        if self.sql.is_some() || self.sql_file.is_some() {
569            return Err(format!(
570                "psql mode currently supports only one -c/--command or -f/--file source; repeated source at {flag}"
571            ));
572        }
573        self.sql_file = Some(path);
574        Ok(())
575    }
576}
577
578fn parse_psql_long_arg(
579    raw: &[String],
580    i: &mut usize,
581    state: &mut PsqlModeState,
582) -> Result<(), String> {
583    let arg = raw[*i].as_str();
584    if arg == "--mode" {
585        let value = take_arg_value(raw, i, "--mode")?;
586        if value != "psql" {
587            return Err(format!(
588                "unsupported psql-mode argument: --mode {value}; only --mode psql is allowed with psql translation"
589            ));
590        }
591        return Ok(());
592    }
593    if let Some(value) = arg.strip_prefix("--mode=") {
594        if value != "psql" {
595            return Err(format!(
596                "unsupported psql-mode argument: {arg}; only --mode=psql is allowed with psql translation"
597            ));
598        }
599        *i += 1;
600        return Ok(());
601    }
602
603    if arg == "--help" || arg.starts_with("--help=") {
604        emit_psql_mode_help();
605        std::process::exit(0);
606    }
607    if arg == "--version" {
608        emit_psql_mode_version();
609        std::process::exit(0);
610    }
611
612    match long_name(arg) {
613        "--command" => {
614            let value = take_long_arg_value(raw, i, "--command")?;
615            state.set_sql(value, "--command")
616        }
617        "--file" => {
618            let value = take_long_arg_value(raw, i, "--file")?;
619            state.set_sql_file(value, "--file")
620        }
621        "--host" => {
622            state.host = Some(take_long_arg_value(raw, i, "--host")?);
623            Ok(())
624        }
625        "--port" => {
626            state.port = Some(parse_port(
627                &take_long_arg_value(raw, i, "--port")?,
628                "--port",
629            )?);
630            Ok(())
631        }
632        "--username" | "--user" => {
633            state.user = Some(take_long_arg_value(raw, i, long_name(arg))?);
634            Ok(())
635        }
636        "--dbname" => {
637            apply_dbname_value(state, take_long_arg_value(raw, i, "--dbname")?);
638            Ok(())
639        }
640        "--set" | "--variable" => {
641            let value = take_long_arg_value(raw, i, long_name(arg))?;
642            add_psql_variable(state, value)
643        }
644        "--list" => {
645            state.list_databases = true;
646            *i += 1;
647            Ok(())
648        }
649        "--no-password"
650        | "--no-psqlrc"
651        | "--no-readline"
652        | "--quiet"
653        | "--echo-all"
654        | "--echo-errors"
655        | "--echo-queries"
656        | "--echo-hidden"
657        | "--no-align"
658        | "--csv"
659        | "--html"
660        | "--tuples-only"
661        | "--expanded"
662        | "--field-separator-zero"
663        | "--record-separator-zero"
664        | "--single-transaction" => {
665            *i += 1;
666            Ok(())
667        }
668        "--field-separator" | "--record-separator" | "--pset" | "--table-attr" => {
669            let _ = take_long_arg_value(raw, i, long_name(arg))?;
670            Ok(())
671        }
672        "--password" => {
673            state.interactive_reason =
674                Some("--password/-W requests an interactive password prompt".to_string());
675            *i += 1;
676            Ok(())
677        }
678        "--single-step" => {
679            state.interactive_reason =
680                Some("--single-step/-s requires interactive command confirmation".to_string());
681            *i += 1;
682            Ok(())
683        }
684        "--single-line" => {
685            state.interactive_reason =
686                Some("--single-line/-S is a human-interactive input mode".to_string());
687            *i += 1;
688            Ok(())
689        }
690        "--dsn-secret" => {
691            state.dsn_secret = Some(take_long_arg_value(raw, i, "--dsn-secret")?);
692            Ok(())
693        }
694        "--dsn-secret-env" => {
695            state.dsn_secret_env = Some(take_long_arg_value(raw, i, "--dsn-secret-env")?);
696            Ok(())
697        }
698        "--conninfo-secret" => {
699            state.conninfo_secret = Some(take_long_arg_value(raw, i, "--conninfo-secret")?);
700            Ok(())
701        }
702        "--password-secret" => {
703            state.password_secret = Some(take_long_arg_value(raw, i, "--password-secret")?);
704            Ok(())
705        }
706        "--password-secret-env" => {
707            state.password_secret_env = Some(take_long_arg_value(raw, i, "--password-secret-env")?);
708            Ok(())
709        }
710        "--output" => {
711            let value = take_long_arg_value(raw, i, "--output")?;
712            if let Ok(format) = parse_output(&value) {
713                state.output = format;
714            } else {
715                state.output_file = Some(value);
716            }
717            Ok(())
718        }
719        "--output-format" => {
720            let value = take_long_arg_value(raw, i, long_name(arg))?;
721            state.output = parse_output(&value)?;
722            Ok(())
723        }
724        "--log" => {
725            let values = take_long_arg_value(raw, i, "--log")?;
726            add_log_entries(state, &values);
727            Ok(())
728        }
729        "--log-file" => {
730            state.log_file = Some(take_long_arg_value(raw, i, "--log-file")?);
731            Ok(())
732        }
733        _ => Err(format!("unsupported psql-mode argument: {arg}")),
734    }
735}
736
737fn parse_psql_short_arg(
738    raw: &[String],
739    i: &mut usize,
740    state: &mut PsqlModeState,
741) -> Result<(), String> {
742    let arg = raw[*i].as_str();
743    let mut offset = 1usize;
744    while offset < arg.len() {
745        let flag = arg.as_bytes()[offset] as char;
746        offset += 1;
747        match flag {
748            '?' => {
749                emit_psql_mode_help();
750                std::process::exit(0);
751            }
752            'V' => {
753                emit_psql_mode_version();
754                std::process::exit(0);
755            }
756            'c' => {
757                let value = take_short_arg_value(raw, i, arg, offset, "-c")?;
758                return state.set_sql(value, "-c");
759            }
760            'f' => {
761                let value = take_short_arg_value(raw, i, arg, offset, "-f")?;
762                return state.set_sql_file(value, "-f");
763            }
764            'h' => {
765                state.host = Some(take_short_arg_value(raw, i, arg, offset, "-h")?);
766                return Ok(());
767            }
768            'p' => {
769                let value = take_short_arg_value(raw, i, arg, offset, "-p")?;
770                state.port = Some(parse_port(&value, "-p")?);
771                return Ok(());
772            }
773            'U' => {
774                state.user = Some(take_short_arg_value(raw, i, arg, offset, "-U")?);
775                return Ok(());
776            }
777            'd' => {
778                apply_dbname_value(state, take_short_arg_value(raw, i, arg, offset, "-d")?);
779                return Ok(());
780            }
781            'v' => {
782                let value = take_short_arg_value(raw, i, arg, offset, "-v")?;
783                return add_psql_variable(state, value);
784            }
785            'F' | 'P' | 'R' | 'T' => {
786                let _ = take_short_arg_value(raw, i, arg, offset, &format!("-{flag}"))?;
787                return Ok(());
788            }
789            'L' => {
790                state.log_file = Some(take_short_arg_value(raw, i, arg, offset, "-L")?);
791                return Ok(());
792            }
793            'o' => {
794                state.output_file = Some(take_short_arg_value(raw, i, arg, offset, "-o")?);
795                return Ok(());
796            }
797            'l' => state.list_databases = true,
798            'W' => {
799                state.interactive_reason =
800                    Some("--password/-W requests an interactive password prompt".to_string());
801            }
802            's' => {
803                state.interactive_reason =
804                    Some("--single-step/-s requires interactive command confirmation".to_string());
805            }
806            'S' => {
807                state.interactive_reason =
808                    Some("--single-line/-S is a human-interactive input mode".to_string());
809            }
810            'a' | 'A' | 'b' | 'e' | 'E' | 'H' | 'n' | 'q' | 't' | 'w' | 'x' | 'X' | 'z' | '0'
811            | '1' => {}
812            _ => return Err(format!("unsupported psql-mode argument: -{flag}")),
813        }
814    }
815    *i += 1;
816    Ok(())
817}
818
819fn long_name(arg: &str) -> &str {
820    arg.split_once('=').map(|(name, _)| name).unwrap_or(arg)
821}
822
823fn take_arg_value(raw: &[String], i: &mut usize, flag: &str) -> Result<String, String> {
824    *i += 1;
825    let value = raw
826        .get(*i)
827        .ok_or_else(|| format!("{flag} requires value"))?
828        .clone();
829    *i += 1;
830    Ok(value)
831}
832
833fn take_long_arg_value(raw: &[String], i: &mut usize, flag: &str) -> Result<String, String> {
834    let arg = raw[*i].as_str();
835    if let Some((_, value)) = arg.split_once('=') {
836        *i += 1;
837        return Ok(value.to_string());
838    }
839    take_arg_value(raw, i, flag)
840}
841
842fn take_short_arg_value(
843    raw: &[String],
844    i: &mut usize,
845    arg: &str,
846    offset: usize,
847    flag: &str,
848) -> Result<String, String> {
849    if offset < arg.len() {
850        let value = arg[offset..].to_string();
851        *i += 1;
852        return Ok(value);
853    }
854    take_arg_value(raw, i, flag)
855}
856
857fn parse_port(value: &str, flag: &str) -> Result<u16, String> {
858    value.parse().map_err(|_| format!("invalid {flag} port"))
859}
860
861fn add_log_entries(state: &mut PsqlModeState, values: &str) {
862    for part in values.split(',') {
863        let trimmed = part.trim();
864        if !trimmed.is_empty() {
865            state.log_entries.push(trimmed.to_string());
866        }
867    }
868}
869
870fn add_psql_variable(state: &mut PsqlModeState, value: String) -> Result<(), String> {
871    let name = value
872        .split_once('=')
873        .map(|(name, _)| name)
874        .unwrap_or(value.as_str());
875    if name.parse::<usize>().is_ok() {
876        if value.contains('=') {
877            state.params_kv.push(value);
878            return Ok(());
879        }
880        return Err(format!("invalid param '{value}', expected N=value"));
881    }
882    if is_psql_behavior_variable(name) {
883        return Ok(());
884    }
885    Err(format!(
886        "invalid or unsupported psql variable '{name}'; afpsql supports numeric -v N=value bind parameters, not client-side :name interpolation"
887    ))
888}
889
890fn is_psql_behavior_variable(name: &str) -> bool {
891    matches!(
892        name.to_ascii_uppercase().as_str(),
893        "ON_ERROR_STOP"
894            | "ON_ERROR_ROLLBACK"
895            | "QUIET"
896            | "ECHO"
897            | "ECHO_HIDDEN"
898            | "FETCH_COUNT"
899            | "VERBOSITY"
900            | "SHOW_CONTEXT"
901            | "HISTCONTROL"
902            | "HISTFILE"
903            | "HISTSIZE"
904            | "IGNOREEOF"
905            | "PAGER"
906            | "COLUMNS"
907    )
908}
909
910fn apply_psql_positionals(state: &mut PsqlModeState) -> Result<(), String> {
911    let positionals = std::mem::take(&mut state.positionals);
912    for value in positionals {
913        if is_postgres_uri(&value) {
914            state.dsn_secret = Some(value);
915            continue;
916        }
917        if looks_like_conninfo(&value) {
918            state.conninfo_secret = Some(value);
919            continue;
920        }
921        if state.dbname.is_none() {
922            state.dbname = Some(value);
923            continue;
924        }
925        if state.user.is_none() {
926            state.user = Some(value);
927            continue;
928        }
929        return Err(format!("too many positional psql arguments: {value}"));
930    }
931    Ok(())
932}
933
934fn apply_dbname_value(state: &mut PsqlModeState, value: String) {
935    if is_postgres_uri(&value) {
936        state.dsn_secret = Some(value);
937    } else if looks_like_conninfo(&value) {
938        state.conninfo_secret = Some(value);
939    } else {
940        state.dbname = Some(value);
941    }
942}
943
944fn is_postgres_uri(value: &str) -> bool {
945    value.starts_with("postgresql://") || value.starts_with("postgres://")
946}
947
948fn looks_like_conninfo(value: &str) -> bool {
949    value.contains('=')
950}
951
952fn psql_list_databases_sql() -> String {
953    "select datname as name from pg_catalog.pg_database where datallowconn order by datname"
954        .to_string()
955}
956
957fn emit_psql_mode_version() {
958    let _ = writeln!(
959        std::io::stdout(),
960        "psql (afpsql wrapper) {}",
961        env!("CARGO_PKG_VERSION")
962    );
963}
964
965fn emit_psql_mode_help() {
966    let _ = writeln!(
967        std::io::stdout(),
968        "psql (afpsql wrapper) {}\n\
969Usage:\n  psql [OPTION]... [DBNAME [USERNAME]]\n\n\
970Supported non-interactive forms:\n  -c, --command=SQL\n  -f, --file=FILE\n  -l, --list\n  -h/-p/-U/-d and --host/--port/--username/--dbname\n  -v N=value, --set N=value for positional bind parameters\n\n\
971Output routing:\n  -o, --output=FILE writes structured output to FILE\n  -L, --log-file=FILE tees structured output to FILE\n  --output-format=json|yaml|plain changes afpsql rendering\n\n\
972Human-interactive psql modes and psql meta-commands are not supported by this wrapper.",
973        env!("CARGO_PKG_VERSION")
974    );
975}
976
977fn top_level_help_requested(raw: &[String]) -> bool {
978    raw.len() == 2 && matches!(raw.get(1).map(String::as_str), Some("--help" | "-h"))
979}
980
981fn top_level_help_markdown_requested(raw: &[String]) -> bool {
982    let mut i = 1usize;
983    while i < raw.len() {
984        let arg = raw[i].as_str();
985        if arg == "--" {
986            break;
987        }
988        if arg == "--help-markdown" {
989            return true;
990        }
991        if arg == "--mode" {
992            i += 2;
993            continue;
994        }
995        if top_level_arg_consumes_value(arg) {
996            i += if arg.contains('=') { 1 } else { 2 };
997            continue;
998        }
999        if arg.starts_with('-') {
1000            i += 1;
1001            continue;
1002        }
1003        break;
1004    }
1005    false
1006}
1007
1008fn psql_admin_action(action: PsqlCliAction) -> PsqlAdminAction {
1009    match action {
1010        PsqlCliAction::Status(args) => PsqlAdminAction::Status {
1011            bin_dir: args.bin_dir,
1012        },
1013        PsqlCliAction::Install(args) => PsqlAdminAction::Install {
1014            bin_dir: args.bin_dir,
1015        },
1016        PsqlCliAction::Uninstall(args) => PsqlAdminAction::Uninstall {
1017            bin_dir: args.bin_dir,
1018        },
1019    }
1020}
1021
1022fn is_psql_mode_requested(raw: &[String]) -> bool {
1023    let mut i = 1usize;
1024    while i < raw.len() {
1025        let arg = raw[i].as_str();
1026        if arg == "--" {
1027            break;
1028        }
1029        if arg == "--mode" {
1030            if let Some(v) = raw.get(i + 1) {
1031                return v == "psql";
1032            }
1033            return false;
1034        }
1035        if arg == "--mode=psql" {
1036            return true;
1037        }
1038        if top_level_arg_consumes_value(arg) {
1039            i += if arg.contains('=') { 1 } else { 2 };
1040            continue;
1041        }
1042        if arg.starts_with('-') {
1043            i += 1;
1044            continue;
1045        }
1046        break;
1047    }
1048    false
1049}
1050
1051fn top_level_arg_consumes_value(arg: &str) -> bool {
1052    let name = arg.split_once('=').map(|(name, _)| name).unwrap_or(arg);
1053    matches!(
1054        name,
1055        "--sql"
1056            | "--sql-file"
1057            | "--param"
1058            | "--batch-rows"
1059            | "--batch-bytes"
1060            | "--statement-timeout-ms"
1061            | "--lock-timeout-ms"
1062            | "--inline-max-rows"
1063            | "--inline-max-bytes"
1064            | "--permission"
1065            | "--dsn-secret"
1066            | "--dsn-secret-env"
1067            | "--conninfo-secret"
1068            | "--host"
1069            | "--port"
1070            | "--user"
1071            | "--dbname"
1072            | "--password-secret"
1073            | "--password-secret-env"
1074            | "--ssh"
1075            | "--ssh-option"
1076            | "--ssh-local-host"
1077            | "--ssh-local-port"
1078            | "--ssh-remote-socket"
1079            | "--ssh-sudo-user"
1080            | "--output"
1081            | "--log"
1082    )
1083}
1084
1085fn load_sql(sql: Option<String>, sql_file: Option<String>) -> Result<String, String> {
1086    match (sql, sql_file) {
1087        (Some(s), None) => validate_sql_size(s),
1088        (None, Some(path)) if path == "-" => {
1089            let stdin = std::io::stdin();
1090            read_limited_sql(stdin.lock(), "read --sql-file -")
1091        }
1092        (None, Some(path)) => {
1093            let metadata =
1094                std::fs::metadata(&path).map_err(|e| format!("read --sql-file failed: {e}"))?;
1095            if metadata.is_file() && metadata.len() > MAX_SQL_BYTES as u64 {
1096                return Err(sql_size_error());
1097            }
1098            let file =
1099                std::fs::File::open(&path).map_err(|e| format!("read --sql-file failed: {e}"))?;
1100            read_limited_sql(file, "read --sql-file")
1101        }
1102        (Some(_), Some(_)) => Err("--sql and --sql-file are mutually exclusive".to_string()),
1103        (None, None) => Err("one of --sql or --sql-file is required".to_string()),
1104    }
1105}
1106
1107fn read_limited_sql<R: Read>(reader: R, context: &str) -> Result<String, String> {
1108    let mut buf = Vec::new();
1109    let mut limited = reader.take(MAX_SQL_BYTES as u64 + 1);
1110    limited
1111        .read_to_end(&mut buf)
1112        .map_err(|e| format!("{context} failed: {e}"))?;
1113    if buf.len() > MAX_SQL_BYTES {
1114        return Err(sql_size_error());
1115    }
1116    String::from_utf8(buf).map_err(|e| format!("{context} failed: {e}"))
1117}
1118
1119fn validate_sql_size(sql: String) -> Result<String, String> {
1120    if sql.len() > MAX_SQL_BYTES {
1121        return Err(sql_size_error());
1122    }
1123    Ok(sql)
1124}
1125
1126fn sql_size_error() -> String {
1127    format!("sql exceeds maximum size; maximum SQL size is {MAX_SQL_BYTES} bytes")
1128}
1129
1130fn parse_output(v: &str) -> Result<OutputFormat, String> {
1131    cli_parse_output(v)
1132}
1133
1134fn parse_permission_arg(v: &str) -> Result<Permission, String> {
1135    v.parse()
1136}
1137
1138fn parse_log_categories(entries: &[String]) -> Vec<String> {
1139    cli_parse_log_filters(entries)
1140}
1141
1142fn startup_requested_from_raw(raw: &[String]) -> bool {
1143    let mut i = 1usize;
1144    while i < raw.len() {
1145        if raw[i] == "--log" {
1146            if let Some(values) = raw.get(i + 1) {
1147                for part in values.split(',') {
1148                    let v = part.trim().to_ascii_lowercase();
1149                    if matches!(v.as_str(), "startup" | "all" | "*") {
1150                        return true;
1151                    }
1152                }
1153            }
1154            i += 2;
1155            continue;
1156        }
1157        if let Some(values) = raw[i].strip_prefix("--log=") {
1158            for part in values.split(',') {
1159                let v = part.trim().to_ascii_lowercase();
1160                if matches!(v.as_str(), "startup" | "all" | "*") {
1161                    return true;
1162                }
1163            }
1164        }
1165        i += 1;
1166    }
1167    false
1168}
1169
1170fn startup_env_snapshot() -> Value {
1171    Value::Array(
1172        STARTUP_ENV_KEYS
1173            .iter()
1174            .map(|key| {
1175                json!({
1176                    "key": key,
1177                    "present": std::env::var_os(key).is_some(),
1178                })
1179            })
1180            .collect(),
1181    )
1182}
1183
1184fn startup_args(
1185    mode: &str,
1186    sql: Option<&str>,
1187    sql_file: Option<&str>,
1188    param_count: usize,
1189) -> Value {
1190    json!({
1191        "mode": mode,
1192        "sql": startup_sql_summary(sql, sql_file),
1193        "param_count": param_count,
1194    })
1195}
1196
1197fn startup_sql_summary(sql: Option<&str>, sql_file: Option<&str>) -> Value {
1198    let Some(sql) = sql else {
1199        return json!({
1200            "present": false,
1201            "source": "none",
1202            "bytes": 0,
1203            "chars": 0,
1204            "operation": null,
1205        });
1206    };
1207    json!({
1208        "present": true,
1209        "source": if sql_file.is_some() { "file" } else { "inline" },
1210        "bytes": sql.len(),
1211        "chars": sql.chars().count(),
1212        "operation": sql_operation(sql),
1213    })
1214}
1215
1216fn sql_operation(sql: &str) -> Option<String> {
1217    let sql = trim_leading_sql_comments(sql);
1218    let token: String = sql
1219        .chars()
1220        .skip_while(|c| c.is_whitespace())
1221        .take_while(|c| c.is_ascii_alphabetic() || *c == '_')
1222        .collect();
1223    if token.is_empty() {
1224        None
1225    } else {
1226        Some(token.to_ascii_lowercase())
1227    }
1228}
1229
1230fn trim_leading_sql_comments(mut sql: &str) -> &str {
1231    loop {
1232        sql = sql.trim_start();
1233        if let Some(rest) = sql.strip_prefix("--") {
1234            sql = rest.split_once('\n').map(|(_, rest)| rest).unwrap_or("");
1235            continue;
1236        }
1237        if let Some(rest) = sql.strip_prefix("/*") {
1238            let Some((_, after)) = rest.split_once("*/") else {
1239                return "";
1240            };
1241            sql = after;
1242            continue;
1243        }
1244        return sql;
1245    }
1246}
1247
1248struct PsqlStartupArgs<'a> {
1249    mode: &'a str,
1250    sql: Option<&'a str>,
1251    sql_file: Option<String>,
1252    param_count: usize,
1253}
1254
1255fn psql_startup_args(args: PsqlStartupArgs<'_>) -> Value {
1256    startup_args(
1257        args.mode,
1258        args.sql,
1259        args.sql_file.as_deref(),
1260        args.param_count,
1261    )
1262}
1263
1264fn resolve_secret_value(
1265    flag_name: &str,
1266    direct: Option<String>,
1267    env_name: Option<&str>,
1268) -> Result<Option<String>, String> {
1269    match (direct, env_name) {
1270        (Some(_), Some(_)) => Err(format!(
1271            "{flag_name} and {flag_name}-env are mutually exclusive"
1272        )),
1273        (Some(value), None) => Ok(Some(value)),
1274        (None, Some(name)) => {
1275            if name.is_empty() {
1276                return Err(format!(
1277                    "{flag_name}-env requires a non-empty variable name"
1278                ));
1279            }
1280            std::env::var(name).map(Some).map_err(|_| {
1281                format!("{flag_name}-env references unset environment variable: {name}")
1282            })
1283        }
1284        (None, None) => Ok(None),
1285    }
1286}
1287
1288pub fn parse_params(entries: &[String]) -> Result<Vec<Value>, String> {
1289    if entries.len() > MAX_PARAMS {
1290        return Err(format!("too many params; maximum params is {MAX_PARAMS}"));
1291    }
1292
1293    let mut by_index: BTreeMap<usize, Value> = BTreeMap::new();
1294    for entry in entries {
1295        let (idx, raw) = split_index_value(entry)?;
1296        if idx == 0 {
1297            return Err("param index must start at 1".to_string());
1298        }
1299        if idx > MAX_PARAMS {
1300            return Err(format!(
1301                "parameter index {idx} exceeds maximum params {MAX_PARAMS}"
1302            ));
1303        }
1304        match by_index.entry(idx) {
1305            Entry::Vacant(slot) => {
1306                slot.insert(parse_param_value(raw));
1307            }
1308            Entry::Occupied(_) => return Err(format!("duplicate parameter index {idx}")),
1309        }
1310    }
1311    if by_index.is_empty() {
1312        return Ok(vec![]);
1313    }
1314    let max = by_index.keys().max().copied().unwrap_or(0);
1315    for i in 1..=max {
1316        if !by_index.contains_key(&i) {
1317            return Err(format!("missing parameter index {i}"));
1318        }
1319    }
1320    Ok(by_index.into_values().collect())
1321}
1322
1323fn split_index_value(entry: &str) -> Result<(usize, &str), String> {
1324    let mut parts = entry.splitn(2, '=');
1325    let left = parts.next().unwrap_or_default();
1326    let right = parts
1327        .next()
1328        .ok_or_else(|| format!("invalid param '{entry}', expected N=value"))?;
1329    let idx = left
1330        .parse::<usize>()
1331        .map_err(|_| format!("invalid param index in '{entry}'"))?;
1332    Ok((idx, right))
1333}
1334
1335fn parse_param_value(v: &str) -> Value {
1336    if v == "null" {
1337        return Value::Null;
1338    }
1339    if v == "true" {
1340        return Value::Bool(true);
1341    }
1342    if v == "false" {
1343        return Value::Bool(false);
1344    }
1345    if let Ok(i) = v.parse::<i64>() {
1346        return Value::Number(i.into());
1347    }
1348    if let Ok(f) = v.parse::<f64>() {
1349        if let Some(n) = serde_json::Number::from_f64(f) {
1350            return Value::Number(n);
1351        }
1352    }
1353    Value::String(v.to_string())
1354}
1355
1356#[cfg(test)]
1357#[path = "../tests/support/unit_cli.rs"]
1358mod tests;