use std::{collections::BTreeMap, time::Duration};
use anyhow::{Context, Result, bail};
use serde::Deserialize;
use super::exec::exec_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 = tokio::time::timeout(PROBE_TIMEOUT, exec_in_container(instance_id, command))
.await
.context("Timed out probing Patroni")??;
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]) -> Option<(String, Vec<PatroniMember>)> {
for instance_id in instance_ids {
if let Ok(members) = probe_cluster(instance_id).await {
return Some((instance_id.clone(), members));
}
}
None
}
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 = format!(
"curl -s --max-time 8 -w '\\nHTTP_STATUS:%{{http_code}}' -X POST localhost:8008/switchover -H 'Content-Type: application/json' -d '{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::*;
#[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());
}
}