Skip to main content

lean_ctx/core/providers/
postgres.rs

1//! PostgreSQL provider — database schema introspection via `psql`.
2//!
3//! Extracts table/column definitions from `information_schema` to make
4//! database structure available as context. Uses `psql` CLI to avoid
5//! adding a native PG driver dependency.
6//!
7//! Configuration via environment variables:
8//!   - `DATABASE_URL`: Full connection string (e.g., "postgres://user:pass@host/db")
9//!   - Or individual: `PGHOST`, `PGPORT`, `PGDATABASE`, `PGUSER`, `PGPASSWORD`
10
11use crate::core::providers::{ContextProvider, ProviderItem, ProviderParams, ProviderResult};
12
13pub struct PostgresProvider {
14    available: bool,
15}
16
17impl Default for PostgresProvider {
18    fn default() -> Self {
19        Self::new()
20    }
21}
22
23impl PostgresProvider {
24    pub fn new() -> Self {
25        let available =
26            std::env::var("DATABASE_URL").is_ok() || std::env::var("PGDATABASE").is_ok();
27        Self { available }
28    }
29}
30
31impl ContextProvider for PostgresProvider {
32    fn id(&self) -> &'static str {
33        "postgres"
34    }
35
36    fn display_name(&self) -> &'static str {
37        "PostgreSQL"
38    }
39
40    fn supported_actions(&self) -> &[&str] {
41        &["schemas", "tables"]
42    }
43
44    fn execute(&self, action: &str, params: &ProviderParams) -> Result<ProviderResult, String> {
45        if !self.available {
46            return Err("PostgreSQL not configured (need DATABASE_URL or PGDATABASE)".into());
47        }
48        match action {
49            "schemas" | "tables" => list_tables(params),
50            _ => Err(format!("Unsupported action: {action}")),
51        }
52    }
53
54    fn cache_ttl_secs(&self) -> u64 {
55        300
56    }
57
58    fn requires_auth(&self) -> bool {
59        true
60    }
61
62    fn is_available(&self) -> bool {
63        self.available
64    }
65}
66
67/// Validates a PostgreSQL identifier before it is interpolated into SQL.
68///
69/// The schema name comes from provider params (agent-controlled), and the query
70/// is executed via `psql -c`, so parameterized queries are not available.
71/// A strict identifier whitelist (`[A-Za-z_][A-Za-z0-9_$]*`, max 63 bytes — the
72/// PostgreSQL `NAMEDATALEN` limit) makes injection impossible.
73fn validate_pg_identifier(name: &str) -> Result<(), String> {
74    let valid_start = name
75        .chars()
76        .next()
77        .is_some_and(|c| c.is_ascii_alphabetic() || c == '_');
78    let valid_rest = name
79        .chars()
80        .skip(1)
81        .all(|c| c.is_ascii_alphanumeric() || c == '_' || c == '$');
82    if name.is_empty() || name.len() > 63 || !valid_start || !valid_rest {
83        return Err(format!(
84            "Invalid PostgreSQL schema identifier: {name:?} (allowed: [A-Za-z_][A-Za-z0-9_$]*, max 63 chars)"
85        ));
86    }
87    Ok(())
88}
89
90fn list_tables(params: &ProviderParams) -> Result<ProviderResult, String> {
91    let schema = params.state.as_deref().unwrap_or("public");
92    validate_pg_identifier(schema)?;
93    let limit = params.limit.unwrap_or(50);
94
95    let query = format!(
96        "SELECT table_name, column_name, data_type, is_nullable \
97         FROM information_schema.columns \
98         WHERE table_schema = '{schema}' \
99         ORDER BY table_name, ordinal_position \
100         LIMIT {limit_cols};",
101        limit_cols = limit * 20, // ~20 columns per table avg
102    );
103
104    let mut cmd = std::process::Command::new("psql");
105
106    if let Ok(url) = std::env::var("DATABASE_URL") {
107        cmd.arg(&url);
108    }
109
110    let output = cmd
111        .args(["-t", "-A", "-F", "|", "-c", &query])
112        .output()
113        .map_err(|e| format!("Failed to run psql: {e}"))?;
114
115    if !output.status.success() {
116        let stderr = String::from_utf8_lossy(&output.stderr);
117        return Err(format!("psql error: {stderr}"));
118    }
119
120    let stdout = String::from_utf8_lossy(&output.stdout);
121    let mut tables: std::collections::BTreeMap<String, Vec<String>> =
122        std::collections::BTreeMap::new();
123
124    for line in stdout.lines() {
125        let parts: Vec<&str> = line.split('|').collect();
126        if parts.len() >= 3 {
127            let table = parts[0].trim();
128            let col = parts[1].trim();
129            let dtype = parts[2].trim();
130            let nullable = parts.get(3).map_or("", |s| s.trim());
131
132            let null_marker = if nullable == "YES" { "?" } else { "" };
133            tables
134                .entry(table.to_string())
135                .or_default()
136                .push(format!("  {col}: {dtype}{null_marker}"));
137        }
138    }
139
140    let items: Vec<ProviderItem> = tables
141        .iter()
142        .take(limit)
143        .map(|(table, columns)| {
144            let body = format!("{schema}.{table}\n{}", columns.join("\n"));
145            ProviderItem {
146                id: table.clone(),
147                title: format!("{schema}.{table}"),
148                state: Some("active".into()),
149                author: None,
150                created_at: None,
151                updated_at: None,
152                url: None,
153                labels: vec![schema.to_string()],
154                body: Some(body),
155                ..Default::default()
156            }
157        })
158        .collect();
159
160    Ok(ProviderResult {
161        provider: "postgres".into(),
162        resource_type: "schemas".into(),
163        items,
164        total_count: Some(tables.len()),
165        truncated: tables.len() > limit,
166    })
167}
168
169#[cfg(test)]
170mod tests {
171    use super::*;
172
173    #[test]
174    fn postgres_provider_unavailable_without_env() {
175        let _env_lock = crate::core::data_dir::test_env_lock();
176        crate::test_env::remove_var("DATABASE_URL");
177        crate::test_env::remove_var("PGDATABASE");
178
179        let provider = PostgresProvider::new();
180        assert!(!provider.is_available());
181        assert_eq!(provider.id(), "postgres");
182        assert!(provider.requires_auth());
183    }
184
185    #[test]
186    fn postgres_provider_supported_actions() {
187        let provider = PostgresProvider::new();
188        assert!(provider.supported_actions().contains(&"schemas"));
189        assert!(provider.supported_actions().contains(&"tables"));
190    }
191
192    // P0-5 (#417): schema names are agent-controlled and interpolated into SQL —
193    // only strict identifiers may pass.
194    #[test]
195    fn valid_pg_identifiers_pass() {
196        for ok in ["public", "my_schema", "_internal", "Schema1", "a$b"] {
197            assert!(validate_pg_identifier(ok).is_ok(), "{ok} should be valid");
198        }
199    }
200
201    #[test]
202    fn sql_injection_payloads_are_rejected() {
203        for evil in [
204            "public' UNION SELECT usename, passwd, '', '' FROM pg_shadow --",
205            "public'; DROP TABLE users; --",
206            "a\"b",
207            "a b",
208            "a;b",
209            "schema\n--",
210            "",
211            "1starts_with_digit",
212        ] {
213            assert!(
214                validate_pg_identifier(evil).is_err(),
215                "{evil:?} must be rejected"
216            );
217        }
218    }
219
220    #[test]
221    fn overlong_identifier_is_rejected() {
222        let too_long = "a".repeat(64);
223        assert!(validate_pg_identifier(&too_long).is_err());
224        let max_ok = "a".repeat(63);
225        assert!(validate_pg_identifier(&max_ok).is_ok());
226    }
227
228    #[test]
229    fn injection_via_params_state_fails_closed() {
230        let params = ProviderParams {
231            state: Some("public' OR '1'='1".into()),
232            ..Default::default()
233        };
234        let err = list_tables(&params).unwrap_err();
235        assert!(err.contains("Invalid PostgreSQL schema identifier"));
236    }
237}