use std::collections::HashMap;
use std::sync::Arc;
use async_trait::async_trait;
use boatramp_core::compute::{ManagedDbEnvResolver, PrivilegeDirective, ReplicaPhase};
use crate::config::ManagedDbPrivilege;
use boatramp_core::deploy::DeployStore;
use boatramp_core::envelope::KeyEnvelope;
use boatramp_core::kv::KvStore;
use boatramp_core::project::ProjectRef;
use boatramp_core::sql::SqlError;
use boatramp_storage::sql_compute::ComputeEndpointResolver;
use boatramp_storage::ExternalSqlKind;
#[cfg_attr(not(feature = "handlers"), allow(dead_code))]
pub fn managed_db_server_env(
kind: ExternalSqlKind,
database: &str,
user: &str,
password: &str,
) -> Vec<(String, String)> {
match kind {
ExternalSqlKind::Postgres => vec![
("POSTGRES_USER".into(), user.into()),
("POSTGRES_PASSWORD".into(), password.into()),
("POSTGRES_DB".into(), database.into()),
],
ExternalSqlKind::Mysql => vec![
("MYSQL_USER".into(), user.into()),
("MYSQL_PASSWORD".into(), password.into()),
("MYSQL_DATABASE".into(), database.into()),
("MYSQL_ROOT_PASSWORD".into(), password.into()),
],
}
}
#[cfg_attr(not(feature = "handlers"), allow(dead_code))]
pub struct ManagedSqlCredentials {
kv: Arc<dyn KvStore>,
envelope: Arc<dyn KeyEnvelope>,
}
impl ManagedSqlCredentials {
#[cfg_attr(not(feature = "handlers"), allow(dead_code))]
pub fn new(kv: Arc<dyn KvStore>, envelope: Arc<dyn KeyEnvelope>) -> Self {
Self { kv, envelope }
}
fn key(project: &str, workload: &str) -> String {
format!("managed-sql-cred/{project}/{workload}")
}
#[cfg_attr(not(feature = "handlers"), allow(dead_code))]
pub async fn password(&self, project: &str, workload: &str) -> Result<String, String> {
let key = Self::key(project, workload);
if let Some(sealed) = self.kv.get(&key).await.map_err(|e| e.to_string())? {
let plain = self
.envelope
.unwrap(&sealed)
.await
.map_err(|e| e.to_string())?;
return String::from_utf8(plain).map_err(|_| {
format!("managed sql credential for {workload:?} is not valid UTF-8")
});
}
let mut bytes = [0u8; 32];
getrandom::getrandom(&mut bytes).map_err(|e| format!("rng: {e}"))?;
let password: String = bytes.iter().map(|b| format!("{b:02x}")).collect();
let sealed = self
.envelope
.wrap(password.as_bytes())
.await
.map_err(|e| e.to_string())?;
self.kv.put(&key, sealed).await.map_err(|e| e.to_string())?;
Ok(password)
}
}
#[cfg_attr(not(feature = "handlers"), allow(dead_code))]
struct ManagedDbSpec {
kind: ExternalSqlKind,
database: String,
user: String,
}
#[cfg_attr(not(feature = "handlers"), allow(dead_code))]
pub struct ManagedDbEnv {
dbs: HashMap<String, ManagedDbSpec>,
creds: ManagedSqlCredentials,
privilege: ManagedDbPrivilege,
}
#[cfg_attr(not(feature = "handlers"), allow(dead_code))]
fn managed_db_default_ids(_kind: ExternalSqlKind) -> (u32, u32) {
(999, 999)
}
#[cfg_attr(not(feature = "handlers"), allow(dead_code))]
fn managed_db_caps() -> Vec<String> {
["CHOWN", "DAC_OVERRIDE", "FOWNER", "SETUID", "SETGID"]
.iter()
.map(|s| (*s).to_string())
.collect()
}
impl ManagedDbEnv {
#[cfg_attr(not(feature = "handlers"), allow(dead_code))]
pub fn from_config(
databases: &std::collections::BTreeMap<String, crate::config::ExternalDatabaseConfig>,
creds: ManagedSqlCredentials,
privilege: ManagedDbPrivilege,
) -> Self {
let mut dbs = HashMap::new();
for db in databases.values() {
if !db.is_managed_credential() {
continue;
}
let (Some(workload), Some(kind), Some(database), Some(user)) = (
db.compute.clone(),
ExternalSqlKind::parse(&db.kind),
db.database.clone(),
db.user.clone(),
) else {
continue;
};
dbs.insert(
workload,
ManagedDbSpec {
kind,
database,
user,
},
);
}
Self {
dbs,
creds,
privilege,
}
}
#[cfg_attr(not(feature = "handlers"), allow(dead_code))]
pub fn is_empty(&self) -> bool {
self.dbs.is_empty()
}
}
#[async_trait]
impl ManagedDbEnvResolver for ManagedDbEnv {
async fn managed_db_env(&self, project: &str, workload: &str) -> Vec<(String, String)> {
let Some(db) = self.dbs.get(workload) else {
return Vec::new();
};
match self.creds.password(project, workload).await {
Ok(password) => managed_db_server_env(db.kind, &db.database, &db.user, &password),
Err(e) => {
tracing::error!(
%workload,
error = %e,
"managed sql: could not resolve the sealed credential; DB launched without managed env"
);
Vec::new()
}
}
}
fn managed_db_privilege(&self, _project: &str, workload: &str) -> Option<PrivilegeDirective> {
let db = self.dbs.get(workload)?;
Some(match self.privilege {
ManagedDbPrivilege::Rootless => {
let (uid, gid) = managed_db_default_ids(db.kind);
PrivilegeDirective::Rootless { uid, gid }
}
ManagedDbPrivilege::Caps => PrivilegeDirective::Caps(managed_db_caps()),
})
}
}
#[cfg_attr(not(feature = "handlers"), allow(dead_code))]
pub struct DeployEndpointResolver {
deploy: DeployStore,
project: String,
}
impl DeployEndpointResolver {
#[cfg_attr(not(feature = "handlers"), allow(dead_code))]
pub fn new(deploy: DeployStore, project: impl Into<String>) -> Self {
Self {
deploy,
project: project.into(),
}
}
}
#[async_trait]
impl ComputeEndpointResolver for DeployEndpointResolver {
async fn endpoints(&self, workload: &str) -> Result<Vec<(String, u16)>, SqlError> {
let states = self
.deploy
.list_replica_states(ProjectRef::new(&self.project), workload)
.await
.map_err(SqlError::other)?;
Ok(states
.into_iter()
.filter(|s| s.phase == ReplicaPhase::Running && s.healthy)
.map(|s| (s.endpoint.host, s.endpoint.port))
.collect())
}
}
#[cfg(test)]
mod tests {
use super::*;
use async_trait::async_trait;
use boatramp_core::envelope::EnvelopeError;
use boatramp_core::kv::MemoryKv;
struct ReverseEnvelope;
#[async_trait]
impl KeyEnvelope for ReverseEnvelope {
async fn wrap(&self, plaintext: &[u8]) -> Result<Vec<u8>, EnvelopeError> {
Ok(plaintext.iter().rev().copied().collect())
}
async fn unwrap(&self, wrapped: &[u8]) -> Result<Vec<u8>, EnvelopeError> {
Ok(wrapped.iter().rev().copied().collect())
}
}
#[tokio::test]
async fn password_is_generated_once_sealed_and_stable() {
let kv: Arc<dyn KvStore> = Arc::new(MemoryKv::new());
let creds = ManagedSqlCredentials::new(kv.clone(), Arc::new(ReverseEnvelope));
let pw = creds.password("default", "pg").await.unwrap();
assert_eq!(pw.len(), 64, "32 random bytes, hex-encoded");
assert_eq!(creds.password("default", "pg").await.unwrap(), pw);
let raw = kv
.get("managed-sql-cred/default/pg")
.await
.unwrap()
.unwrap();
assert_ne!(
raw,
pw.as_bytes(),
"the stored blob is sealed, not the password"
);
assert_eq!(
raw.iter().rev().copied().collect::<Vec<u8>>(),
pw.as_bytes()
);
let after_restart = ManagedSqlCredentials::new(kv, Arc::new(ReverseEnvelope));
assert_eq!(after_restart.password("default", "pg").await.unwrap(), pw);
assert_ne!(creds.password("default", "other").await.unwrap(), pw);
}
#[test]
fn server_env_recipe_per_engine() {
let pg = managed_db_server_env(ExternalSqlKind::Postgres, "analytics", "app", "pw");
assert_eq!(
pg,
vec![
("POSTGRES_USER".into(), "app".into()),
("POSTGRES_PASSWORD".into(), "pw".into()),
("POSTGRES_DB".into(), "analytics".into()),
]
);
let my = managed_db_server_env(ExternalSqlKind::Mysql, "shop", "app", "pw");
assert!(my.contains(&("MYSQL_USER".into(), "app".into())));
assert!(my.contains(&("MYSQL_DATABASE".into(), "shop".into())));
assert!(my.contains(&("MYSQL_ROOT_PASSWORD".into(), "pw".into())));
}
use crate::config::ExternalDatabaseConfig;
use std::collections::BTreeMap;
fn db(
kind: &str,
compute: Option<&str>,
url_env: &str,
pw_env: Option<&str>,
) -> ExternalDatabaseConfig {
ExternalDatabaseConfig {
kind: kind.into(),
url_env: url_env.into(),
compute: compute.map(Into::into),
database: compute.map(|_| "analytics".into()),
user: compute.map(|_| "app".into()),
password_env: pw_env.map(Into::into),
..Default::default()
}
}
#[tokio::test]
async fn managed_db_env_only_covers_managed_workloads() {
let mut dbs = BTreeMap::new();
dbs.insert(
"analytics".to_string(),
db("postgres", Some("pg"), "", None),
);
dbs.insert(
"byo".to_string(),
db("postgres", Some("pg2"), "", Some("PG2_PW")),
);
dbs.insert("ext".to_string(), db("mysql", None, "MYSQL_URL", None));
let kv: Arc<dyn KvStore> = Arc::new(MemoryKv::new());
let creds = ManagedSqlCredentials::new(kv, Arc::new(ReverseEnvelope));
let env = ManagedDbEnv::from_config(&dbs, creds, ManagedDbPrivilege::default());
assert!(!env.is_empty());
assert_eq!(
env.managed_db_privilege("default", "pg"),
Some(PrivilegeDirective::Rootless { uid: 999, gid: 999 })
);
assert_eq!(env.managed_db_privilege("default", "nope"), None);
let pg = env.managed_db_env("default", "pg").await;
assert!(pg.contains(&("POSTGRES_USER".into(), "app".into())));
assert!(pg.contains(&("POSTGRES_DB".into(), "analytics".into())));
let password = pg
.iter()
.find(|(k, _)| k == "POSTGRES_PASSWORD")
.map(|(_, v)| v.clone())
.expect("password present");
assert_eq!(password.len(), 64, "managed 32-byte hex password");
let pg2 = env.managed_db_env("default", "pg").await;
assert_eq!(pg, pg2);
assert!(env.managed_db_env("default", "pg2").await.is_empty());
assert!(env.managed_db_env("default", "nope").await.is_empty());
}
use boatramp_core::{ByteStream, GetObject, ObjectMeta, PutMeta, Storage, StorageError};
struct NullStorage;
#[async_trait]
impl Storage for NullStorage {
async fn get(&self, _: &str) -> Result<GetObject, StorageError> {
Err(StorageError::NotFound(String::new()))
}
async fn get_range(
&self,
_: &str,
_: u64,
_: Option<u64>,
) -> Result<GetObject, StorageError> {
Err(StorageError::NotFound(String::new()))
}
async fn put(
&self,
_: &str,
_: ByteStream,
_: PutMeta,
) -> Result<ObjectMeta, StorageError> {
Err(StorageError::unsupported("null"))
}
async fn head(&self, _: &str) -> Result<ObjectMeta, StorageError> {
Err(StorageError::NotFound(String::new()))
}
async fn delete(&self, _: &str) -> Result<(), StorageError> {
Ok(())
}
async fn list(&self, _: &str) -> Result<Vec<ObjectMeta>, StorageError> {
Ok(Vec::new())
}
}
fn replica(
workload: &str,
replica: u32,
host: &str,
port: u16,
healthy: bool,
phase: ReplicaPhase,
) -> boatramp_core::compute::ObservedInstance {
use boatramp_core::compute::{Endpoint, InstanceHandle, Scheme};
boatramp_core::compute::ObservedInstance {
handle: InstanceHandle {
workload: workload.into(),
replica,
backend_ref: String::new(),
},
node: 0,
backend: "fake".into(),
endpoint: Endpoint {
scheme: Scheme::Http,
host: host.into(),
port,
},
region: None,
healthy,
phase,
snapshot: None,
}
}
#[tokio::test]
async fn endpoint_resolver_returns_only_healthy_running_replicas() {
let deploy = DeployStore::new(Arc::new(NullStorage), Arc::new(MemoryKv::new()));
let p = ProjectRef::DEFAULT;
deploy
.set_replica_state(
p,
&replica("pg", 0, "10.0.0.1", 5432, true, ReplicaPhase::Running),
)
.await
.unwrap();
deploy
.set_replica_state(
p,
&replica("pg", 1, "10.0.0.2", 5432, true, ReplicaPhase::Running),
)
.await
.unwrap();
deploy
.set_replica_state(
p,
&replica("pg", 2, "10.0.0.3", 5432, false, ReplicaPhase::Running),
)
.await
.unwrap();
deploy
.set_replica_state(
p,
&replica("pg", 3, "10.0.0.4", 5432, false, ReplicaPhase::Zero),
)
.await
.unwrap();
let resolver = DeployEndpointResolver::new(deploy, "default");
let mut eps = resolver.endpoints("pg").await.unwrap();
eps.sort();
assert_eq!(
eps,
vec![
("10.0.0.1".to_string(), 5432),
("10.0.0.2".to_string(), 5432)
],
"only the healthy running replicas, unhealthy + Zero filtered out"
);
assert!(resolver.endpoints("absent").await.unwrap().is_empty());
}
}