use serde::Deserialize;
use std::time::Duration;
const PROBE_TIMEOUT: Duration = Duration::from_secs(8);
const OSV_ENDPOINT: &str = "https://api.osv.dev/v1/querybatch";
const OSV_BATCH_LIMIT: usize = 500;
const NPM_DOWNLOADS_BASE: &str = "https://api.npmjs.org/downloads/point/last-week";
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct MaliciousAdvisory {
pub package: String,
pub advisory_id: String,
}
#[derive(Debug, thiserror::Error)]
pub enum SupplyChainError {
#[error("supply-chain probe HTTP error: {0}")]
Http(#[from] reqwest::Error),
#[error("supply-chain probe JSON decode failed: {0}")]
Decode(#[from] serde_json::Error),
#[error("supply-chain probe returned non-success status: {0}")]
Status(reqwest::StatusCode),
#[error("OSV returned {got} results for {expected} queries — truncated response")]
TruncatedOsvResponse { expected: usize, got: usize },
}
#[derive(Debug, serde::Serialize)]
struct OsvQuery<'a> {
package: OsvPackage<'a>,
}
#[derive(Debug, serde::Serialize)]
struct OsvPackage<'a> {
name: &'a str,
ecosystem: &'a str,
}
#[derive(Debug, serde::Serialize)]
struct OsvBatchRequest<'a> {
queries: Vec<OsvQuery<'a>>,
}
#[derive(Debug, Deserialize, Default)]
struct OsvBatchResponse {
#[serde(default)]
results: Vec<OsvResult>,
}
#[derive(Debug, Deserialize, Default)]
struct OsvResult {
#[serde(default)]
vulns: Vec<OsvVuln>,
}
#[derive(Debug, Deserialize)]
struct OsvVuln {
id: String,
}
#[derive(Debug, Deserialize)]
struct NpmDownloadsResponse {
#[serde(default)]
downloads: Option<u64>,
#[serde(default)]
error: Option<String>,
}
pub fn build_probe_client() -> Result<reqwest::Client, SupplyChainError> {
Ok(reqwest::Client::builder().timeout(PROBE_TIMEOUT).build()?)
}
pub async fn fetch_malicious_advisories(
client: &reqwest::Client,
names: &[String],
) -> Result<Vec<MaliciousAdvisory>, SupplyChainError> {
if names.is_empty() {
return Ok(Vec::new());
}
let mut hits = Vec::new();
for chunk in names.chunks(OSV_BATCH_LIMIT) {
hits.extend(fetch_malicious_advisories_chunk(client, chunk).await?);
}
Ok(hits)
}
async fn fetch_malicious_advisories_chunk(
client: &reqwest::Client,
names: &[String],
) -> Result<Vec<MaliciousAdvisory>, SupplyChainError> {
let body = OsvBatchRequest {
queries: names
.iter()
.map(|n| OsvQuery {
package: OsvPackage {
name: n.as_str(),
ecosystem: "npm",
},
})
.collect(),
};
let resp = client.post(OSV_ENDPOINT).json(&body).send().await?;
if !resp.status().is_success() {
return Err(SupplyChainError::Status(resp.status()));
}
let bytes = resp.bytes().await?;
let parsed: OsvBatchResponse = serde_json::from_slice(&bytes)?;
if parsed.results.len() != names.len() {
return Err(SupplyChainError::TruncatedOsvResponse {
expected: names.len(),
got: parsed.results.len(),
});
}
Ok(extract_malicious(names, &parsed))
}
fn extract_malicious(names: &[String], resp: &OsvBatchResponse) -> Vec<MaliciousAdvisory> {
let mut hits = Vec::new();
for (name, result) in names.iter().zip(resp.results.iter()) {
for vuln in &result.vulns {
if vuln.id.starts_with("MAL-") {
hits.push(MaliciousAdvisory {
package: name.clone(),
advisory_id: vuln.id.clone(),
});
}
}
}
hits
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum DownloadCount {
Known(u64),
Unknown,
}
pub async fn fetch_weekly_downloads_with(
client: &reqwest::Client,
name: &str,
) -> Result<DownloadCount, SupplyChainError> {
let encoded = name.replace('/', "%2F");
let url = format!("{NPM_DOWNLOADS_BASE}/{encoded}");
let resp = client.get(&url).send().await?;
let status = resp.status();
if status == reqwest::StatusCode::NOT_FOUND {
return Ok(DownloadCount::Unknown);
}
if !status.is_success() {
return Err(SupplyChainError::Status(status));
}
let bytes = resp.bytes().await?;
let parsed: NpmDownloadsResponse = serde_json::from_slice(&bytes)?;
Ok(parse_downloads(&parsed))
}
fn parse_downloads(resp: &NpmDownloadsResponse) -> DownloadCount {
if resp.error.is_some() {
return DownloadCount::Unknown;
}
match resp.downloads {
Some(n) => DownloadCount::Known(n),
None => DownloadCount::Unknown,
}
}
pub fn advisory_url(advisory_id: &str) -> String {
format!("https://osv.dev/vulnerability/{advisory_id}")
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn extract_malicious_filters_non_mal_ids() {
let names = vec!["evil-pkg".to_string(), "fine-pkg".to_string()];
let resp = OsvBatchResponse {
results: vec![
OsvResult {
vulns: vec![
OsvVuln {
id: "MAL-2026-3652".to_string(),
},
OsvVuln {
id: "GHSA-xxxx".to_string(),
},
],
},
OsvResult {
vulns: vec![OsvVuln {
id: "CVE-2024-9999".to_string(),
}],
},
],
};
let hits = extract_malicious(&names, &resp);
assert_eq!(hits.len(), 1);
assert_eq!(hits[0].package, "evil-pkg");
assert_eq!(hits[0].advisory_id, "MAL-2026-3652");
}
#[test]
fn truncated_osv_response_carries_lengths_in_error() {
let err = SupplyChainError::TruncatedOsvResponse {
expected: 3,
got: 1,
};
let rendered = err.to_string();
assert!(rendered.contains("3"), "expected count missing: {rendered}");
assert!(rendered.contains("1"), "got count missing: {rendered}");
assert!(
rendered.contains("truncated"),
"category word missing: {rendered}"
);
}
#[test]
fn parse_downloads_treats_error_body_as_unknown() {
let resp = NpmDownloadsResponse {
downloads: None,
error: Some("package @scope/name not found".to_string()),
};
assert_eq!(parse_downloads(&resp), DownloadCount::Unknown);
}
#[test]
fn parse_downloads_reads_known_count() {
let resp = NpmDownloadsResponse {
downloads: Some(42_000_000),
error: None,
};
assert_eq!(parse_downloads(&resp), DownloadCount::Known(42_000_000));
}
#[test]
fn advisory_url_uses_osv_domain() {
assert_eq!(
advisory_url("MAL-2026-3652"),
"https://osv.dev/vulnerability/MAL-2026-3652"
);
}
}