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 std::env::remove_var("DATABASE_URL");
176 std::env::remove_var("PGDATABASE");
177
178 let provider = PostgresProvider::new();
179 assert!(!provider.is_available());
180 assert_eq!(provider.id(), "postgres");
181 assert!(provider.requires_auth());
182 }
183
184 #[test]
185 fn postgres_provider_supported_actions() {
186 let provider = PostgresProvider::new();
187 assert!(provider.supported_actions().contains(&"schemas"));
188 assert!(provider.supported_actions().contains(&"tables"));
189 }
190
191 #[test]
194 fn valid_pg_identifiers_pass() {
195 for ok in ["public", "my_schema", "_internal", "Schema1", "a$b"] {
196 assert!(validate_pg_identifier(ok).is_ok(), "{ok} should be valid");
197 }
198 }
199
200 #[test]
201 fn sql_injection_payloads_are_rejected() {
202 for evil in [
203 "public' UNION SELECT usename, passwd, '', '' FROM pg_shadow --",
204 "public'; DROP TABLE users; --",
205 "a\"b",
206 "a b",
207 "a;b",
208 "schema\n--",
209 "",
210 "1starts_with_digit",
211 ] {
212 assert!(
213 validate_pg_identifier(evil).is_err(),
214 "{evil:?} must be rejected"
215 );
216 }
217 }
218
219 #[test]
220 fn overlong_identifier_is_rejected() {
221 let too_long = "a".repeat(64);
222 assert!(validate_pg_identifier(&too_long).is_err());
223 let max_ok = "a".repeat(63);
224 assert!(validate_pg_identifier(&max_ok).is_ok());
225 }
226
227 #[test]
228 fn injection_via_params_state_fails_closed() {
229 let params = ProviderParams {
230 state: Some("public' OR '1'='1".into()),
231 ..Default::default()
232 };
233 let err = list_tables(¶ms).unwrap_err();
234 assert!(err.contains("Invalid PostgreSQL schema identifier"));
235 }
236}