use serde::Deserialize;
use std::collections::HashMap;
use std::time::Duration;
const PROBE_TIMEOUT: Duration = Duration::from_secs(8);
const OSV_ENDPOINT: &str = "https://api.osv.dev/v1/querybatch";
const OSV_VULN_BASE: &str = "https://api.osv.dev/v1/vulns";
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,
pub version: Option<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>,
#[serde(skip_serializing_if = "Option::is_none")]
version: Option<&'a str>,
}
#[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, Clone, Deserialize, Default)]
struct OsvVulnDetails {
#[serde(default)]
affected: Vec<OsvAffected>,
}
#[derive(Debug, Clone, Deserialize, Default)]
struct OsvAffected {
#[serde(default)]
package: OsvAffectedPackage,
#[serde(default)]
versions: Vec<String>,
}
#[derive(Debug, Clone, Deserialize, Default)]
struct OsvAffectedPackage {
#[serde(default)]
name: String,
#[serde(default)]
ecosystem: 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)
}
pub async fn fetch_malicious_advisories_versioned(
client: &reqwest::Client,
pairs: &[(String, String)],
) -> Result<Vec<MaliciousAdvisory>, SupplyChainError> {
if pairs.is_empty() {
return Ok(Vec::new());
}
let mut hits = Vec::new();
for chunk in pairs.chunks(OSV_BATCH_LIMIT) {
hits.extend(fetch_malicious_advisories_versioned_chunk(client, chunk).await?);
}
filter_malicious_versioned_hits(client, hits).await
}
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",
},
version: None,
})
.collect(),
};
let parsed = post_osv_batch(client, &body, names.len()).await?;
Ok(extract_malicious(names, &parsed))
}
async fn fetch_malicious_advisories_versioned_chunk(
client: &reqwest::Client,
pairs: &[(String, String)],
) -> Result<Vec<MaliciousAdvisory>, SupplyChainError> {
let body = OsvBatchRequest {
queries: pairs
.iter()
.map(|(name, version)| OsvQuery {
package: OsvPackage {
name: name.as_str(),
ecosystem: "npm",
},
version: Some(version.as_str()),
})
.collect(),
};
let parsed = post_osv_batch(client, &body, pairs.len()).await?;
Ok(extract_malicious_versioned(pairs, &parsed))
}
async fn post_osv_batch(
client: &reqwest::Client,
body: &OsvBatchRequest<'_>,
expected: usize,
) -> Result<OsvBatchResponse, SupplyChainError> {
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() != expected {
return Err(SupplyChainError::TruncatedOsvResponse {
expected,
got: parsed.results.len(),
});
}
Ok(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(),
version: None,
});
}
}
}
hits
}
fn extract_malicious_versioned(
pairs: &[(String, String)],
resp: &OsvBatchResponse,
) -> Vec<MaliciousAdvisory> {
let mut hits = Vec::new();
for ((name, version), result) in pairs.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(),
version: Some(version.clone()),
});
}
}
}
hits
}
async fn filter_malicious_versioned_hits(
client: &reqwest::Client,
hits: Vec<MaliciousAdvisory>,
) -> Result<Vec<MaliciousAdvisory>, SupplyChainError> {
let mut details = HashMap::new();
let mut tasks = tokio::task::JoinSet::new();
for hit in &hits {
if hit.version.is_none() || details.contains_key(&hit.advisory_id) {
continue;
}
details.insert(hit.advisory_id.clone(), None);
let client = client.clone();
let advisory_id = hit.advisory_id.clone();
tasks.spawn(async move {
let details = fetch_osv_vuln_details(&client, &advisory_id).await.ok();
(advisory_id, details)
});
}
while let Some(joined) = tasks.join_next().await {
if let Ok((advisory_id, fetched)) = joined {
details.insert(advisory_id, fetched);
}
}
let mut filtered = Vec::new();
for hit in hits {
let affects = details
.get(&hit.advisory_id)
.and_then(Option::as_ref)
.is_none_or(|d| {
let Some(version) = hit.version.as_deref() else {
return true;
};
versioned_hit_affects_resolved_version(d, &hit.package, version)
});
if affects {
filtered.push(hit);
}
}
Ok(filtered)
}
async fn fetch_osv_vuln_details(
client: &reqwest::Client,
advisory_id: &str,
) -> Result<OsvVulnDetails, SupplyChainError> {
let url = format!("{OSV_VULN_BASE}/{advisory_id}");
let resp = client.get(url).send().await?;
if !resp.status().is_success() {
return Err(SupplyChainError::Status(resp.status()));
}
let bytes = resp.bytes().await?;
Ok(serde_json::from_slice(&bytes)?)
}
fn versioned_hit_affects_resolved_version(
details: &OsvVulnDetails,
package: &str,
version: &str,
) -> bool {
let mut matched_package = false;
for affected in &details.affected {
if !affected.package.ecosystem.eq_ignore_ascii_case("npm")
|| affected.package.name != package
{
continue;
}
matched_package = true;
if affected.versions.is_empty() || affected.versions.iter().any(|v| v == version) {
return true;
}
}
!matched_package
}
#[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");
assert_eq!(hits[0].version, None);
}
#[test]
fn extract_malicious_versioned_carries_resolved_version() {
let pairs = vec![("ansi-regex".to_string(), "6.2.1".to_string())];
let resp = OsvBatchResponse {
results: vec![OsvResult {
vulns: vec![OsvVuln {
id: "MAL-2025-46966".to_string(),
}],
}],
};
let hits = extract_malicious_versioned(&pairs, &resp);
assert_eq!(hits.len(), 1);
assert_eq!(hits[0].package, "ansi-regex");
assert_eq!(hits[0].version.as_deref(), Some("6.2.1"));
assert_eq!(hits[0].advisory_id, "MAL-2025-46966");
}
#[test]
fn versioned_hit_prefers_explicit_affected_versions_over_broad_range() {
let details = OsvVulnDetails {
affected: vec![OsvAffected {
package: OsvAffectedPackage {
name: "@mistralai/mistralai".to_string(),
ecosystem: "npm".to_string(),
},
versions: vec![
"2.2.4".to_string(),
"2.2.3".to_string(),
"2.2.2".to_string(),
],
}],
};
assert!(!versioned_hit_affects_resolved_version(
&details,
"@mistralai/mistralai",
"2.2.1",
));
assert!(versioned_hit_affects_resolved_version(
&details,
"@mistralai/mistralai",
"2.2.2",
));
}
#[test]
fn versioned_hit_without_explicit_versions_still_blocks() {
let details = OsvVulnDetails {
affected: vec![OsvAffected {
package: OsvAffectedPackage {
name: "evil-pkg".to_string(),
ecosystem: "npm".to_string(),
},
versions: Vec::new(),
}],
};
assert!(versioned_hit_affects_resolved_version(
&details, "evil-pkg", "1.0.0",
));
}
#[test]
fn versioned_hit_without_matching_package_still_blocks() {
let details = OsvVulnDetails {
affected: vec![OsvAffected {
package: OsvAffectedPackage {
name: "evil-pkg".to_string(),
ecosystem: "PyPI".to_string(),
},
versions: vec!["1.0.0".to_string()],
}],
};
assert!(versioned_hit_affects_resolved_version(
&details, "evil-pkg", "1.0.0",
));
}
#[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"
);
}
}