Skip to main content

gitcortex_indexer/parser/
python.rs

1use std::{
2    collections::HashMap,
3    path::{Path, PathBuf},
4};
5
6use gitcortex_core::{
7    error::{GitCortexError, Result},
8    graph::{Edge, Node, NodeId, NodeMetadata, Span},
9    schema::{EdgeKind, NodeKind, Visibility},
10};
11use tree_sitter::{Node as TsNode, Parser};
12
13use super::{LanguageParser, ParseResult};
14
15pub struct PythonParser {
16    language: tree_sitter::Language,
17}
18
19impl PythonParser {
20    pub fn new() -> Self {
21        Self {
22            language: tree_sitter_python::LANGUAGE.into(),
23        }
24    }
25}
26
27impl Default for PythonParser {
28    fn default() -> Self {
29        Self::new()
30    }
31}
32
33impl LanguageParser for PythonParser {
34    fn extensions(&self) -> &[&str] {
35        &["py"]
36    }
37
38    fn parse(&self, path: &Path, source: &str) -> Result<ParseResult> {
39        let mut parser = Parser::new();
40        parser
41            .set_language(&self.language)
42            .map_err(|e| GitCortexError::Parse {
43                file: path.to_owned(),
44                message: e.to_string(),
45            })?;
46
47        let tree = parser
48            .parse(source, None)
49            .ok_or_else(|| GitCortexError::Parse {
50                file: path.to_owned(),
51                message: "tree-sitter returned no parse tree".into(),
52            })?;
53
54        let mut visitor = FileVisitor::new(path, source);
55        visitor.collect_names(tree.root_node());
56        visitor.visit_module(tree.root_node());
57        visitor.collect_imports(tree.root_node());
58
59        Ok(ParseResult {
60            nodes: visitor.nodes,
61            edges: visitor.edges,
62            deferred_calls: visitor.deferred_calls,
63            deferred_uses: visitor.deferred_uses,
64            deferred_implements: visitor.deferred_implements,
65            deferred_imports: visitor.deferred_imports,
66            deferred_inherits: Vec::new(),
67            deferred_throws: Vec::new(),
68            deferred_annotated: visitor.deferred_annotated,
69        })
70    }
71}
72
73// ── Internal visitor ──────────────────────────────────────────────────────────
74
75struct FileVisitor<'src> {
76    source: &'src [u8],
77    file: PathBuf,
78    /// NodeId of the file-level Module node (anchor for Imports edges).
79    module_id: NodeId,
80    nodes: Vec<Node>,
81    edges: Vec<Edge>,
82    /// class name → NodeId (pass 1)
83    class_index: HashMap<String, NodeId>,
84    /// function/method name → NodeId (pass 1)
85    fn_index: HashMap<String, NodeId>,
86    deferred_calls: Vec<(NodeId, String)>,
87    deferred_uses: Vec<(NodeId, String)>,
88    deferred_implements: Vec<(NodeId, String)>,
89    deferred_imports: Vec<(NodeId, String)>,
90    deferred_annotated: Vec<(NodeId, String)>,
91}
92
93impl<'src> FileVisitor<'src> {
94    fn new(file: &Path, source: &'src str) -> Self {
95        let module_id = NodeId::new();
96        // Derive module name from the file stem (e.g. "auth" from "auth.py").
97        let module_name = file
98            .file_stem()
99            .and_then(|s| s.to_str())
100            .unwrap_or("__init__")
101            .to_owned();
102        let module_node = Node {
103            id: module_id.clone(),
104            qualified_name: module_name.clone(),
105            kind: NodeKind::Module,
106            name: module_name,
107            file: file.to_owned(),
108            span: Span {
109                start_line: 1,
110                end_line: 1,
111            },
112            metadata: NodeMetadata {
113                loc: source.lines().count() as u32,
114                visibility: Visibility::Pub,
115                is_async: false,
116                is_unsafe: false,
117                ..Default::default()
118            },
119        };
120        let nodes = vec![module_node];
121        Self {
122            source: source.as_bytes(),
123            file: file.to_owned(),
124            module_id,
125            nodes,
126            edges: Vec::new(),
127            class_index: HashMap::new(),
128            fn_index: HashMap::new(),
129            deferred_calls: Vec::new(),
130            deferred_uses: Vec::new(),
131            deferred_implements: Vec::new(),
132            deferred_imports: Vec::new(),
133            deferred_annotated: Vec::new(),
134        }
135    }
136
137    fn text<'t>(&self, node: TsNode<'t>) -> &'src str {
138        node.utf8_text(self.source).unwrap_or("")
139    }
140
141    fn span(node: TsNode<'_>) -> Span {
142        Span {
143            start_line: node.start_position().row as u32 + 1,
144            end_line: node.end_position().row as u32 + 1,
145        }
146    }
147
148    /// In Python, public = not starting with `_`.
149    fn visibility(name: &str) -> Visibility {
150        if name.starts_with('_') {
151            Visibility::Private
152        } else {
153            Visibility::Pub
154        }
155    }
156
157    fn qualified(scope: &[String], name: &str) -> String {
158        if scope.is_empty() {
159            name.to_owned()
160        } else {
161            format!("{}.{name}", scope.join("."))
162        }
163    }
164
165    fn make_node(
166        &self,
167        id: NodeId,
168        kind: NodeKind,
169        name: String,
170        scope: &[String],
171        ts_node: TsNode<'_>,
172        is_async: bool,
173    ) -> Node {
174        let vis = Self::visibility(&name);
175        Node {
176            id,
177            qualified_name: Self::qualified(scope, &name),
178            kind,
179            name,
180            file: self.file.clone(),
181            span: Self::span(ts_node),
182            metadata: NodeMetadata {
183                loc: (ts_node.end_position().row - ts_node.start_position().row + 1) as u32,
184                visibility: vis,
185                is_async,
186                is_unsafe: false,
187                ..Default::default()
188            },
189        }
190    }
191
192    // ── Pass 1: pre-allocate NodeIds ──────────────────────────────────────────
193
194    fn collect_names(&mut self, node: TsNode<'_>) {
195        let mut cursor = node.walk();
196        let children: Vec<TsNode<'_>> = node.named_children(&mut cursor).collect();
197        for child in children {
198            match child.kind() {
199                "class_definition" => {
200                    if let Some(name_node) = child.child_by_field_name("name") {
201                        let name = self.text(name_node).to_owned();
202                        self.class_index.entry(name).or_default();
203                    }
204                    // Recurse into class body to index nested classes and methods.
205                    if let Some(body) = child.child_by_field_name("body") {
206                        self.collect_names(body);
207                    }
208                }
209                "function_definition" => {
210                    if let Some(name_node) = child.child_by_field_name("name") {
211                        let name = self.text(name_node).to_owned();
212                        self.fn_index.entry(name).or_default();
213                    }
214                }
215                "decorated_definition" => {
216                    let def = child.child_by_field_name("definition");
217                    if let Some(def) = def {
218                        match def.kind() {
219                            "function_definition" => {
220                                if let Some(name_node) = def.child_by_field_name("name") {
221                                    let name = self.text(name_node).to_owned();
222                                    self.fn_index.entry(name).or_default();
223                                }
224                            }
225                            "class_definition" => {
226                                if let Some(name_node) = def.child_by_field_name("name") {
227                                    let name = self.text(name_node).to_owned();
228                                    self.class_index.entry(name).or_default();
229                                }
230                                // Recurse into decorated nested class body.
231                                if let Some(body) = def.child_by_field_name("body") {
232                                    self.collect_names(body);
233                                }
234                            }
235                            _ => {}
236                        }
237                    }
238                }
239                _ => {}
240            }
241        }
242    }
243
244    // ── Pass 2: emit nodes + edges ────────────────────────────────────────────
245
246    fn visit_module(&mut self, node: TsNode<'_>) {
247        let mut cursor = node.walk();
248        let children: Vec<TsNode<'_>> = node.named_children(&mut cursor).collect();
249        for child in children {
250            self.visit_top_level(child, &[]);
251        }
252    }
253
254    fn visit_top_level(&mut self, node: TsNode<'_>, scope: &[String]) {
255        match node.kind() {
256            "function_definition" => {
257                let is_async = Self::fn_is_async(node);
258                self.visit_function(node, scope, None, is_async, &[]);
259            }
260            "decorated_definition" => {
261                let decorators = self.collect_decorators(node);
262                let is_async = node
263                    .child_by_field_name("definition")
264                    .map(Self::fn_is_async)
265                    .unwrap_or(false);
266                if let Some(def) = node.child_by_field_name("definition") {
267                    match def.kind() {
268                        "function_definition" => {
269                            self.visit_function(def, scope, None, is_async, &decorators)
270                        }
271                        "class_definition" => self.visit_class(def, scope, &decorators),
272                        _ => {}
273                    }
274                }
275            }
276            "class_definition" => self.visit_class(node, scope, &[]),
277            "expression_statement" => self.maybe_visit_constant(node, scope),
278            _ => {}
279        }
280    }
281
282    fn visit_function(
283        &mut self,
284        node: TsNode<'_>,
285        scope: &[String],
286        container_id: Option<NodeId>,
287        is_async: bool,
288        decorators: &[String],
289    ) {
290        let Some(name_node) = node.child_by_field_name("name") else {
291            return;
292        };
293        let name = self.text(name_node).to_owned();
294        let id = self
295            .fn_index
296            .get(&name)
297            .cloned()
298            .unwrap_or_else(NodeId::new);
299
300        // Determine kind and metadata flags from decorators.
301        let has_property = decorators.iter().any(|d| d == "property");
302        let has_staticmethod = decorators.iter().any(|d| d == "staticmethod");
303        let has_classmethod = decorators.iter().any(|d| d == "classmethod");
304
305        let kind = if has_property {
306            NodeKind::Property
307        } else if container_id.is_some() {
308            NodeKind::Method
309        } else {
310            NodeKind::Function
311        };
312
313        // Check if the body contains a `yield` or `yield_from` → generator.
314        let is_generator = node
315            .child_by_field_name("body")
316            .map(|body| Self::body_has_yield(body))
317            .unwrap_or(false);
318
319        let mut graph_node = self.make_node(id.clone(), kind, name, scope, node, is_async);
320        if has_property {
321            graph_node.metadata.is_property = true;
322        }
323        if has_staticmethod || has_classmethod {
324            graph_node.metadata.is_static = true;
325        }
326        if is_generator {
327            graph_node.metadata.is_generator = true;
328        }
329
330        if let Some(cid) = container_id {
331            self.edges.push(Edge {
332                src: cid,
333                dst: id.clone(),
334                kind: EdgeKind::Contains,
335            });
336        }
337        self.nodes.push(graph_node);
338
339        // Type annotations → Uses edges
340        self.extract_param_types(node, &id);
341        self.extract_return_type(node, &id);
342
343        // Decorator names → Uses edges (e.g. @property, @staticmethod, @dataclass)
344        // and → deferred_annotated edges.
345        for dec in decorators {
346            self.deferred_uses.push((id.clone(), dec.clone()));
347            self.deferred_annotated.push((id.clone(), dec.clone()));
348        }
349
350        if let Some(body) = node.child_by_field_name("body") {
351            self.collect_calls(body, &id);
352        }
353    }
354
355    fn visit_class(&mut self, node: TsNode<'_>, scope: &[String], decorators: &[String]) {
356        let Some(name_node) = node.child_by_field_name("name") else {
357            return;
358        };
359        let name = self.text(name_node).to_owned();
360        let id = self
361            .class_index
362            .get(&name)
363            .cloned()
364            .unwrap_or_else(NodeId::new);
365
366        // Determine whether this class inherits from Protocol → Interface.
367        let mut is_protocol = false;
368        if let Some(bases) = node.child_by_field_name("superclasses") {
369            let mut c = bases.walk();
370            for base in bases.named_children(&mut c) {
371                let base_name = match base.kind() {
372                    "identifier" => Some(self.text(base).to_owned()),
373                    "attribute" => base
374                        .child_by_field_name("attribute")
375                        .map(|n| self.text(n).to_owned()),
376                    _ => None,
377                };
378                if base_name.as_deref() == Some("Protocol") {
379                    is_protocol = true;
380                }
381            }
382        }
383
384        let class_kind = if is_protocol {
385            NodeKind::Interface
386        } else {
387            NodeKind::Struct
388        };
389
390        let mut graph_node =
391            self.make_node(id.clone(), class_kind, name.clone(), scope, node, false);
392        if is_protocol {
393            graph_node.metadata.is_abstract = true;
394        }
395        self.nodes.push(graph_node);
396
397        // Base classes → Implements edges
398        if let Some(bases) = node.child_by_field_name("superclasses") {
399            let mut c = bases.walk();
400            for base in bases.named_children(&mut c) {
401                let base_name = match base.kind() {
402                    "identifier" => Some(self.text(base).to_owned()),
403                    "attribute" => base
404                        .child_by_field_name("attribute")
405                        .map(|n| self.text(n).to_owned()),
406                    _ => None,
407                };
408                if let Some(b) = base_name {
409                    self.deferred_implements.push((id.clone(), b));
410                }
411            }
412        }
413
414        // Decorator names → Uses edges (e.g. @dataclass) and → deferred_annotated.
415        for dec in decorators {
416            self.deferred_uses.push((id.clone(), dec.clone()));
417            self.deferred_annotated.push((id.clone(), dec.clone()));
418        }
419
420        let mut class_scope = scope.to_vec();
421        class_scope.push(name.clone());
422
423        if let Some(body) = node.child_by_field_name("body") {
424            let mut cursor = body.walk();
425            let children: Vec<TsNode<'_>> = body.named_children(&mut cursor).collect();
426            for child in children {
427                match child.kind() {
428                    "function_definition" => {
429                        let is_async = Self::fn_is_async(child);
430                        self.visit_function(child, &class_scope, Some(id.clone()), is_async, &[]);
431                    }
432                    "decorated_definition" => {
433                        let method_decorators = self.collect_decorators(child);
434                        let is_async = child
435                            .child_by_field_name("definition")
436                            .map(Self::fn_is_async)
437                            .unwrap_or(false);
438                        if let Some(def) = child.child_by_field_name("definition") {
439                            match def.kind() {
440                                "function_definition" => {
441                                    self.visit_function(
442                                        def,
443                                        &class_scope,
444                                        Some(id.clone()),
445                                        is_async,
446                                        &method_decorators,
447                                    );
448                                }
449                                "class_definition" => {
450                                    self.visit_class(def, &class_scope, &method_decorators);
451                                    // Add Contains edge from parent class to nested class.
452                                    if let Some(nested_name_node) = def.child_by_field_name("name")
453                                    {
454                                        let nested_name = self.text(nested_name_node).to_owned();
455                                        if let Some(nested_id) =
456                                            self.class_index.get(&nested_name).cloned()
457                                        {
458                                            self.edges.push(Edge {
459                                                src: id.clone(),
460                                                dst: nested_id,
461                                                kind: EdgeKind::Contains,
462                                            });
463                                        }
464                                    }
465                                }
466                                _ => {}
467                            }
468                        }
469                    }
470                    "class_definition" => {
471                        self.visit_class(child, &class_scope, &[]);
472                        // Add Contains edge from parent class to nested class.
473                        if let Some(nested_name_node) = child.child_by_field_name("name") {
474                            let nested_name = self.text(nested_name_node).to_owned();
475                            if let Some(nested_id) = self.class_index.get(&nested_name).cloned() {
476                                self.edges.push(Edge {
477                                    src: id.clone(),
478                                    dst: nested_id,
479                                    kind: EdgeKind::Contains,
480                                });
481                            }
482                        }
483                    }
484                    _ => {}
485                }
486            }
487        }
488    }
489
490    fn maybe_visit_constant(&mut self, node: TsNode<'_>, scope: &[String]) {
491        let mut cursor = node.walk();
492        for child in node.named_children(&mut cursor) {
493            if child.kind() == "assignment" {
494                if let Some(left) = child.child_by_field_name("left") {
495                    if left.kind() == "identifier" {
496                        let name = self.text(left).to_owned();
497                        if name
498                            .chars()
499                            .all(|c| c.is_uppercase() || c == '_' || c.is_ascii_digit())
500                            && name.len() > 1
501                            && !name.starts_with('_')
502                        {
503                            let id = NodeId::new();
504                            let graph_node =
505                                self.make_node(id, NodeKind::Constant, name, scope, node, false);
506                            self.nodes.push(graph_node);
507                        }
508                    }
509                }
510            }
511        }
512    }
513
514    // ── Pass 3: collect import statements ────────────────────────────────────
515
516    fn collect_imports(&mut self, node: TsNode<'_>) {
517        let mut cursor = node.walk();
518        let children: Vec<TsNode<'_>> = node.named_children(&mut cursor).collect();
519        for child in children {
520            match child.kind() {
521                "import_statement" => {
522                    // `import foo`, `import foo.bar`, `import foo as f`
523                    let mut c = child.walk();
524                    for name_node in child.named_children(&mut c) {
525                        let leaf = match name_node.kind() {
526                            "dotted_name" => {
527                                let text = self.text(name_node);
528                                text.split('.').next_back().map(|s| s.to_owned())
529                            }
530                            "aliased_import" => name_node
531                                .child_by_field_name("name")
532                                .map(|n| self.text(n))
533                                .map(|t| t.split('.').next_back().unwrap_or(t).to_owned()),
534                            _ => None,
535                        };
536                        if let Some(name) = leaf {
537                            self.deferred_imports.push((self.module_id.clone(), name));
538                        }
539                    }
540                }
541                "import_from_statement" => {
542                    // `from foo import bar, baz`
543                    // Named children: first is the source module (dotted_name or
544                    // relative_import), the rest are the imported names.
545                    let mut c = child.walk();
546                    let all_children: Vec<TsNode<'_>> = child.named_children(&mut c).collect();
547                    for name_node in all_children.iter().skip(1) {
548                        let leaf = match name_node.kind() {
549                            "dotted_name" => {
550                                let text = self.text(*name_node);
551                                Some(text.split('.').next_back().unwrap_or(text).to_owned())
552                            }
553                            "aliased_import" => name_node
554                                .child_by_field_name("name")
555                                .map(|n| self.text(n))
556                                .map(|t| t.split('.').next_back().unwrap_or(t).to_owned()),
557                            "wildcard_import" => None,
558                            _ => None,
559                        };
560                        if let Some(name) = leaf {
561                            self.deferred_imports.push((self.module_id.clone(), name));
562                        }
563                    }
564                }
565                _ => {}
566            }
567        }
568    }
569
570    // ── Call collection ───────────────────────────────────────────────────────
571
572    fn collect_calls(&mut self, node: TsNode<'_>, caller_id: &NodeId) {
573        let mut cursor = node.walk();
574        let children: Vec<TsNode<'_>> = node.named_children(&mut cursor).collect();
575        for child in children {
576            if child.kind() == "call" {
577                if let Some(callee) = self.callee_name(child) {
578                    self.record_call(caller_id.clone(), callee);
579                }
580                if let Some(args) = child.child_by_field_name("arguments") {
581                    self.collect_calls(args, caller_id);
582                }
583            } else {
584                self.collect_calls(child, caller_id);
585            }
586        }
587    }
588
589    fn callee_name(&self, call_node: TsNode<'_>) -> Option<String> {
590        let func = call_node.child_by_field_name("function")?;
591        match func.kind() {
592            "identifier" => Some(self.text(func).to_owned()),
593            "attribute" => func
594                .child_by_field_name("attribute")
595                .map(|n| self.text(n).to_owned()),
596            _ => None,
597        }
598    }
599
600    fn record_call(&mut self, caller_id: NodeId, callee_name: String) {
601        if callee_name.is_empty() {
602            return;
603        }
604        if let Some(callee_id) = self.fn_index.get(&callee_name).cloned() {
605            let edge = Edge {
606                src: caller_id,
607                dst: callee_id,
608                kind: EdgeKind::Calls,
609            };
610            if !self.edges.contains(&edge) {
611                self.edges.push(edge);
612            }
613        } else if !self
614            .deferred_calls
615            .iter()
616            .any(|(c, n)| c == &caller_id && n == &callee_name)
617        {
618            self.deferred_calls.push((caller_id, callee_name));
619        }
620    }
621
622    // ── Helpers ───────────────────────────────────────────────────────────────
623
624    /// Returns true if this `function_definition` node is `async def`.
625    fn fn_is_async(node: TsNode<'_>) -> bool {
626        let mut c = node.walk();
627        // `async` appears as an anonymous child (keyword) before `def`
628        let result = node.children(&mut c).any(|n| n.kind() == "async");
629        result
630    }
631
632    /// Returns true if the body subtree contains a `yield` or `yield_from` expression.
633    fn body_has_yield(node: TsNode<'_>) -> bool {
634        if node.kind() == "yield" || node.kind() == "yield_from" {
635            return true;
636        }
637        // Don't descend into nested function definitions — their yields are not
638        // generators of the outer function.
639        if node.kind() == "function_definition" {
640            return false;
641        }
642        let mut c = node.walk();
643        let found = node.named_children(&mut c).any(Self::body_has_yield);
644        found
645    }
646
647    /// Extract decorator names from a `decorated_definition` node.
648    fn collect_decorators(&self, node: TsNode<'_>) -> Vec<String> {
649        let mut c = node.walk();
650        node.named_children(&mut c)
651            .filter(|n| n.kind() == "decorator")
652            .filter_map(|d| self.decorator_name(d))
653            .collect()
654    }
655
656    /// Get the callable name from a `decorator` node.
657    fn decorator_name(&self, decorator: TsNode<'_>) -> Option<String> {
658        let mut c = decorator.walk();
659        let child = decorator.named_children(&mut c).next()?;
660        match child.kind() {
661            "identifier" => Some(self.text(child).to_owned()),
662            "attribute" => child
663                .child_by_field_name("attribute")
664                .map(|n| self.text(n).to_owned()),
665            "call" => child
666                .child_by_field_name("function")
667                .and_then(|f| match f.kind() {
668                    "identifier" => Some(self.text(f).to_owned()),
669                    "attribute" => f
670                        .child_by_field_name("attribute")
671                        .map(|n| self.text(n).to_owned()),
672                    _ => None,
673                }),
674            _ => None,
675        }
676    }
677
678    /// Extract type identifiers from a type-annotation node.
679    /// Records all identifiers found (e.g. `List[MyType]` → ["List", "MyType"]).
680    fn extract_param_types(&mut self, fn_node: TsNode<'_>, fn_id: &NodeId) {
681        let Some(params) = fn_node.child_by_field_name("parameters") else {
682            return;
683        };
684        let mut c = params.walk();
685        for param in params.named_children(&mut c) {
686            let type_node = match param.kind() {
687                "typed_parameter" | "typed_default_parameter" => param.child_by_field_name("type"),
688                _ => None,
689            };
690            if let Some(t) = type_node {
691                for name in self.collect_type_names(t) {
692                    self.deferred_uses.push((fn_id.clone(), name));
693                }
694            }
695        }
696    }
697
698    fn extract_return_type(&mut self, fn_node: TsNode<'_>, fn_id: &NodeId) {
699        if let Some(ret) = fn_node.child_by_field_name("return_type") {
700            for name in self.collect_type_names(ret) {
701                self.deferred_uses.push((fn_id.clone(), name));
702            }
703        }
704    }
705
706    /// Walk a type-annotation subtree and collect all identifier names.
707    fn collect_type_names(&self, node: TsNode<'_>) -> Vec<String> {
708        let mut names = Vec::new();
709        self.walk_type_names(node, &mut names);
710        names
711    }
712
713    fn walk_type_names(&self, node: TsNode<'_>, out: &mut Vec<String>) {
714        match node.kind() {
715            "identifier" => {
716                let name = self.text(node).to_owned();
717                if !is_builtin_type(&name) {
718                    out.push(name);
719                }
720            }
721            _ => {
722                let mut c = node.walk();
723                for child in node.named_children(&mut c) {
724                    self.walk_type_names(child, out);
725                }
726            }
727        }
728    }
729}
730
731/// Returns true for built-in Python types and typing-module generics that do
732/// not correspond to user-defined symbols.
733fn is_builtin_type(name: &str) -> bool {
734    matches!(
735        name,
736        "int"
737            | "str"
738            | "bool"
739            | "float"
740            | "complex"
741            | "bytes"
742            | "bytearray"
743            | "None"
744            | "list"
745            | "dict"
746            | "set"
747            | "frozenset"
748            | "tuple"
749            | "type"
750            | "object"
751            | "Any"
752            | "Optional"
753            | "Union"
754            | "List"
755            | "Dict"
756            | "Set"
757            | "FrozenSet"
758            | "Tuple"
759            | "Callable"
760            | "Type"
761            | "ClassVar"
762            | "Final"
763            | "Literal"
764            | "TypeVar"
765            | "Generic"
766            | "Protocol"
767            | "Sequence"
768            | "Iterable"
769            | "Iterator"
770            | "Generator"
771            | "Coroutine"
772            | "Awaitable"
773            | "AsyncIterator"
774            | "AsyncGenerator"
775            | "NoReturn"
776            | "Never"
777            | "Self"
778            | "Annotated"
779            | "TypeAlias"
780            | "ParamSpec"
781            | "TypeVarTuple"
782            | "overload"
783            | "abstractmethod"
784            | "staticmethod"
785            | "classmethod"
786    )
787}
788
789// ── Tests ─────────────────────────────────────────────────────────────────────
790
791#[cfg(test)]
792mod tests {
793    use super::PythonParser;
794    use crate::parser::LanguageParser;
795    use gitcortex_core::schema::{EdgeKind, NodeKind};
796    use std::path::Path;
797
798    fn parse(
799        src: &str,
800    ) -> (
801        Vec<gitcortex_core::graph::Node>,
802        Vec<gitcortex_core::graph::Edge>,
803    ) {
804        let r = PythonParser::new()
805            .parse(Path::new("test.py"), src)
806            .unwrap();
807        (r.nodes, r.edges)
808    }
809
810    #[allow(clippy::type_complexity)]
811    fn parse_full(
812        src: &str,
813    ) -> (
814        Vec<gitcortex_core::graph::Node>,
815        Vec<gitcortex_core::graph::Edge>,
816        Vec<(gitcortex_core::graph::NodeId, String)>,
817        Vec<(gitcortex_core::graph::NodeId, String)>,
818        Vec<(gitcortex_core::graph::NodeId, String)>,
819        Vec<(gitcortex_core::graph::NodeId, String)>,
820    ) {
821        let r = PythonParser::new()
822            .parse(Path::new("test.py"), src)
823            .unwrap();
824        (
825            r.nodes,
826            r.edges,
827            r.deferred_calls,
828            r.deferred_uses,
829            r.deferred_implements,
830            r.deferred_imports,
831        )
832    }
833
834    #[test]
835    fn parses_free_function() {
836        let (nodes, _) = parse("def greet(name):\n    return name\n");
837        // Module node + Function node
838        assert_eq!(nodes.len(), 2);
839        let fns: Vec<_> = nodes
840            .iter()
841            .filter(|n| n.kind == NodeKind::Function)
842            .collect();
843        assert_eq!(fns.len(), 1);
844        assert_eq!(fns[0].name, "greet");
845    }
846
847    #[test]
848    fn parses_class_and_method() {
849        let src = "class Person:\n    def greet(self):\n        pass\n";
850        let (nodes, edges) = parse(src);
851        let classes: Vec<_> = nodes
852            .iter()
853            .filter(|n| n.kind == NodeKind::Struct)
854            .collect();
855        let methods: Vec<_> = nodes
856            .iter()
857            .filter(|n| n.kind == NodeKind::Method)
858            .collect();
859        assert_eq!(classes.len(), 1);
860        assert_eq!(methods.len(), 1);
861        let contains: Vec<_> = edges
862            .iter()
863            .filter(|e| e.kind == EdgeKind::Contains)
864            .collect();
865        assert!(!contains.is_empty());
866    }
867
868    #[test]
869    fn detects_call_edges() {
870        let src = "def caller():\n    callee()\ndef callee():\n    pass\n";
871        let (_, edges) = parse(src);
872        let calls: Vec<_> = edges.iter().filter(|e| e.kind == EdgeKind::Calls).collect();
873        assert_eq!(calls.len(), 1);
874    }
875
876    #[test]
877    fn detects_base_class_implements() {
878        let src = "class Base:\n    pass\nclass Child(Base):\n    pass\n";
879        let (_, _, _, _, implements, _) = parse_full(src);
880        assert!(
881            implements.iter().any(|(_, name)| name == "Base"),
882            "expected Implements edge to Base, got: {implements:?}"
883        );
884    }
885
886    #[test]
887    fn detects_type_annotation_uses() {
888        let src = "class Foo:\n    pass\ndef bar(x: Foo) -> Foo:\n    pass\n";
889        let (_, _, _, uses, _, _) = parse_full(src);
890        let uses_foo: Vec<_> = uses.iter().filter(|(_, n)| n == "Foo").collect();
891        assert!(
892            uses_foo.len() >= 2,
893            "expected at least 2 Uses edges to Foo, got: {uses_foo:?}"
894        );
895    }
896
897    #[test]
898    fn detects_decorator_uses() {
899        let src = "class Foo:\n    @property\n    def name(self):\n        return self._name\n";
900        let (_, _, _, uses, _, _) = parse_full(src);
901        assert!(
902            uses.iter().any(|(_, n)| n == "property"),
903            "expected Uses edge to 'property' decorator, got: {uses:?}"
904        );
905    }
906
907    #[test]
908    fn detects_import_statement() {
909        let src = "import os\nimport sys\n\ndef main():\n    pass\n";
910        let (_, _, _, _, _, imports) = parse_full(src);
911        assert!(
912            imports.iter().any(|(_, n)| n == "os"),
913            "expected import 'os', got: {imports:?}"
914        );
915        assert!(
916            imports.iter().any(|(_, n)| n == "sys"),
917            "expected import 'sys', got: {imports:?}"
918        );
919    }
920
921    #[test]
922    fn detects_from_import_statement() {
923        let src = "from os.path import join, exists\n\ndef main():\n    pass\n";
924        let (_, _, _, _, _, imports) = parse_full(src);
925        assert!(
926            imports.iter().any(|(_, n)| n == "join"),
927            "expected import 'join', got: {imports:?}"
928        );
929        assert!(
930            imports.iter().any(|(_, n)| n == "exists"),
931            "expected import 'exists', got: {imports:?}"
932        );
933    }
934
935    #[test]
936    fn module_node_is_emitted() {
937        let (nodes, _) = parse("x = 1\n");
938        let modules: Vec<_> = nodes
939            .iter()
940            .filter(|n| n.kind == NodeKind::Module)
941            .collect();
942        assert_eq!(modules.len(), 1);
943        assert_eq!(modules[0].name, "test");
944    }
945
946    #[test]
947    fn async_function_flagged() {
948        let src = "async def fetch():\n    pass\n";
949        let (nodes, _) = parse(src);
950        let fns: Vec<_> = nodes
951            .iter()
952            .filter(|n| n.kind == NodeKind::Function)
953            .collect();
954        assert_eq!(fns.len(), 1);
955        assert!(fns[0].metadata.is_async);
956    }
957}