Skip to main content

agentshield/adapter/
langchain.rs

1use std::path::Path;
2
3use crate::config::ScanPathFilter;
4use crate::error::Result;
5use crate::ir::taint_builder::build_data_surface;
6use crate::ir::*;
7
8/// LangChain framework adapter.
9///
10/// Detects LangChain projects by looking for:
11/// - `pyproject.toml` with `langchain` dependency
12/// - `requirements.txt` containing `langchain` or `langgraph`
13/// - `langgraph.json` configuration file
14/// - Python files importing `from langchain` / `from langchain_core` / `from langgraph`
15pub struct LangChainAdapter;
16
17impl super::Adapter for LangChainAdapter {
18    fn framework(&self) -> Framework {
19        Framework::LangChain
20    }
21
22    fn detect(&self, root: &Path) -> bool {
23        // Check pyproject.toml for langchain dependency
24        let pyproject = root.join("pyproject.toml");
25        if pyproject.exists() {
26            if let Ok(content) = std::fs::read_to_string(&pyproject) {
27                if content.contains("langchain") || content.contains("langgraph") {
28                    return true;
29                }
30            }
31        }
32
33        // Check requirements.txt for langchain/langgraph
34        let requirements = root.join("requirements.txt");
35        if requirements.exists() {
36            if let Ok(content) = std::fs::read_to_string(&requirements) {
37                if content.lines().any(|l| {
38                    let trimmed = l.trim();
39                    trimmed.starts_with("langchain") || trimmed.starts_with("langgraph")
40                }) {
41                    return true;
42                }
43            }
44        }
45
46        // Check for langgraph.json configuration file
47        if root.join("langgraph.json").exists() {
48            return true;
49        }
50
51        // Check package.json for @langchain dependencies
52        let package_json = root.join("package.json");
53        if package_json.exists() {
54            if let Ok(content) = std::fs::read_to_string(&package_json) {
55                if content.contains("@langchain/")
56                    || content.contains("\"langchain\"")
57                    || content.contains("@langchain/core")
58                {
59                    return true;
60                }
61            }
62        }
63
64        if super::mcp::has_recursive_python_import(
65            root,
66            &[
67                "from langchain",
68                "import langchain",
69                "from langgraph",
70                "import langgraph",
71            ],
72        ) {
73            return true;
74        }
75
76        false
77    }
78
79    fn load(&self, root: &Path, ignore_tests: bool) -> Result<Vec<ScanTarget>> {
80        let filter = ScanPathFilter::for_ignore_tests(ignore_tests);
81        self.load_with_filter(root, &filter)
82    }
83
84    fn load_with_filter(&self, root: &Path, filter: &ScanPathFilter) -> Result<Vec<ScanTarget>> {
85        let name = root
86            .file_name()
87            .map(|n| n.to_string_lossy().to_string())
88            .unwrap_or_else(|| "langchain-project".into());
89
90        let mut source_files = Vec::new();
91        // Phase 0: Collect source files (reuses MCP adapter's walker)
92        super::mcp::collect_source_files_with_filter(root, filter, &mut source_files)?;
93
94        // Retain Python and TypeScript/JavaScript source files for LangChain
95        source_files.retain(|sf| {
96            matches!(
97                sf.language,
98                Language::Python | Language::TypeScript | Language::JavaScript
99            )
100        });
101
102        let execution = super::pipeline::build_execution_surface(&source_files);
103
104        // Parse dependencies from pyproject.toml / requirements.txt / package.json
105        let dependencies = super::mcp::parse_dependencies(root, filter);
106
107        // Parse provenance from pyproject.toml / package.json
108        let provenance = super::mcp::parse_provenance(root, filter);
109
110        let tools = vec![];
111        let data = build_data_surface(&tools, &execution);
112
113        Ok(vec![ScanTarget {
114            name,
115            framework: Framework::LangChain,
116            root_path: root.to_path_buf(),
117            tools,
118            execution,
119            data,
120            dependencies,
121            provenance,
122            source_files,
123        }])
124    }
125}
126
127#[cfg(test)]
128mod tests {
129    use super::*;
130    use crate::adapter::Adapter;
131    use std::io::Write;
132    use std::path::PathBuf;
133    use tempfile::TempDir;
134
135    #[test]
136    fn test_detect_langchain_via_pyproject() {
137        let dir =
138            PathBuf::from(env!("CARGO_MANIFEST_DIR")).join("tests/fixtures/langchain_project");
139        let adapter = LangChainAdapter;
140        assert!(adapter.detect(&dir));
141    }
142
143    #[test]
144    fn test_detect_langchain_via_langgraph_json() {
145        let dir =
146            PathBuf::from(env!("CARGO_MANIFEST_DIR")).join("tests/fixtures/langchain_project");
147        let adapter = LangChainAdapter;
148        // The fixture has pyproject.toml, but langgraph.json also triggers detection
149        assert!(adapter.detect(&dir));
150    }
151
152    #[test]
153    fn test_detect_non_langchain_project() {
154        let dir = PathBuf::from(env!("CARGO_MANIFEST_DIR"))
155            .join("tests/fixtures/mcp_servers/safe_calculator");
156        let adapter = LangChainAdapter;
157        assert!(!adapter.detect(&dir));
158    }
159
160    #[test]
161    fn test_load_langchain_finds_cmd_injection() {
162        let dir =
163            PathBuf::from(env!("CARGO_MANIFEST_DIR")).join("tests/fixtures/langchain_project");
164        let adapter = LangChainAdapter;
165        let targets = adapter.load(&dir, false).unwrap();
166        assert_eq!(targets.len(), 1);
167
168        let target = &targets[0];
169        assert_eq!(target.framework, Framework::LangChain);
170        assert_eq!(target.name, "langchain_project");
171
172        // Should find command injection in shell_tool.py
173        assert!(
174            !target.execution.commands.is_empty(),
175            "expected command execution findings from shell_tool.py"
176        );
177        // Should find tainted command args
178        assert!(
179            target
180                .execution
181                .commands
182                .iter()
183                .any(|c| c.command_arg.is_tainted()),
184            "expected tainted command source from subprocess.run with user input"
185        );
186    }
187
188    #[test]
189    fn test_load_langchain_finds_ssrf() {
190        let dir =
191            PathBuf::from(env!("CARGO_MANIFEST_DIR")).join("tests/fixtures/langchain_project");
192        let adapter = LangChainAdapter;
193        let targets = adapter.load(&dir, false).unwrap();
194        let target = &targets[0];
195
196        // Should find network operations in fetch_tool.py
197        assert!(
198            !target.execution.network_operations.is_empty(),
199            "expected network operation findings from fetch_tool.py"
200        );
201    }
202
203    #[test]
204    fn test_load_langchain_only_python_files() {
205        let dir =
206            PathBuf::from(env!("CARGO_MANIFEST_DIR")).join("tests/fixtures/langchain_project");
207        let adapter = LangChainAdapter;
208        let targets = adapter.load(&dir, false).unwrap();
209        let target = &targets[0];
210
211        // All source files should be Python
212        for sf in &target.source_files {
213            assert_eq!(
214                sf.language,
215                Language::Python,
216                "non-Python file found: {:?}",
217                sf.path
218            );
219        }
220    }
221
222    #[test]
223    fn test_detect_langchain_nested_python_import() {
224        let tmp_root = TempDir::new().unwrap();
225        let nested_dir = tmp_root.path().join("src/pkg");
226        std::fs::create_dir_all(&nested_dir).unwrap();
227
228        let mut file = std::fs::File::create(nested_dir.join("tool.py")).unwrap();
229        writeln!(file, "from langchain_core import chat_models").unwrap();
230
231        let adapter = LangChainAdapter;
232        assert!(adapter.detect(tmp_root.path()));
233    }
234}