Skip to main content

dbx_tools_databricks_auth/storage/
keyring.rs

1use std::{path::PathBuf, time::Duration};
2
3use async_trait::async_trait;
4use base64::{engine::general_purpose::STANDARD, Engine as _};
5use keyring::Entry;
6
7use super::{CredentialStore, FileStore, StorageLock};
8use crate::{Error, Result, Token};
9
10const SERVICE: &str = "databricks-cli";
11const PROBE_ACCOUNT_PREFIX: &str = "__probe_";
12const GO_KEYRING_BASE64_PREFIX: &str = "go-keyring-base64:";
13const GO_KEYRING_HEX_PREFIX: &str = "go-keyring-encoded:";
14
15#[derive(serde::Serialize, serde::Deserialize)]
16struct KeyringEntry {
17    token: Token,
18}
19
20pub struct KeyringStore {
21    lock_store: FileStore,
22}
23
24impl KeyringStore {
25    fn new(cache_dir: PathBuf) -> Result<Self> {
26        Ok(Self {
27            lock_store: FileStore::new(cache_dir.join("locks"))?,
28        })
29    }
30
31    pub async fn open_for_read(cache_dir: PathBuf) -> Result<Self> {
32        let store = Self::new(cache_dir)?;
33        let account = format!("{PROBE_ACCOUNT_PREFIX}{}", uuid::Uuid::new_v4());
34        tokio::task::spawn_blocking(move || match Self::entry(&account)?.get_password() {
35            Ok(_) | Err(keyring::Error::NoEntry) => Ok(()),
36            Err(error) => Err(keyring_error(error)),
37        })
38        .await
39        .map_err(|error| Error::Storage(format!("keyring read probe task failed: {error}")))??;
40        Ok(store)
41    }
42
43    async fn prepare_write_probe(&self) -> Result<()> {
44        let account = format!("{PROBE_ACCOUNT_PREFIX}{}", uuid::Uuid::new_v4());
45        tokio::task::spawn_blocking(move || {
46            let entry = Entry::new(SERVICE, &account).map_err(keyring_error)?;
47            entry.set_password("probe").map_err(keyring_error)?;
48            entry.delete_credential().map_err(keyring_error)?;
49            Ok::<_, Error>(())
50        })
51        .await
52        .map_err(|error| Error::Storage(format!("keyring probe task failed: {error}")))??;
53        Ok(())
54    }
55
56    fn entry(profile: &str) -> Result<Entry> {
57        Entry::new(SERVICE, profile).map_err(keyring_error)
58    }
59}
60
61fn decode_keyring_value(raw: &str) -> Result<String> {
62    if let Some(value) = raw.strip_prefix(GO_KEYRING_BASE64_PREFIX) {
63        return STANDARD
64            .decode(value)
65            .map_err(|error| Error::Storage(format!("invalid Databricks keyring base64: {error}")))
66            .and_then(|value| {
67                String::from_utf8(value).map_err(|error| {
68                    Error::Storage(format!("invalid Databricks keyring UTF-8: {error}"))
69                })
70            });
71    }
72    if let Some(value) = raw.strip_prefix(GO_KEYRING_HEX_PREFIX) {
73        if value.len() % 2 != 0 {
74            return Err(Error::Storage(
75                "invalid Databricks keyring hex: odd length".into(),
76            ));
77        }
78        let bytes = (0..value.len())
79            .step_by(2)
80            .map(|index| {
81                u8::from_str_radix(&value[index..index + 2], 16).map_err(|error| {
82                    Error::Storage(format!("invalid Databricks keyring hex: {error}"))
83                })
84            })
85            .collect::<Result<Vec<_>>>()?;
86        return String::from_utf8(bytes)
87            .map_err(|error| Error::Storage(format!("invalid Databricks keyring UTF-8: {error}")));
88    }
89    Ok(raw.to_owned())
90}
91
92#[async_trait]
93impl CredentialStore for KeyringStore {
94    async fn load(&self, profile: &str) -> Result<Option<Token>> {
95        let profile = profile.to_owned();
96        tokio::task::spawn_blocking(move || match Self::entry(&profile)?.get_password() {
97            Ok(raw) => {
98                let raw = decode_keyring_value(&raw)?;
99                Ok(Some(serde_json::from_str::<KeyringEntry>(&raw)?.token))
100            }
101            Err(keyring::Error::NoEntry) => Ok(None),
102            Err(error) => Err(keyring_error(error)),
103        })
104        .await
105        .map_err(|error| Error::Storage(format!("keyring read task failed: {error}")))?
106    }
107
108    async fn prepare_write(&self) -> Result<()> {
109        self.prepare_write_probe().await
110    }
111
112    async fn save(&self, profile: &str, token: &Token) -> Result<()> {
113        let profile = profile.to_owned();
114        let raw = serde_json::to_string(&KeyringEntry {
115            token: token.clone(),
116        })?;
117        tokio::task::spawn_blocking(move || {
118            Self::entry(&profile)?
119                .set_password(&raw)
120                .map_err(keyring_error)
121        })
122        .await
123        .map_err(|error| Error::Storage(format!("keyring write task failed: {error}")))?
124    }
125
126    async fn delete(&self, profile: &str) -> Result<()> {
127        let profile = profile.to_owned();
128        tokio::task::spawn_blocking(move || match Self::entry(&profile)?.delete_credential() {
129            Ok(()) | Err(keyring::Error::NoEntry) => Ok(()),
130            Err(error) => Err(keyring_error(error)),
131        })
132        .await
133        .map_err(|error| Error::Storage(format!("keyring delete task failed: {error}")))?
134    }
135
136    async fn lock(&self, profile: &str, timeout: Duration) -> Result<Box<dyn StorageLock>> {
137        Ok(Box::new(
138            self.lock_store.acquire_file_lock(profile, timeout).await?,
139        ))
140    }
141
142    fn name(&self) -> &'static str {
143        "keyring"
144    }
145}
146
147fn keyring_error(error: keyring::Error) -> Error {
148    Error::Storage(format!("OS keyring: {error}"))
149}
150
151#[cfg(test)]
152mod tests {
153    use super::decode_keyring_value;
154    use base64::{engine::general_purpose::STANDARD, Engine as _};
155
156    #[test]
157    fn decodes_go_keyring_base64_values() {
158        let json = r#"{"token":{"access_token":"value"}}"#;
159        let encoded = format!("go-keyring-base64:{}", STANDARD.encode(json));
160        assert_eq!(decode_keyring_value(&encoded).unwrap(), json);
161    }
162
163    #[test]
164    fn decodes_go_keyring_hex_values() {
165        assert_eq!(
166            decode_keyring_value("go-keyring-encoded:7b22746f6b656e223a7b7d7d").unwrap(),
167            r#"{"token":{}}"#
168        );
169    }
170
171    #[test]
172    fn leaves_native_keyring_values_unchanged() {
173        let json = r#"{"token":{"access_token":"value"}}"#;
174        assert_eq!(decode_keyring_value(json).unwrap(), json);
175    }
176}