1use std::collections::BTreeMap;
39use std::net::Ipv4Addr;
40use std::path::{Path, PathBuf};
41use std::sync::Arc;
42use std::time::{Duration, Instant};
43
44use anyhow::{Context, Result};
45use kamaji::native::NativeRuntime;
46use kamaji::{Kamaji, MeshAssignment, MeshIdent};
47use tracing::{info, warn};
48
49use super::native_support::{native_spec, sanitize_ident};
50use crate::capability::Capability;
51use crate::config::{MirrorConfig, Provider, ServiceWithMirrors};
52
53pub const PG_DRIVER_IDENT: &str = "yah-pg-dev";
57
58pub const PG_DEV_BIN_ENV: &str = "YAH_PG_DEV_BIN";
61
62const DEV_ENV: &str = "dev";
65
66#[derive(Debug, Clone, Default)]
68pub struct PgDriverOptions {
69 pub binary: Option<PathBuf>,
72 pub ready_timeout: Option<Duration>,
77}
78
79impl PgDriverOptions {
80 fn resolved_binary(&self) -> PathBuf {
81 if let Some(ref p) = self.binary {
82 return p.clone();
83 }
84 if let Some(p) = std::env::var_os(PG_DEV_BIN_ENV) {
85 return PathBuf::from(p);
86 }
87 PathBuf::from("yah-pg-dev")
88 }
89
90 fn ready_timeout(&self) -> Duration {
91 self.ready_timeout.unwrap_or(Duration::from_secs(180))
92 }
93}
94
95pub struct RunningPgDriver {
97 pub port: u16,
99 pub databases: Vec<String>,
101 runtime: Arc<NativeRuntime>,
102 ident: MeshIdent,
103}
104
105impl RunningPgDriver {
106 pub async fn teardown(&self) {
109 self.runtime.teardown_workload(&self.ident).await.ok();
110 }
111}
112
113pub fn declared_pg_databases(services: &BTreeMap<String, ServiceWithMirrors>) -> Vec<String> {
122 let mut out = Vec::new();
123 for (name, svc) in services {
124 let Some(mirror) = svc.mirrors.get(DEV_ENV) else {
125 continue;
126 };
127 if binds_local_pg_dev(mirror) {
128 out.push(yah_pg_dev_database_name(name));
129 }
130 }
131 out.sort();
132 out.dedup();
133 out
134}
135
136fn binds_local_pg_dev(mirror: &MirrorConfig) -> bool {
138 mirror
139 .driver(Capability::Pg)
140 .and_then(|slot| slot.inline_kind())
141 == Some(Provider::LocalPgDev)
142}
143
144fn yah_pg_dev_database_name(service: &str) -> String {
149 let sanitize = |s: &str| -> String {
150 s.chars()
151 .map(|c| {
152 if c.is_ascii_alphanumeric() {
153 c.to_ascii_lowercase()
154 } else {
155 '_'
156 }
157 })
158 .collect::<String>()
159 };
160 let mut name = format!("svc_{}_{}", sanitize(service), sanitize(DEV_ENV));
161 name.truncate(63);
162 name
163}
164
165pub fn coords_path(workspace_root: &Path) -> PathBuf {
167 workspace_root.join(".yah/infra/state/dev/pg/coords.json")
168}
169
170pub async fn up_pg_driver(
177 workspace_root: &Path,
178 databases: Vec<String>,
179 opts: &PgDriverOptions,
180) -> Result<RunningPgDriver> {
181 let binary = opts.resolved_binary();
182 let ident_str = sanitize_ident(PG_DRIVER_IDENT);
183 let ident = MeshIdent(ident_str.clone());
184
185 let mut argv: Vec<String> = vec![
186 binary.display().to_string(),
187 "serve".to_string(),
188 "--workspace".to_string(),
189 workspace_root.display().to_string(),
190 ];
191 for db in &databases {
192 argv.push("--database".to_string());
193 argv.push(db.clone());
194 }
195
196 let coords = coords_path(workspace_root);
200 let _ = std::fs::remove_file(&coords);
201
202 let spec = native_spec(&ident_str, argv, Vec::new());
203 let state_dir = workspace_root.join(".yah/jit/native");
204 let runtime = Arc::new(NativeRuntime::new(&state_dir));
205 let mesh = MeshAssignment::inlined(Ipv4Addr::LOCALHOST);
206
207 info!(
208 binary = %binary.display(),
209 databases = databases.len(),
210 ident = %ident_str,
211 "spawning yah-pg-dev (kamaji native backend)",
212 );
213
214 runtime
215 .deploy_workload(&spec, &mesh)
216 .await
217 .with_context(|| {
218 format!(
219 "deploying the dev-tier pg driver via kamaji — install it with \
220 `cargo install --path crates/yah/pg-dev` or point {PG_DEV_BIN_ENV} \
221 at the binary ({})",
222 binary.display(),
223 )
224 })?;
225
226 let timeout = opts.ready_timeout();
227 let Some(port) = wait_for_coords(&coords, timeout).await else {
228 warn!(timeout = ?timeout, "yah-pg-dev did not publish coords; tearing down");
229 runtime.teardown_workload(&ident).await.ok();
230 let (_out, err) = super::native_support::capture_paths(&state_dir, &ident_str);
231 anyhow::bail!(
232 "the dev-tier pg driver did not become ready within {timeout:?} — \
233 check {} for why",
234 err.display(),
235 );
236 };
237
238 info!(
239 port,
240 databases = databases.len(),
241 "dev-tier pg driver ready"
242 );
243 Ok(RunningPgDriver {
244 port,
245 databases,
246 runtime,
247 ident,
248 })
249}
250
251async fn wait_for_coords(path: &Path, timeout: Duration) -> Option<u16> {
257 let deadline = Instant::now() + timeout;
258 while Instant::now() < deadline {
259 if let Some(port) = read_coords_port(path) {
260 return Some(port);
261 }
262 tokio::time::sleep(Duration::from_millis(100)).await;
263 }
264 None
265}
266
267fn read_coords_port(path: &Path) -> Option<u16> {
270 let bytes = std::fs::read(path).ok()?;
271 let v: serde_json::Value = serde_json::from_slice(&bytes).ok()?;
272 let port = u16::try_from(v.get("port")?.as_u64()?).ok()?;
273 (port != 0).then_some(port)
274}
275
276#[cfg(test)]
277mod tests {
278 use super::*;
279 use crate::config::ServiceConfig;
280
281 fn service(name: &str, dev_mirror: Option<&str>) -> ServiceWithMirrors {
282 let service: ServiceConfig = toml::from_str(&format!(
283 "schema_version = 1\nname = \"{name}\"\ndomain = \"{name}.example\"\n"
284 ))
285 .expect("parse service");
286 let mut mirrors = BTreeMap::new();
287 if let Some(src) = dev_mirror {
288 mirrors.insert(
289 DEV_ENV.to_string(),
290 toml::from_str::<MirrorConfig>(src).expect("parse mirror"),
291 );
292 }
293 ServiceWithMirrors {
294 service,
295 mirrors,
296 component_transform_recipes: BTreeMap::new(),
297 }
298 }
299
300 const BINDS_PG: &str = r#"
301schema_version = 1
302shape = "local"
303[drivers.pg]
304kind = "local-pg-dev"
305"#;
306
307 const NO_DRIVERS: &str = r#"
308schema_version = 1
309shape = "local"
310[providers.static]
311kind = "local-static"
312port = 4324
313"#;
314
315 fn services(entries: Vec<(&str, ServiceWithMirrors)>) -> BTreeMap<String, ServiceWithMirrors> {
316 entries
317 .into_iter()
318 .map(|(n, s)| (n.to_string(), s))
319 .collect()
320 }
321
322 #[test]
323 fn only_services_binding_the_pg_driver_get_a_database() {
324 let svcs = services(vec![
325 ("scrabcake", service("scrabcake", Some(BINDS_PG))),
326 ("yah-dashboard", service("yah-dashboard", Some(NO_DRIVERS))),
327 ("yah-cloud", service("yah-cloud", None)),
328 ]);
329 assert_eq!(
330 declared_pg_databases(&svcs),
331 vec!["svc_scrabcake_dev".to_string()]
332 );
333 }
334
335 #[test]
336 fn no_binding_anywhere_means_no_driver_to_spawn() {
337 let svcs = services(vec![(
338 "yah-dashboard",
339 service("yah-dashboard", Some(NO_DRIVERS)),
340 )]);
341 assert!(declared_pg_databases(&svcs).is_empty());
342 }
343
344 #[test]
347 fn a_non_dev_mirror_binding_pg_is_ignored() {
348 let mut svc = service("scrabcake", None);
349 svc.mirrors.insert(
350 "pond".to_string(),
351 toml::from_str::<MirrorConfig>(BINDS_PG).expect("parse mirror"),
352 );
353 assert!(declared_pg_databases(&services(vec![("scrabcake", svc)])).is_empty());
354 }
355
356 #[test]
357 fn database_names_match_the_drivers_own_convention() {
358 assert_eq!(
361 yah_pg_dev_database_name("yah-dashboard"),
362 "svc_yah_dashboard_dev"
363 );
364 assert_eq!(yah_pg_dev_database_name("Scrabcake"), "svc_scrabcake_dev");
365 }
366
367 #[test]
368 fn binary_resolution_prefers_explicit_over_env_over_path() {
369 let explicit = PgDriverOptions {
370 binary: Some(PathBuf::from("/opt/yah-pg-dev")),
371 ..Default::default()
372 };
373 assert_eq!(explicit.resolved_binary(), PathBuf::from("/opt/yah-pg-dev"));
374 if std::env::var_os(PG_DEV_BIN_ENV).is_none() {
378 assert_eq!(
379 PgDriverOptions::default().resolved_binary(),
380 PathBuf::from("yah-pg-dev")
381 );
382 }
383 }
384
385 #[tokio::test]
386 async fn wait_for_coords_times_out_on_a_missing_or_portless_file() {
387 let tmp = tempfile::tempdir().unwrap();
388 let path = tmp.path().join("coords.json");
389 assert_eq!(
390 wait_for_coords(&path, Duration::from_millis(150)).await,
391 None
392 );
393 std::fs::write(&path, br#"{"port":0}"#).unwrap();
395 assert_eq!(
396 wait_for_coords(&path, Duration::from_millis(150)).await,
397 None
398 );
399 }
400
401 #[tokio::test]
402 async fn wait_for_coords_returns_the_published_port() {
403 let tmp = tempfile::tempdir().unwrap();
404 let path = tmp.path().join("coords.json");
405 std::fs::write(&path, br#"{"port":25432,"username":"postgres"}"#).unwrap();
406 assert_eq!(
407 wait_for_coords(&path, Duration::from_secs(1)).await,
408 Some(25432)
409 );
410 }
411}