use std::collections::{BTreeMap, HashMap, HashSet};
use std::path::{Path, PathBuf};
use std::sync::{Mutex, PoisonError};
use anyhow::{bail, Context, Result};
use serde::{Deserialize, Serialize};
use serde_json::Value;
const GH_BIN_ENV: &str = "OMNI_DEV_GH_BIN";
const GH_BINARY_CANDIDATES: &[&str] = &[
"/opt/homebrew/bin/gh",
"/usr/local/bin/gh",
"/home/linuxbrew/.linuxbrew/bin/gh",
"/usr/bin/gh",
];
const FAILURE_STATES: &[&str] = &[
"FAILURE",
"ERROR",
"CANCELLED",
"TIMED_OUT",
"ACTION_REQUIRED",
"STARTUP_FAILURE",
"STALE",
];
const SUCCESS_STATES: &[&str] = &["SUCCESS", "NEUTRAL", "SKIPPED"];
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
#[serde(rename_all = "lowercase")]
pub enum PrCheckState {
Success,
Failure,
Pending,
None,
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize)]
pub struct PrBadge {
pub number: u64,
#[serde(rename = "isDraft")]
pub is_draft: bool,
pub checks: PrCheckState,
pub url: String,
#[serde(skip)]
pub head_oid: String,
}
impl PrBadge {
#[must_use]
pub fn is_stale_for(&self, head_sha: Option<&str>) -> bool {
head_sha.is_some_and(|sha| sha != self.head_oid)
}
}
#[derive(Debug, Clone, PartialEq, Eq, Hash, PartialOrd, Ord, Serialize, Deserialize)]
pub struct PrTarget {
pub owner: String,
pub name: String,
pub branch: String,
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum PrResolution {
Pr(PrBadge),
NoPr,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
enum EntryState {
Failure,
Pending,
Success,
}
fn check_entry_state(entry: &Value) -> EntryState {
let status = entry.get("status").and_then(Value::as_str).unwrap_or("");
if !status.is_empty() && !status.eq_ignore_ascii_case("COMPLETED") {
return EntryState::Pending;
}
let raw = entry
.get("conclusion")
.and_then(Value::as_str)
.filter(|s| !s.is_empty())
.or_else(|| {
entry
.get("state")
.and_then(Value::as_str)
.filter(|s| !s.is_empty())
})
.unwrap_or("")
.to_ascii_uppercase();
if FAILURE_STATES.contains(&raw.as_str()) {
return EntryState::Failure;
}
if SUCCESS_STATES.contains(&raw.as_str()) {
return EntryState::Success;
}
EntryState::Pending
}
fn rollup_check_state(contexts: &[Value]) -> PrCheckState {
if contexts.is_empty() {
return PrCheckState::None;
}
let mut saw_pending = false;
for entry in contexts {
match check_entry_state(entry) {
EntryState::Failure => return PrCheckState::Failure,
EntryState::Pending => saw_pending = true,
EntryState::Success => {}
}
}
if saw_pending {
return PrCheckState::Pending;
}
if contexts.iter().any(suite_still_running) {
return PrCheckState::Pending;
}
PrCheckState::Success
}
fn suite_still_running(entry: &Value) -> bool {
entry
.get("checkSuite")
.and_then(|suite| suite.get("status"))
.and_then(Value::as_str)
.is_some_and(|status| !status.is_empty() && !status.eq_ignore_ascii_case("COMPLETED"))
}
const ROLLUP_FRAGMENT: &str = "statusCheckRollup{ contexts(first:100){ nodes{
__typename
...on CheckRun{ status conclusion checkSuite{ status } }
...on StatusContext{ state }
} } }";
fn branch_fragment(alias: &str, branch: &str) -> String {
let qualified = Value::String(format!("refs/heads/{branch}"));
format!(
r"{alias}: ref(qualifiedName:{qualified}){{
target{{ ...on Commit{{ oid
{ROLLUP_FRAGMENT}
}} }}
associatedPullRequests(first:1, states:OPEN){{ nodes{{ number isDraft url }} }}
}}"
)
}
fn merge_branch_fragment(alias: &str, branch: &str) -> String {
let qualified = Value::String(format!("refs/heads/{branch}"));
format!(
r"{alias}: ref(qualifiedName:{qualified}){{
target{{ ...on Commit{{ oid
{ROLLUP_FRAGMENT}
}} }}
associatedPullRequests(first:1, states:OPEN){{ nodes{{
id number isDraft url headRefOid mergeStateStatus
mergeQueueEntry{{ state }}
}} }}
}}"
)
}
type QueryIndex = HashMap<(usize, usize), PrTarget>;
fn build_query(targets: &[PrTarget]) -> Option<(String, QueryIndex)> {
if targets.is_empty() {
return None;
}
let mut by_repo: BTreeMap<(&str, &str), Vec<&PrTarget>> = BTreeMap::new();
for t in targets {
by_repo
.entry((t.owner.as_str(), t.name.as_str()))
.or_default()
.push(t);
}
let mut index = HashMap::new();
let mut repos = Vec::new();
for (ri, ((owner, name), branches)) in by_repo.iter().enumerate() {
let mut frags = Vec::new();
for (bi, target) in branches.iter().enumerate() {
frags.push(branch_fragment(&format!("b{bi}"), &target.branch));
index.insert((ri, bi), (*target).clone());
}
let owner = Value::String((*owner).to_string());
let name = Value::String((*name).to_string());
repos.push(format!(
"r{ri}: repository(owner:{owner}, name:{name}){{\n{}\n}}",
frags.join("\n")
));
}
Some((
format!(
"query{{\nrateLimit{{ limit cost remaining used resetAt }}\n{}\n}}",
repos.join("\n")
),
index,
))
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct QueryRateLimit {
pub cost: u64,
pub used: u64,
pub limit: u64,
pub remaining: u64,
pub reset: i64,
}
fn parse_rate_limit(body: &Value) -> Option<QueryRateLimit> {
let rl = body.get("data")?.get("rateLimit")?;
let reset = rl
.get("resetAt")
.and_then(Value::as_str)
.and_then(|s| chrono::DateTime::parse_from_rfc3339(s).ok())
.map_or(0, |dt| dt.timestamp());
Some(QueryRateLimit {
cost: rl.get("cost")?.as_u64()?,
used: rl.get("used")?.as_u64()?,
limit: rl.get("limit")?.as_u64()?,
remaining: rl.get("remaining")?.as_u64()?,
reset,
})
}
fn badge_from_ref(node: &Value) -> Option<PrBadge> {
let pr = node
.get("associatedPullRequests")?
.get("nodes")?
.as_array()?
.first()?;
let contexts = node
.get("target")
.and_then(|t| t.get("statusCheckRollup"))
.and_then(|r| r.get("contexts"))
.and_then(|c| c.get("nodes"))
.and_then(Value::as_array)
.cloned()
.unwrap_or_default();
Some(PrBadge {
number: pr.get("number").and_then(Value::as_u64)?,
is_draft: pr.get("isDraft").and_then(Value::as_bool).unwrap_or(false),
checks: rollup_check_state(&contexts),
url: pr
.get("url")
.and_then(Value::as_str)
.unwrap_or_default()
.to_string(),
head_oid: node
.get("target")
.and_then(|t| t.get("oid"))
.and_then(Value::as_str)
.unwrap_or_default()
.to_string(),
})
}
fn resolution_from_ref(node: &Value) -> Option<PrResolution> {
let prs = node
.get("associatedPullRequests")?
.get("nodes")?
.as_array()?;
if prs.is_empty() {
return Some(PrResolution::NoPr);
}
badge_from_ref(node).map(PrResolution::Pr)
}
fn parse_response(body: &Value, index: &QueryIndex) -> HashMap<PrTarget, PrResolution> {
let mut out = HashMap::new();
let Some(data) = body.get("data") else {
return out;
};
for ((ri, bi), target) in index {
let Some(repo) = data.get(format!("r{ri}")).filter(|r| !r.is_null()) else {
continue;
};
let Some(node) = repo.get(format!("b{bi}")) else {
continue;
};
if node.is_null() {
out.insert(target.clone(), PrResolution::NoPr);
continue;
}
if let Some(resolution) = resolution_from_ref(node) {
out.insert(target.clone(), resolution);
}
}
out
}
#[must_use]
pub fn resolve_gh_binary() -> PathBuf {
resolve_gh_binary_from(std::env::var_os(GH_BIN_ENV), GH_BINARY_CANDIDATES)
}
fn resolve_gh_binary_from(
env_override: Option<std::ffi::OsString>,
candidates: &[&str],
) -> PathBuf {
if let Some(path) = env_override.filter(|p| !p.is_empty()) {
return PathBuf::from(path);
}
for candidate in candidates {
let path = Path::new(candidate);
if path.exists() {
return path.to_path_buf();
}
}
PathBuf::from("gh")
}
fn run_gh_graphql(bin: &Path, query: &str) -> Result<Value> {
let query_arg = format!("query={query}");
let output = crate::github_metrics::run_gh(
bin,
["api", "graphql", "-f", query_arg.as_str()],
"api graphql",
None,
)
.with_context(|| {
format!(
"failed to run {} (is the GitHub CLI installed?)",
bin.display()
)
})?;
if !output.status.success() {
let stderr = String::from_utf8_lossy(&output.stderr);
bail!("gh api graphql failed: {}", stderr.trim());
}
serde_json::from_slice(&output.stdout).context("gh api graphql returned invalid JSON")
}
pub fn resolve_with(bin: &Path, targets: &[PrTarget]) -> Result<HashMap<PrTarget, PrResolution>> {
resolve_inner(bin, targets).map(|(resolutions, _budget)| resolutions)
}
pub fn resolve_with_budget(
bin: &Path,
targets: &[PrTarget],
) -> Result<(HashMap<PrTarget, PrResolution>, Option<QueryRateLimit>)> {
resolve_inner(bin, targets)
}
fn resolve_inner(
bin: &Path,
targets: &[PrTarget],
) -> Result<(HashMap<PrTarget, PrResolution>, Option<QueryRateLimit>)> {
let Some((query, index)) = build_query(targets) else {
return Ok((HashMap::new(), None));
};
let body = run_gh_graphql(bin, &query)?;
if let Some(errors) = body.get("errors").and_then(Value::as_array) {
if !errors.is_empty() {
bail!("gh api graphql returned errors: {errors:?}");
}
}
Ok((parse_response(&body, &index), parse_rate_limit(&body)))
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct MergeInfo {
pub pr_id: String,
pub number: u64,
pub url: String,
pub is_draft: bool,
pub head_oid: String,
pub merge_state: Option<String>,
pub already_queued: bool,
pub checks: PrCheckState,
}
fn build_merge_query(targets: &[PrTarget]) -> Option<(String, QueryIndex)> {
if targets.is_empty() {
return None;
}
let mut by_repo: BTreeMap<(&str, &str), Vec<&PrTarget>> = BTreeMap::new();
for t in targets {
by_repo
.entry((t.owner.as_str(), t.name.as_str()))
.or_default()
.push(t);
}
let mut index = HashMap::new();
let mut repos = Vec::new();
for (ri, ((owner, name), branches)) in by_repo.iter().enumerate() {
let mut frags = Vec::new();
for (bi, target) in branches.iter().enumerate() {
frags.push(merge_branch_fragment(&format!("b{bi}"), &target.branch));
index.insert((ri, bi), (*target).clone());
}
let owner = Value::String((*owner).to_string());
let name = Value::String((*name).to_string());
repos.push(format!(
"r{ri}: repository(owner:{owner}, name:{name}){{\n{}\n}}",
frags.join("\n")
));
}
Some((format!("query{{\n{}\n}}", repos.join("\n")), index))
}
fn merge_info_from_ref(node: &Value) -> Option<MergeInfo> {
let pr = node
.get("associatedPullRequests")?
.get("nodes")?
.as_array()?
.first()?;
let contexts = node
.get("target")
.and_then(|t| t.get("statusCheckRollup"))
.and_then(|r| r.get("contexts"))
.and_then(|c| c.get("nodes"))
.and_then(Value::as_array)
.cloned()
.unwrap_or_default();
Some(MergeInfo {
pr_id: pr.get("id").and_then(Value::as_str)?.to_string(),
number: pr.get("number").and_then(Value::as_u64)?,
url: pr
.get("url")
.and_then(Value::as_str)
.unwrap_or_default()
.to_string(),
is_draft: pr.get("isDraft").and_then(Value::as_bool).unwrap_or(false),
head_oid: pr
.get("headRefOid")
.and_then(Value::as_str)
.unwrap_or_default()
.to_string(),
merge_state: pr
.get("mergeStateStatus")
.and_then(Value::as_str)
.map(str::to_ascii_uppercase),
already_queued: pr
.get("mergeQueueEntry")
.is_some_and(|entry| !entry.is_null()),
checks: rollup_check_state(&contexts),
})
}
fn parse_merge_response(body: &Value, index: &QueryIndex) -> HashMap<PrTarget, MergeInfo> {
let mut out = HashMap::new();
let Some(data) = body.get("data") else {
return out;
};
for ((ri, bi), target) in index {
let Some(repo) = data.get(format!("r{ri}")).filter(|r| !r.is_null()) else {
continue;
};
let Some(node) = repo.get(format!("b{bi}")).filter(|n| !n.is_null()) else {
continue;
};
if let Some(info) = merge_info_from_ref(node) {
out.insert(target.clone(), info);
}
}
out
}
pub fn resolve_merge_targets(
bin: &Path,
targets: &[PrTarget],
) -> Result<HashMap<PrTarget, MergeInfo>> {
let Some((query, index)) = build_merge_query(targets) else {
return Ok(HashMap::new());
};
let body = run_gh_graphql(bin, &query)?;
if let Some(errors) = body.get("errors").and_then(Value::as_array) {
if !errors.is_empty() {
bail!("gh api graphql returned errors: {errors:?}");
}
}
Ok(parse_merge_response(&body, &index))
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum EnqueueOutcome {
Queued(Option<String>),
Rejected(String),
}
fn enqueue_error_message(errors: &[Value]) -> String {
let joined = errors
.iter()
.filter_map(|e| e.get("message").and_then(Value::as_str))
.collect::<Vec<_>>()
.join("; ");
if joined.is_empty() {
"enqueue rejected by GitHub".to_string()
} else {
joined
}
}
pub fn enqueue_pull_request(bin: &Path, pr_node_id: &str) -> Result<EnqueueOutcome> {
let id_lit = Value::String(pr_node_id.to_string());
let mutation =
format!("mutation{{ enqueuePullRequest(input:{{pullRequestId:{id_lit}}}){{ mergeQueueEntry{{ state }} }} }}");
let query_arg = format!("query={mutation}");
let output = crate::github_metrics::run_gh(
bin,
["api", "graphql", "-f", query_arg.as_str()],
"api graphql",
None,
)
.with_context(|| {
format!(
"failed to run {} (is the GitHub CLI installed?)",
bin.display()
)
})?;
if !output.status.success() {
let msg = serde_json::from_slice::<Value>(&output.stdout)
.ok()
.and_then(|body| {
body.get("errors")
.and_then(Value::as_array)
.filter(|errs| !errs.is_empty())
.map(|errs| enqueue_error_message(errs))
})
.unwrap_or_else(|| String::from_utf8_lossy(&output.stderr).trim().to_string());
return Ok(EnqueueOutcome::Rejected(msg));
}
let body: Value =
serde_json::from_slice(&output.stdout).context("gh api graphql returned invalid JSON")?;
if let Some(errors) = body.get("errors").and_then(Value::as_array) {
if !errors.is_empty() {
return Ok(EnqueueOutcome::Rejected(enqueue_error_message(errors)));
}
}
let state = body
.pointer("/data/enqueuePullRequest/mergeQueueEntry/state")
.and_then(Value::as_str)
.map(str::to_string);
Ok(EnqueueOutcome::Queued(state))
}
#[derive(Debug, Default)]
pub struct PrStatusCache {
resolutions: Mutex<HashMap<PrTarget, PrResolution>>,
}
impl PrStatusCache {
#[must_use]
pub fn new() -> Self {
Self::default()
}
#[must_use]
pub fn get(&self, owner: &str, name: &str, branch: &str) -> Option<PrResolution> {
let key = PrTarget {
owner: owner.to_string(),
name: name.to_string(),
branch: branch.to_string(),
};
self.lock().get(&key).cloned()
}
pub fn replace(&self, next: HashMap<PrTarget, PrResolution>) -> bool {
let mut guard = self.lock();
if *guard == next {
return false;
}
*guard = next;
true
}
pub fn retain_targets(&self, keep: &HashSet<PrTarget>) -> bool {
let mut guard = self.lock();
let before = guard.len();
guard.retain(|target, _| keep.contains(target));
guard.len() != before
}
#[must_use]
pub fn entries(&self) -> Vec<(PrTarget, PrResolution)> {
self.lock()
.iter()
.map(|(t, r)| (t.clone(), r.clone()))
.collect()
}
pub fn seed(&self, entries: impl IntoIterator<Item = (PrTarget, PrResolution)>) {
let mut guard = self.lock();
for (target, resolution) in entries {
guard.insert(target, resolution);
}
}
#[must_use]
pub fn any_pending(&self) -> bool {
self.lock()
.values()
.any(|r| matches!(r, PrResolution::Pr(b) if b.checks == PrCheckState::Pending))
}
fn lock(&self) -> std::sync::MutexGuard<'_, HashMap<PrTarget, PrResolution>> {
self.resolutions
.lock()
.unwrap_or_else(PoisonError::into_inner)
}
}
#[cfg(test)]
#[allow(clippy::unwrap_used, clippy::expect_used)]
mod tests {
use super::*;
use crate::test_support::shim::{retry_on_etxtbsy, shim_lock, write_exec_script};
use serde_json::json;
use std::sync::MutexGuard;
fn target(branch: &str) -> PrTarget {
PrTarget {
owner: "rust-works".into(),
name: "omni-dev".into(),
branch: branch.into(),
}
}
#[test]
fn rollup_is_none_for_an_empty_rollup() {
assert_eq!(rollup_check_state(&[]), PrCheckState::None);
}
#[test]
fn rollup_reads_completed_check_run_conclusions() {
for (conclusion, want) in [
("SUCCESS", PrCheckState::Success),
("NEUTRAL", PrCheckState::Success),
("SKIPPED", PrCheckState::Success),
("FAILURE", PrCheckState::Failure),
("CANCELLED", PrCheckState::Failure),
("TIMED_OUT", PrCheckState::Failure),
("ACTION_REQUIRED", PrCheckState::Failure),
("STARTUP_FAILURE", PrCheckState::Failure),
("STALE", PrCheckState::Failure),
] {
let entry =
json!({"__typename":"CheckRun","status":"COMPLETED","conclusion":conclusion});
assert_eq!(
rollup_check_state(&[entry]),
want,
"conclusion {conclusion} should be {want:?}"
);
}
}
#[test]
fn rollup_treats_an_incomplete_check_run_as_pending() {
for status in ["IN_PROGRESS", "QUEUED", "WAITING", "PENDING"] {
let entry = json!({"__typename":"CheckRun","status":status,"conclusion":null});
assert_eq!(rollup_check_state(&[entry]), PrCheckState::Pending);
}
}
#[test]
fn rollup_reads_status_context_states() {
for (state, want) in [
("SUCCESS", PrCheckState::Success),
("FAILURE", PrCheckState::Failure),
("ERROR", PrCheckState::Failure),
("PENDING", PrCheckState::Pending),
("EXPECTED", PrCheckState::Pending),
] {
let entry = json!({"__typename":"StatusContext","state":state});
assert_eq!(rollup_check_state(&[entry]), want, "state {state}");
}
}
#[test]
fn rollup_never_reads_an_unknown_value_as_a_pass() {
let entry = json!({"__typename":"CheckRun","status":"COMPLETED","conclusion":"WAT"});
assert_eq!(rollup_check_state(&[entry]), PrCheckState::Pending);
let entry = json!({"__typename":"CheckRun","status":"COMPLETED","conclusion":""});
assert_eq!(rollup_check_state(&[entry]), PrCheckState::Pending);
}
#[test]
fn rollup_precedence_is_failure_then_pending_then_success() {
let ok = json!({"__typename":"CheckRun","status":"COMPLETED","conclusion":"SUCCESS"});
let bad = json!({"__typename":"CheckRun","status":"COMPLETED","conclusion":"FAILURE"});
let run = json!({"__typename":"CheckRun","status":"IN_PROGRESS","conclusion":null});
assert_eq!(
rollup_check_state(&[ok.clone(), run.clone(), bad.clone()]),
PrCheckState::Failure
);
assert_eq!(
rollup_check_state(&[bad, ok.clone()]),
PrCheckState::Failure
);
assert_eq!(
rollup_check_state(&[ok.clone(), run]),
PrCheckState::Pending
);
assert_eq!(rollup_check_state(&[ok]), PrCheckState::Success);
}
fn run(conclusion: &str, suite: Option<&str>) -> Value {
let mut e = json!({
"__typename": "CheckRun",
"status": "COMPLETED",
"conclusion": conclusion,
});
if let Some(s) = suite {
e["checkSuite"] = json!({ "status": s });
}
e
}
#[test]
fn rollup_is_pending_while_a_suite_is_still_creating_jobs() {
let contexts = vec![
run("SUCCESS", Some("IN_PROGRESS")),
run("SUCCESS", Some("IN_PROGRESS")),
];
assert_eq!(rollup_check_state(&contexts), PrCheckState::Pending);
assert_eq!(
rollup_check_state(&[run("SUCCESS", Some("QUEUED"))]),
PrCheckState::Pending
);
}
#[test]
fn rollup_is_success_once_every_backing_suite_is_terminal() {
let contexts = vec![
run("SUCCESS", Some("COMPLETED")),
run("SKIPPED", Some("COMPLETED")),
];
assert_eq!(rollup_check_state(&contexts), PrCheckState::Success);
}
#[test]
fn rollup_ignores_a_zero_run_zombie_suite() {
let contexts = vec![run("SUCCESS", Some("COMPLETED"))];
assert_eq!(rollup_check_state(&contexts), PrCheckState::Success);
}
#[test]
fn rollup_tolerates_entries_without_suite_information() {
assert_eq!(
rollup_check_state(&[json!({"__typename":"StatusContext","state":"SUCCESS"})]),
PrCheckState::Success
);
assert_eq!(
rollup_check_state(&[run("SUCCESS", None)]),
PrCheckState::Success
);
assert_eq!(
rollup_check_state(&[run("SUCCESS", Some(""))]),
PrCheckState::Success
);
}
#[test]
fn rollup_failure_still_dominates_a_running_suite() {
let contexts = vec![
run("FAILURE", Some("IN_PROGRESS")),
run("SUCCESS", Some("IN_PROGRESS")),
];
assert_eq!(rollup_check_state(&contexts), PrCheckState::Failure);
}
#[test]
fn build_query_asks_for_the_backing_suite_status() {
let (query, _) = build_query(&[target("main")]).unwrap();
assert!(query.contains("checkSuite{ status }"), "{query}");
}
#[test]
fn build_query_folds_in_the_rate_limit_block() {
let (query, _) = build_query(&[target("main")]).unwrap();
assert!(
query.contains("rateLimit{ limit cost remaining used resetAt }"),
"{query}"
);
}
#[test]
fn parse_rate_limit_reads_the_folded_budget_block() {
let body = json!({"data":{"rateLimit":
{"limit":5000,"cost":1,"remaining":4970,"used":30,"resetAt":"2026-07-21T16:00:00Z"}}});
let rl = parse_rate_limit(&body).expect("a complete block should parse");
assert_eq!(rl.cost, 1);
assert_eq!(rl.used, 30);
assert_eq!(rl.limit, 5000);
assert_eq!(rl.remaining, 4970);
assert!(rl.reset > 0, "resetAt should parse to an epoch");
assert!(parse_rate_limit(&json!({"data":{"r0":{}}})).is_none());
assert!(parse_rate_limit(
&json!({"data":{"rateLimit":{"limit":5000,"cost":1,"remaining":4970}}})
)
.is_none());
}
#[test]
fn build_query_is_none_without_targets() {
assert!(build_query(&[]).is_none());
}
#[test]
fn build_query_groups_branches_under_one_repo_alias() {
let (query, index) = build_query(&[target("main"), target("feature")]).unwrap();
assert_eq!(query.matches("repository(owner:").count(), 1);
assert_eq!(query.matches(": ref(qualifiedName:").count(), 2);
assert!(query.contains(r#""refs/heads/main""#), "{query}");
assert!(query.contains(r#""refs/heads/feature""#), "{query}");
assert_eq!(index.len(), 2);
}
#[test]
fn build_query_reads_open_prs_off_the_ref_not_the_commit() {
let (query, _) = build_query(&[target("main")]).unwrap();
assert!(
query.contains("associatedPullRequests(first:1, states:OPEN)"),
"{query}"
);
let commit_block = query.find("...on Commit").unwrap();
let pr_lookup = query.find("associatedPullRequests").unwrap();
assert!(
pr_lookup > commit_block,
"PR lookup must be on the Ref, after the Commit block"
);
}
#[test]
fn build_query_separates_distinct_repos() {
let a = PrTarget {
owner: "o1".into(),
name: "r1".into(),
branch: "main".into(),
};
let b = PrTarget {
owner: "o2".into(),
name: "r2".into(),
branch: "main".into(),
};
let (query, index) = build_query(&[a, b]).unwrap();
assert_eq!(query.matches("repository(owner:").count(), 2);
assert_eq!(index.len(), 2);
}
#[test]
fn build_query_escapes_branch_names() {
let (query, _) = build_query(&[target(r#"we"ird"#)]).unwrap();
assert!(query.contains(r#"refs/heads/we\"ird"#), "{query}");
}
#[test]
fn parse_response_reads_badges_and_resolves_absent_refs_as_negatives() {
let targets = vec![target("feature"), target("unpushed"), target("no-pr")];
let (_, index) = build_query(&targets).unwrap();
let body = json!({"data":{"r0":{
"b0": {
"target": {"oid":"abc","statusCheckRollup":{"contexts":{"nodes":[
{"__typename":"CheckRun","status":"COMPLETED","conclusion":"SUCCESS"}
]}}},
"associatedPullRequests":{"nodes":[{"number":65,"isDraft":true,"url":"u"}]}
},
"b1": null,
"b2": {
"target": {"oid":"def","statusCheckRollup":null},
"associatedPullRequests":{"nodes":[]}
}
}}});
let out = parse_response(&body, &index);
assert_eq!(out.len(), 3, "{out:?}");
let Some(PrResolution::Pr(badge)) = out.get(&target("feature")) else {
panic!("expected a badge for feature: {out:?}");
};
assert_eq!(badge.number, 65);
assert!(badge.is_draft);
assert_eq!(badge.checks, PrCheckState::Success);
assert_eq!(badge.url, "u");
assert_eq!(out.get(&target("unpushed")), Some(&PrResolution::NoPr));
assert_eq!(out.get(&target("no-pr")), Some(&PrResolution::NoPr));
}
#[test]
fn parse_response_reports_a_pushed_branch_with_no_pr_as_a_negative() {
let (_, index) = build_query(&[target("quiet")]).unwrap();
let body = json!({"data":{"r0":{"b0":{
"target": {"oid":"abc","statusCheckRollup":null},
"associatedPullRequests":{"nodes":[]}
}}}});
let out = parse_response(&body, &index);
assert_eq!(out.get(&target("quiet")), Some(&PrResolution::NoPr));
}
#[test]
fn parse_response_reports_an_unpushed_branch_as_a_negative() {
let (_, index) = build_query(&[target("local-only")]).unwrap();
let body = json!({"data":{"r0":{"b0":null}}});
let out = parse_response(&body, &index);
assert_eq!(out.get(&target("local-only")), Some(&PrResolution::NoPr));
}
#[test]
fn parse_response_leaves_a_missing_alias_unresolved_rather_than_negative() {
let (_, index) = build_query(&[target("x")]).unwrap();
assert!(parse_response(&json!({"data":{"r0":{}}}), &index).is_empty());
assert!(parse_response(&json!({"data":{"r0":null}}), &index).is_empty());
}
#[test]
fn parse_response_leaves_a_malformed_pr_node_unresolved_rather_than_negative() {
let (_, index) = build_query(&[target("x")]).unwrap();
let body = json!({"data":{"r0":{"b0":{
"target": {"oid":"abc","statusCheckRollup":null},
"associatedPullRequests":{"nodes":[{"isDraft":false,"url":"u"}]}
}}}});
assert!(parse_response(&body, &index).is_empty());
}
#[test]
fn parse_response_reads_a_pr_with_no_checks_as_none() {
let targets = vec![target("feature")];
let (_, index) = build_query(&targets).unwrap();
let body = json!({"data":{"r0":{"b0":{
"target": {"oid":"abc","statusCheckRollup":null},
"associatedPullRequests":{"nodes":[{"number":7,"isDraft":false,"url":"u"}]}
}}}});
let out = parse_response(&body, &index);
let Some(PrResolution::Pr(badge)) = out.get(&target("feature")) else {
panic!("expected a badge: {out:?}");
};
assert_eq!(badge.checks, PrCheckState::None);
}
#[test]
fn parse_response_tolerates_a_missing_data_block() {
let (_, index) = build_query(&[target("x")]).unwrap();
assert!(parse_response(&json!({}), &index).is_empty());
}
#[test]
fn badge_serializes_is_draft_as_camel_case() {
let badge = PrBadge {
number: 65,
is_draft: true,
checks: PrCheckState::Pending,
url: "u".into(),
head_oid: String::new(),
};
let v = serde_json::to_value(&badge).unwrap();
assert_eq!(v["isDraft"], json!(true));
assert_eq!(v["checks"], json!("pending"));
assert!(v.get("is_draft").is_none(), "{v}");
}
#[test]
fn check_state_serializes_lowercase() {
for (state, want) in [
(PrCheckState::Success, "success"),
(PrCheckState::Failure, "failure"),
(PrCheckState::Pending, "pending"),
(PrCheckState::None, "none"),
] {
assert_eq!(serde_json::to_value(state).unwrap(), json!(want));
}
}
#[test]
fn resolve_gh_binary_from_prefers_env_then_candidate_then_fallback() {
assert_eq!(
resolve_gh_binary_from(Some("/custom/gh".into()), &["/usr/bin/gh"]),
PathBuf::from("/custom/gh")
);
let existing = tempfile::NamedTempFile::new().unwrap();
let existing_path = existing.path().to_str().unwrap();
assert_eq!(
resolve_gh_binary_from(None, &["/no/such/gh/xyzzy", existing_path]),
PathBuf::from(existing_path)
);
assert_eq!(
resolve_gh_binary_from(None, &["/no/such/gh/xyzzy"]),
PathBuf::from("gh")
);
assert_eq!(
resolve_gh_binary_from(Some("".into()), &["/no/such/gh/xyzzy"]),
PathBuf::from("gh")
);
let _ = resolve_gh_binary();
}
fn fake_gh(dir: &Path, stdout: &str, code: i32) -> (PathBuf, MutexGuard<'static, ()>) {
let guard = shim_lock();
let path = dir.join("fake-gh");
write_exec_script(
&path,
&format!("#!/bin/sh\ncat <<'JSON'\n{stdout}\nJSON\nexit {code}\n"),
);
(path, guard)
}
#[test]
fn resolve_with_asks_nothing_for_no_targets() {
let out = resolve_with(Path::new("/no/such/gh/xyzzy"), &[]).unwrap();
assert!(out.is_empty());
}
#[test]
fn resolve_with_errors_when_gh_is_missing() {
let err = resolve_with(Path::new("/no/such/gh/xyzzy"), &[target("main")]).unwrap_err();
let msg = format!("{err:#}");
assert!(msg.contains("failed to run"), "{msg}");
assert!(msg.contains("GitHub CLI"), "{msg}");
}
#[test]
fn resolve_with_errors_on_a_nonzero_exit() {
let dir = tempfile::tempdir().unwrap();
let (bin, _shim) = fake_gh(dir.path(), "", 1);
let err = retry_on_etxtbsy(|| resolve_with(&bin, &[target("main")])).unwrap_err();
assert!(
format!("{err:#}").contains("gh api graphql failed"),
"{err:#}"
);
}
#[test]
fn resolve_with_errors_on_unparseable_output() {
let dir = tempfile::tempdir().unwrap();
let (bin, _shim) = fake_gh(dir.path(), "not json at all", 0);
let err = retry_on_etxtbsy(|| resolve_with(&bin, &[target("main")])).unwrap_err();
assert!(format!("{err:#}").contains("invalid JSON"), "{err:#}");
}
#[test]
fn resolve_with_surfaces_graphql_errors_rather_than_reporting_no_badges() {
let dir = tempfile::tempdir().unwrap();
let (bin, _shim) = fake_gh(
dir.path(),
r#"{"data":null,"errors":[{"message":"API rate limit exceeded"}]}"#,
0,
);
let err = retry_on_etxtbsy(|| resolve_with(&bin, &[target("main")])).unwrap_err();
let msg = format!("{err:#}");
assert!(msg.contains("returned errors"), "{msg}");
assert!(msg.contains("rate limit"), "{msg}");
}
#[test]
fn resolve_with_ignores_an_empty_errors_array() {
let dir = tempfile::tempdir().unwrap();
let (bin, _shim) = fake_gh(
dir.path(),
r#"{"errors":[],"data":{"r0":{"b0":{
"target":{"oid":"a","statusCheckRollup":null},
"associatedPullRequests":{"nodes":[{"number":9,"isDraft":false,"url":"u"}]}}}}}"#,
0,
);
let out = retry_on_etxtbsy(|| resolve_with(&bin, &[target("main")])).unwrap();
let Some(PrResolution::Pr(badge)) = out.get(&target("main")) else {
panic!("expected a badge: {out:?}");
};
assert_eq!(badge.number, 9);
}
#[test]
fn resolve_with_reads_a_real_reply_end_to_end() {
let dir = tempfile::tempdir().unwrap();
let (bin, _shim) = fake_gh(
dir.path(),
r#"{"data":{"r0":{"b0":{
"target":{"oid":"a","statusCheckRollup":{"contexts":{"nodes":[
{"__typename":"CheckRun","status":"COMPLETED","conclusion":"FAILURE"}
]}}},
"associatedPullRequests":{"nodes":[{"number":42,"isDraft":true,"url":"u42"}]}}}}}"#,
0,
);
let out = retry_on_etxtbsy(|| resolve_with(&bin, &[target("main")])).unwrap();
let Some(PrResolution::Pr(badge)) = out.get(&target("main")) else {
panic!("expected a badge: {out:?}");
};
assert_eq!(badge.number, 42);
assert!(badge.is_draft);
assert_eq!(badge.checks, PrCheckState::Failure);
}
#[test]
fn resolve_with_reports_negatives_alongside_badges_end_to_end() {
let dir = tempfile::tempdir().unwrap();
let (bin, _shim) = fake_gh(
dir.path(),
r#"{"data":{"r0":{
"b0":{
"target":{"oid":"a","statusCheckRollup":null},
"associatedPullRequests":{"nodes":[{"number":7,"isDraft":false,"url":"u"}]}},
"b1":null,
"b2":{
"target":{"oid":"b","statusCheckRollup":null},
"associatedPullRequests":{"nodes":[]}}}}}"#,
0,
);
let targets = vec![target("main"), target("unpushed"), target("quiet")];
let out = retry_on_etxtbsy(|| resolve_with(&bin, &targets)).unwrap();
assert_eq!(out.len(), 3, "{out:?}");
assert!(
matches!(out.get(&target("main")), Some(PrResolution::Pr(b)) if b.number == 7),
"{out:?}"
);
assert_eq!(out.get(&target("unpushed")), Some(&PrResolution::NoPr));
assert_eq!(out.get(&target("quiet")), Some(&PrResolution::NoPr));
}
fn pending_pr(number: u64) -> PrResolution {
PrResolution::Pr(PrBadge {
number,
is_draft: false,
checks: PrCheckState::Pending,
url: "u".into(),
head_oid: String::new(),
})
}
#[test]
fn cache_replace_reports_whether_anything_changed() {
let cache = PrStatusCache::new();
let mut map = HashMap::new();
map.insert(target("f"), pending_pr(1));
assert!(cache.replace(map.clone()));
assert!(!cache.replace(map.clone()));
let mut moved = map.clone();
if let Some(PrResolution::Pr(badge)) = moved.get_mut(&target("f")) {
badge.checks = PrCheckState::Success;
}
assert!(cache.replace(moved));
assert!(cache.replace(HashMap::new()));
assert!(!cache.replace(HashMap::new()));
}
#[test]
fn cache_replace_counts_a_new_negative_as_a_change() {
let cache = PrStatusCache::new();
let mut negatives = HashMap::new();
negatives.insert(target("f"), PrResolution::NoPr);
assert!(cache.replace(negatives.clone()));
assert!(!cache.replace(negatives.clone()));
let mut opened = HashMap::new();
opened.insert(target("f"), pending_pr(1));
assert!(cache.replace(opened));
assert!(cache.replace(negatives));
}
#[test]
fn cache_any_pending_ignores_negative_resolutions() {
let cache = PrStatusCache::new();
let mut map = HashMap::new();
map.insert(target("f"), PrResolution::NoPr);
cache.replace(map);
assert!(!cache.any_pending());
}
#[test]
fn cache_get_and_any_pending() {
let cache = PrStatusCache::new();
assert!(cache.get("rust-works", "omni-dev", "f").is_none());
assert!(!cache.any_pending());
let mut map = HashMap::new();
map.insert(target("f"), pending_pr(1));
cache.replace(map.clone());
assert!(
matches!(
cache.get("rust-works", "omni-dev", "f"),
Some(PrResolution::Pr(b)) if b.number == 1
),
"expected the cached badge back"
);
assert!(cache.any_pending());
assert!(cache.get("rust-works", "omni-dev", "other").is_none());
if let Some(PrResolution::Pr(badge)) = map.get_mut(&target("f")) {
badge.checks = PrCheckState::Success;
}
cache.replace(map);
assert!(!cache.any_pending());
}
#[test]
fn cache_retain_targets_drops_only_the_vanished_entries() {
let cache = PrStatusCache::new();
let mut map = HashMap::new();
map.insert(target("keep"), pending_pr(1));
map.insert(target("drop"), pending_pr(2));
cache.replace(map);
let keep: HashSet<PrTarget> = std::iter::once(target("keep")).collect();
assert!(cache.retain_targets(&keep), "a target was removed");
assert!(cache.get("rust-works", "omni-dev", "keep").is_some());
assert!(cache.get("rust-works", "omni-dev", "drop").is_none());
assert!(!cache.retain_targets(&keep));
}
#[test]
fn cache_seed_then_entries_round_trips_without_a_bump() {
let cache = PrStatusCache::new();
cache.seed([
(target("f"), pending_pr(7)),
(target("g"), PrResolution::NoPr),
]);
assert!(matches!(
cache.get("rust-works", "omni-dev", "f"),
Some(PrResolution::Pr(b)) if b.number == 7
));
assert!(matches!(
cache.get("rust-works", "omni-dev", "g"),
Some(PrResolution::NoPr)
));
let mut entries = cache.entries();
entries.sort_by(|a, b| a.0.cmp(&b.0));
assert_eq!(entries.len(), 2);
assert_eq!(entries[0].0, target("f"));
assert_eq!(entries[1].0, target("g"));
}
#[test]
fn merge_query_asks_for_the_enqueue_fields() {
let (query, index) = build_merge_query(&[target("main")]).unwrap();
for field in [
"id",
"headRefOid",
"mergeStateStatus",
"mergeQueueEntry{ state }",
] {
assert!(query.contains(field), "missing {field}: {query}");
}
assert!(
query.contains("associatedPullRequests(first:1, states:OPEN)"),
"{query}"
);
assert!(query.contains("checkSuite{ status }"), "{query}");
assert!(!query.contains("rateLimit"), "{query}");
assert_eq!(index.len(), 1);
}
#[test]
fn merge_query_is_none_without_targets() {
assert!(build_merge_query(&[]).is_none());
}
#[test]
fn parse_merge_response_reads_pr_facts_and_already_queued() {
let (_, index) = build_merge_query(&[target("ready"), target("queued")]).unwrap();
let body = json!({"data":{"r0":{
"b0":{
"target":{"oid":"sha-ready","statusCheckRollup":{"contexts":{"nodes":[
{"__typename":"CheckRun","status":"COMPLETED","conclusion":"SUCCESS"}
]}}},
"associatedPullRequests":{"nodes":[{
"id":"PR_ready","number":10,"isDraft":false,
"url":"https://github.com/rust-works/omni-dev/pull/10",
"headRefOid":"sha-ready","mergeStateStatus":"CLEAN","mergeQueueEntry":null
}]}
},
"b1":{
"target":{"oid":"sha-q","statusCheckRollup":null},
"associatedPullRequests":{"nodes":[{
"id":"PR_queued","number":11,"isDraft":true,
"url":"https://github.com/rust-works/omni-dev/pull/11",
"headRefOid":"sha-q","mergeStateStatus":"BLOCKED",
"mergeQueueEntry":{"state":"QUEUED"}
}]}
}
}}});
let out = parse_merge_response(&body, &index);
let ready = out.get(&target("ready")).expect("ready resolved");
assert_eq!(ready.pr_id, "PR_ready");
assert_eq!(ready.number, 10);
assert_eq!(ready.head_oid, "sha-ready");
assert_eq!(ready.merge_state.as_deref(), Some("CLEAN"));
assert!(!ready.is_draft);
assert!(!ready.already_queued);
assert_eq!(ready.checks, PrCheckState::Success);
let queued = out.get(&target("queued")).expect("queued resolved");
assert!(queued.is_draft);
assert!(queued.already_queued);
assert_eq!(queued.merge_state.as_deref(), Some("BLOCKED"));
assert_eq!(queued.checks, PrCheckState::None);
}
#[test]
fn parse_merge_response_treats_absent_or_prless_refs_as_skips() {
let (_, index) = build_merge_query(&[target("null-ref"), target("no-pr")]).unwrap();
let body = json!({"data":{"r0":{
"b0": null,
"b1": {
"target":{"oid":"x","statusCheckRollup":null},
"associatedPullRequests":{"nodes":[]}
}
}}});
let out = parse_merge_response(&body, &index);
assert!(out.is_empty(), "{out:?}");
}
#[test]
fn merge_info_requires_an_id_and_number() {
let node = json!({
"target":{"oid":"x","statusCheckRollup":null},
"associatedPullRequests":{"nodes":[{"number":9,"url":"u"}]}
});
assert!(merge_info_from_ref(&node).is_none());
}
#[test]
fn enqueue_error_message_joins_or_falls_back() {
let errs = vec![
json!({"message":"Merge queue is not enabled"}),
json!({"message":"Pull request is not mergeable"}),
];
assert_eq!(
enqueue_error_message(&errs),
"Merge queue is not enabled; Pull request is not mergeable"
);
assert_eq!(
enqueue_error_message(&[json!({"code":"X"})]),
"enqueue rejected by GitHub"
);
}
#[test]
fn parse_merge_response_tolerates_a_missing_data_block() {
let (_, index) = build_merge_query(&[target("main")]).unwrap();
assert!(parse_merge_response(&json!({}), &index).is_empty());
}
#[test]
fn resolve_merge_targets_asks_nothing_for_no_targets() {
let out = resolve_merge_targets(Path::new("/no/such/gh/xyzzy"), &[]).unwrap();
assert!(out.is_empty());
}
#[test]
fn resolve_merge_targets_reads_a_reply_end_to_end() {
let dir = tempfile::tempdir().unwrap();
let (bin, _shim) = fake_gh(
dir.path(),
r#"{"data":{"r0":{"b0":{
"target":{"oid":"a","statusCheckRollup":{"contexts":{"nodes":[
{"__typename":"CheckRun","status":"COMPLETED","conclusion":"SUCCESS"}
]}}},
"associatedPullRequests":{"nodes":[{
"id":"PR_1","number":42,"isDraft":false,"url":"u42",
"headRefOid":"a","mergeStateStatus":"CLEAN","mergeQueueEntry":null
}]}}}}}"#,
0,
);
let out = retry_on_etxtbsy(|| resolve_merge_targets(&bin, &[target("main")])).unwrap();
let info = out.get(&target("main")).expect("resolved");
assert_eq!(info.pr_id, "PR_1");
assert_eq!(info.number, 42);
assert_eq!(info.head_oid, "a");
assert_eq!(info.merge_state.as_deref(), Some("CLEAN"));
assert!(!info.already_queued);
assert_eq!(info.checks, PrCheckState::Success);
}
#[test]
fn resolve_merge_targets_errors_on_a_graphql_errors_body() {
let dir = tempfile::tempdir().unwrap();
let (bin, _shim) = fake_gh(
dir.path(),
r#"{"errors":[{"message":"Something is wrong"}]}"#,
0,
);
let err = retry_on_etxtbsy(|| resolve_merge_targets(&bin, &[target("main")])).unwrap_err();
assert!(format!("{err:#}").contains("returned errors"), "{err:#}");
}
#[test]
fn enqueue_pull_request_reports_queued_on_success() {
let dir = tempfile::tempdir().unwrap();
let (bin, _shim) = fake_gh(
dir.path(),
r#"{"data":{"enqueuePullRequest":{"mergeQueueEntry":{"state":"QUEUED"}}}}"#,
0,
);
let out = retry_on_etxtbsy(|| enqueue_pull_request(&bin, "PR_1")).unwrap();
assert_eq!(out, EnqueueOutcome::Queued(Some("QUEUED".to_string())));
}
#[test]
fn enqueue_pull_request_maps_a_nonzero_exit_to_rejected() {
let dir = tempfile::tempdir().unwrap();
let (bin, _shim) = fake_gh(
dir.path(),
r#"{"errors":[{"message":"Merge queue is not enabled on this repository"}]}"#,
1,
);
let out = retry_on_etxtbsy(|| enqueue_pull_request(&bin, "PR_1")).unwrap();
assert_eq!(
out,
EnqueueOutcome::Rejected("Merge queue is not enabled on this repository".to_string())
);
}
#[test]
fn enqueue_pull_request_maps_a_200_with_errors_to_rejected() {
let dir = tempfile::tempdir().unwrap();
let (bin, _shim) = fake_gh(
dir.path(),
r#"{"errors":[{"message":"Pull request is not mergeable"}]}"#,
0,
);
let out = retry_on_etxtbsy(|| enqueue_pull_request(&bin, "PR_1")).unwrap();
assert_eq!(
out,
EnqueueOutcome::Rejected("Pull request is not mergeable".to_string())
);
}
}