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