Skip to main content

omni_dev/daemon/services/
snowflake.rs

1//! The Snowflake daemon service.
2//!
3//! A thin adapter that hosts the account-agnostic [`SnowflakeEngine`] under the
4//! daemon's lifecycle and exposes query/sessions/disconnect over the control
5//! socket, plus a tray submenu.
6//!
7//! All real work (lazy multiplexed auth, per-query `USE …`, heartbeats, the
8//! arbitrary-schema → JSON mapping) lives in [`crate::snowflake`]; this adapter
9//! only routes ops and renders the menu/status. Unlike the bridge it persists no
10//! secret to disk — sessions live only in memory.
11
12use anyhow::{anyhow, bail, Context, Result};
13use async_trait::async_trait;
14use chrono::Utc;
15use serde_json::{json, Value};
16
17use crate::daemon::service::{DaemonService, MenuAction, MenuItem, MenuSnapshot, ServiceStatus};
18use crate::snowflake::session::SessionInfo;
19use crate::snowflake::{QueryRequest, SnowflakeEngine, SnowflakeEngineConfig};
20
21/// The Snowflake service name (the control-socket routing key).
22pub const SERVICE_NAME: &str = "snowflake";
23
24/// Hosts a [`SnowflakeEngine`] as a [`DaemonService`].
25pub struct SnowflakeService {
26    engine: SnowflakeEngine,
27}
28
29impl SnowflakeService {
30    /// Creates the service and starts the engine's background keep-alive
31    /// heartbeat. Cheap — no eager auth or I/O; each `(account, user)` session
32    /// is authenticated lazily on its first query, and the heartbeat only
33    /// touches sessions that exist.
34    #[must_use]
35    pub fn new(config: SnowflakeEngineConfig) -> Self {
36        let engine = SnowflakeEngine::new(config);
37        engine.start_heartbeat();
38        Self { engine }
39    }
40
41    /// Handles the `disconnect` op's three mutually-exclusive selectors — every
42    /// pool (`{ all: true }`), a pool by numeric id (`{ id }`), or the
43    /// `(account, user)` pair (`{ account, user }`, the original contract). The
44    /// bulk and by-id forms were previously reachable only from the macOS tray
45    /// (#1228). Every reply carries `count` (pools evicted) plus `disconnected`
46    /// (`count > 0`), so the by-id and pair forms stay wire-compatible with the
47    /// pre-#1228 `{ disconnected }` reply.
48    fn handle_disconnect(&self, payload: &Value) -> Result<Value> {
49        if payload.get("all").and_then(Value::as_bool) == Some(true) {
50            let count = self.engine.disconnect_all();
51            return Ok(json!({ "disconnected": count > 0, "count": count }));
52        }
53        if let Some(id) = payload.get("id").and_then(Value::as_u64) {
54            let removed = self.engine.disconnect_by_id(id);
55            return Ok(json!({ "disconnected": removed, "count": u64::from(removed) }));
56        }
57        let account = payload
58            .get("account")
59            .and_then(Value::as_str)
60            .ok_or_else(|| anyhow!("`disconnect` requires `account`, `id`, or `all`"))?;
61        let user = payload
62            .get("user")
63            .and_then(Value::as_str)
64            .ok_or_else(|| anyhow!("`disconnect` requires `user`"))?;
65        let removed = self.engine.disconnect(account, user);
66        Ok(json!({ "disconnected": removed, "count": u64::from(removed) }))
67    }
68
69    /// Handles the `cancel` op: abort the running query on a target pool without
70    /// evicting the session. Selectors mirror `disconnect` — every pool
71    /// (`{ all: true }`), a pool by id (`{ id }`), or the `(account, user)` pair —
72    /// plus an optional `member` (a numeric member id from `sessions`) to target
73    /// one authenticated session's query rather than every busy one in the pool
74    /// (ignored with `all`). The reply's `cancelled` is how many running
75    /// statements an abort was issued for.
76    async fn handle_cancel(&self, payload: &Value) -> Result<Value> {
77        let member = payload.get("member").and_then(Value::as_u64);
78        let cancelled = if payload.get("all").and_then(Value::as_bool) == Some(true) {
79            self.engine.cancel_all().await
80        } else if let Some(id) = payload.get("id").and_then(Value::as_u64) {
81            self.engine.cancel_by_id(id, member).await
82        } else {
83            let account = payload
84                .get("account")
85                .and_then(Value::as_str)
86                .ok_or_else(|| anyhow!("`cancel` requires `account`, `id`, or `all`"))?;
87            let user = payload
88                .get("user")
89                .and_then(Value::as_str)
90                .ok_or_else(|| anyhow!("`cancel` requires `user`"))?;
91            self.engine.cancel(account, user, member).await
92        };
93        Ok(json!({ "cancelled": cancelled }))
94    }
95}
96
97#[async_trait]
98impl DaemonService for SnowflakeService {
99    fn name(&self) -> &'static str {
100        SERVICE_NAME
101    }
102
103    async fn handle(&self, op: &str, payload: Value) -> Result<Value> {
104        match op {
105            "query" => {
106                let req: QueryRequest =
107                    serde_json::from_value(payload).context("invalid `query` payload")?;
108                if req.sql.trim().is_empty() {
109                    bail!("`query` requires a non-empty `sql`");
110                }
111                self.engine.query(req).await
112            }
113            "sessions" => Ok(json!({ "sessions": self.engine.sessions() })),
114            "cancel" => self.handle_cancel(&payload).await,
115            "disconnect" => self.handle_disconnect(&payload),
116            other => bail!("unknown snowflake op: {other}"),
117        }
118    }
119
120    fn menu(&self) -> MenuSnapshot {
121        let sessions = self.engine.sessions();
122        let items = if sessions.is_empty() {
123            vec![MenuItem::Label("No sessions".to_string())]
124        } else {
125            session_menu_items(&sessions)
126        };
127        MenuSnapshot {
128            title: "Snowflake".to_string(),
129            items,
130        }
131    }
132
133    async fn menu_action(&self, action_id: &str) -> Result<()> {
134        if action_id == "disconnect-all" {
135            self.engine.disconnect_all();
136            return Ok(());
137        }
138        if let Some(id) = action_id.strip_prefix("disconnect:") {
139            let id: u64 = id
140                .parse()
141                .with_context(|| format!("invalid session id in action {action_id}"))?;
142            self.engine.disconnect_by_id(id);
143            return Ok(());
144        }
145        bail!("unknown snowflake menu action: {action_id}")
146    }
147
148    async fn status(&self) -> ServiceStatus {
149        let sessions = self.engine.sessions();
150        let live: usize = sessions.iter().map(|s| s.sessions).sum();
151        ServiceStatus {
152            name: SERVICE_NAME.to_string(),
153            healthy: true,
154            summary: format!("{} pool(s), {live} session(s)", sessions.len()),
155            detail: json!({ "sessions": sessions }),
156        }
157    }
158
159    async fn shutdown(&self) {
160        self.engine.shutdown().await;
161    }
162}
163
164/// Builds the tray items for a non-empty session list: a label per pool, an
165/// indented label per authenticated session (with what it's doing), a separator,
166/// and the per-pool + "Disconnect all" actions.
167fn session_menu_items(sessions: &[SessionInfo]) -> Vec<MenuItem> {
168    let mut items = Vec::new();
169    for session in sessions {
170        items.push(MenuItem::Label(format!(
171            "{} · {} · {}/{} sessions · {} queries",
172            session.account,
173            session.user,
174            session.sessions,
175            session.max_sessions,
176            session.query_count
177        )));
178        // One line per individual authenticated session (auth), with what it's
179        // doing: the running query + elapsed when busy, else idle time.
180        for member in &session.members {
181            let state = if let Some(running) = &member.running {
182                let secs = (Utc::now() - running.started_at).num_seconds().max(0);
183                format!("running {secs}s: {}", running.sql)
184            } else if member.busy {
185                "busy".to_string()
186            } else {
187                let idle = (Utc::now() - member.last_used).num_seconds().max(0);
188                format!("idle {idle}s · {} queries", member.query_count)
189            };
190            items.push(MenuItem::Label(format!(
191                "    #{} {} · {state}",
192                member.id,
193                member.context.summary(),
194            )));
195        }
196    }
197    items.push(MenuItem::Separator);
198    for session in sessions {
199        items.push(MenuItem::Action(MenuAction {
200            id: format!("disconnect:{}", session.id),
201            label: format!("Disconnect {} · {}", session.account, session.user),
202            enabled: true,
203        }));
204    }
205    items.push(MenuItem::Action(MenuAction {
206        id: "disconnect-all".to_string(),
207        label: "Disconnect all".to_string(),
208        enabled: true,
209    }));
210    items
211}
212
213#[cfg(test)]
214#[allow(clippy::unwrap_used, clippy::expect_used)]
215mod tests {
216    use super::*;
217
218    /// A service with no resolvable defaults, so query resolution fails before
219    /// any network/auth — keeping these tests offline and deterministic.
220    fn offline_service() -> SnowflakeService {
221        SnowflakeService::new(SnowflakeEngineConfig::default())
222    }
223
224    #[tokio::test]
225    async fn name_and_unknown_op() {
226        let svc = offline_service();
227        assert_eq!(svc.name(), "snowflake");
228        assert!(svc.handle("frobnicate", Value::Null).await.is_err());
229    }
230
231    #[tokio::test]
232    async fn sessions_op_is_empty_initially() {
233        let svc = offline_service();
234        let payload = svc.handle("sessions", Value::Null).await.unwrap();
235        assert_eq!(payload, json!({ "sessions": [] }));
236    }
237
238    #[tokio::test]
239    async fn empty_sql_is_rejected_before_auth() {
240        let svc = offline_service();
241        assert!(svc.handle("query", json!({ "sql": "   " })).await.is_err());
242    }
243
244    #[tokio::test]
245    async fn query_without_account_errors_not_panics() {
246        let svc = offline_service();
247        // Non-empty SQL but no resolvable account: errors on resolution, no auth.
248        let err = svc
249            .handle("query", json!({ "sql": "SELECT 1" }))
250            .await
251            .unwrap_err();
252        assert!(err.to_string().contains("account"));
253    }
254
255    #[tokio::test]
256    async fn disconnect_requires_account_and_user() {
257        let svc = offline_service();
258        assert!(svc.handle("disconnect", json!({})).await.is_err());
259        assert!(svc
260            .handle("disconnect", json!({ "account": "ACCT" }))
261            .await
262            .is_err());
263        // Both present: evicts nothing on an empty engine, but succeeds.
264        let payload = svc
265            .handle("disconnect", json!({ "account": "ACCT", "user": "me" }))
266            .await
267            .unwrap();
268        assert_eq!(payload, json!({ "disconnected": false, "count": 0 }));
269    }
270
271    #[tokio::test]
272    async fn disconnect_by_id_and_all_selectors() {
273        let svc = offline_service();
274        // Nothing to evict on an empty engine, but both bulk forms succeed and
275        // report a zero count.
276        let by_id = svc.handle("disconnect", json!({ "id": 7 })).await.unwrap();
277        assert_eq!(by_id, json!({ "disconnected": false, "count": 0 }));
278        let all = svc
279            .handle("disconnect", json!({ "all": true }))
280            .await
281            .unwrap();
282        assert_eq!(all, json!({ "disconnected": false, "count": 0 }));
283    }
284
285    #[tokio::test]
286    async fn cancel_requires_a_selector_and_is_zero_on_an_empty_engine() {
287        let svc = offline_service();
288        // No selector at all is an error.
289        assert!(svc.handle("cancel", json!({})).await.is_err());
290        // `account` without `user` is an error.
291        assert!(svc
292            .handle("cancel", json!({ "account": "ACCT" }))
293            .await
294            .is_err());
295        // Each valid selector succeeds and cancels nothing on an empty engine.
296        for payload in [
297            json!({ "account": "ACCT", "user": "me" }),
298            json!({ "account": "ACCT", "user": "me", "member": 2 }),
299            json!({ "id": 7 }),
300            json!({ "all": true }),
301        ] {
302            let reply = svc.handle("cancel", payload.clone()).await.unwrap();
303            assert_eq!(reply, json!({ "cancelled": 0 }), "for {payload}");
304        }
305        // Cancel must not create a pool.
306        let sessions = svc.handle("sessions", Value::Null).await.unwrap();
307        assert_eq!(sessions, json!({ "sessions": [] }));
308    }
309
310    #[tokio::test]
311    async fn menu_and_status_shape_with_no_sessions() {
312        let svc = offline_service();
313        let menu = svc.menu();
314        assert_eq!(menu.title, "Snowflake");
315        assert!(matches!(
316            menu.items.first(),
317            Some(MenuItem::Label(text)) if text == "No sessions"
318        ));
319        let status = svc.status().await;
320        assert_eq!(status.name, "snowflake");
321        assert!(status.healthy);
322        assert_eq!(status.summary, "0 pool(s), 0 session(s)");
323    }
324
325    #[tokio::test]
326    async fn menu_actions_route_and_reject_unknown() {
327        let svc = offline_service();
328        // Both forms are no-ops on an empty engine but must not error.
329        svc.menu_action("disconnect-all").await.unwrap();
330        svc.menu_action("disconnect:7").await.unwrap();
331        assert!(svc.menu_action("disconnect:not-a-number").await.is_err());
332        assert!(svc.menu_action("bogus").await.is_err());
333        svc.shutdown().await;
334    }
335
336    #[test]
337    fn session_menu_items_render_each_member_state_and_actions() {
338        use crate::snowflake::session::{MemberInfo, QueryContext, RunningQuery};
339
340        let now = Utc::now();
341        let wh_ctx = QueryContext {
342            warehouse: Some("WH".to_string()),
343            role: Some("R".to_string()),
344            ..QueryContext::default()
345        };
346        let sessions = vec![SessionInfo {
347            id: 5,
348            account: "ACME".to_string(),
349            user: "me".to_string(),
350            created_at: now,
351            last_used: now,
352            query_count: 9,
353            sessions: 3,
354            max_sessions: 4,
355            members: vec![
356                MemberInfo {
357                    id: 1,
358                    busy: true,
359                    context: wh_ctx,
360                    last_used: now,
361                    query_count: 3,
362                    running: Some(RunningQuery {
363                        sql: "SELECT 42".to_string(),
364                        started_at: now,
365                    }),
366                },
367                MemberInfo {
368                    id: 2,
369                    busy: true,
370                    context: QueryContext::default(),
371                    last_used: now,
372                    query_count: 1,
373                    running: None,
374                },
375                MemberInfo {
376                    id: 3,
377                    busy: false,
378                    context: QueryContext::default(),
379                    last_used: now,
380                    query_count: 0,
381                    running: None,
382                },
383            ],
384        }];
385
386        let items = session_menu_items(&sessions);
387        let labels: Vec<&str> = items
388            .iter()
389            .filter_map(|i| match i {
390                MenuItem::Label(t) => Some(t.as_str()),
391                _ => None,
392            })
393            .collect();
394        assert!(labels
395            .iter()
396            .any(|l| l.contains("ACME · me · 3/4 sessions · 9 queries")));
397        assert!(labels
398            .iter()
399            .any(|l| l.contains("running") && l.contains("SELECT 42") && l.contains("WH/R")));
400        assert!(labels.iter().any(|l| l.contains("busy")));
401        assert!(labels
402            .iter()
403            .any(|l| l.contains("idle") && l.contains("(default)")));
404
405        assert!(items.iter().any(|i| matches!(i, MenuItem::Separator)));
406        let action_ids: Vec<&str> = items
407            .iter()
408            .filter_map(|i| match i {
409                MenuItem::Action(a) => Some(a.id.as_str()),
410                _ => None,
411            })
412            .collect();
413        assert!(action_ids.contains(&"disconnect:5"));
414        assert!(action_ids.contains(&"disconnect-all"));
415    }
416}