Skip to main content

cac_webhook/
provider.rs

1use crate::payload::ProviderKind;
2use crate::retry::{retry, with_timeout, RetryPolicy};
3use reqwest::Client;
4use serde::Serialize;
5use std::time::Duration;
6use thiserror::Error;
7
8#[derive(Debug, Error)]
9pub enum ProviderError {
10    #[error("http error: {0}")]
11    Http(#[from] reqwest::Error),
12    #[error("api error ({status}): {body}")]
13    Api { status: u16, body: String },
14    #[error("missing token for {0}")]
15    MissingToken(&'static str),
16    #[error("serialization error: {0}")]
17    Json(#[from] serde_json::Error),
18}
19
20#[derive(Debug, Clone, Serialize)]
21struct CommitStatus<'a> {
22    state: &'a str,
23    context: &'a str,
24    description: String,
25    #[serde(skip_serializing_if = "Option::is_none")]
26    target_url: Option<String>,
27}
28
29pub struct ProviderClient {
30    http: Client,
31    github_token: Option<String>,
32    codeberg_token: Option<String>,
33}
34
35impl ProviderClient {
36    pub fn new(github_token: Option<String>, codeberg_token: Option<String>) -> Self {
37        let http = Client::builder()
38            .connect_timeout(Duration::from_secs(5))
39            .timeout(Duration::from_secs(30))
40            .build()
41            .unwrap_or_else(|_| Client::new());
42        Self {
43            http,
44            github_token,
45            codeberg_token,
46        }
47    }
48
49    pub async fn set_pending(
50        &self,
51        provider: ProviderKind,
52        owner: &str,
53        repo: &str,
54        sha: &str,
55        context: &str,
56    ) -> Result<(), ProviderError> {
57        self.set_status(
58            provider,
59            owner,
60            repo,
61            sha,
62            context,
63            "pending",
64            "Compliance scan in progress".into(),
65            None,
66        )
67        .await
68    }
69
70    pub async fn set_result(
71        &self,
72        provider: ProviderKind,
73        owner: &str,
74        repo: &str,
75        sha: &str,
76        context: &str,
77        passed: bool,
78        description: String,
79        target_url: Option<String>,
80    ) -> Result<(), ProviderError> {
81        self.set_status(
82            provider,
83            owner,
84            repo,
85            sha,
86            context,
87            if passed { "success" } else { "failure" },
88            description,
89            target_url,
90        )
91        .await
92    }
93
94    async fn set_status(
95        &self,
96        provider: ProviderKind,
97        owner: &str,
98        repo: &str,
99        sha: &str,
100        context: &str,
101        state: &str,
102        description: String,
103        target_url: Option<String>,
104    ) -> Result<(), ProviderError> {
105        let body = CommitStatus {
106            state,
107            context,
108            description,
109            target_url,
110        };
111
112        match provider {
113            ProviderKind::GitHub => {
114                let token = self
115                    .github_token
116                    .as_ref()
117                    .ok_or(ProviderError::MissingToken("GitHub"))?;
118                let url = format!("https://api.github.com/repos/{owner}/{repo}/statuses/{sha}");
119                self.post_json(&url, token, "Bearer", &body).await
120            }
121            ProviderKind::Gitea => {
122                let token = self
123                    .codeberg_token
124                    .as_ref()
125                    .ok_or(ProviderError::MissingToken("Codeberg/Gitea"))?;
126                let url = format!("https://codeberg.org/api/v1/repos/{owner}/{repo}/statuses/{sha}");
127                self.post_json(&url, token, "token", &body).await
128            }
129        }
130    }
131
132    pub async fn comment_on_pr(
133        &self,
134        provider: ProviderKind,
135        owner: &str,
136        repo: &str,
137        pr_number: u64,
138        body: &str,
139    ) -> Result<(), ProviderError> {
140        #[derive(Serialize)]
141        struct Comment<'a> {
142            body: &'a str,
143        }
144
145        match provider {
146            ProviderKind::GitHub => {
147                let token = self
148                    .github_token
149                    .as_ref()
150                    .ok_or(ProviderError::MissingToken("GitHub"))?;
151                let url =
152                    format!("https://api.github.com/repos/{owner}/{repo}/issues/{pr_number}/comments");
153                self.post_json(&url, token, "Bearer", &Comment { body }).await
154            }
155            ProviderKind::Gitea => {
156                let token = self
157                    .codeberg_token
158                    .as_ref()
159                    .ok_or(ProviderError::MissingToken("Codeberg/Gitea"))?;
160                let url = format!(
161                    "https://codeberg.org/api/v1/repos/{owner}/{repo}/issues/{pr_number}/comments"
162                );
163                self.post_json(&url, token, "token", &Comment { body }).await
164            }
165        }
166    }
167
168    pub async fn open_fix_pr(
169        &self,
170        provider: ProviderKind,
171        owner: &str,
172        repo: &str,
173        head_branch: &str,
174        base_branch: &str,
175        title: &str,
176        body: &str,
177    ) -> Result<String, ProviderError> {
178        #[derive(Serialize)]
179        struct NewPr<'a> {
180            title: &'a str,
181            head: &'a str,
182            base: &'a str,
183            body: &'a str,
184        }
185
186        match provider {
187            ProviderKind::GitHub => {
188                let token = self
189                    .github_token
190                    .as_ref()
191                    .ok_or(ProviderError::MissingToken("GitHub"))?;
192                let url = format!("https://api.github.com/repos/{owner}/{repo}/pulls");
193                let resp = self
194                    .http
195                    .post(&url)
196                    .header("Authorization", format!("Bearer {token}"))
197                    .header("Accept", "application/vnd.github+json")
198                    .header("User-Agent", "compliance-as-code-agent")
199                    .json(&NewPr {
200                        title,
201                        head: head_branch,
202                        base: base_branch,
203                        body,
204                    })
205                    .send()
206                    .await?;
207                let status = resp.status();
208                let text = resp.text().await?;
209                if !status.is_success() {
210                    return Err(ProviderError::Api {
211                        status: status.as_u16(),
212                        body: text,
213                    });
214                }
215                let json: serde_json::Value = serde_json::from_str(&text).unwrap_or_default();
216                Ok(json["html_url"].as_str().unwrap_or("").to_string())
217            }
218            ProviderKind::Gitea => {
219                let token = self
220                    .codeberg_token
221                    .as_ref()
222                    .ok_or(ProviderError::MissingToken("Codeberg/Gitea"))?;
223                let url = format!("https://codeberg.org/api/v1/repos/{owner}/{repo}/pulls");
224                let resp = self
225                    .http
226                    .post(&url)
227                    .header("Authorization", format!("token {token}"))
228                    .header("User-Agent", "compliance-as-code-agent")
229                    .json(&NewPr {
230                        title,
231                        head: head_branch,
232                        base: base_branch,
233                        body,
234                    })
235                    .send()
236                    .await?;
237                let status = resp.status();
238                let text = resp.text().await?;
239                if !status.is_success() {
240                    return Err(ProviderError::Api {
241                        status: status.as_u16(),
242                        body: text,
243                    });
244                }
245                let json: serde_json::Value = serde_json::from_str(&text).unwrap_or_default();
246                Ok(json["html_url"].as_str().unwrap_or("").to_string())
247            }
248        }
249    }
250
251    async fn post_json<T: Serialize>(
252        &self,
253        url: &str,
254        token: &str,
255        auth_prefix: &str,
256        body: &T,
257    ) -> Result<(), ProviderError> {
258        // Serialize once; reqwest's RequestBuilder is not re-runnable, so each
259        // retry attempt rebuilds the request from these owned values.
260        let auth = format!("{auth_prefix} {token}");
261        let payload = serde_json::to_vec(body)?;
262
263        // Up to 3 attempts with exponential backoff + full jitter, retrying only
264        // on transient transport errors and 429/5xx responses.
265        let policy = RetryPolicy::with_max_attempts(3);
266        let result = retry(
267            || {
268                let req = self
269                    .http
270                    .post(url)
271                    .header("Authorization", &auth)
272                    .header("Accept", "application/json")
273                    .header("User-Agent", "compliance-as-code-agent")
274                    .header("Content-Type", "application/json")
275                    .body(payload.clone());
276                async move {
277                    let resp = with_timeout(
278                        async { req.send().await.map_err(ProviderError::from) },
279                        Duration::from_secs(30),
280                    )
281                    .await
282                    .map_err(|e| match e.into_source() {
283                        Some(err) => err,
284                        None => ProviderError::Api {
285                            status: 0,
286                            body: "request timed out".into(),
287                        },
288                    })?;
289                    let status = resp.status();
290                    if status.is_success() {
291                        Ok(())
292                    } else {
293                        Err(ProviderError::Api {
294                            status: status.as_u16(),
295                            body: resp.text().await.unwrap_or_default(),
296                        })
297                    }
298                }
299            },
300            &policy,
301            is_retryable,
302        )
303        .await;
304
305        result.map_err(|e| e.into_source().unwrap_or(ProviderError::Api {
306            status: 0,
307            body: "request timed out".into(),
308        }))
309    }
310}
311
312/// Transient failures worth retrying: network/transport errors and rate-limit
313/// (429) or server (5xx) responses. Client errors (4xx other than 429) are
314/// terminal.
315fn is_retryable(err: &ProviderError) -> bool {
316    match err {
317        ProviderError::Http(_) => true,
318        ProviderError::Api { status, .. } => *status == 429 || (500..600).contains(status),
319        ProviderError::MissingToken(_) => false,
320        ProviderError::Json(_) => false,
321    }
322}