#![allow(unsafe_code)]
use std::sync::Arc;
use async_trait::async_trait;
use linkme::distributed_slice;
use secrecy::{ExposeSecret, SecretString};
use tokio::io::AsyncRead;
use crate::config::ReleaseSourceConfig;
use crate::http;
use crate::release::{
ProviderError, ProviderFactory, ProviderRegistration, RegisteredProvider, Release,
ReleaseAsset, ReleaseProvider, RELEASE_PROVIDERS,
};
pub struct GithubProvider {
client: reqwest::Client,
scheme: &'static str,
host: String,
owner: String,
repo: String,
token: Option<SecretString>,
}
fn normalise_host(raw: &str) -> String {
let stripped =
raw.trim_end_matches('/').trim_start_matches("https://").trim_start_matches("http://");
if stripped == "api.github.com" || stripped.ends_with("/api/v3") || stripped.contains("/api/") {
stripped.to_string()
} else if let Some(rest) = stripped.strip_prefix("api.") {
format!("{rest}/api/v3")
} else {
format!("{stripped}/api/v3")
}
}
pub fn factory(
cfg: &ReleaseSourceConfig,
token: Option<SecretString>,
) -> Result<Arc<dyn ReleaseProvider>, ProviderError> {
let ReleaseSourceConfig::Github(params) = cfg else {
return Err(ProviderError::InvalidConfig(format!(
"github factory called with non-github config: source_type={}",
cfg.source_type()
)));
};
if params.host.trim().is_empty() {
return Err(ProviderError::InvalidConfig("github host must not be empty".to_string()));
}
if params.host.starts_with("http://") {
return Err(ProviderError::InvalidConfig(format!(
"github host must be https; got {}",
params.host
)));
}
if params.owner.trim().is_empty() || params.repo.trim().is_empty() {
return Err(ProviderError::InvalidConfig(
"github owner and repo must not be empty".to_string(),
));
}
let host = if params.allow_insecure_base_url {
params.host.trim_end_matches('/').to_string()
} else {
normalise_host(¶ms.host)
};
let client = http::build_client(params.timeout_seconds, params.allow_insecure_base_url)?;
let scheme = http::scheme_for(params.allow_insecure_base_url);
Ok(Arc::new(GithubProvider {
client,
scheme,
host,
owner: params.owner.clone(),
repo: params.repo.clone(),
token,
}))
}
#[distributed_slice(RELEASE_PROVIDERS)]
fn __register_github() -> Box<dyn ProviderRegistration> {
Box::new(RegisteredProvider { source_type: "github", factory: factory as ProviderFactory })
}
#[async_trait]
impl ReleaseProvider for GithubProvider {
async fn latest_release(&self) -> Result<Release, ProviderError> {
let url = format!(
"{scheme}://{host}/repos/{owner}/{repo}/releases/latest",
scheme = self.scheme,
host = self.host,
owner = self.owner,
repo = self.repo
);
let resp = self.send(&url).await?;
let dto: ApiRelease = http::parse_json(resp).await?;
Ok(dto.into_release())
}
async fn release_by_tag(&self, tag: &str) -> Result<Release, ProviderError> {
let url = format!(
"{scheme}://{host}/repos/{owner}/{repo}/releases/tags/{tag}",
scheme = self.scheme,
host = self.host,
owner = self.owner,
repo = self.repo,
tag = http::urlencode(tag),
);
let resp = self.send(&url).await?;
let dto: ApiRelease = http::parse_json(resp).await?;
Ok(dto.into_release())
}
async fn list_releases(&self, limit: usize) -> Result<Vec<Release>, ProviderError> {
let per_page = limit.clamp(1, 100);
let url = format!(
"{scheme}://{host}/repos/{owner}/{repo}/releases?per_page={per_page}",
scheme = self.scheme,
host = self.host,
owner = self.owner,
repo = self.repo
);
let resp = self.send(&url).await?;
let dtos: Vec<ApiRelease> = http::parse_json(resp).await?;
Ok(dtos.into_iter().take(limit).map(ApiRelease::into_release).collect())
}
async fn download_asset(
&self,
asset: &ReleaseAsset,
) -> Result<(Box<dyn AsyncRead + Send + Unpin>, u64), ProviderError> {
let mut req = self
.client
.get(&asset.download_url)
.header("Accept", "application/octet-stream")
.header("X-GitHub-Api-Version", "2022-11-28");
if let Some(tok) = &self.token {
req = req.bearer_auth(tok.expose_secret());
}
let resp = req.send().await.map_err(|e| ProviderError::Transport(e.to_string()))?;
check_status(&resp, &self.host)?;
Ok(http::stream_body(resp))
}
}
impl GithubProvider {
async fn send(&self, url: &str) -> Result<reqwest::Response, ProviderError> {
let mut req = self
.client
.get(url)
.header("Accept", "application/vnd.github+json")
.header("X-GitHub-Api-Version", "2022-11-28");
if let Some(tok) = &self.token {
req = req.bearer_auth(tok.expose_secret());
}
let resp = req.send().await.map_err(|e| ProviderError::Transport(e.to_string()))?;
check_status(&resp, &self.host)?;
Ok(resp)
}
}
fn check_status(resp: &reqwest::Response, host: &str) -> Result<(), ProviderError> {
let extra_rate_limit_signal = resp.status() == reqwest::StatusCode::FORBIDDEN
&& http::header_str(resp.headers(), "x-ratelimit-remaining") == Some("0");
http::map_status_to_error(resp, host, extra_rate_limit_signal)
}
#[derive(Debug, serde::Deserialize)]
struct ApiRelease {
#[serde(default)]
id: u64,
#[serde(default)]
name: Option<String>,
tag_name: String,
#[serde(default)]
body: Option<String>,
#[serde(default)]
draft: bool,
#[serde(default)]
prerelease: bool,
created_at: String,
#[serde(default)]
published_at: Option<String>,
#[serde(default)]
assets: Vec<ApiAsset>,
}
#[derive(Debug, serde::Deserialize)]
struct ApiAsset {
id: u64,
name: String,
#[serde(default)]
size: u64,
#[serde(default)]
content_type: Option<String>,
#[serde(rename = "browser_download_url")]
download_url: String,
}
impl ApiRelease {
fn into_release(self) -> Release {
let created_at =
parse_iso8601(&self.created_at).unwrap_or(time::OffsetDateTime::UNIX_EPOCH);
let published_at = self.published_at.as_deref().and_then(parse_iso8601);
let name = self.name.unwrap_or_else(|| self.tag_name.clone());
let body = self.body.unwrap_or_default();
let tag = self.tag_name.clone();
let _ = self.id;
let mut release = Release::new(name, tag, created_at);
release.body = body;
release.draft = self.draft;
release.prerelease = self.prerelease;
release.published_at = published_at;
release.assets = self.assets.into_iter().map(ApiAsset::into_asset).collect();
release
}
}
impl ApiAsset {
fn into_asset(self) -> ReleaseAsset {
let mut a = ReleaseAsset::new(self.id.to_string(), self.name, self.download_url);
a.size = self.size;
a.content_type = self.content_type;
a
}
}
fn parse_iso8601(s: &str) -> Option<time::OffsetDateTime> {
time::OffsetDateTime::parse(s, &time::format_description::well_known::Rfc3339).ok()
}
#[cfg(test)]
mod tests {
use super::normalise_host;
#[test]
fn normalises_api_github_com() {
assert_eq!(normalise_host("api.github.com"), "api.github.com");
assert_eq!(normalise_host("https://api.github.com"), "api.github.com");
assert_eq!(normalise_host("https://api.github.com/"), "api.github.com");
}
#[test]
fn promotes_bare_enterprise_host() {
assert_eq!(normalise_host("github.example.com"), "github.example.com/api/v3");
assert_eq!(normalise_host("https://github.example.com/"), "github.example.com/api/v3");
}
#[test]
fn promotes_api_prefixed_enterprise_host() {
assert_eq!(normalise_host("api.github.example.com"), "github.example.com/api/v3");
}
#[test]
fn preserves_explicit_api_path() {
assert_eq!(normalise_host("github.example.com/api/v3"), "github.example.com/api/v3");
}
}