use std::{collections::BTreeMap, time::Duration};
use anyhow::{Context, Result, bail};
use serde::Deserialize;
use super::exec::{exec_in_container, exec_probe_in_container};
use super::project::{ServiceContext, find_service_instance, get_environment_instances};
const PROBE_TIMEOUT: Duration = Duration::from_secs(5);
const SWITCHOVER_CURL_MAX_TIME_SECS: u64 = 30;
const SWITCHOVER_TIMEOUT: Duration = Duration::from_secs(40);
#[derive(Debug, Clone, Default, Deserialize)]
#[serde(default)]
pub struct PatroniMember {
pub name: String,
pub role: String,
pub state: String,
pub lag: Option<serde_json::Value>,
pub timeline: Option<i64>,
}
#[derive(Debug, Clone, Default, Deserialize)]
#[serde(default)]
struct PatroniClusterResponse {
members: Vec<PatroniMember>,
}
pub async fn probe_cluster(instance_id: &str) -> Result<Vec<PatroniMember>> {
let command = "curl -s --max-time 4 localhost:8008/cluster";
let output = exec_probe_in_container(instance_id, command, PROBE_TIMEOUT)
.await
.context("Probing Patroni failed")?;
let parsed: PatroniClusterResponse = serde_json::from_str(output.trim())
.with_context(|| format!("Unexpected response from Patroni: {}", output.trim()))?;
Ok(parsed.members)
}
pub async fn probe_any(
instance_ids: &[String],
) -> Result<(String, Vec<PatroniMember>), Vec<(String, String)>> {
let mut failures = Vec::with_capacity(instance_ids.len());
for instance_id in instance_ids {
match probe_cluster(instance_id).await {
Ok(members) => return Ok((instance_id.clone(), members)),
Err(e) => failures.push((instance_id.clone(), format!("{e:#}"))),
}
}
Err(failures)
}
const RESTAPI_AUTH_PRELUDE: &str = concat!(
r#"PATRONI_REST_PW="${PATRONI_RESTAPI_PASSWORD:-${PATRONI_SUPERUSER_PASSWORD:-${PGPASSWORD:-${POSTGRES_PASSWORD:-}}}}"; "#,
r#"PATRONI_REST_USER="${PATRONI_RESTAPI_USERNAME:-${PATRONI_SUPERUSER_USERNAME:-${PGUSER:-${POSTGRES_USER:-postgres}}}}"; "#,
r#"curl_cfg_quote() { s=$1; q=; nl=$(printf '\nx'); nl=${nl%x}; "#,
r#"while [ -n "$s" ]; do r=${s#?}; c=${s%"$r"}; s=$r; "#,
r#"case $c in \\) q="$q\\\\";; \") q="$q\\\"";; "$nl") q="$q\\n";; *) q="$q$c";; esac; done; "#,
r#"printf '%s' "$q"; }; "#,
r#"if [ -n "$PATRONI_REST_PW" ]; then PATRONI_REST_CFG="user = \"$(curl_cfg_quote "$PATRONI_REST_USER:$PATRONI_REST_PW")\""; set -- -K -; else PATRONI_REST_CFG=; set --; fi; "#,
);
fn switchover_command(body: &str) -> String {
format!(
r#"{prelude}printf '%s\n' "$PATRONI_REST_CFG" | curl -s --max-time {max_time} -w '\nHTTP_STATUS:%{{http_code}}' "$@" -X POST localhost:8008/switchover -H 'Content-Type: application/json' -d '{body}'"#,
prelude = RESTAPI_AUTH_PRELUDE,
max_time = SWITCHOVER_CURL_MAX_TIME_SECS,
)
}
pub async fn switchover(instance_id: &str, leader: &str, candidate: &str) -> Result<String> {
let body = serde_json::json!({ "leader": leader, "candidate": candidate }).to_string();
let command = switchover_command(&body);
let output = match tokio::time::timeout(
SWITCHOVER_TIMEOUT,
exec_in_container(instance_id, &command),
)
.await
{
Ok(Ok(output)) => output,
Ok(Err(err)) => return map_switchover_exec_error(err),
Err(_elapsed) => bail!(
"Timed out requesting switchover. The failover may still be in progress — check `ha status`."
),
};
parse_switchover_response(&output)
}
fn map_switchover_exec_error(err: anyhow::Error) -> Result<String> {
let detail = format!("{err:#}");
if detail.contains("exit code 28") || detail.contains("HTTP_STATUS:000") {
bail!(
"Patroni did not answer the switchover in time. \
The failover may still be in progress — check `ha status`."
);
}
Err(err)
}
fn parse_switchover_response(output: &str) -> Result<String> {
let (response_body, status) = match output.rsplit_once("HTTP_STATUS:") {
Some((body, status)) => (body.trim().to_string(), status.trim().parse::<u16>().ok()),
None => (output.trim().to_string(), None),
};
match status {
Some(200..=299) => Ok(response_body),
Some(0) => bail!(
"Patroni did not answer the switchover in time. \
The failover may still be in progress — check `ha status`."
),
Some(code) => bail!("Patroni switchover failed ({code}): {response_body}"),
None => bail!("Patroni switchover returned an unexpected response: {response_body}"),
}
}
pub async fn resolve_instance_ids(
ctx: &ServiceContext,
service_ids: &[String],
) -> Result<BTreeMap<String, String>> {
let instances = get_environment_instances(
&ctx.client,
&ctx.configs,
&ctx.project_id,
&ctx.environment_id,
)
.await?;
Ok(service_ids
.iter()
.filter_map(|id| {
find_service_instance(&instances, id).map(|si| (id.clone(), si.id.clone()))
})
.collect())
}
#[derive(Debug, Clone, Default)]
pub struct MemberProbe {
pub reachable: bool,
pub self_view: Option<PatroniMember>,
}
pub async fn probe_members(
ctx: &ServiceContext,
members: &[(String, String)],
) -> Result<BTreeMap<String, MemberProbe>> {
let service_ids: Vec<String> = members.iter().map(|(id, _)| id.clone()).collect();
let instance_ids = resolve_instance_ids(ctx, &service_ids).await?;
let probes = members.iter().map(|(service_id, service_name)| {
let instance_id = instance_ids.get(service_id).cloned();
let name_lower = service_name.to_ascii_lowercase();
let service_id = service_id.clone();
async move {
let Some(instance_id) = instance_id else {
return (service_id, MemberProbe::default());
};
match probe_cluster(&instance_id).await {
Ok(cluster_members) => {
let self_view = cluster_members
.into_iter()
.find(|m| m.name.to_ascii_lowercase() == name_lower);
(
service_id,
MemberProbe {
reachable: true,
self_view,
},
)
}
Err(_) => (service_id, MemberProbe::default()),
}
}
});
Ok(futures::future::join_all(probes)
.await
.into_iter()
.collect())
}
#[cfg(test)]
mod tests {
use super::*;
#[cfg(unix)]
struct CurlCall {
argv: Vec<String>,
config: Option<String>,
}
#[cfg(unix)]
fn curl_call_for(env: &[(&str, &str)]) -> CurlCall {
use std::io::Write;
let dir = std::env::temp_dir().join(format!(
"cli-patroni-auth-{}-{:?}",
std::process::id(),
std::thread::current().id()
));
let _ = std::fs::remove_dir_all(&dir);
std::fs::create_dir_all(&dir).unwrap();
let shim = dir.join("curl");
std::fs::write(
&shim,
concat!(
"#!/bin/sh\n",
"printf '%s\\n' \"$@\"\n",
"prev=\n",
"for a in \"$@\"; do\n",
" if [ \"$prev\" = -K ] && [ \"$a\" = - ]; then printf '%s\\n' '@@CONFIG@@'; cat; fi\n",
" prev=$a\n",
"done\n",
),
)
.unwrap();
#[cfg(unix)]
{
use std::os::unix::fs::PermissionsExt;
std::fs::set_permissions(&shim, std::fs::Permissions::from_mode(0o755)).unwrap();
}
let mut child = std::process::Command::new("sh")
.arg("-s")
.env(
"PATH",
format!(
"{}:{}",
dir.display(),
std::env::var("PATH").unwrap_or_default()
),
)
.env_remove("PATRONI_RESTAPI_PASSWORD")
.env_remove("PATRONI_SUPERUSER_PASSWORD")
.env_remove("PGPASSWORD")
.env_remove("POSTGRES_PASSWORD")
.env_remove("PATRONI_RESTAPI_USERNAME")
.env_remove("PATRONI_SUPERUSER_USERNAME")
.env_remove("PGUSER")
.env_remove("POSTGRES_USER")
.envs(env.iter().copied())
.stdin(std::process::Stdio::piped())
.stdout(std::process::Stdio::piped())
.stderr(std::process::Stdio::piped())
.spawn()
.unwrap();
child
.stdin
.take()
.unwrap()
.write_all(
switchover_command(r#"{"leader":"postgres-1","candidate":"postgres-2"}"#)
.as_bytes(),
)
.unwrap();
let out = child.wait_with_output().unwrap();
let _ = std::fs::remove_dir_all(&dir);
assert!(
out.status.success(),
"shell rejected the switchover text: {}",
String::from_utf8_lossy(&out.stderr)
);
let stdout = String::from_utf8_lossy(&out.stdout).into_owned();
let (argv, config) = match stdout.split_once("@@CONFIG@@\n") {
Some((argv, config)) => (argv.to_string(), Some(config.to_string())),
None => (stdout, None),
};
CurlCall {
argv: argv.lines().map(str::to_string).collect(),
config,
}
}
#[cfg(unix)]
fn curl_config_value(line: &str) -> String {
let (_, quoted) = line.split_once('"').expect("a double-quoted value");
let mut out = String::new();
let mut chars = quoted.chars();
while let Some(c) = chars.next() {
match c {
'"' => break,
'\\' => match chars.next() {
Some('t') => out.push('\t'),
Some('n') => out.push('\n'),
Some('r') => out.push('\r'),
Some('v') => out.push('\u{0B}'),
Some(other) => out.push(other),
None => break,
},
c => out.push(c),
}
}
out
}
#[test]
fn switchover_command_reads_the_credential_by_name() {
let cmd = switchover_command(r#"{"leader":"postgres-1","candidate":"postgres-2"}"#);
assert!(cmd.contains("${PATRONI_RESTAPI_PASSWORD:-${PATRONI_SUPERUSER_PASSWORD:-"));
assert!(cmd.contains(
r#"PATRONI_REST_CFG="user = \"$(curl_cfg_quote "$PATRONI_REST_USER:$PATRONI_REST_PW")\""; set -- -K -"#
));
assert!(cmd.contains("else PATRONI_REST_CFG=; set --; fi"));
assert!(
!cmd.contains(" -u "),
"the credential must never be a curl argument"
);
let pipe = cmd
.find(r#"printf '%s\n' "$PATRONI_REST_CFG" | curl "#)
.expect("curl reads the config document from stdin");
let post = cmd.find("-X POST").expect("the POST survives");
assert!(pipe < post, "the credential must precede the request");
assert!(
cmd.contains(&format!("--max-time {SWITCHOVER_CURL_MAX_TIME_SECS}")),
"switchover curl budget too short: {cmd}"
);
assert!(SWITCHOVER_CURL_MAX_TIME_SECS >= 20);
assert!(SWITCHOVER_TIMEOUT.as_secs() > SWITCHOVER_CURL_MAX_TIME_SECS);
}
#[cfg(unix)]
#[test]
fn switchover_authenticates_from_the_members_own_env() {
let call = curl_call_for(&[
("PATRONI_RESTAPI_PASSWORD", "rest-pw"),
("PATRONI_SUPERUSER_PASSWORD", "super-pw"),
("POSTGRES_USER", "rw"),
]);
assert_eq!(call.config.as_deref(), Some("user = \"rw:rest-pw\"\n"));
assert!(
call.argv
.iter()
.all(|a| !a.contains("rest-pw") && !a.contains("super-pw")),
"credential in curl's argv: {:?}",
call.argv
);
assert!(!call.argv.iter().any(|a| a == "-u"), "{:?}", call.argv);
assert!(call.argv.iter().any(|a| a == "localhost:8008/switchover"));
assert!(
call.argv
.iter()
.any(|a| a == r#"{"leader":"postgres-1","candidate":"postgres-2"}"#)
);
}
#[cfg(unix)]
#[test]
fn switchover_falls_back_to_the_superuser_credential() {
let call = curl_call_for(&[("PATRONI_SUPERUSER_PASSWORD", "super-pw")]);
assert_eq!(
call.config.as_deref(),
Some("user = \"postgres:super-pw\"\n")
);
assert!(
call.argv.iter().all(|a| !a.contains("super-pw")),
"credential in curl's argv: {:?}",
call.argv
);
}
#[cfg(unix)]
#[test]
fn switchover_stays_bare_when_the_member_has_no_password() {
let call = curl_call_for(&[]);
assert!(
call.config.is_none(),
"config document without a password: {:?}",
call.config
);
assert!(
!call.argv.iter().any(|a| a == "-K" || a == "-u"),
"bare POST expected, got {:?}",
call.argv
);
assert!(call.argv.iter().any(|a| a == "localhost:8008/switchover"));
}
#[cfg(unix)]
#[test]
fn switchover_escapes_the_credential_for_curls_config_parser() {
let password = "p\"a\\s$s' w#rd\nnext\ttab";
let call = curl_call_for(&[
("PATRONI_RESTAPI_PASSWORD", password),
("POSTGRES_USER", "rw"),
]);
let config = call.config.expect("config document piped to curl");
assert_eq!(config, "user = \"rw:p\\\"a\\\\s$s' w#rd\\nnext\ttab\"\n");
assert_eq!(curl_config_value(&config), format!("rw:{password}"));
assert!(
call.argv.iter().all(|a| !a.contains("w#rd")),
"credential in curl's argv: {:?}",
call.argv
);
}
#[test]
fn parses_cluster_response_with_partial_fields() {
let raw = r#"{"members": [
{"name": "postgres-1", "role": "leader", "state": "running", "timeline": 3},
{"name": "postgres-replica-1", "role": "replica", "state": "streaming", "lag": 0}
]}"#;
let parsed: PatroniClusterResponse = serde_json::from_str(raw).unwrap();
assert_eq!(parsed.members.len(), 2);
assert_eq!(parsed.members[0].role, "leader");
assert_eq!(parsed.members[1].lag, Some(serde_json::json!(0)));
}
#[test]
fn parses_cluster_response_tolerates_missing_fields() {
let raw = r#"{"members": [{"name": "postgres-1"}]}"#;
let parsed: PatroniClusterResponse = serde_json::from_str(raw).unwrap();
assert_eq!(parsed.members.len(), 1);
assert_eq!(parsed.members[0].role, "");
assert!(parsed.members[0].lag.is_none());
}
#[test]
fn switchover_response_accepts_2xx_with_body() {
let ok =
parse_switchover_response("Successfully switched over to \"pg-2\"\nHTTP_STATUS:200")
.unwrap();
assert_eq!(ok, "Successfully switched over to \"pg-2\"");
}
#[test]
fn switchover_response_surfaces_patronis_rejection_body() {
let err = parse_switchover_response(
"candidate name does not match with the switchover candidate\nHTTP_STATUS:412",
)
.unwrap_err();
let message = err.to_string();
assert!(message.contains("412"));
assert!(message.contains("candidate name does not match"));
}
#[test]
fn switchover_response_rejects_missing_status_marker() {
let err = parse_switchover_response("curl: (7) connection refused").unwrap_err();
assert!(err.to_string().contains("unexpected response"));
}
#[test]
fn switchover_response_rejects_unparseable_status_code() {
assert!(parse_switchover_response("body\nHTTP_STATUS:abc").is_err());
}
#[test]
fn switchover_response_treats_curl_timeout_as_maybe_in_progress() {
let err = parse_switchover_response("HTTP_STATUS:000").unwrap_err();
let msg = format!("{err:#}");
assert!(msg.contains("may still be in progress"), "{msg}");
assert!(msg.contains("ha status"), "{msg}");
}
#[test]
fn switchover_exec_timeout_exit_is_maybe_in_progress() {
let err = map_switchover_exec_error(anyhow::anyhow!(
"SSH command failed (exit code 28): HTTP_STATUS:000"
))
.unwrap_err();
let msg = format!("{err:#}");
assert!(msg.contains("may still be in progress"), "{msg}");
assert!(!msg.contains("exit code 28"), "{msg}");
}
#[test]
fn switchover_exec_other_failures_pass_through() {
let err = map_switchover_exec_error(anyhow::anyhow!(
"SSH command failed (exit code 127): sh: curl: command not found"
))
.unwrap_err();
assert!(format!("{err:#}").contains("command not found"));
}
}