lean_ctx/core/providers/
postgres.rs1use 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
67fn 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, );
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 #[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(¶ms).unwrap_err();
235 assert!(err.contains("Invalid PostgreSQL schema identifier"));
236 }
237}