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 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 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 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 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 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}