From c413303ae01589e216fa8c60fb51bdb35a146e4d Mon Sep 17 00:00:00 2001 From: Jasper Date: Sun, 2 Aug 2026 07:48:50 -0500 Subject: [PATCH] fix: parse PyPI requirements for OSV checks (#10510) --- .../src/agents/extension_malware_check.rs | 320 +++++++++++++++++- 1 file changed, 304 insertions(+), 16 deletions(-) diff --git a/crates/goose/src/agents/extension_malware_check.rs b/crates/goose/src/agents/extension_malware_check.rs index 29b8d566c..6191122e3 100644 --- a/crates/goose/src/agents/extension_malware_check.rs +++ b/crates/goose/src/agents/extension_malware_check.rs @@ -131,24 +131,121 @@ fn parse_npm_token(token: &str) -> Option<(String, Option)> { } fn parse_pypi_token(token: &str) -> Option<(String, Option)> { - // Accept exact pins: - // package==1.2.3 - // package[extra]==1.2.3 - // Treat "latest" as None. Ignore other specifiers (>=, <=, ~=, !=) for pinning. - let lowered = token.to_ascii_lowercase(); - if let Some(idx) = lowered.find("==") { - let (name, ver) = token.split_at(idx); - let ver = ver.trim_start_matches('=').trim_start_matches('='); - let name = name.trim(); - if name.is_empty() { + let token = token.trim(); + let name_end = token + .char_indices() + .find_map(|(index, character)| { + (!character.is_ascii_alphanumeric() && !matches!(character, '-' | '_' | '.')) + .then_some(index) + }) + .unwrap_or(token.len()); + let (name, requirement) = token.split_at(name_end); + let name = name.trim(); + if name.is_empty() + || !name + .chars() + .next() + .is_some_and(|character| character.is_ascii_alphanumeric()) + || !name + .chars() + .next_back() + .is_some_and(|character| character.is_ascii_alphanumeric()) + { + return None; + } + + let mut raw_requirement = requirement; + let mut requirement = raw_requirement.trim_start(); + if !requirement.is_empty() + && !requirement + .chars() + .next() + .is_some_and(|character| matches!(character, '[' | '@' | ';' | '(')) + && !["===", "~=", "==", "!=", "<=", ">=", "<", ">"] + .iter() + .any(|operator| requirement.starts_with(operator)) + { + return None; + } + if let Some(extras) = requirement.strip_prefix('[') { + let (_, remaining) = extras.split_once(']')?; + raw_requirement = remaining; + requirement = remaining.trim_start(); + } + + let version = parse_pypi_exact_version(requirement).or_else(|| { + raw_requirement + .strip_prefix('@') + .and_then(parse_concrete_version) + }); + + Some((normalize_pypi_name(name), version)) +} + +fn parse_pypi_exact_version(requirement: &str) -> Option { + let requirement = requirement.trim(); + let comparison = if let Some(parenthesized) = requirement + .strip_prefix('(') + .and_then(|requirement| requirement.strip_suffix(')')) + { + let comparison = parenthesized.trim(); + if comparison.contains(['(', ')']) { return None; } - if ver.is_empty() || ver.eq_ignore_ascii_case("latest") { - return Some((name.to_string(), None)); + comparison + } else { + if requirement.contains(['(', ')']) { + return None; } - return Some((name.to_string(), Some(ver.to_string()))); + requirement + }; + + let version = comparison.strip_prefix("==")?.trim(); + if version.starts_with('=') { + return None; } - Some((token.to_string(), None)) + parse_concrete_version(version) +} + +fn parse_concrete_version(version: &str) -> Option { + let version = version.trim(); + let numeric_version = version + .strip_prefix('v') + .or_else(|| version.strip_prefix('V')) + .unwrap_or(version); + if version.is_empty() + || version.eq_ignore_ascii_case("latest") + || !numeric_version + .chars() + .next() + .is_some_and(|character| character.is_ascii_digit()) + || !version.chars().all(|character| { + character.is_ascii_alphanumeric() || matches!(character, '.' | '-' | '_' | '+' | '!') + }) + { + return None; + } + + Some(version.to_string()) +} + +fn normalize_pypi_name(name: &str) -> String { + let mut normalized = String::with_capacity(name.len()); + let mut pending_separator = false; + + for character in name.chars() { + if matches!(character, '-' | '_' | '.') { + pending_separator = true; + } else { + if pending_separator && !normalized.is_empty() { + normalized.push('-'); + } + normalized.push(character.to_ascii_lowercase()); + pending_separator = false; + } + } + + normalized } const DEFAULT_OSV_ENDPOINT: &str = "https://api.osv.dev/v1/query"; @@ -288,7 +385,7 @@ mod tests { use serde_json::json; use serial_test; use tokio; - use wiremock::matchers::{method, path}; + use wiremock::matchers::{body_json, method, path}; use wiremock::{Mock, MockServer, ResponseTemplate}; fn checker_for(server: &MockServer) -> OsvChecker { @@ -383,6 +480,70 @@ mod tests { assert!(res.is_ok()); } + #[tokio::test] + #[serial_test::serial] + async fn cmd_args_pypi_normalizes_requirement_before_osv_query() { + let server = MockServer::start().await; + Mock::given(method("POST")) + .and(path("/v1/query")) + .and(body_json(json!({ + "package": { + "name": "evil-pkg", + "ecosystem": "PyPI" + } + }))) + .respond_with(ResponseTemplate::new(200).set_body_json(json!({ + "vulns": [ { "id": "MAL-240", "summary": "Malicious package" } ], + "next_page_token": null + }))) + .mount(&server) + .await; + + let _env = TempEnvVar::set("OSV_ENDPOINT", &format!("{}/v1/query", server.uri())); + let result = deny_if_malicious_cmd_args("uvx", &["Evil_Pkg[cli]>=0".to_string()]).await; + + assert!(result.is_err()); + assert!(format!("{:?}", result.unwrap_err()).contains("MAL-240")); + } + + #[tokio::test] + #[serial_test::serial] + async fn cmd_args_pypi_preserves_uvx_exact_version() { + let server = MockServer::start().await; + Mock::given(method("POST")) + .and(path("/v1/query")) + .and(body_json(json!({ + "version": "0.3.0", + "package": { + "name": "ruff", + "ecosystem": "PyPI" + } + }))) + .respond_with(ResponseTemplate::new(200).set_body_json(json!({ + "vulns": [], + "next_page_token": null + }))) + .mount(&server) + .await; + + let _env = TempEnvVar::set("OSV_ENDPOINT", &format!("{}/v1/query", server.uri())); + let uvx_result = deny_if_malicious_cmd_args("uvx", &["Ruff@0.3.0".to_string()]).await; + let pep_508_result = deny_if_malicious_cmd_args("uvx", &["Ruff== 0.3.0".to_string()]).await; + let parenthesized_result = deny_if_malicious_cmd_args( + "uvx", + &[ + "--from".to_string(), + "Ruff (==0.3.0)".to_string(), + "ruff".to_string(), + ], + ) + .await; + + assert!(uvx_result.is_ok()); + assert!(pep_508_result.is_ok()); + assert!(parenthesized_result.is_ok()); + } + #[tokio::test] #[serial_test::serial] async fn cmd_args_npm_scoped_malicious() { @@ -498,7 +659,7 @@ mod tests { } #[test] - fn parse_pypi_exact_pin_and_latest() { + fn parse_pypi_requirement_uses_canonical_project_name() { assert_eq!( super::parse_pypi_token("requests==2.32.3"), Some(("requests".into(), Some("2.32.3".into()))) @@ -507,5 +668,132 @@ mod tests { super::parse_pypi_token("requests==latest"), Some(("requests".into(), None)) ); + + for requirement in [ + "requests>=2.32", + "requests<=2.32", + "requests~=2.32", + "requests!=2.32", + "requests>2.32", + "requests<3", + "requests===custom", + "requests==2.32.*", + "requests==2.32,!=2.32.1", + "requests@https://example.invalid/requests.whl", + ] { + assert_eq!( + super::parse_pypi_token(requirement), + Some(("requests".into(), None)), + "failed to parse {requirement}" + ); + } + + assert_eq!( + super::parse_pypi_token("requests[security]==2.32.3"), + Some(("requests".into(), Some("2.32.3".into()))) + ); + for requirement in [ + "requests== 2.32.3", + "requests == 2.32.3", + "requests[security] == 2.32.3", + "requests (==2.32.3)", + "requests( == 2.32.3 )", + "requests[security] (==2.32.3)", + ] { + assert_eq!( + super::parse_pypi_token(requirement), + Some(("requests".into(), Some("2.32.3".into()))), + "failed to preserve exact pin in {requirement}" + ); + } + assert_eq!( + super::parse_pypi_token("requests[security,socks]>=2.32"), + Some(("requests".into(), None)) + ); + assert_eq!( + super::parse_pypi_token("zope.interface"), + Some(("zope-interface".into(), None)) + ); + assert_eq!( + super::parse_pypi_token("Friendly-._.-Bard"), + Some(("friendly-bard".into(), None)) + ); + assert_eq!( + super::parse_pypi_token("Evil_Pkg[cli]>=0"), + Some(("evil-pkg".into(), None)) + ); + assert_eq!( + super::parse_pypi_token("evilpkg @ https://example.invalid/pkg.whl"), + Some(("evilpkg".into(), None)) + ); + assert_eq!( + super::parse_pypi_token("evilpkg; python_version < '3.12'"), + Some(("evilpkg".into(), None)) + ); + assert_eq!( + super::parse_pypi_token("evilpkg (>=1.0)"), + Some(("evilpkg".into(), None)) + ); + assert_eq!(super::parse_pypi_token("[security]>=2.32"), None); + assert_eq!( + super::parse_pypi_token("ruff@0.3.0"), + Some(("ruff".into(), Some("0.3.0".into()))) + ); + assert_eq!( + super::parse_pypi_token("ruff@v0.3.0"), + Some(("ruff".into(), Some("v0.3.0".into()))) + ); + for direct_reference in [ + "ruff@latest", + "ruff@https://example.invalid/ruff.whl", + "ruff@git+https://example.invalid/ruff.git", + "ruff@file:///tmp/ruff.whl", + "ruff@../ruff", + "ruff @ 0.3.0", + "ruff=== 0.3.0", + "ruff== 0.3.*", + "ruff== 0.3.0,!=0.3.1", + "ruff== 0.3.0 ; python_version < '3.12'", + "ruff (==0.3.0,!=0.3.1)", + "ruff (>=0.3.0)", + "ruff (==0.3.*)", + "ruff (==0.3.0); python_version < '3.12'", + "ruff (===0.3.0)", + "ruff (==0.3.0", + "ruff ==0.3.0)", + "ruff ((==0.3.0))", + "ruff (==0.3.0) trailing", + ] { + assert_eq!( + super::parse_pypi_token(direct_reference), + Some(("ruff".into(), None)), + "unexpected pin in {direct_reference}" + ); + } + for invalid in [ + "", + " ", + "==1.0", + "https://example.invalid/pkg.whl", + "git+https://example.invalid/pkg.git", + "./local", + "-invalid", + "invalid-", + "requests[security", + ] { + assert_eq!( + super::parse_pypi_token(invalid), + None, + "unexpected package name in {invalid}" + ); + } + } + + #[test] + fn parse_first_pypi_argument_normalizes_name_before_osv_query() { + assert_eq!( + super::parse_first_package_arg("PyPI", &["Evil_Pkg[cli]>=0".to_string()]), + Some(("evil-pkg".into(), None)) + ); } }