Skip to main content

nexql_tools/
session.rs

1//! Session state: resolved profiles + active pool + optional schema index.
2
3use std::collections::HashMap;
4use std::path::PathBuf;
5use std::sync::Arc;
6
7use deadpool_postgres::{Object, Pool};
8use nexql_conn::{
9    ConnectionParams, PoolOptions, ProfileConfig, ResolvedConnection, checkout_guarded, create_pool,
10};
11use nexql_index::IndexStore;
12use nexql_policy::{AccessMode, PolicyCaps, PolicyFilter};
13use tokio::sync::RwLock;
14
15use crate::error::ToolError;
16
17/// Index root: `NEXQL_MCP_INDEX_DIR`, else `~/.local/share/nexql-mcp` (same as CLI).
18pub fn default_index_root() -> PathBuf {
19    if let Ok(p) = std::env::var("NEXQL_MCP_INDEX_DIR") {
20        return PathBuf::from(p);
21    }
22    let home = std::env::var("HOME").unwrap_or_else(|_| ".".into());
23    PathBuf::from(home).join(".local/share/nexql-mcp")
24}
25
26#[derive(Debug, Clone)]
27pub struct ConnectionInfo {
28    pub id: String,
29    pub name: String,
30    pub host: Option<String>,
31    pub port: Option<u16>,
32    pub database: Option<String>,
33    pub params: ConnectionParams,
34}
35
36pub struct ToolSession {
37    pub connections: Vec<ConnectionInfo>,
38    pub access_mode: AccessMode,
39    pub caps: PolicyCaps,
40    pub filter: PolicyFilter,
41    pub pool_opts: PoolOptions,
42    /// Schema index root; `None` disables Phase 3 tools with an actionable error.
43    pub index_store: Option<IndexStore>,
44    inner: RwLock<SessionInner>,
45}
46
47fn filter_from_profile(profile: &ProfileConfig) -> PolicyFilter {
48    PolicyFilter {
49        allow_schemas: profile.schemas.clone(),
50        deny_schemas: profile.deny_schemas.clone(),
51        deny_tables: profile.deny_tables.clone(),
52        pii_columns: profile.pii_columns.clone(),
53    }
54}
55
56struct SessionInner {
57    active_id: String,
58    database: String,
59    pools: HashMap<String, Pool>,
60}
61
62fn pool_key(connection_id: &str, database: &str) -> String {
63    format!("{connection_id}\0{database}")
64}
65
66fn params_for_database(base: &ConnectionParams, database: &str) -> ConnectionParams {
67    let mut params = base.clone();
68    params.dbname = Some(database.to_string());
69    // `to_url()` prefers `url` over `dbname`; drop stale URL path when host fields exist.
70    if params.host.is_some() {
71        params.url = None;
72    }
73    params
74}
75
76impl ToolSession {
77    pub async fn from_resolved(
78        resolved: ResolvedConnection,
79        access_mode: AccessMode,
80        caps: PolicyCaps,
81    ) -> Result<Arc<Self>, ToolError> {
82        let id = resolved
83            .profile_name
84            .clone()
85            .unwrap_or_else(|| "default".into());
86        let filter = resolved
87            .profile
88            .as_ref()
89            .map(filter_from_profile)
90            .unwrap_or_default();
91        let info = ConnectionInfo {
92            id: id.clone(),
93            name: id.clone(),
94            host: resolved.params.host.clone(),
95            port: resolved.params.port,
96            database: resolved.params.dbname.clone(),
97            params: resolved.params,
98        };
99        Self::from_connections_with_filter(vec![info], access_mode, caps, Some(id), filter).await
100    }
101
102    pub async fn from_connections(
103        connections: Vec<ConnectionInfo>,
104        access_mode: AccessMode,
105        caps: PolicyCaps,
106        active_id: Option<String>,
107    ) -> Result<Arc<Self>, ToolError> {
108        Self::from_connections_with_filter(
109            connections,
110            access_mode,
111            caps,
112            active_id,
113            PolicyFilter::default(),
114        )
115        .await
116    }
117
118    pub async fn from_connections_with_filter(
119        connections: Vec<ConnectionInfo>,
120        access_mode: AccessMode,
121        caps: PolicyCaps,
122        active_id: Option<String>,
123        filter: PolicyFilter,
124    ) -> Result<Arc<Self>, ToolError> {
125        if connections.is_empty() {
126            return Err(ToolError::Execution("no connections configured".into()));
127        }
128        let pool_opts = PoolOptions {
129            read_only: !access_mode.allows_writes(),
130            ..Default::default()
131        };
132        let active = active_id.unwrap_or_else(|| connections[0].id.clone());
133        let mut pools = HashMap::new();
134        for c in &connections {
135            let database = c.database.as_deref().unwrap_or("postgres");
136            let pool = create_pool(&c.params, &pool_opts).await?;
137            pools.insert(pool_key(&c.id, database), pool);
138        }
139        let database = connections
140            .iter()
141            .find(|c| c.id == active)
142            .and_then(|c| c.database.clone())
143            .unwrap_or_else(|| "postgres".into());
144        Ok(Arc::new(Self {
145            connections,
146            access_mode,
147            caps,
148            filter,
149            pool_opts,
150            index_store: Some(IndexStore::new(default_index_root())),
151            inner: RwLock::new(SessionInner {
152                active_id: active,
153                database,
154                pools,
155            }),
156        }))
157    }
158
159    pub async fn active_context(&self) -> (String, String) {
160        let g = self.inner.read().await;
161        (g.active_id.clone(), g.database.clone())
162    }
163
164    pub async fn switch(
165        &self,
166        connection_id: &str,
167        database: Option<String>,
168    ) -> Result<(), ToolError> {
169        let conn = self
170            .connections
171            .iter()
172            .find(|c| c.id == connection_id)
173            .ok_or_else(|| {
174                ToolError::Execution(format!(
175                    "Connection not found for ID: {connection_id} — call list_connections"
176                ))
177            })?;
178        let target_db = database
179            .or_else(|| conn.database.clone())
180            .unwrap_or_else(|| "postgres".into());
181        let key = pool_key(connection_id, &target_db);
182        let mut g = self.inner.write().await;
183        if let std::collections::hash_map::Entry::Vacant(e) = g.pools.entry(key) {
184            let params = params_for_database(&conn.params, &target_db);
185            let pool = create_pool(&params, &self.pool_opts).await?;
186            e.insert(pool);
187        }
188        g.active_id = connection_id.to_string();
189        g.database = target_db;
190        Ok(())
191    }
192
193    pub async fn checkout(&self) -> Result<Object, ToolError> {
194        let g = self.inner.read().await;
195        let key = pool_key(&g.active_id, &g.database);
196        let pool = g
197            .pools
198            .get(&key)
199            .ok_or_else(|| ToolError::Execution("no active pool".into()))?;
200        Ok(checkout_guarded(pool, &self.pool_opts).await?)
201    }
202
203    /// Test helper: session with no live pools (index tools only).
204    #[cfg(test)]
205    pub fn for_tests(
206        connections: Vec<ConnectionInfo>,
207        filter: PolicyFilter,
208        index_store: Option<IndexStore>,
209    ) -> Arc<Self> {
210        assert!(!connections.is_empty());
211        let active = connections[0].id.clone();
212        let database = connections[0]
213            .database
214            .clone()
215            .unwrap_or_else(|| "postgres".into());
216        Arc::new(Self {
217            connections,
218            access_mode: AccessMode::Read,
219            caps: PolicyCaps::default(),
220            filter,
221            pool_opts: PoolOptions::default(),
222            index_store,
223            inner: RwLock::new(SessionInner {
224                active_id: active,
225                database,
226                pools: HashMap::new(),
227            }),
228        })
229    }
230}
231
232#[cfg(test)]
233mod tests {
234    use super::*;
235
236    #[test]
237    fn filter_maps_profile_fields() {
238        let profile = ProfileConfig {
239            schemas: vec!["public".into()],
240            deny_schemas: vec!["pgboss".into()],
241            deny_tables: vec!["auth.*".into()],
242            pii_columns: vec!["public.users.ssn".into()],
243            ..Default::default()
244        };
245        let f = filter_from_profile(&profile);
246        assert_eq!(f.allow_schemas, vec!["public"]);
247        assert_eq!(f.deny_schemas, vec!["pgboss"]);
248        assert_eq!(f.deny_tables, vec!["auth.*"]);
249        assert_eq!(f.pii_columns, vec!["public.users.ssn"]);
250        assert!(f.allows_schema("public"));
251        assert!(!f.allows_schema("pgboss"));
252        assert!(!f.allows_table("auth", "sessions"));
253    }
254}