Skip to main content

agentshield/fix/
dependencies.rs

1use std::path::Path;
2
3use regex::Regex;
4
5use super::{AppliedFix, FixOutput};
6
7/// Fix unpinned dependency specifications in requirements.txt or package.json (SHIELD-009).
8pub fn fix_unpinned_dependencies(content: &str, path: &Path) -> Option<FixOutput> {
9    let file_name = path.file_name().and_then(|f| f.to_str()).unwrap_or("");
10
11    if file_name == "requirements.txt" || file_name.ends_with(".requirements.txt") {
12        fix_requirements_txt(content)
13    } else if file_name == "package.json" {
14        fix_package_json(content)
15    } else {
16        None
17    }
18}
19
20fn fix_requirements_txt(content: &str) -> Option<FixOutput> {
21    let unpinned_re =
22        Regex::new(r"^([a-zA-Z0-9_.\-]+)\s*(?:>=|~=|>|\^)\s*([0-9a-zA-Z_.\-]+)(.*)$").ok()?;
23
24    let mut modified_lines = Vec::new();
25    let mut fixes = Vec::new();
26    let mut made_changes = false;
27
28    for (line_idx, line) in content.lines().enumerate() {
29        let trimmed = line.trim();
30        if trimmed.starts_with('#') || trimmed.is_empty() {
31            modified_lines.push(line.to_string());
32            continue;
33        }
34
35        if let Some(caps) = unpinned_re.captures(trimmed) {
36            let pkg = &caps[1];
37            let ver = &caps[2];
38            let rest = &caps[3];
39
40            let new_line = format!("{pkg}=={ver}{rest}");
41            fixes.push(AppliedFix {
42                rule_id: "SHIELD-009".into(),
43                description: format!("Pinned '{pkg}' to exact version '=={ver}'"),
44                line_number: line_idx + 1,
45            });
46            modified_lines.push(new_line);
47            made_changes = true;
48        } else {
49            modified_lines.push(line.to_string());
50        }
51    }
52
53    if !made_changes {
54        return None;
55    }
56
57    let mut output = modified_lines.join("\n");
58    if content.ends_with('\n') {
59        output.push('\n');
60    }
61
62    Some(FixOutput {
63        content: output,
64        fixes,
65    })
66}
67
68fn fix_package_json(content: &str) -> Option<FixOutput> {
69    let npm_unpinned_re =
70        Regex::new(r#"^(\s*"[^"]+"\s*:\s*)"[\^~>=]+([0-9a-zA-Z_.\-]+)"(.*)$"#).ok()?;
71
72    let mut modified_lines = Vec::new();
73    let mut fixes = Vec::new();
74    let mut made_changes = false;
75
76    for (line_idx, line) in content.lines().enumerate() {
77        if let Some(caps) = npm_unpinned_re.captures(line) {
78            let prefix = &caps[1];
79            let ver = &caps[2];
80            let suffix = &caps[3];
81
82            let new_line = format!("{prefix}\"{ver}\"{suffix}");
83            fixes.push(AppliedFix {
84                rule_id: "SHIELD-009".into(),
85                description: format!("Pinned npm dependency to exact version '{ver}'"),
86                line_number: line_idx + 1,
87            });
88            modified_lines.push(new_line);
89            made_changes = true;
90        } else {
91            modified_lines.push(line.to_string());
92        }
93    }
94
95    if !made_changes {
96        return None;
97    }
98
99    let mut output = modified_lines.join("\n");
100    if content.ends_with('\n') {
101        output.push('\n');
102    }
103
104    Some(FixOutput {
105        content: output,
106        fixes,
107    })
108}
109
110#[cfg(test)]
111mod tests {
112    use super::*;
113
114    #[test]
115    fn test_fix_requirements_txt_unpinned() {
116        let reqs = "requests>=2.31.0\nfastapi~=0.100.0\npytest==8.0.0\n";
117        let res = fix_unpinned_dependencies(reqs, Path::new("requirements.txt")).unwrap();
118        assert_eq!(
119            res.content,
120            "requests==2.31.0\nfastapi==0.100.0\npytest==8.0.0\n"
121        );
122        assert_eq!(res.fixes.len(), 2);
123    }
124
125    #[test]
126    fn test_fix_package_json_unpinned() {
127        let pkg = r#"{
128  "dependencies": {
129    "@modelcontextprotocol/sdk": "^1.0.0",
130    "express": "~4.18.2"
131  }
132}"#;
133        let res = fix_unpinned_dependencies(pkg, Path::new("package.json")).unwrap();
134        assert!(
135            res.content
136                .contains("\"@modelcontextprotocol/sdk\": \"1.0.0\"")
137        );
138        assert!(res.content.contains("\"express\": \"4.18.2\""));
139        assert_eq!(res.fixes.len(), 2);
140    }
141}