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);
#[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#"if [ -n "$PATRONI_REST_PW" ]; then set -- -u "$PATRONI_REST_USER:$PATRONI_REST_PW"; else set --; fi; "#,
);
fn switchover_command(body: &str) -> String {
format!(
"{prelude}curl -s --max-time 8 -w '\\nHTTP_STATUS:%{{http_code}}' \"$@\" -X POST localhost:8008/switchover -H 'Content-Type: application/json' -d '{body}'",
prelude = RESTAPI_AUTH_PRELUDE,
)
}
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 = tokio::time::timeout(
Duration::from_secs(10),
exec_in_container(instance_id, &command),
)
.await
.context("Timed out requesting switchover")??;
parse_switchover_response(&output)
}
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(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)]
fn curl_argv_for(env: &[(&str, &str)]) -> Vec<String> {
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, "#!/bin/sh\nprintf '%s\\n' \"$@\"\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)
);
String::from_utf8_lossy(&out.stdout)
.lines()
.map(str::to_string)
.collect()
}
#[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#"set -- -u "$PATRONI_REST_USER:$PATRONI_REST_PW""#));
assert!(cmd.contains("else set --; fi"));
let at = cmd
.find(r#""$@""#)
.expect("curl carries the resolved credential");
let post = cmd.find("-X POST").expect("the POST survives");
assert!(at < post, "the credential must precede the request");
}
#[cfg(unix)]
#[test]
fn switchover_authenticates_from_the_members_own_env() {
let argv = curl_argv_for(&[
("PATRONI_RESTAPI_PASSWORD", "rest-pw"),
("PATRONI_SUPERUSER_PASSWORD", "super-pw"),
("POSTGRES_USER", "rw"),
]);
let u = argv
.iter()
.position(|a| a == "-u")
.expect("credential passed to curl");
assert_eq!(argv[u + 1], "rw:rest-pw");
assert!(argv.iter().any(|a| a == "localhost:8008/switchover"));
assert!(
argv.iter()
.any(|a| a == r#"{"leader":"postgres-1","candidate":"postgres-2"}"#)
);
}
#[cfg(unix)]
#[test]
fn switchover_falls_back_to_the_superuser_credential() {
let argv = curl_argv_for(&[("PATRONI_SUPERUSER_PASSWORD", "super-pw")]);
let u = argv
.iter()
.position(|a| a == "-u")
.expect("credential passed to curl");
assert_eq!(argv[u + 1], "postgres:super-pw");
}
#[cfg(unix)]
#[test]
fn switchover_stays_bare_when_the_member_has_no_password() {
let argv = curl_argv_for(&[]);
assert!(
!argv.iter().any(|a| a == "-u"),
"bare POST expected, got {argv:?}"
);
assert!(argv.iter().any(|a| a == "localhost:8008/switchover"));
}
#[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());
}
}