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 let auth = format!("{auth_prefix} {token}");
261 let payload = serde_json::to_vec(body)?;
262
263 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
312fn 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}