Skip to main content

github_actions_maintainer/
github.rs

1use std::{thread, time::Duration};
2
3use anyhow::{Context, Result, anyhow, bail};
4use reqwest::{
5    StatusCode,
6    blocking::{Client, RequestBuilder, Response},
7    header::{AUTHORIZATION, HeaderMap, HeaderValue, RETRY_AFTER, USER_AGENT},
8};
9use serde::{Deserialize, Serialize};
10
11#[derive(Debug, Clone)]
12pub struct GitHubClient {
13    base_url: String,
14    token: Option<String>,
15    client: Client,
16    max_retries: u32,
17    retry_delay: Duration,
18    max_retry_delay: Duration,
19}
20
21#[derive(Debug, Deserialize)]
22struct CommitResponse {
23    sha: String,
24}
25
26#[derive(Debug, Clone, Eq, PartialEq)]
27pub struct LatestReference {
28    pub version: String,
29    pub sha: String,
30}
31
32#[derive(Debug, Deserialize)]
33struct LatestReleaseResponse {
34    tag_name: String,
35}
36
37#[derive(Debug, Deserialize)]
38struct TagResponse {
39    name: String,
40    commit: Option<CommitResponse>,
41}
42
43#[derive(Debug, Clone)]
44pub struct GitHubClientOptions {
45    pub base_url: String,
46    pub token: Option<String>,
47    pub timeout: Duration,
48    pub max_retries: u32,
49    pub retry_delay: Duration,
50    pub max_retry_delay: Duration,
51}
52
53impl Default for GitHubClientOptions {
54    fn default() -> Self {
55        Self {
56            base_url: String::from("https://api.github.com"),
57            token: None,
58            timeout: Duration::from_secs(30),
59            max_retries: 3,
60            retry_delay: Duration::from_secs(1),
61            max_retry_delay: Duration::from_mins(1),
62        }
63    }
64}
65
66#[derive(Debug, Clone, Eq, PartialEq)]
67pub struct TreeEntry {
68    pub path: String,
69    pub sha: String,
70}
71
72#[derive(Debug, Clone, Eq, PartialEq)]
73pub struct PullRequestInfo {
74    pub number: u64,
75    pub url: String,
76}
77
78#[derive(Debug, Clone, Eq, PartialEq)]
79pub struct TagInfo {
80    pub name: String,
81    pub sha: Option<String>,
82}
83
84#[derive(Debug, Clone, Eq, PartialEq)]
85pub struct CommitInfo {
86    pub sha: String,
87    pub message: String,
88    pub is_merge: bool,
89}
90
91#[derive(Debug, Clone, Eq, PartialEq)]
92pub struct CommitRange {
93    pub commits: Vec<CommitInfo>,
94    pub truncated: bool,
95}
96
97#[derive(Debug, Clone, Eq, PartialEq)]
98pub struct ReleaseInfo {
99    pub url: String,
100}
101
102#[derive(Debug, Deserialize)]
103struct RepositoryResponse {
104    default_branch: String,
105}
106
107#[derive(Debug, Deserialize)]
108struct ReferenceResponse {
109    object: ReferenceObject,
110}
111
112#[derive(Debug, Deserialize)]
113struct ReferenceObject {
114    sha: String,
115}
116
117#[derive(Debug, Deserialize)]
118struct CommitTreeResponse {
119    tree: CommitTreeObject,
120}
121
122#[derive(Debug, Deserialize)]
123struct CommitTreeObject {
124    sha: String,
125}
126
127#[derive(Debug, Deserialize)]
128struct BlobResponse {
129    sha: String,
130}
131
132#[derive(Debug, Deserialize)]
133struct TreeResponse {
134    sha: String,
135}
136
137#[derive(Debug, Deserialize)]
138struct CreatedCommitResponse {
139    sha: String,
140}
141
142#[derive(Debug, Deserialize)]
143struct PullRequestResponse {
144    number: u64,
145    html_url: String,
146}
147
148#[derive(Debug, Deserialize)]
149struct UserResponse {
150    login: Option<String>,
151}
152
153#[derive(Debug, Deserialize)]
154struct CompareResponse {
155    total_commits: u64,
156    commits: Vec<RepoCommitResponse>,
157}
158
159#[derive(Debug, Deserialize)]
160struct RepoCommitResponse {
161    sha: String,
162    commit: RepoCommitDetail,
163    parents: Vec<CommitParent>,
164}
165
166#[derive(Debug, Deserialize)]
167struct RepoCommitDetail {
168    message: String,
169}
170
171#[derive(Debug, Deserialize)]
172struct CommitParent {}
173
174#[derive(Debug, Deserialize)]
175struct CreatedReleaseResponse {
176    html_url: String,
177}
178
179#[derive(Debug, Serialize)]
180struct CreateReferenceRequest<'a> {
181    #[serde(rename = "ref")]
182    reference: &'a str,
183    sha: &'a str,
184}
185
186#[derive(Debug, Serialize)]
187struct CreateBlobRequest<'a> {
188    content: &'a str,
189    encoding: &'a str,
190}
191
192#[derive(Debug, Serialize)]
193struct CreateTreeRequest<'a> {
194    base_tree: &'a str,
195    tree: Vec<CreateTreeEntry<'a>>,
196}
197
198#[derive(Debug, Serialize)]
199struct CreateTreeEntry<'a> {
200    path: &'a str,
201    mode: &'a str,
202    #[serde(rename = "type")]
203    object_type: &'a str,
204    sha: &'a str,
205}
206
207#[derive(Debug, Serialize)]
208struct CreateCommitRequest<'a> {
209    message: &'a str,
210    tree: &'a str,
211    parents: Vec<&'a str>,
212}
213
214#[derive(Debug, Serialize)]
215struct UpdateReferenceRequest<'a> {
216    sha: &'a str,
217    force: bool,
218}
219
220#[derive(Debug, Serialize)]
221struct CreatePullRequestRequest<'a> {
222    title: &'a str,
223    body: &'a str,
224    head: &'a str,
225    base: &'a str,
226}
227
228#[derive(Debug, Serialize)]
229struct AddLabelsRequest<'a> {
230    labels: &'a [String],
231}
232
233#[derive(Debug, Serialize)]
234struct CreateReleaseRequest<'a> {
235    tag_name: &'a str,
236    target_commitish: &'a str,
237    name: &'a str,
238    body: &'a str,
239}
240
241impl GitHubClient {
242    pub fn new(base_url: impl Into<String>, token: Option<String>) -> Result<Self> {
243        Self::with_options(GitHubClientOptions {
244            base_url: base_url.into(),
245            token,
246            ..GitHubClientOptions::default()
247        })
248    }
249
250    pub fn with_options(options: GitHubClientOptions) -> Result<Self> {
251        let GitHubClientOptions {
252            base_url,
253            token,
254            timeout,
255            max_retries,
256            retry_delay,
257            max_retry_delay,
258        } = options;
259        let mut headers = HeaderMap::new();
260        headers.insert(USER_AGENT, HeaderValue::from_static("github-actions-maintainer"));
261
262        let client = Client::builder()
263            .default_headers(headers)
264            .timeout(timeout)
265            .build()
266            .context("failed to build GitHub HTTP client")?;
267
268        Ok(Self {
269            base_url: base_url.trim_end_matches('/').to_owned(),
270            token: token.as_deref().and_then(normalize_token),
271            client,
272            max_retries,
273            retry_delay,
274            max_retry_delay,
275        })
276    }
277
278    pub fn latest_reference(&self, owner: &str, repository: &str) -> Result<LatestReference> {
279        if let Some(version) = self.latest_release_tag(owner, repository)? {
280            let sha = self.resolve_reference(owner, repository, &version)?;
281            return Ok(LatestReference { version, sha });
282        }
283
284        let tags = self.send_with_retry(
285            || self.get(&format!("/repos/{owner}/{repository}/tags?per_page=1")),
286            || format!("fetch tags for {owner}/{repository}"),
287        )?;
288        let mut tags = tags
289            .json::<Vec<TagResponse>>()
290            .with_context(|| format!("failed to decode tags response for {owner}/{repository}"))?;
291
292        let tag = tags.pop().ok_or_else(|| {
293            anyhow::anyhow!("GitHub did not return any tags for {owner}/{repository}")
294        })?;
295        let sha = if let Some(commit) = tag.commit {
296            commit.sha
297        } else {
298            self.resolve_reference(owner, repository, &tag.name)?
299        };
300
301        Ok(LatestReference { version: tag.name, sha })
302    }
303
304    pub fn resolve_reference(
305        &self,
306        owner: &str,
307        repository: &str,
308        reference: &str,
309    ) -> Result<String> {
310        let encoded_reference = urlencoding::encode(reference);
311        let response = self.send_with_retry(
312            || self.get(&format!("/repos/{owner}/{repository}/commits/{encoded_reference}")),
313            || format!("resolve {owner}/{repository}@{reference}"),
314        )?;
315        let commit = response.json::<CommitResponse>().with_context(|| {
316            format!("failed to decode commit response for {owner}/{repository}@{reference}")
317        })?;
318
319        Ok(commit.sha)
320    }
321
322    fn latest_release_tag(&self, owner: &str, repository: &str) -> Result<Option<String>> {
323        let response = self.send_with_retry_allowing_not_found(
324            || self.get(&format!("/repos/{owner}/{repository}/releases/latest")),
325            || format!("fetch latest release for {owner}/{repository}"),
326        )?;
327        let Some(response) = response else {
328            return Ok(None);
329        };
330        let release = response.json::<LatestReleaseResponse>().with_context(|| {
331            format!("failed to decode release response for {owner}/{repository}")
332        })?;
333
334        Ok(Some(release.tag_name))
335    }
336
337    pub fn validate_token_scopes(&self) -> Result<()> {
338        let token = self
339            .token
340            .as_deref()
341            .ok_or_else(|| anyhow!("a GitHub token is required for remote PR creation"))?;
342
343        let response = self.send_with_retry(
344            || self.get("/user").header(AUTHORIZATION, format!("Bearer {token}")),
345            || String::from("validate GitHub token scopes"),
346        )?;
347        let headers = response.headers().clone();
348        let user =
349            response.json::<UserResponse>().context("failed to decode GitHub user response")?;
350
351        if user.login.is_none() {
352            bail!("failed to validate GitHub token: authenticated user is missing");
353        }
354
355        let Some(scopes) = headers.get("x-oauth-scopes").and_then(|value| value.to_str().ok())
356        else {
357            return Ok(());
358        };
359
360        let has_repo_scope = scopes.contains("repo") || scopes.contains("public_repo");
361        if !has_repo_scope {
362            bail!("GitHub token is missing the repo or public_repo scope");
363        }
364        if !scopes.contains("workflow") {
365            bail!("GitHub token is missing the workflow scope");
366        }
367
368        Ok(())
369    }
370
371    pub fn default_branch(&self, owner: &str, repository: &str) -> Result<String> {
372        let response = self.send_with_retry(
373            || self.get(&format!("/repos/{owner}/{repository}")),
374            || format!("fetch repository metadata for {owner}/{repository}"),
375        )?;
376        let repository = response.json::<RepositoryResponse>().with_context(|| {
377            format!("failed to decode repository response for {owner}/{repository}")
378        })?;
379        Ok(repository.default_branch)
380    }
381
382    pub fn branch_head_sha(&self, owner: &str, repository: &str, branch: &str) -> Result<String> {
383        let response = self.send_with_retry(
384            || self.get(&format!("/repos/{owner}/{repository}/git/ref/heads/{branch}")),
385            || format!("fetch branch ref for {owner}/{repository}:{branch}"),
386        )?;
387        let reference = response.json::<ReferenceResponse>().with_context(|| {
388            format!("failed to decode branch ref for {owner}/{repository}:{branch}")
389        })?;
390        Ok(reference.object.sha)
391    }
392
393    pub fn commit_tree_sha(
394        &self,
395        owner: &str,
396        repository: &str,
397        commit_sha: &str,
398    ) -> Result<String> {
399        let response = self.send_with_retry(
400            || self.get(&format!("/repos/{owner}/{repository}/git/commits/{commit_sha}")),
401            || format!("fetch commit tree for {owner}/{repository}@{commit_sha}"),
402        )?;
403        let commit = response.json::<CommitTreeResponse>().with_context(|| {
404            format!("failed to decode commit tree for {owner}/{repository}@{commit_sha}")
405        })?;
406        Ok(commit.tree.sha)
407    }
408
409    pub fn create_branch(
410        &self,
411        owner: &str,
412        repository: &str,
413        branch: &str,
414        base_sha: &str,
415    ) -> Result<()> {
416        self.create_ref(owner, repository, &format!("heads/{branch}"), base_sha)
417    }
418
419    pub fn create_ref(
420        &self,
421        owner: &str,
422        repository: &str,
423        ref_path: &str,
424        sha: &str,
425    ) -> Result<()> {
426        let reference = format!("refs/{ref_path}");
427        let payload = CreateReferenceRequest { reference: &reference, sha };
428        self.post_json(&format!("/repos/{owner}/{repository}/git/refs"), &payload, || {
429            format!("create ref {ref_path} for {owner}/{repository}")
430        })?;
431        Ok(())
432    }
433
434    pub fn reference_sha(
435        &self,
436        owner: &str,
437        repository: &str,
438        ref_path: &str,
439    ) -> Result<Option<String>> {
440        let response = self.send_with_retry_allowing_not_found(
441            || self.get(&format!("/repos/{owner}/{repository}/git/ref/{ref_path}")),
442            || format!("fetch ref {ref_path} for {owner}/{repository}"),
443        )?;
444        let Some(response) = response else {
445            return Ok(None);
446        };
447        let reference = response
448            .json::<ReferenceResponse>()
449            .with_context(|| format!("failed to decode ref {ref_path} for {owner}/{repository}"))?;
450        Ok(Some(reference.object.sha))
451    }
452
453    pub fn create_blob(&self, owner: &str, repository: &str, content: &str) -> Result<String> {
454        let payload = CreateBlobRequest { content, encoding: "utf-8" };
455        let response =
456            self.post_json(&format!("/repos/{owner}/{repository}/git/blobs"), &payload, || {
457                format!("create blob for {owner}/{repository}")
458            })?;
459        let blob = response
460            .json::<BlobResponse>()
461            .with_context(|| format!("failed to decode blob response for {owner}/{repository}"))?;
462        Ok(blob.sha)
463    }
464
465    pub fn create_tree(
466        &self,
467        owner: &str,
468        repository: &str,
469        base_tree_sha: &str,
470        entries: &[TreeEntry],
471    ) -> Result<String> {
472        let payload = CreateTreeRequest {
473            base_tree: base_tree_sha,
474            tree: entries
475                .iter()
476                .map(|entry| CreateTreeEntry {
477                    path: &entry.path,
478                    mode: "100644",
479                    object_type: "blob",
480                    sha: &entry.sha,
481                })
482                .collect(),
483        };
484        let response =
485            self.post_json(&format!("/repos/{owner}/{repository}/git/trees"), &payload, || {
486                format!("create tree for {owner}/{repository}")
487            })?;
488        let tree = response
489            .json::<TreeResponse>()
490            .with_context(|| format!("failed to decode tree response for {owner}/{repository}"))?;
491        Ok(tree.sha)
492    }
493
494    pub fn create_commit(
495        &self,
496        owner: &str,
497        repository: &str,
498        message: &str,
499        tree_sha: &str,
500        parent_sha: &str,
501    ) -> Result<String> {
502        let payload = CreateCommitRequest { message, tree: tree_sha, parents: vec![parent_sha] };
503        let response =
504            self.post_json(&format!("/repos/{owner}/{repository}/git/commits"), &payload, || {
505                format!("create commit for {owner}/{repository}")
506            })?;
507        let commit = response.json::<CreatedCommitResponse>().with_context(|| {
508            format!("failed to decode commit response for {owner}/{repository}")
509        })?;
510        Ok(commit.sha)
511    }
512
513    pub fn update_branch(
514        &self,
515        owner: &str,
516        repository: &str,
517        branch: &str,
518        commit_sha: &str,
519    ) -> Result<()> {
520        self.update_ref(owner, repository, &format!("heads/{branch}"), commit_sha, false)
521    }
522
523    pub fn update_ref(
524        &self,
525        owner: &str,
526        repository: &str,
527        ref_path: &str,
528        sha: &str,
529        force: bool,
530    ) -> Result<()> {
531        let payload = UpdateReferenceRequest { sha, force };
532        self.patch_json(
533            &format!("/repos/{owner}/{repository}/git/refs/{ref_path}"),
534            &payload,
535            || format!("update ref {ref_path} for {owner}/{repository}"),
536        )?;
537        Ok(())
538    }
539
540    /// Fast-forward `ref_path` to `sha`; returns `Ok(false)` when GitHub rejects
541    /// the update because it is not a fast forward (HTTP 422).
542    pub fn update_ref_fast_forward(
543        &self,
544        owner: &str,
545        repository: &str,
546        ref_path: &str,
547        sha: &str,
548    ) -> Result<bool> {
549        let payload = UpdateReferenceRequest { sha, force: false };
550        let response = self.send_with_retry_allowing_non_fast_forward(
551            || {
552                self.client
553                    .patch(format!(
554                        "{}/repos/{owner}/{repository}/git/refs/{ref_path}",
555                        self.base_url
556                    ))
557                    .with_auth(self)
558                    .json(&payload)
559            },
560            || format!("fast-forward ref {ref_path} for {owner}/{repository}"),
561        )?;
562        Ok(response.is_some())
563    }
564
565    pub fn list_tags(&self, owner: &str, repository: &str, max_pages: u32) -> Result<Vec<TagInfo>> {
566        let mut tags = Vec::new();
567        for page in 1..=max_pages {
568            let response = self.send_with_retry(
569                || self.get(&format!("/repos/{owner}/{repository}/tags?per_page=100&page={page}")),
570                || format!("list tags for {owner}/{repository}"),
571            )?;
572            let page_tags = response.json::<Vec<TagResponse>>().with_context(|| {
573                format!("failed to decode tags response for {owner}/{repository}")
574            })?;
575            let page_len = page_tags.len();
576            tags.extend(
577                page_tags.into_iter().map(|tag| TagInfo {
578                    name: tag.name,
579                    sha: tag.commit.map(|commit| commit.sha),
580                }),
581            );
582            if page_len < 100 {
583                break;
584            }
585        }
586        Ok(tags)
587    }
588
589    /// Latest tag whose name is `prefix` followed by a semver version,
590    /// scanning up to 1000 tags.
591    pub fn latest_semver_tag(
592        &self,
593        owner: &str,
594        repository: &str,
595        prefix: &str,
596    ) -> Result<Option<TagInfo>> {
597        const MAX_TAG_PAGES: u32 = 10;
598        let tags = self.list_tags(owner, repository, MAX_TAG_PAGES)?;
599        Ok(tags
600            .into_iter()
601            .filter_map(|tag| {
602                let version = tag.name.strip_prefix(prefix)?;
603                let parsed = semver::Version::parse(version).ok()?;
604                Some((parsed, tag))
605            })
606            .max_by(|(left, _), (right, _)| left.cmp(right))
607            .map(|(_, tag)| tag))
608    }
609
610    pub fn compare_commits(
611        &self,
612        owner: &str,
613        repository: &str,
614        base: &str,
615        head: &str,
616        max_pages: u32,
617    ) -> Result<CommitRange> {
618        let encoded_base = urlencoding::encode(base);
619        let encoded_head = urlencoding::encode(head);
620        let mut commits = Vec::new();
621        let mut total_commits = 0usize;
622        for page in 1..=max_pages {
623            let response = self.send_with_retry(
624                || {
625                    self.get(&format!(
626                        "/repos/{owner}/{repository}/compare/{encoded_base}...{encoded_head}?per_page=100&page={page}"
627                    ))
628                },
629                || format!("compare {base}...{head} for {owner}/{repository}"),
630            )?;
631            let compare = response.json::<CompareResponse>().with_context(|| {
632                format!("failed to decode compare response for {owner}/{repository}")
633            })?;
634            total_commits = usize::try_from(compare.total_commits).unwrap_or(usize::MAX);
635            // The compare endpoint caps the commit list; once pages come back
636            // empty, further requests cannot make progress.
637            if compare.commits.is_empty() {
638                break;
639            }
640            commits.extend(compare.commits.into_iter().map(commit_info_from_response));
641            if commits.len() >= total_commits {
642                break;
643            }
644        }
645        let truncated = commits.len() < total_commits;
646        Ok(CommitRange { commits, truncated })
647    }
648
649    pub fn list_commits(
650        &self,
651        owner: &str,
652        repository: &str,
653        head_sha: &str,
654        max_pages: u32,
655    ) -> Result<CommitRange> {
656        let encoded_head = urlencoding::encode(head_sha);
657        let mut commits = Vec::new();
658        let mut last_page_full = false;
659        for page in 1..=max_pages {
660            let response = self.send_with_retry(
661                || {
662                    self.get(&format!(
663                        "/repos/{owner}/{repository}/commits?sha={encoded_head}&per_page=100&page={page}"
664                    ))
665                },
666                || format!("list commits for {owner}/{repository}"),
667            )?;
668            let page_commits = response.json::<Vec<RepoCommitResponse>>().with_context(|| {
669                format!("failed to decode commits response for {owner}/{repository}")
670            })?;
671            last_page_full = page_commits.len() == 100;
672            commits.extend(page_commits.into_iter().map(commit_info_from_response));
673            if !last_page_full {
674                break;
675            }
676        }
677        Ok(CommitRange { commits, truncated: last_page_full })
678    }
679
680    pub fn create_release(
681        &self,
682        owner: &str,
683        repository: &str,
684        tag_name: &str,
685        name: &str,
686        body: &str,
687        target_commitish: &str,
688    ) -> Result<ReleaseInfo> {
689        let payload = CreateReleaseRequest { tag_name, target_commitish, name, body };
690        let response =
691            self.post_json(&format!("/repos/{owner}/{repository}/releases"), &payload, || {
692                format!("create release {tag_name} for {owner}/{repository}")
693            })?;
694        let release = response.json::<CreatedReleaseResponse>().with_context(|| {
695            format!("failed to decode release response for {owner}/{repository}")
696        })?;
697        Ok(ReleaseInfo { url: release.html_url })
698    }
699
700    pub fn ensure_token(&self) -> Result<()> {
701        if self.token.is_none() {
702            bail!("a GitHub token is required to create releases; provide --token or GITHUB_TOKEN");
703        }
704        Ok(())
705    }
706
707    pub fn create_pull_request(
708        &self,
709        owner: &str,
710        repository: &str,
711        title: &str,
712        body: &str,
713        head: &str,
714        base: &str,
715    ) -> Result<PullRequestInfo> {
716        let payload = CreatePullRequestRequest { title, body, head, base };
717        let response =
718            self.post_json(&format!("/repos/{owner}/{repository}/pulls"), &payload, || {
719                format!("create pull request for {owner}/{repository}")
720            })?;
721        let pull_request = response.json::<PullRequestResponse>().with_context(|| {
722            format!("failed to decode pull request response for {owner}/{repository}")
723        })?;
724        Ok(PullRequestInfo { number: pull_request.number, url: pull_request.html_url })
725    }
726
727    pub fn add_labels(
728        &self,
729        owner: &str,
730        repository: &str,
731        issue_number: u64,
732        labels: &[String],
733    ) -> Result<()> {
734        let payload = AddLabelsRequest { labels };
735        self.post_json(
736            &format!("/repos/{owner}/{repository}/issues/{issue_number}/labels"),
737            &payload,
738            || format!("add labels to issue {issue_number} for {owner}/{repository}"),
739        )?;
740        Ok(())
741    }
742
743    fn get(&self, path: &str) -> RequestBuilder {
744        let mut request = self.client.get(format!("{}{}", self.base_url, path));
745        if let Some(token) = self.token.as_deref() {
746            request = request.header(AUTHORIZATION, format!("Bearer {token}"));
747        }
748        request
749    }
750
751    fn post_json<T: Serialize, F>(&self, path: &str, payload: &T, describe: F) -> Result<Response>
752    where
753        F: Fn() -> String,
754    {
755        self.send_with_retry(
756            || self.client.post(format!("{}{}", self.base_url, path)).with_auth(self).json(payload),
757            describe,
758        )
759    }
760
761    fn patch_json<T: Serialize, F>(&self, path: &str, payload: &T, describe: F) -> Result<Response>
762    where
763        F: Fn() -> String,
764    {
765        self.send_with_retry(
766            || {
767                self.client
768                    .patch(format!("{}{}", self.base_url, path))
769                    .with_auth(self)
770                    .json(payload)
771            },
772            describe,
773        )
774    }
775
776    /// Send with retry/backoff and return the first non-retryable response,
777    /// whatever its status. Status-specific handling lives in the wrappers.
778    fn send_raw_with_retry<F, D>(&self, mut build_request: F, describe: &D) -> Result<Response>
779    where
780        F: FnMut() -> RequestBuilder,
781        D: Fn() -> String,
782    {
783        let mut attempt = 0u32;
784
785        loop {
786            match build_request().send() {
787                Ok(response) => {
788                    if Self::should_retry_response(&response) && attempt < self.max_retries {
789                        self.sleep_for_retry(response.headers(), attempt);
790                        attempt += 1;
791                        continue;
792                    }
793                    return Ok(response);
794                }
795                Err(error) => {
796                    if (error.is_timeout() || error.is_connect()) && attempt < self.max_retries {
797                        thread::sleep(self.calculate_backoff(attempt));
798                        attempt += 1;
799                        continue;
800                    }
801                    return Err(error).with_context(describe);
802                }
803            }
804        }
805    }
806
807    fn send_with_retry<F, D>(&self, build_request: F, describe: D) -> Result<Response>
808    where
809        F: FnMut() -> RequestBuilder,
810        D: Fn() -> String,
811    {
812        let response = self.send_raw_with_retry(build_request, &describe)?;
813        if response.status().is_success() {
814            return Ok(response);
815        }
816        self.error_from_response(response, &describe())
817    }
818
819    fn send_with_retry_allowing_not_found<D>(
820        &self,
821        build_request: impl FnMut() -> RequestBuilder,
822        describe: D,
823    ) -> Result<Option<Response>>
824    where
825        D: Fn() -> String,
826    {
827        let response = self.send_raw_with_retry(build_request, &describe)?;
828        if response.status() == StatusCode::NOT_FOUND {
829            return Ok(None);
830        }
831        if response.status().is_success() {
832            return Ok(Some(response));
833        }
834        self.error_from_response(response, &describe()).map(Some)
835    }
836
837    /// Like `send_with_retry`, but a 422 whose body reports a non-fast-forward
838    /// ref update returns `Ok(None)`. Any other 422 is still an error so
839    /// validation failures (bad SHA, invalid ref name) surface loudly.
840    fn send_with_retry_allowing_non_fast_forward<D>(
841        &self,
842        build_request: impl FnMut() -> RequestBuilder,
843        describe: D,
844    ) -> Result<Option<Response>>
845    where
846        D: Fn() -> String,
847    {
848        let response = self.send_raw_with_retry(build_request, &describe)?;
849        if response.status() == StatusCode::UNPROCESSABLE_ENTITY {
850            let body =
851                response.text().unwrap_or_else(|_| String::from("<response body unavailable>"));
852            if body.to_ascii_lowercase().contains("fast forward") {
853                return Ok(None);
854            }
855            bail!("{}: GitHub API returned 422 Unprocessable Entity ({body})", describe());
856        }
857        if response.status().is_success() {
858            return Ok(Some(response));
859        }
860        self.error_from_response(response, &describe()).map(Some)
861    }
862
863    fn should_retry_response(response: &Response) -> bool {
864        if response.status() == StatusCode::TOO_MANY_REQUESTS || response.status().is_server_error()
865        {
866            return true;
867        }
868
869        response.status() == StatusCode::FORBIDDEN
870            && (response
871                .headers()
872                .get("x-ratelimit-remaining")
873                .and_then(|value| value.to_str().ok())
874                == Some("0")
875                || response.headers().contains_key(RETRY_AFTER))
876    }
877
878    fn sleep_for_retry(&self, headers: &HeaderMap, attempt: u32) {
879        let delay = retry_delay_from_headers(headers)
880            .filter(|delay| *delay > Duration::ZERO && *delay <= self.max_retry_delay * 10)
881            .unwrap_or_else(|| self.calculate_backoff(attempt));
882        thread::sleep(delay);
883    }
884
885    fn calculate_backoff(&self, attempt: u32) -> Duration {
886        let shift = attempt.min(10);
887        let candidate = self.retry_delay.saturating_mul(1u32 << shift);
888        candidate.min(self.max_retry_delay)
889    }
890
891    fn error_from_response(&self, response: Response, context: &str) -> Result<Response> {
892        let status = response.status();
893        let body = response.text().unwrap_or_else(|_| String::from("<response body unavailable>"));
894
895        if status == StatusCode::FORBIDDEN
896            && body.to_ascii_lowercase().contains("rate limit")
897            && self.token.is_none()
898        {
899            bail!(
900                "{context}: GitHub API rate limit exceeded. Provide --token or GITHUB_TOKEN for higher limits."
901            )
902        }
903        if status == StatusCode::NOT_FOUND {
904            bail!("{context}: resource not found ({body})");
905        }
906
907        bail!("{context}: GitHub API returned {status} ({body})")
908    }
909}
910
911fn commit_info_from_response(commit: RepoCommitResponse) -> CommitInfo {
912    CommitInfo {
913        sha: commit.sha,
914        message: commit.commit.message,
915        is_merge: commit.parents.len() > 1,
916    }
917}
918
919fn normalize_token(token: &str) -> Option<String> {
920    let trimmed = token.trim();
921    if trimmed.is_empty() { None } else { Some(trimmed.to_owned()) }
922}
923
924fn retry_delay_from_headers(headers: &HeaderMap) -> Option<Duration> {
925    if let Some(retry_after) = headers.get(RETRY_AFTER).and_then(|value| value.to_str().ok())
926        && let Ok(seconds) = retry_after.parse::<u64>()
927    {
928        return Some(Duration::from_secs(seconds));
929    }
930
931    let remaining = headers
932        .get("x-ratelimit-remaining")
933        .and_then(|value| value.to_str().ok())
934        .and_then(|value| value.parse::<u64>().ok());
935    let reset = headers
936        .get("x-ratelimit-reset")
937        .and_then(|value| value.to_str().ok())
938        .and_then(|value| value.parse::<u64>().ok());
939
940    if remaining == Some(0)
941        && let Some(reset) = reset
942    {
943        let now =
944            std::time::SystemTime::now().duration_since(std::time::UNIX_EPOCH).ok()?.as_secs();
945        if reset > now {
946            return Some(Duration::from_secs(reset - now) + Duration::from_millis(100));
947        }
948    }
949
950    None
951}
952
953trait RequestBuilderAuthExt {
954    fn with_auth(self, client: &GitHubClient) -> Self;
955}
956
957impl RequestBuilderAuthExt for RequestBuilder {
958    fn with_auth(self, client: &GitHubClient) -> Self {
959        if let Some(token) = client.token.as_deref() {
960            self.header(AUTHORIZATION, format!("Bearer {token}"))
961        } else {
962            self
963        }
964    }
965}
966
967#[cfg(test)]
968#[allow(clippy::significant_drop_tightening)]
969mod tests {
970    use mockito::{Matcher, Server};
971    use std::time::Duration;
972
973    use super::{GitHubClient, GitHubClientOptions};
974
975    #[test]
976    fn resolve_reference_returns_commit_sha() {
977        let mut server = Server::new();
978        let _mock = server
979            .mock("GET", "/repos/actions/checkout/commits/v4")
980            .match_header("user-agent", "github-actions-maintainer")
981            .with_status(200)
982            .with_body(r#"{"sha":"de0fac2e4500dabe0009e67214ff5f5447ce83dd"}"#)
983            .create();
984
985        let client = GitHubClient::new(server.url(), None).expect("github client");
986        let sha = client.resolve_reference("actions", "checkout", "v4").expect("resolve reference");
987
988        assert_eq!(sha, "de0fac2e4500dabe0009e67214ff5f5447ce83dd");
989    }
990
991    #[test]
992    fn resolve_reference_sends_authorization_when_token_is_present() {
993        let mut server = Server::new();
994        let _mock = server
995            .mock("GET", "/repos/actions/cache/commits/v4")
996            .match_header("authorization", Matcher::Regex(r"^Bearer\s+ghp_testtoken$".into()))
997            .with_status(200)
998            .with_body(r#"{"sha":"668228422ae6a00e4ad889ee87cd7109ec5666a7"}"#)
999            .create();
1000
1001        let client = GitHubClient::new(server.url(), Some(String::from("ghp_testtoken")))
1002            .expect("github client");
1003        let sha = client.resolve_reference("actions", "cache", "v4").expect("resolve reference");
1004
1005        assert_eq!(sha, "668228422ae6a00e4ad889ee87cd7109ec5666a7");
1006    }
1007
1008    #[test]
1009    fn resolve_reference_retries_after_rate_limit() {
1010        let mut server = Server::new();
1011        let now = std::time::SystemTime::now()
1012            .duration_since(std::time::UNIX_EPOCH)
1013            .expect("system time")
1014            .as_secs();
1015
1016        let _rate_limited = server
1017            .mock("GET", "/repos/actions/checkout/commits/v4")
1018            .expect(1)
1019            .with_status(403)
1020            .with_header("x-ratelimit-remaining", "0")
1021            .with_header("x-ratelimit-reset", &now.to_string())
1022            .with_body(r#"{"message":"API rate limit exceeded"}"#)
1023            .create();
1024        let _success = server
1025            .mock("GET", "/repos/actions/checkout/commits/v4")
1026            .expect(1)
1027            .with_status(200)
1028            .with_body(r#"{"sha":"de0fac2e4500dabe0009e67214ff5f5447ce83dd"}"#)
1029            .create();
1030
1031        let client = GitHubClient::with_options(GitHubClientOptions {
1032            base_url: server.url(),
1033            token: None,
1034            timeout: Duration::from_secs(5),
1035            max_retries: 1,
1036            retry_delay: Duration::from_millis(1),
1037            max_retry_delay: Duration::from_millis(5),
1038        })
1039        .expect("github client");
1040
1041        let sha = client.resolve_reference("actions", "checkout", "v4").expect("resolve reference");
1042
1043        assert_eq!(sha, "de0fac2e4500dabe0009e67214ff5f5447ce83dd");
1044    }
1045
1046    #[test]
1047    fn resolve_reference_retries_when_retry_after_is_present() {
1048        let mut server = Server::new();
1049
1050        let _rate_limited = server
1051            .mock("GET", "/repos/actions/cache/commits/v4")
1052            .expect(1)
1053            .with_status(403)
1054            .with_header("retry-after", "0")
1055            .with_body(r#"{"message":"You have exceeded a secondary rate limit"}"#)
1056            .create();
1057        let _success = server
1058            .mock("GET", "/repos/actions/cache/commits/v4")
1059            .expect(1)
1060            .with_status(200)
1061            .with_body(r#"{"sha":"668228422ae6a00e4ad889ee87cd7109ec5666a7"}"#)
1062            .create();
1063
1064        let client = GitHubClient::with_options(GitHubClientOptions {
1065            base_url: server.url(),
1066            token: Some(String::from("ghp_testtoken")),
1067            timeout: Duration::from_secs(5),
1068            max_retries: 1,
1069            retry_delay: Duration::from_millis(1),
1070            max_retry_delay: Duration::from_millis(5),
1071        })
1072        .expect("github client");
1073
1074        let sha = client.resolve_reference("actions", "cache", "v4").expect("resolve reference");
1075
1076        assert_eq!(sha, "668228422ae6a00e4ad889ee87cd7109ec5666a7");
1077    }
1078
1079    #[test]
1080    fn reference_sha_returns_none_when_ref_is_missing() {
1081        let mut server = Server::new();
1082        let _mock = server
1083            .mock("GET", "/repos/acme/demo/git/ref/tags/v1.2.3")
1084            .with_status(404)
1085            .with_body(r#"{"message":"Not Found"}"#)
1086            .create();
1087
1088        let client = GitHubClient::new(server.url(), Some(String::from("ghp_testtoken")))
1089            .expect("github client");
1090        let sha = client.reference_sha("acme", "demo", "tags/v1.2.3").expect("reference sha");
1091
1092        assert_eq!(sha, None);
1093    }
1094
1095    #[test]
1096    fn reference_sha_returns_object_sha() {
1097        let mut server = Server::new();
1098        let _mock = server
1099            .mock("GET", "/repos/acme/demo/git/ref/tags/v1.2.3")
1100            .with_status(200)
1101            .with_body(r#"{"ref":"refs/tags/v1.2.3","object":{"sha":"tagsha"}}"#)
1102            .create();
1103
1104        let client = GitHubClient::new(server.url(), Some(String::from("ghp_testtoken")))
1105            .expect("github client");
1106        let sha = client.reference_sha("acme", "demo", "tags/v1.2.3").expect("reference sha");
1107
1108        assert_eq!(sha.as_deref(), Some("tagsha"));
1109    }
1110
1111    #[test]
1112    fn create_ref_posts_fully_qualified_tag_reference() {
1113        let mut server = Server::new();
1114        let _mock = server
1115            .mock("POST", "/repos/acme/demo/git/refs")
1116            .match_body(Matcher::Regex(r#""ref":"refs/tags/v1\.2\.3""#.into()))
1117            .with_status(201)
1118            .with_body(r#"{"ref":"refs/tags/v1.2.3"}"#)
1119            .create();
1120
1121        let client = GitHubClient::new(server.url(), Some(String::from("ghp_testtoken")))
1122            .expect("github client");
1123        client.create_ref("acme", "demo", "tags/v1.2.3", "commitsha").expect("create ref");
1124    }
1125
1126    #[test]
1127    fn update_ref_serializes_force_flag() {
1128        let mut server = Server::new();
1129        let _mock = server
1130            .mock("PATCH", "/repos/acme/demo/git/refs/tags/v1")
1131            .match_body(Matcher::Regex(r#""force":true"#.into()))
1132            .with_status(200)
1133            .with_body(r#"{"ref":"refs/tags/v1"}"#)
1134            .create();
1135
1136        let client = GitHubClient::new(server.url(), Some(String::from("ghp_testtoken")))
1137            .expect("github client");
1138        client.update_ref("acme", "demo", "tags/v1", "commitsha", true).expect("update ref");
1139    }
1140
1141    #[test]
1142    fn update_ref_fast_forward_reports_non_fast_forward_updates() {
1143        let mut server = Server::new();
1144        let _mock = server
1145            .mock("PATCH", "/repos/acme/demo/git/refs/heads/main")
1146            .match_body(Matcher::Regex(r#""force":false"#.into()))
1147            .with_status(422)
1148            .with_body(r#"{"message":"Update is not a fast forward"}"#)
1149            .create();
1150
1151        let client = GitHubClient::new(server.url(), Some(String::from("ghp_testtoken")))
1152            .expect("github client");
1153        let advanced = client
1154            .update_ref_fast_forward("acme", "demo", "heads/main", "commitsha")
1155            .expect("fast-forward ref");
1156
1157        assert!(!advanced);
1158    }
1159
1160    #[test]
1161    fn update_ref_fast_forward_errors_on_unrelated_validation_failures() {
1162        let mut server = Server::new();
1163        let _mock = server
1164            .mock("PATCH", "/repos/acme/demo/git/refs/heads/main")
1165            .with_status(422)
1166            .with_body(r#"{"message":"Object does not exist"}"#)
1167            .create();
1168
1169        let client = GitHubClient::new(server.url(), Some(String::from("ghp_testtoken")))
1170            .expect("github client");
1171        let error = client
1172            .update_ref_fast_forward("acme", "demo", "heads/main", "commitsha")
1173            .expect_err("validation failure");
1174
1175        assert!(error.to_string().contains("Object does not exist"), "{error}");
1176    }
1177
1178    #[test]
1179    fn update_ref_fast_forward_succeeds_when_ref_is_current() {
1180        let mut server = Server::new();
1181        let _mock = server
1182            .mock("PATCH", "/repos/acme/demo/git/refs/heads/main")
1183            .with_status(200)
1184            .with_body(r#"{"ref":"refs/heads/main"}"#)
1185            .create();
1186
1187        let client = GitHubClient::new(server.url(), Some(String::from("ghp_testtoken")))
1188            .expect("github client");
1189        let advanced = client
1190            .update_ref_fast_forward("acme", "demo", "heads/main", "commitsha")
1191            .expect("fast-forward ref");
1192
1193        assert!(advanced);
1194    }
1195
1196    #[test]
1197    fn list_tags_paginates_until_a_short_page() {
1198        let mut server = Server::new();
1199        let full_page: Vec<String> = (0..100)
1200            .map(|index| format!(r#"{{"name":"v0.0.{index}","commit":{{"sha":"{index:040}"}}}}"#))
1201            .collect();
1202        let _first = server
1203            .mock("GET", "/repos/acme/demo/tags?per_page=100&page=1")
1204            .expect(1)
1205            .with_status(200)
1206            .with_body(format!("[{}]", full_page.join(",")))
1207            .create();
1208        let _second = server
1209            .mock("GET", "/repos/acme/demo/tags?per_page=100&page=2")
1210            .expect(1)
1211            .with_status(200)
1212            .with_body(r#"[{"name":"v1.0.0","commit":{"sha":"lasttagsha"}}]"#)
1213            .create();
1214
1215        let client = GitHubClient::new(server.url(), Some(String::from("ghp_testtoken")))
1216            .expect("github client");
1217        let tags = client.list_tags("acme", "demo", 5).expect("list tags");
1218
1219        assert_eq!(tags.len(), 101);
1220        assert_eq!(tags[100].name, "v1.0.0");
1221        assert_eq!(tags[100].sha.as_deref(), Some("lasttagsha"));
1222    }
1223
1224    #[test]
1225    fn compare_commits_paginates_and_flags_merge_commits() {
1226        let mut server = Server::new();
1227        let first_page: Vec<String> = (0..100)
1228            .map(|index| {
1229                format!(
1230                    r#"{{"sha":"{index:040}","commit":{{"message":"feat: change {index}"}},"parents":[{{}}]}}"#
1231                )
1232            })
1233            .collect();
1234        let _first = server
1235            .mock("GET", "/repos/acme/demo/compare/v0.1.0...headsha?per_page=100&page=1")
1236            .expect(1)
1237            .with_status(200)
1238            .with_body(format!(r#"{{"total_commits":101,"commits":[{}]}}"#, first_page.join(",")))
1239            .create();
1240        let _second = server
1241            .mock("GET", "/repos/acme/demo/compare/v0.1.0...headsha?per_page=100&page=2")
1242            .expect(1)
1243            .with_status(200)
1244            .with_body(
1245                r#"{"total_commits":101,"commits":[{"sha":"mergesha","commit":{"message":"Merge pull request #1"},"parents":[{},{}]}]}"#,
1246            )
1247            .create();
1248
1249        let client = GitHubClient::new(server.url(), Some(String::from("ghp_testtoken")))
1250            .expect("github client");
1251        let range = client
1252            .compare_commits("acme", "demo", "v0.1.0", "headsha", 5)
1253            .expect("compare commits");
1254
1255        assert_eq!(range.commits.len(), 101);
1256        assert!(!range.truncated);
1257        assert!(!range.commits[0].is_merge);
1258        assert!(range.commits[100].is_merge);
1259        assert_eq!(range.commits[100].sha, "mergesha");
1260    }
1261
1262    #[test]
1263    fn compare_commits_marks_truncation_at_the_page_cap() {
1264        let mut server = Server::new();
1265        let first_page: Vec<String> = (0..100)
1266            .map(|index| {
1267                format!(r#"{{"sha":"{index:040}","commit":{{"message":"fix: {index}"}},"parents":[{{}}]}}"#)
1268            })
1269            .collect();
1270        let _first = server
1271            .mock("GET", "/repos/acme/demo/compare/v0.1.0...headsha?per_page=100&page=1")
1272            .expect(1)
1273            .with_status(200)
1274            .with_body(format!(r#"{{"total_commits":150,"commits":[{}]}}"#, first_page.join(",")))
1275            .create();
1276
1277        let client = GitHubClient::new(server.url(), Some(String::from("ghp_testtoken")))
1278            .expect("github client");
1279        let range = client
1280            .compare_commits("acme", "demo", "v0.1.0", "headsha", 1)
1281            .expect("compare commits");
1282
1283        assert_eq!(range.commits.len(), 100);
1284        assert!(range.truncated);
1285    }
1286
1287    #[test]
1288    fn compare_commits_stops_when_pages_run_dry() {
1289        let mut server = Server::new();
1290        let first_page: Vec<String> = (0..100)
1291            .map(|index| {
1292                format!(
1293                    r#"{{"sha":"{index:040}","commit":{{"message":"fix: {index}"}},"parents":[{{}}]}}"#
1294                )
1295            })
1296            .collect();
1297        let _first = server
1298            .mock("GET", "/repos/acme/demo/compare/v0.1.0...headsha?per_page=100&page=1")
1299            .expect(1)
1300            .with_status(200)
1301            .with_body(format!(r#"{{"total_commits":300,"commits":[{}]}}"#, first_page.join(",")))
1302            .create();
1303        let _second = server
1304            .mock("GET", "/repos/acme/demo/compare/v0.1.0...headsha?per_page=100&page=2")
1305            .expect(1)
1306            .with_status(200)
1307            .with_body(r#"{"total_commits":300,"commits":[]}"#)
1308            .create();
1309        let _third = server
1310            .mock("GET", "/repos/acme/demo/compare/v0.1.0...headsha?per_page=100&page=3")
1311            .expect(0)
1312            .create();
1313
1314        let client = GitHubClient::new(server.url(), Some(String::from("ghp_testtoken")))
1315            .expect("github client");
1316        let range =
1317            client.compare_commits("acme", "demo", "v0.1.0", "headsha", 10).expect("compare");
1318
1319        assert_eq!(range.commits.len(), 100);
1320        assert!(range.truncated);
1321    }
1322
1323    #[test]
1324    fn list_commits_stops_on_a_short_page() {
1325        let mut server = Server::new();
1326        let _first = server
1327            .mock("GET", "/repos/acme/demo/commits?sha=headsha&per_page=100&page=1")
1328            .expect(1)
1329            .with_status(200)
1330            .with_body(r#"[{"sha":"onlysha","commit":{"message":"feat: initial"},"parents":[]}]"#)
1331            .create();
1332
1333        let client = GitHubClient::new(server.url(), Some(String::from("ghp_testtoken")))
1334            .expect("github client");
1335        let range = client.list_commits("acme", "demo", "headsha", 3).expect("list commits");
1336
1337        assert_eq!(range.commits.len(), 1);
1338        assert!(!range.truncated);
1339        assert_eq!(range.commits[0].message, "feat: initial");
1340    }
1341
1342    #[test]
1343    fn create_release_posts_tag_and_returns_url() {
1344        let mut server = Server::new();
1345        let _mock = server
1346            .mock("POST", "/repos/acme/demo/releases")
1347            .match_body(Matcher::AllOf(vec![
1348                Matcher::Regex(r#""tag_name":"v1\.2\.3""#.into()),
1349                Matcher::Regex(r#""target_commitish":"commitsha""#.into()),
1350                Matcher::Regex(r#""name":"Release v1\.2\.3""#.into()),
1351            ]))
1352            .with_status(201)
1353            .with_body(r#"{"html_url":"https://github.com/acme/demo/releases/tag/v1.2.3"}"#)
1354            .create();
1355
1356        let client = GitHubClient::new(server.url(), Some(String::from("ghp_testtoken")))
1357            .expect("github client");
1358        let release = client
1359            .create_release("acme", "demo", "v1.2.3", "Release v1.2.3", "notes", "commitsha")
1360            .expect("create release");
1361
1362        assert_eq!(release.url, "https://github.com/acme/demo/releases/tag/v1.2.3");
1363    }
1364
1365    #[test]
1366    fn ensure_token_requires_a_token() {
1367        let client = GitHubClient::new("https://api.github.com", None).expect("github client");
1368        let error = client.ensure_token().expect_err("missing token");
1369
1370        assert!(error.to_string().contains("GitHub token"));
1371    }
1372
1373    #[test]
1374    fn validate_token_scopes_requires_workflow_scope() {
1375        let mut server = Server::new();
1376        let _user = server
1377            .mock("GET", "/user")
1378            .match_header("authorization", Matcher::Regex("^Bearer\\s+ghp_testtoken$".into()))
1379            .with_status(200)
1380            .with_header("x-oauth-scopes", "repo")
1381            .with_body(r#"{"login":"octocat"}"#)
1382            .create();
1383
1384        let client = GitHubClient::new(server.url(), Some(String::from("ghp_testtoken")))
1385            .expect("github client");
1386        let error = client.validate_token_scopes().expect_err("missing workflow scope");
1387
1388        assert!(error.to_string().contains("workflow scope"));
1389    }
1390}