use std::collections::HashMap;
use schemars::JsonSchema;
use serde::{Deserialize, Serialize};
#[derive(Debug, Clone, PartialEq, Eq, Default, Serialize, Deserialize, JsonSchema)]
#[serde(tag = "type", content = "config", rename_all = "snake_case")]
pub enum DeltaCredentials {
#[default]
Default,
Aws(AwsCredentials),
Azure(AzureCredentials),
Gcp(GcpCredentials),
}
impl DeltaCredentials {
pub fn apply(&self, options: &mut HashMap<String, String>) {
let derived = self.storage_options();
for (k, v) in derived {
options.entry(k).or_insert(v);
}
}
pub fn storage_options(&self) -> HashMap<String, String> {
let mut m = HashMap::new();
match self {
DeltaCredentials::Default => {}
DeltaCredentials::Aws(a) => {
if let Some(v) = &a.access_key_id {
m.insert("AWS_ACCESS_KEY_ID".into(), v.clone());
}
if let Some(v) = &a.secret_access_key {
m.insert("AWS_SECRET_ACCESS_KEY".into(), v.clone());
}
if let Some(v) = &a.session_token {
m.insert("AWS_SESSION_TOKEN".into(), v.clone());
}
if let Some(v) = &a.region {
m.insert("AWS_REGION".into(), v.clone());
}
if let Some(v) = &a.endpoint_url {
m.insert("AWS_ENDPOINT_URL".into(), v.clone());
}
if let Some(v) = a.allow_http {
m.insert("AWS_ALLOW_HTTP".into(), v.to_string());
}
}
DeltaCredentials::Azure(a) => {
if let Some(v) = &a.account_name {
m.insert("AZURE_STORAGE_ACCOUNT_NAME".into(), v.clone());
}
if let Some(v) = &a.access_key {
m.insert("AZURE_STORAGE_ACCESS_KEY".into(), v.clone());
}
if let Some(v) = &a.sas_token {
m.insert("AZURE_STORAGE_SAS_KEY".into(), v.clone());
}
}
DeltaCredentials::Gcp(g) => {
if let Some(v) = &g.service_account_path {
m.insert("GOOGLE_SERVICE_ACCOUNT".into(), v.clone());
}
if let Some(v) = &g.service_account_key {
m.insert("GOOGLE_SERVICE_ACCOUNT_KEY".into(), v.clone());
}
}
}
m
}
}
#[derive(Debug, Clone, PartialEq, Eq, Default, Serialize, Deserialize, JsonSchema)]
pub struct AwsCredentials {
#[serde(default, skip_serializing_if = "Option::is_none")]
pub access_key_id: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub secret_access_key: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub session_token: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub region: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub endpoint_url: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub allow_http: Option<bool>,
}
#[derive(Debug, Clone, PartialEq, Eq, Default, Serialize, Deserialize, JsonSchema)]
pub struct AzureCredentials {
#[serde(default, skip_serializing_if = "Option::is_none")]
pub account_name: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub access_key: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub sas_token: Option<String>,
}
#[derive(Debug, Clone, PartialEq, Eq, Default, Serialize, Deserialize, JsonSchema)]
pub struct GcpCredentials {
#[serde(default, skip_serializing_if = "Option::is_none")]
pub service_account_path: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub service_account_key: Option<String>,
}
#[cfg(test)]
mod tests {
use super::*;
use serde_json::json;
#[test]
fn default_injects_nothing() {
let c = DeltaCredentials::default();
assert!(c.storage_options().is_empty());
}
#[test]
fn aws_maps_all_fields() {
let c = DeltaCredentials::Aws(AwsCredentials {
access_key_id: Some("AK".into()),
secret_access_key: Some("SK".into()),
session_token: Some("ST".into()),
region: Some("us-east-1".into()),
endpoint_url: Some("http://localhost:9000".into()),
allow_http: Some(true),
});
let m = c.storage_options();
assert_eq!(m["AWS_ACCESS_KEY_ID"], "AK");
assert_eq!(m["AWS_SECRET_ACCESS_KEY"], "SK");
assert_eq!(m["AWS_SESSION_TOKEN"], "ST");
assert_eq!(m["AWS_REGION"], "us-east-1");
assert_eq!(m["AWS_ENDPOINT_URL"], "http://localhost:9000");
assert_eq!(m["AWS_ALLOW_HTTP"], "true");
}
#[test]
fn azure_and_gcp_partial() {
let az = DeltaCredentials::Azure(AzureCredentials {
account_name: Some("acct".into()),
access_key: None,
sas_token: Some("sas".into()),
});
let m = az.storage_options();
assert_eq!(m["AZURE_STORAGE_ACCOUNT_NAME"], "acct");
assert_eq!(m["AZURE_STORAGE_SAS_KEY"], "sas");
assert!(!m.contains_key("AZURE_STORAGE_ACCESS_KEY"));
let gcp = DeltaCredentials::Gcp(GcpCredentials {
service_account_path: Some("/keys/sa.json".into()),
service_account_key: None,
});
assert_eq!(
gcp.storage_options()["GOOGLE_SERVICE_ACCOUNT"],
"/keys/sa.json"
);
let az_key = DeltaCredentials::Azure(AzureCredentials {
account_name: None,
access_key: Some("shared-key".into()),
sas_token: None,
});
assert_eq!(
az_key.storage_options()["AZURE_STORAGE_ACCESS_KEY"],
"shared-key"
);
let gcp_key = DeltaCredentials::Gcp(GcpCredentials {
service_account_path: None,
service_account_key: Some("{\"type\":\"service_account\"}".into()),
});
assert_eq!(
gcp_key.storage_options()["GOOGLE_SERVICE_ACCOUNT_KEY"],
"{\"type\":\"service_account\"}"
);
}
#[test]
fn explicit_options_win_over_credentials() {
let c = DeltaCredentials::Aws(AwsCredentials {
region: Some("us-east-1".into()),
..Default::default()
});
let mut opts = HashMap::new();
opts.insert("AWS_REGION".to_string(), "eu-west-1".to_string());
c.apply(&mut opts);
assert_eq!(opts["AWS_REGION"], "eu-west-1");
}
#[test]
fn tagged_shape_round_trips() {
let v = json!({ "type": "aws", "config": { "region": "us-east-1" } });
let c: DeltaCredentials = serde_json::from_value(v).unwrap();
assert_eq!(
c,
DeltaCredentials::Aws(AwsCredentials {
region: Some("us-east-1".into()),
..Default::default()
})
);
let d: DeltaCredentials = serde_json::from_value(json!({ "type": "default" })).unwrap();
assert_eq!(d, DeltaCredentials::Default);
}
}