use chrono::{DateTime, Utc};
use serde::{Deserialize, Serialize};
use uuid::Uuid;
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
pub enum Engine {
Postgres,
MySql,
Sqlite,
}
impl Engine {
pub const ALL: [Engine; 3] = [Engine::Postgres, Engine::MySql, Engine::Sqlite];
pub fn label(self) -> &'static str {
match self {
Engine::Postgres => "PostgreSQL",
Engine::MySql => "MySQL",
Engine::Sqlite => "SQLite",
}
}
pub fn default_port(self) -> u16 {
match self {
Engine::Postgres => 5432,
Engine::MySql => 3306,
Engine::Sqlite => 0,
}
}
pub fn is_file_based(self) -> bool {
matches!(self, Engine::Sqlite)
}
}
#[derive(Debug, Clone, Copy, Default, PartialEq, Eq, Serialize, Deserialize)]
pub enum SafetyMode {
ReadOnly,
ConfirmWrites,
#[default]
Staged,
AutoApply,
}
impl SafetyMode {
pub const ALL: [SafetyMode; 4] = [
SafetyMode::ReadOnly,
SafetyMode::ConfirmWrites,
SafetyMode::Staged,
SafetyMode::AutoApply,
];
pub fn label(self) -> &'static str {
match self {
SafetyMode::ReadOnly => "Read-only",
SafetyMode::ConfirmWrites => "Confirm writes",
SafetyMode::Staged => "Staged edits",
SafetyMode::AutoApply => "Auto-apply",
}
}
pub fn description(self) -> &'static str {
match self {
SafetyMode::ReadOnly => "Refuses anything that writes",
SafetyMode::ConfirmWrites => "Shows every write before it runs",
SafetyMode::Staged => "Edits wait until you apply them",
SafetyMode::AutoApply => "Edits are written when you leave the row",
}
}
pub fn is_read_only(self) -> bool {
matches!(self, SafetyMode::ReadOnly)
}
pub fn confirms_writes(self) -> bool {
matches!(self, SafetyMode::ConfirmWrites)
}
pub fn auto_applies(self) -> bool {
matches!(self, SafetyMode::AutoApply)
}
}
#[derive(Debug, Clone, Copy, Default, PartialEq, Eq, Serialize, Deserialize)]
#[serde(rename_all = "kebab-case")]
pub enum SslMode {
Disable,
#[default]
Prefer,
Require,
VerifyCa,
VerifyFull,
}
impl SslMode {
pub const ALL: [SslMode; 5] = [
SslMode::Disable,
SslMode::Prefer,
SslMode::Require,
SslMode::VerifyCa,
SslMode::VerifyFull,
];
pub fn label(self) -> &'static str {
match self {
SslMode::Disable => "Disable",
SslMode::Prefer => "Prefer",
SslMode::Require => "Require",
SslMode::VerifyCa => "Verify CA",
SslMode::VerifyFull => "Verify full",
}
}
pub fn description(self) -> &'static str {
match self {
SslMode::Disable => "Never encrypt the connection",
SslMode::Prefer => "Encrypt when the server offers it",
SslMode::Require => "Always encrypt, without checking the certificate",
SslMode::VerifyCa => "Always encrypt, and check the certificate's authority",
SslMode::VerifyFull => "Always encrypt, and check the authority and the host name",
}
}
pub fn uses_files(self) -> bool {
matches!(
self,
SslMode::Require | SslMode::VerifyCa | SslMode::VerifyFull
)
}
}
#[derive(Debug, Clone, Default, PartialEq, Eq, Serialize, Deserialize)]
#[serde(default)]
pub struct SslConfig {
pub mode: SslMode,
pub ca_cert: String,
pub client_cert: String,
pub client_key: String,
}
impl SslConfig {
pub fn invalid(&self) -> Option<&'static str> {
if !self.mode.uses_files() {
return None;
}
match (
self.client_cert.trim().is_empty(),
self.client_key.trim().is_empty(),
) {
(false, true) => Some("A client certificate needs its private key too."),
(true, false) => Some("A client key needs its certificate too."),
_ => None,
}
}
pub(crate) fn path(value: &str) -> Option<&str> {
Some(value.trim()).filter(|value| !value.is_empty())
}
}
#[derive(Debug, Clone, Copy, Default, PartialEq, Eq, Serialize, Deserialize)]
#[serde(rename_all = "kebab-case")]
pub enum SshAuth {
#[default]
Agent,
PrivateKey,
Password,
}
impl SshAuth {
pub const ALL: [SshAuth; 3] = [SshAuth::Agent, SshAuth::PrivateKey, SshAuth::Password];
pub fn label(self) -> &'static str {
match self {
SshAuth::Agent => "SSH agent",
SshAuth::PrivateKey => "Key file",
SshAuth::Password => "Password",
}
}
pub fn secret_label(self) -> Option<&'static str> {
match self {
SshAuth::Agent => None,
SshAuth::PrivateKey => Some("Key passphrase"),
SshAuth::Password => Some("SSH password"),
}
}
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
#[serde(default)]
pub struct SshConfig {
pub enabled: bool,
pub host: String,
pub port: u16,
pub username: String,
pub auth: SshAuth,
pub key_path: String,
}
impl Default for SshConfig {
fn default() -> Self {
Self {
enabled: false,
host: String::new(),
port: 22,
username: String::new(),
auth: SshAuth::default(),
key_path: String::new(),
}
}
}
impl SshConfig {
pub fn invalid(&self) -> Option<&'static str> {
if !self.enabled {
return None;
}
if self.host.trim().is_empty() {
return Some("Enter the SSH host to tunnel through.");
}
if self.port == 0 {
return Some("The SSH port must be a number from 1 to 65535.");
}
if self.username.trim().is_empty() {
return Some("Enter the SSH user.");
}
if self.auth == SshAuth::PrivateKey && self.key_path.trim().is_empty() {
return Some("Enter the path to the SSH private key.");
}
None
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
#[serde(rename_all = "kebab-case")]
pub enum TagColor {
Red,
Orange,
Yellow,
Green,
Teal,
Blue,
Purple,
Pink,
Gray,
}
impl TagColor {
pub const ALL: [TagColor; 9] = [
TagColor::Red,
TagColor::Orange,
TagColor::Yellow,
TagColor::Green,
TagColor::Teal,
TagColor::Blue,
TagColor::Purple,
TagColor::Pink,
TagColor::Gray,
];
pub fn label(self) -> &'static str {
match self {
TagColor::Red => "Red",
TagColor::Orange => "Orange",
TagColor::Yellow => "Yellow",
TagColor::Green => "Green",
TagColor::Teal => "Teal",
TagColor::Blue => "Blue",
TagColor::Purple => "Purple",
TagColor::Pink => "Pink",
TagColor::Gray => "Gray",
}
}
pub fn key(self) -> &'static str {
match self {
TagColor::Red => "red",
TagColor::Orange => "orange",
TagColor::Yellow => "yellow",
TagColor::Green => "green",
TagColor::Teal => "teal",
TagColor::Blue => "blue",
TagColor::Purple => "purple",
TagColor::Pink => "pink",
TagColor::Gray => "gray",
}
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ConnectionConfig {
pub id: Uuid,
pub name: String,
pub engine: Engine,
pub host: String,
pub port: u16,
pub username: String,
pub database: String,
#[serde(default)]
pub safety: SafetyMode,
#[serde(default)]
pub color: Option<TagColor>,
#[serde(default)]
pub last_connected: Option<DateTime<Utc>>,
#[serde(default)]
pub ssl: SslConfig,
#[serde(default)]
pub ssh: SshConfig,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub statement_timeout: Option<u32>,
}
impl ConnectionConfig {
pub fn new(engine: Engine) -> Self {
Self {
id: Uuid::new_v4(),
name: String::new(),
engine,
host: "localhost".into(),
port: engine.default_port(),
username: String::new(),
database: String::new(),
safety: SafetyMode::default(),
color: None,
last_connected: None,
ssl: SslConfig::default(),
ssh: SshConfig::default(),
statement_timeout: None,
}
}
pub fn display_name(&self) -> String {
if !self.name.trim().is_empty() {
return self.name.clone();
}
if self.engine.is_file_based() {
return file_name(&self.database);
}
self.display_target()
}
pub fn display_target(&self) -> String {
if self.engine.is_file_based() {
return self.database.clone();
}
format!("{}:{}/{}", self.host, self.port, self.database)
}
pub fn is_risky_auto_apply(&self) -> bool {
is_risky_auto_apply(&self.name, self.color, self.safety)
}
}
pub fn is_risky_auto_apply(name: &str, color: Option<TagColor>, safety: SafetyMode) -> bool {
safety.auto_applies() && (color == Some(TagColor::Red) || name.to_lowercase().contains("prod"))
}
impl Default for ConnectionConfig {
fn default() -> Self {
Self::new(Engine::Postgres)
}
}
pub(crate) fn file_name(path: &str) -> String {
std::path::Path::new(path)
.file_name()
.map(|name| name.to_string_lossy().into_owned())
.unwrap_or_else(|| path.to_string())
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn a_client_certificate_needs_its_key() {
let half = |cert: &str, key: &str| SslConfig {
mode: SslMode::Require,
client_cert: cert.into(),
client_key: key.into(),
..SslConfig::default()
};
assert!(half("client.crt", "").invalid().is_some());
assert!(half("", "client.key").invalid().is_some());
assert_eq!(half("client.crt", "client.key").invalid(), None);
assert_eq!(half("", "").invalid(), None);
for mode in [SslMode::Disable, SslMode::Prefer] {
assert_eq!(
SslConfig {
mode,
..half("client.crt", "")
}
.invalid(),
None
);
}
}
#[test]
fn a_connection_saved_before_ssl_settings_reads_back_as_prefer() {
let config: ConnectionConfig = serde_json::from_str(
r#"{"id":"6f1c1f0e-8a5c-4d36-9a49-1b1f3c0f2a10","name":"","engine":"Postgres",
"host":"localhost","port":5432,"username":"","database":""}"#,
)
.unwrap();
assert_eq!(config.ssl, SslConfig::default());
assert_eq!(config.ssl.mode, SslMode::Prefer);
}
#[test]
fn an_ssh_tunnel_needs_a_host_a_user_and_a_key_file_for_key_auth() {
let tunnel = SshConfig {
enabled: true,
host: "bastion".into(),
username: "deploy".into(),
..SshConfig::default()
};
assert_eq!(tunnel.invalid(), None);
assert!(
SshConfig {
host: " ".into(),
..tunnel.clone()
}
.invalid()
.is_some()
);
assert!(
SshConfig {
username: String::new(),
..tunnel.clone()
}
.invalid()
.is_some()
);
assert!(
SshConfig {
auth: SshAuth::PrivateKey,
..tunnel.clone()
}
.invalid()
.is_some()
);
assert_eq!(SshConfig::default().invalid(), None);
}
#[test]
fn a_connection_saved_before_tunnels_reads_back_without_one() {
let config: ConnectionConfig = serde_json::from_str(
r#"{"id":"6f1c1f0e-8a5c-4d36-9a49-1b1f3c0f2a10","name":"","engine":"MySql",
"host":"localhost","port":3306,"username":"","database":""}"#,
)
.unwrap();
assert!(!config.ssh.enabled);
assert_eq!(config.ssh.port, 22);
}
#[test]
fn ssl_settings_round_trip_through_json() {
let config = ConnectionConfig {
ssl: SslConfig {
mode: SslMode::VerifyFull,
ca_cert: "/etc/ssl/ca.pem".into(),
..SslConfig::default()
},
..ConnectionConfig::default()
};
let json = serde_json::to_string(&config).unwrap();
assert!(json.contains(r#""mode":"verify-full""#), "{json}");
let read: ConnectionConfig = serde_json::from_str(&json).unwrap();
assert_eq!(read.ssl, config.ssl);
}
#[test]
fn a_production_connection_with_auto_apply_is_risky() {
assert!(is_risky_auto_apply(
"App",
Some(TagColor::Red),
SafetyMode::AutoApply
));
assert!(is_risky_auto_apply(
"prod-east",
None,
SafetyMode::AutoApply
));
}
#[test]
fn a_production_connection_is_not_risky_under_a_safer_mode() {
for safety in [
SafetyMode::ReadOnly,
SafetyMode::ConfirmWrites,
SafetyMode::Staged,
] {
assert!(!is_risky_auto_apply(
"Production",
Some(TagColor::Red),
safety
));
}
}
#[test]
fn auto_apply_is_not_risky_without_a_production_marker() {
assert!(!is_risky_auto_apply("", None, SafetyMode::AutoApply));
assert!(!is_risky_auto_apply(
"Staging",
Some(TagColor::Orange),
SafetyMode::AutoApply
));
}
}