use std::collections::{BTreeMap, HashMap};
use std::path::{Path, PathBuf};
use std::process::Command;
use std::sync::{Mutex, PoisonError};
use anyhow::{bail, Context, Result};
use serde::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)]
#[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)]
pub struct PrTarget {
pub owner: String,
pub name: String,
pub branch: String,
}
#[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"))
}
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
statusCheckRollup{{ contexts(first:100){{ nodes{{
__typename
...on CheckRun{{ status conclusion checkSuite{{ status }} }}
...on StatusContext{{ state }}
}} }} }}
}} }}
associatedPullRequests(first:1, states:OPEN){{ nodes{{ number isDraft url }} }}
}}"
)
}
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{{\n{}\n}}", repos.join("\n")), index))
}
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 parse_response(body: &Value, index: &QueryIndex) -> HashMap<PrTarget, PrBadge> {
let mut out = HashMap::new();
let Some(data) = body.get("data") else {
return out;
};
for ((ri, bi), target) in index {
let node = data
.get(format!("r{ri}"))
.and_then(|r| r.get(format!("b{bi}")));
let Some(node) = node.filter(|n| !n.is_null()) else {
continue;
};
if let Some(badge) = badge_from_ref(node) {
out.insert(target.clone(), badge);
}
}
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 output = Command::new(bin)
.args(["api", "graphql", "-f"])
.arg(format!("query={query}"))
.output()
.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, PrBadge>> {
let Some((query, index)) = build_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_response(&body, &index))
}
#[derive(Debug, Default)]
pub struct PrStatusCache {
badges: Mutex<HashMap<PrTarget, PrBadge>>,
}
impl PrStatusCache {
#[must_use]
pub fn new() -> Self {
Self::default()
}
#[must_use]
pub fn get(&self, owner: &str, name: &str, branch: &str) -> Option<PrBadge> {
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, PrBadge>) -> bool {
let mut guard = self.lock();
if *guard == next {
return false;
}
*guard = next;
true
}
#[must_use]
pub fn any_pending(&self) -> bool {
self.lock()
.values()
.any(|b| b.checks == PrCheckState::Pending)
}
fn lock(&self) -> std::sync::MutexGuard<'_, HashMap<PrTarget, PrBadge>> {
self.badges.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::{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_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_skips_absent_refs() {
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(), 1, "{out:?}");
let badge = out.get(&target("feature")).unwrap();
assert_eq!(badge.number, 65);
assert!(badge.is_draft);
assert_eq!(badge.checks, PrCheckState::Success);
assert_eq!(badge.url, "u");
}
#[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);
assert_eq!(
out.get(&target("feature")).unwrap().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 = 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 = 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 = 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 = resolve_with(&bin, &[target("main")]).unwrap();
assert_eq!(out.get(&target("main")).unwrap().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 = resolve_with(&bin, &[target("main")]).unwrap();
let badge = out.get(&target("main")).unwrap();
assert_eq!(badge.number, 42);
assert!(badge.is_draft);
assert_eq!(badge.checks, PrCheckState::Failure);
}
#[test]
fn cache_replace_reports_whether_anything_changed() {
let cache = PrStatusCache::new();
let badge = PrBadge {
number: 1,
is_draft: false,
checks: PrCheckState::Pending,
url: "u".into(),
head_oid: String::new(),
};
let mut map = HashMap::new();
map.insert(target("f"), badge);
assert!(cache.replace(map.clone()));
assert!(!cache.replace(map.clone()));
let mut moved = map.clone();
moved.get_mut(&target("f")).unwrap().checks = PrCheckState::Success;
assert!(cache.replace(moved));
assert!(cache.replace(HashMap::new()));
assert!(!cache.replace(HashMap::new()));
}
#[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"),
PrBadge {
number: 1,
is_draft: false,
checks: PrCheckState::Pending,
url: "u".into(),
head_oid: String::new(),
},
);
cache.replace(map.clone());
assert_eq!(cache.get("rust-works", "omni-dev", "f").unwrap().number, 1);
assert!(cache.any_pending());
assert!(cache.get("rust-works", "omni-dev", "other").is_none());
map.get_mut(&target("f")).unwrap().checks = PrCheckState::Success;
cache.replace(map);
assert!(!cache.any_pending());
}
}