1use std::collections::HashMap;
7use std::path::{Path, PathBuf};
8use std::sync::{Arc, RwLock};
9
10use deadpool_postgres::{Object, Pool};
11use nexql_conn::{
12 ConnectionParams, PoolOptions, ProfileConfig, ResolvedConnection, checkout_guarded, create_pool,
13};
14use nexql_index::IndexStore;
15use nexql_policy::{AccessMode, PolicyCaps, PolicyFilter};
16use tokio::sync::RwLock as AsyncRwLock;
17
18use crate::error::ToolError;
19
20pub fn default_index_root() -> PathBuf {
22 if let Ok(p) = std::env::var("NEXQL_MCP_INDEX_DIR") {
23 return PathBuf::from(p);
24 }
25 let home = std::env::var("HOME").unwrap_or_else(|_| ".".into());
26 PathBuf::from(home).join(".local/share/nexql-mcp")
27}
28
29pub fn resolve_index_root(
32 workspace_root: Option<&Path>,
33 project_config: Option<&nexql_conn::ProjectConfigFile>,
34) -> PathBuf {
35 if let Ok(p) = std::env::var("NEXQL_MCP_INDEX_DIR") {
36 return PathBuf::from(p);
37 }
38 if let (Some(root), Some(cfg)) = (workspace_root, project_config) {
39 if let Some(ref index_dir) = cfg.index_dir {
40 return root.join(".nexql").join(index_dir);
41 }
42 }
43 default_index_root()
44}
45
46#[derive(Debug, Clone)]
47pub struct ConnectionPolicy {
48 pub access_mode: AccessMode,
49 pub caps: PolicyCaps,
50 pub filter: PolicyFilter,
51 pub environment: Option<String>,
52}
53
54#[derive(Debug, Clone)]
55pub struct ConnectionInfo {
56 pub id: String,
57 pub name: String,
58 pub host: Option<String>,
59 pub port: Option<u16>,
60 pub database: Option<String>,
61 pub params: ConnectionParams,
62 pub policy: ConnectionPolicy,
63}
64
65#[derive(Debug, Clone)]
66struct ActivePolicy {
67 access_mode: AccessMode,
68 caps: PolicyCaps,
69 filter: PolicyFilter,
70 pool_opts: PoolOptions,
71}
72
73pub struct ToolSession {
74 pub connections: Vec<ConnectionInfo>,
75 policy: RwLock<ActivePolicy>,
76 pub index_store: Option<IndexStore>,
78 inner: AsyncRwLock<SessionInner>,
79}
80
81fn filter_from_profile(profile: &ProfileConfig) -> PolicyFilter {
82 PolicyFilter {
83 allow_schemas: profile.schemas.clone(),
84 deny_schemas: profile.deny_schemas.clone(),
85 deny_tables: profile.deny_tables.clone(),
86 pii_columns: profile.pii_columns.clone(),
87 }
88}
89
90pub fn policy_from_profile(
91 profile: Option<&ProfileConfig>,
92 default_mode: AccessMode,
93 default_caps: PolicyCaps,
94) -> ConnectionPolicy {
95 let access_mode = profile
96 .and_then(|p| p.access_mode.as_deref())
97 .and_then(|m| m.parse::<AccessMode>().ok())
98 .unwrap_or(default_mode);
99 let mut caps = default_caps;
100 if let Some(n) = profile.and_then(|p| p.max_rows) {
101 caps = caps.with_max_rows(n);
102 }
103 let filter = profile.map(filter_from_profile).unwrap_or_default();
104 ConnectionPolicy {
105 access_mode,
106 caps,
107 filter,
108 environment: None,
109 }
110}
111
112struct SessionInner {
113 active_id: String,
114 database: String,
115 pools: HashMap<String, Pool>,
116}
117
118fn pool_key(connection_id: &str, database: &str) -> String {
119 format!("{connection_id}\0{database}")
120}
121
122fn params_for_database(base: &ConnectionParams, database: &str) -> ConnectionParams {
123 let mut params = base.clone();
124 params.dbname = Some(database.to_string());
125 if params.host.is_some() {
127 params.url = None;
128 }
129 params
130}
131
132fn active_policy_from(policy: &ConnectionPolicy) -> ActivePolicy {
133 ActivePolicy {
134 access_mode: policy.access_mode,
135 caps: policy.caps.clone(),
136 filter: policy.filter.clone(),
137 pool_opts: PoolOptions {
138 read_only: !policy.access_mode.allows_writes(),
139 ..Default::default()
140 },
141 }
142}
143
144impl ToolSession {
145 pub fn access_mode(&self) -> AccessMode {
146 self.policy
147 .read()
148 .map(|p| p.access_mode)
149 .unwrap_or(AccessMode::Read)
150 }
151
152 pub fn caps(&self) -> PolicyCaps {
153 self.policy
154 .read()
155 .map(|p| p.caps.clone())
156 .unwrap_or_default()
157 }
158
159 pub fn filter(&self) -> PolicyFilter {
160 self.policy
161 .read()
162 .map(|p| p.filter.clone())
163 .unwrap_or_default()
164 }
165
166 pub fn pool_opts(&self) -> PoolOptions {
167 self.policy
168 .read()
169 .map(|p| p.pool_opts.clone())
170 .unwrap_or_default()
171 }
172
173 fn apply_policy(&self, policy: &ConnectionPolicy) {
174 if let Ok(mut active) = self.policy.write() {
175 *active = active_policy_from(policy);
176 }
177 }
178
179 pub async fn from_resolved(
180 resolved: ResolvedConnection,
181 access_mode: AccessMode,
182 caps: PolicyCaps,
183 ) -> Result<Arc<Self>, ToolError> {
184 let id = resolved
185 .profile_name
186 .clone()
187 .unwrap_or_else(|| "default".into());
188 let conn_policy = policy_from_profile(resolved.profile.as_ref(), access_mode, caps);
189 let info = ConnectionInfo {
190 id: id.clone(),
191 name: id.clone(),
192 host: resolved.params.host.clone(),
193 port: resolved.params.port,
194 database: resolved.params.dbname.clone(),
195 params: resolved.params,
196 policy: conn_policy,
197 };
198 Self::from_connections(vec![info], Some(id)).await
199 }
200
201 pub async fn from_connections(
202 connections: Vec<ConnectionInfo>,
203 active_id: Option<String>,
204 ) -> Result<Arc<Self>, ToolError> {
205 if connections.is_empty() {
206 return Err(ToolError::Execution("no connections configured".into()));
207 }
208 let active = active_id.unwrap_or_else(|| connections[0].id.clone());
209 let active_conn = connections
210 .iter()
211 .find(|c| c.id == active)
212 .unwrap_or(&connections[0]);
213 let active_policy = active_policy_from(&active_conn.policy);
214 let pool_opts = active_policy.pool_opts.clone();
215 let mut pools = HashMap::new();
216 for c in &connections {
217 let database = c.database.as_deref().unwrap_or("postgres");
218 let pool = create_pool(&c.params, &pool_opts).await?;
219 pools.insert(pool_key(&c.id, database), pool);
220 }
221 let database = active_conn
222 .database
223 .clone()
224 .unwrap_or_else(|| "postgres".into());
225 Ok(Arc::new(Self {
226 connections,
227 policy: RwLock::new(active_policy),
228 index_store: Some(IndexStore::new(default_index_root())),
229 inner: AsyncRwLock::new(SessionInner {
230 active_id: active,
231 database,
232 pools,
233 }),
234 }))
235 }
236
237 pub async fn from_connections_with_filter(
238 connections: Vec<ConnectionInfo>,
239 access_mode: AccessMode,
240 caps: PolicyCaps,
241 active_id: Option<String>,
242 filter: PolicyFilter,
243 ) -> Result<Arc<Self>, ToolError> {
244 let connections = connections
245 .into_iter()
246 .map(|mut c| {
247 c.policy.access_mode = access_mode;
248 c.policy.caps = caps.clone();
249 c.policy.filter = filter.clone();
250 c
251 })
252 .collect();
253 Self::from_connections(connections, active_id).await
254 }
255
256 pub async fn active_context(&self) -> (String, String) {
257 let g = self.inner.read().await;
258 (g.active_id.clone(), g.database.clone())
259 }
260
261 pub async fn switch(
262 &self,
263 connection_id: &str,
264 database: Option<String>,
265 ) -> Result<(), ToolError> {
266 let conn = self
267 .connections
268 .iter()
269 .find(|c| c.id == connection_id)
270 .ok_or_else(|| {
271 ToolError::Execution(format!(
272 "Connection not found for ID: {connection_id} — call list_connections"
273 ))
274 })?;
275 let target_db = database
276 .or_else(|| conn.database.clone())
277 .unwrap_or_else(|| "postgres".into());
278 self.apply_policy(&conn.policy);
279 let pool_opts = self.pool_opts();
280 let key = pool_key(connection_id, &target_db);
281 let mut g = self.inner.write().await;
282 if let std::collections::hash_map::Entry::Vacant(e) = g.pools.entry(key) {
283 let params = params_for_database(&conn.params, &target_db);
284 let pool = create_pool(¶ms, &pool_opts).await?;
285 e.insert(pool);
286 }
287 g.active_id = connection_id.to_string();
288 g.database = target_db;
289 Ok(())
290 }
291
292 pub async fn checkout(&self) -> Result<Object, ToolError> {
293 let g = self.inner.read().await;
294 let key = pool_key(&g.active_id, &g.database);
295 let pool = g
296 .pools
297 .get(&key)
298 .ok_or_else(|| ToolError::Execution("no active pool".into()))?;
299 let pool_opts = self.pool_opts();
300 Ok(checkout_guarded(pool, &pool_opts).await?)
301 }
302
303 #[cfg(test)]
305 pub fn for_tests(
306 connections: Vec<ConnectionInfo>,
307 filter: PolicyFilter,
308 index_store: Option<IndexStore>,
309 ) -> Arc<Self> {
310 assert!(!connections.is_empty());
311 let active = connections[0].id.clone();
312 let database = connections[0]
313 .database
314 .clone()
315 .unwrap_or_else(|| "postgres".into());
316 let policy = ConnectionPolicy {
317 access_mode: AccessMode::Read,
318 caps: PolicyCaps::default(),
319 filter,
320 environment: None,
321 };
322 Arc::new(Self {
323 connections,
324 policy: RwLock::new(active_policy_from(&policy)),
325 index_store,
326 inner: AsyncRwLock::new(SessionInner {
327 active_id: active,
328 database,
329 pools: HashMap::new(),
330 }),
331 })
332 }
333}
334
335#[cfg(test)]
336mod tests {
337 use super::*;
338
339 #[test]
340 fn filter_maps_profile_fields() {
341 let profile = ProfileConfig {
342 schemas: vec!["public".into()],
343 deny_schemas: vec!["pgboss".into()],
344 deny_tables: vec!["auth.*".into()],
345 pii_columns: vec!["public.users.ssn".into()],
346 ..Default::default()
347 };
348 let f = filter_from_profile(&profile);
349 assert_eq!(f.allow_schemas, vec!["public"]);
350 assert_eq!(f.deny_schemas, vec!["pgboss"]);
351 assert_eq!(f.deny_tables, vec!["auth.*"]);
352 assert_eq!(f.pii_columns, vec!["public.users.ssn"]);
353 assert!(f.allows_schema("public"));
354 assert!(!f.allows_schema("pgboss"));
355 assert!(!f.allows_table("auth", "sessions"));
356 }
357}