use std::collections::BTreeMap;
use std::net::Ipv4Addr;
use std::path::{Path, PathBuf};
use std::sync::Arc;
use std::time::{Duration, Instant};
use anyhow::{Context, Result};
use kamaji::native::NativeRuntime;
use kamaji::{Kamaji, MeshAssignment, MeshIdent};
use tracing::{info, warn};
use super::native_support::{native_spec, sanitize_ident};
use crate::capability::Capability;
use crate::config::{MirrorConfig, Provider, ServiceWithMirrors};
pub const PG_DRIVER_IDENT: &str = "yah-pg-dev";
pub const PG_DEV_BIN_ENV: &str = "YAH_PG_DEV_BIN";
const DEV_ENV: &str = "dev";
#[derive(Debug, Clone, Default)]
pub struct PgDriverOptions {
pub binary: Option<PathBuf>,
pub ready_timeout: Option<Duration>,
}
impl PgDriverOptions {
fn resolved_binary(&self) -> PathBuf {
if let Some(ref p) = self.binary {
return p.clone();
}
if let Some(p) = std::env::var_os(PG_DEV_BIN_ENV) {
return PathBuf::from(p);
}
PathBuf::from("yah-pg-dev")
}
fn ready_timeout(&self) -> Duration {
self.ready_timeout.unwrap_or(Duration::from_secs(180))
}
}
pub struct RunningPgDriver {
pub port: u16,
pub databases: Vec<String>,
runtime: Arc<NativeRuntime>,
ident: MeshIdent,
}
impl RunningPgDriver {
pub async fn teardown(&self) {
self.runtime.teardown_workload(&self.ident).await.ok();
}
}
pub fn declared_pg_databases(services: &BTreeMap<String, ServiceWithMirrors>) -> Vec<String> {
let mut out = Vec::new();
for (name, svc) in services {
let Some(mirror) = svc.mirrors.get(DEV_ENV) else {
continue;
};
if binds_local_pg_dev(mirror) {
out.push(yah_pg_dev_database_name(name));
}
}
out.sort();
out.dedup();
out
}
fn binds_local_pg_dev(mirror: &MirrorConfig) -> bool {
mirror
.driver(Capability::Pg)
.and_then(|slot| slot.inline_kind())
== Some(Provider::LocalPgDev)
}
fn yah_pg_dev_database_name(service: &str) -> String {
let sanitize = |s: &str| -> String {
s.chars()
.map(|c| {
if c.is_ascii_alphanumeric() {
c.to_ascii_lowercase()
} else {
'_'
}
})
.collect::<String>()
};
let mut name = format!("svc_{}_{}", sanitize(service), sanitize(DEV_ENV));
name.truncate(63);
name
}
pub fn coords_path(workspace_root: &Path) -> PathBuf {
workspace_root.join(".yah/infra/state/dev/pg/coords.json")
}
pub async fn up_pg_driver(
workspace_root: &Path,
databases: Vec<String>,
opts: &PgDriverOptions,
) -> Result<RunningPgDriver> {
let binary = opts.resolved_binary();
let ident_str = sanitize_ident(PG_DRIVER_IDENT);
let ident = MeshIdent(ident_str.clone());
let mut argv: Vec<String> = vec![
binary.display().to_string(),
"serve".to_string(),
"--workspace".to_string(),
workspace_root.display().to_string(),
];
for db in &databases {
argv.push("--database".to_string());
argv.push(db.clone());
}
let coords = coords_path(workspace_root);
let _ = std::fs::remove_file(&coords);
let spec = native_spec(&ident_str, argv, Vec::new());
let state_dir = workspace_root.join(".yah/jit/native");
let runtime = Arc::new(NativeRuntime::new(&state_dir));
let mesh = MeshAssignment::inlined(Ipv4Addr::LOCALHOST);
info!(
binary = %binary.display(),
databases = databases.len(),
ident = %ident_str,
"spawning yah-pg-dev (kamaji native backend)",
);
runtime
.deploy_workload(&spec, &mesh)
.await
.with_context(|| {
format!(
"deploying the dev-tier pg driver via kamaji — install it with \
`cargo install --path crates/yah/pg-dev` or point {PG_DEV_BIN_ENV} \
at the binary ({})",
binary.display(),
)
})?;
let timeout = opts.ready_timeout();
let Some(port) = wait_for_coords(&coords, timeout).await else {
warn!(timeout = ?timeout, "yah-pg-dev did not publish coords; tearing down");
runtime.teardown_workload(&ident).await.ok();
let (_out, err) = super::native_support::capture_paths(&state_dir, &ident_str);
anyhow::bail!(
"the dev-tier pg driver did not become ready within {timeout:?} — \
check {} for why",
err.display(),
);
};
info!(
port,
databases = databases.len(),
"dev-tier pg driver ready"
);
Ok(RunningPgDriver {
port,
databases,
runtime,
ident,
})
}
async fn wait_for_coords(path: &Path, timeout: Duration) -> Option<u16> {
let deadline = Instant::now() + timeout;
while Instant::now() < deadline {
if let Some(port) = read_coords_port(path) {
return Some(port);
}
tokio::time::sleep(Duration::from_millis(100)).await;
}
None
}
fn read_coords_port(path: &Path) -> Option<u16> {
let bytes = std::fs::read(path).ok()?;
let v: serde_json::Value = serde_json::from_slice(&bytes).ok()?;
let port = u16::try_from(v.get("port")?.as_u64()?).ok()?;
(port != 0).then_some(port)
}
#[cfg(test)]
mod tests {
use super::*;
use crate::config::ServiceConfig;
fn service(name: &str, dev_mirror: Option<&str>) -> ServiceWithMirrors {
let service: ServiceConfig = toml::from_str(&format!(
"schema_version = 1\nname = \"{name}\"\n[address]\nkind = \"front-door\"\ndomain = \"{name}.example\"\n"
))
.expect("parse service");
let mut mirrors = BTreeMap::new();
if let Some(src) = dev_mirror {
mirrors.insert(
DEV_ENV.to_string(),
toml::from_str::<MirrorConfig>(src).expect("parse mirror"),
);
}
ServiceWithMirrors {
service,
mirrors,
component_transform_recipes: BTreeMap::new(),
passway_machines: BTreeMap::new(),
}
}
const BINDS_PG: &str = r#"
schema_version = 1
shape = "local"
[drivers.pg]
kind = "local-pg-dev"
"#;
const NO_DRIVERS: &str = r#"
schema_version = 1
shape = "local"
[providers.static]
kind = "miniflare-native"
port = 4324
"#;
fn services(entries: Vec<(&str, ServiceWithMirrors)>) -> BTreeMap<String, ServiceWithMirrors> {
entries
.into_iter()
.map(|(n, s)| (n.to_string(), s))
.collect()
}
#[test]
fn only_services_binding_the_pg_driver_get_a_database() {
let svcs = services(vec![
("scrabcake", service("scrabcake", Some(BINDS_PG))),
("yah-dashboard", service("yah-dashboard", Some(NO_DRIVERS))),
("yah-cloud", service("yah-cloud", None)),
]);
assert_eq!(
declared_pg_databases(&svcs),
vec!["svc_scrabcake_dev".to_string()]
);
}
#[test]
fn no_binding_anywhere_means_no_driver_to_spawn() {
let svcs = services(vec![(
"yah-dashboard",
service("yah-dashboard", Some(NO_DRIVERS)),
)]);
assert!(declared_pg_databases(&svcs).is_empty());
}
#[test]
fn a_non_dev_mirror_binding_pg_is_ignored() {
let mut svc = service("scrabcake", None);
svc.mirrors.insert(
"pond".to_string(),
toml::from_str::<MirrorConfig>(BINDS_PG).expect("parse mirror"),
);
assert!(declared_pg_databases(&services(vec![("scrabcake", svc)])).is_empty());
}
#[test]
fn database_names_match_the_drivers_own_convention() {
assert_eq!(
yah_pg_dev_database_name("yah-dashboard"),
"svc_yah_dashboard_dev"
);
assert_eq!(yah_pg_dev_database_name("Scrabcake"), "svc_scrabcake_dev");
}
#[test]
fn binary_resolution_prefers_explicit_over_env_over_path() {
let explicit = PgDriverOptions {
binary: Some(PathBuf::from("/opt/yah-pg-dev")),
..Default::default()
};
assert_eq!(explicit.resolved_binary(), PathBuf::from("/opt/yah-pg-dev"));
if std::env::var_os(PG_DEV_BIN_ENV).is_none() {
assert_eq!(
PgDriverOptions::default().resolved_binary(),
PathBuf::from("yah-pg-dev")
);
}
}
#[tokio::test]
async fn wait_for_coords_times_out_on_a_missing_or_portless_file() {
let tmp = tempfile::tempdir().unwrap();
let path = tmp.path().join("coords.json");
assert_eq!(
wait_for_coords(&path, Duration::from_millis(150)).await,
None
);
std::fs::write(&path, br#"{"port":0}"#).unwrap();
assert_eq!(
wait_for_coords(&path, Duration::from_millis(150)).await,
None
);
}
#[tokio::test]
async fn wait_for_coords_returns_the_published_port() {
let tmp = tempfile::tempdir().unwrap();
let path = tmp.path().join("coords.json");
std::fs::write(&path, br#"{"port":25432,"username":"postgres"}"#).unwrap();
assert_eq!(
wait_for_coords(&path, Duration::from_secs(1)).await,
Some(25432)
);
}
}