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::{capture_definition, 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                definition: capture_definition(self.source, ts_node),
188                ..Default::default()
189            },
190        }
191    }
192
193    // ── Pass 1: pre-allocate NodeIds ──────────────────────────────────────────
194
195    fn collect_names(&mut self, node: TsNode<'_>) {
196        let mut cursor = node.walk();
197        let children: Vec<TsNode<'_>> = node.named_children(&mut cursor).collect();
198        for child in children {
199            match child.kind() {
200                "class_definition" => {
201                    if let Some(name_node) = child.child_by_field_name("name") {
202                        let name = self.text(name_node).to_owned();
203                        self.class_index.entry(name).or_default();
204                    }
205                    // Recurse into class body to index nested classes and methods.
206                    if let Some(body) = child.child_by_field_name("body") {
207                        self.collect_names(body);
208                    }
209                }
210                "function_definition" => {
211                    if let Some(name_node) = child.child_by_field_name("name") {
212                        let name = self.text(name_node).to_owned();
213                        self.fn_index.entry(name).or_default();
214                    }
215                }
216                "decorated_definition" => {
217                    let def = child.child_by_field_name("definition");
218                    if let Some(def) = def {
219                        match def.kind() {
220                            "function_definition" => {
221                                if let Some(name_node) = def.child_by_field_name("name") {
222                                    let name = self.text(name_node).to_owned();
223                                    self.fn_index.entry(name).or_default();
224                                }
225                            }
226                            "class_definition" => {
227                                if let Some(name_node) = def.child_by_field_name("name") {
228                                    let name = self.text(name_node).to_owned();
229                                    self.class_index.entry(name).or_default();
230                                }
231                                // Recurse into decorated nested class body.
232                                if let Some(body) = def.child_by_field_name("body") {
233                                    self.collect_names(body);
234                                }
235                            }
236                            _ => {}
237                        }
238                    }
239                }
240                _ => {}
241            }
242        }
243    }
244
245    // ── Pass 2: emit nodes + edges ────────────────────────────────────────────
246
247    fn visit_module(&mut self, node: TsNode<'_>) {
248        let mut cursor = node.walk();
249        let children: Vec<TsNode<'_>> = node.named_children(&mut cursor).collect();
250        for child in children {
251            self.visit_top_level(child, &[]);
252        }
253    }
254
255    fn visit_top_level(&mut self, node: TsNode<'_>, scope: &[String]) {
256        match node.kind() {
257            "function_definition" => {
258                let is_async = Self::fn_is_async(node);
259                self.visit_function(node, scope, None, is_async, &[]);
260            }
261            "decorated_definition" => {
262                let decorators = self.collect_decorators(node);
263                let is_async = node
264                    .child_by_field_name("definition")
265                    .map(Self::fn_is_async)
266                    .unwrap_or(false);
267                if let Some(def) = node.child_by_field_name("definition") {
268                    match def.kind() {
269                        "function_definition" => {
270                            self.visit_function(def, scope, None, is_async, &decorators)
271                        }
272                        "class_definition" => self.visit_class(def, scope, &decorators),
273                        _ => {}
274                    }
275                }
276            }
277            "class_definition" => self.visit_class(node, scope, &[]),
278            "expression_statement" => self.maybe_visit_constant(node, scope),
279            _ => {}
280        }
281    }
282
283    fn visit_function(
284        &mut self,
285        node: TsNode<'_>,
286        scope: &[String],
287        container_id: Option<NodeId>,
288        is_async: bool,
289        decorators: &[String],
290    ) {
291        let Some(name_node) = node.child_by_field_name("name") else {
292            return;
293        };
294        let name = self.text(name_node).to_owned();
295        let id = self
296            .fn_index
297            .get(&name)
298            .cloned()
299            .unwrap_or_else(NodeId::new);
300
301        // Determine kind and metadata flags from decorators.
302        let has_property = decorators.iter().any(|d| d == "property");
303        let has_staticmethod = decorators.iter().any(|d| d == "staticmethod");
304        let has_classmethod = decorators.iter().any(|d| d == "classmethod");
305
306        let kind = if has_property {
307            NodeKind::Property
308        } else if container_id.is_some() {
309            NodeKind::Method
310        } else {
311            NodeKind::Function
312        };
313
314        // Check if the body contains a `yield` or `yield_from` → generator.
315        let is_generator = node
316            .child_by_field_name("body")
317            .map(|body| Self::body_has_yield(body))
318            .unwrap_or(false);
319
320        let mut graph_node = self.make_node(id.clone(), kind, name, scope, node, is_async);
321        if has_property {
322            graph_node.metadata.is_property = true;
323        }
324        if has_staticmethod || has_classmethod {
325            graph_node.metadata.is_static = true;
326        }
327        if is_generator {
328            graph_node.metadata.is_generator = true;
329        }
330
331        if let Some(cid) = container_id {
332            self.edges.push(Edge {
333                src: cid,
334                dst: id.clone(),
335                kind: EdgeKind::Contains,
336            });
337        }
338        self.nodes.push(graph_node);
339
340        // Type annotations → Uses edges
341        self.extract_param_types(node, &id);
342        self.extract_return_type(node, &id);
343
344        // Decorator names → Uses edges (e.g. @property, @staticmethod, @dataclass)
345        // and → deferred_annotated edges.
346        for dec in decorators {
347            self.deferred_uses.push((id.clone(), dec.clone()));
348            self.deferred_annotated.push((id.clone(), dec.clone()));
349        }
350
351        if let Some(body) = node.child_by_field_name("body") {
352            self.collect_calls(body, &id);
353        }
354    }
355
356    fn visit_class(&mut self, node: TsNode<'_>, scope: &[String], decorators: &[String]) {
357        let Some(name_node) = node.child_by_field_name("name") else {
358            return;
359        };
360        let name = self.text(name_node).to_owned();
361        let id = self
362            .class_index
363            .get(&name)
364            .cloned()
365            .unwrap_or_else(NodeId::new);
366
367        // Determine whether this class inherits from Protocol → Interface.
368        let mut is_protocol = false;
369        if let Some(bases) = node.child_by_field_name("superclasses") {
370            let mut c = bases.walk();
371            for base in bases.named_children(&mut c) {
372                let base_name = match base.kind() {
373                    "identifier" => Some(self.text(base).to_owned()),
374                    "attribute" => base
375                        .child_by_field_name("attribute")
376                        .map(|n| self.text(n).to_owned()),
377                    _ => None,
378                };
379                if base_name.as_deref() == Some("Protocol") {
380                    is_protocol = true;
381                }
382            }
383        }
384
385        let class_kind = if is_protocol {
386            NodeKind::Interface
387        } else {
388            NodeKind::Struct
389        };
390
391        let mut graph_node =
392            self.make_node(id.clone(), class_kind, name.clone(), scope, node, false);
393        if is_protocol {
394            graph_node.metadata.is_abstract = true;
395        }
396        self.nodes.push(graph_node);
397
398        // Base classes → Implements edges
399        if let Some(bases) = node.child_by_field_name("superclasses") {
400            let mut c = bases.walk();
401            for base in bases.named_children(&mut c) {
402                let base_name = match base.kind() {
403                    "identifier" => Some(self.text(base).to_owned()),
404                    "attribute" => base
405                        .child_by_field_name("attribute")
406                        .map(|n| self.text(n).to_owned()),
407                    _ => None,
408                };
409                if let Some(b) = base_name {
410                    self.deferred_implements.push((id.clone(), b));
411                }
412            }
413        }
414
415        // Decorator names → Uses edges (e.g. @dataclass) and → deferred_annotated.
416        for dec in decorators {
417            self.deferred_uses.push((id.clone(), dec.clone()));
418            self.deferred_annotated.push((id.clone(), dec.clone()));
419        }
420
421        let mut class_scope = scope.to_vec();
422        class_scope.push(name.clone());
423
424        if let Some(body) = node.child_by_field_name("body") {
425            let mut cursor = body.walk();
426            let children: Vec<TsNode<'_>> = body.named_children(&mut cursor).collect();
427            for child in children {
428                match child.kind() {
429                    "function_definition" => {
430                        let is_async = Self::fn_is_async(child);
431                        self.visit_function(child, &class_scope, Some(id.clone()), is_async, &[]);
432                    }
433                    "decorated_definition" => {
434                        let method_decorators = self.collect_decorators(child);
435                        let is_async = child
436                            .child_by_field_name("definition")
437                            .map(Self::fn_is_async)
438                            .unwrap_or(false);
439                        if let Some(def) = child.child_by_field_name("definition") {
440                            match def.kind() {
441                                "function_definition" => {
442                                    self.visit_function(
443                                        def,
444                                        &class_scope,
445                                        Some(id.clone()),
446                                        is_async,
447                                        &method_decorators,
448                                    );
449                                }
450                                "class_definition" => {
451                                    self.visit_class(def, &class_scope, &method_decorators);
452                                    // Add Contains edge from parent class to nested class.
453                                    if let Some(nested_name_node) = def.child_by_field_name("name")
454                                    {
455                                        let nested_name = self.text(nested_name_node).to_owned();
456                                        if let Some(nested_id) =
457                                            self.class_index.get(&nested_name).cloned()
458                                        {
459                                            self.edges.push(Edge {
460                                                src: id.clone(),
461                                                dst: nested_id,
462                                                kind: EdgeKind::Contains,
463                                            });
464                                        }
465                                    }
466                                }
467                                _ => {}
468                            }
469                        }
470                    }
471                    "class_definition" => {
472                        self.visit_class(child, &class_scope, &[]);
473                        // Add Contains edge from parent class to nested class.
474                        if let Some(nested_name_node) = child.child_by_field_name("name") {
475                            let nested_name = self.text(nested_name_node).to_owned();
476                            if let Some(nested_id) = self.class_index.get(&nested_name).cloned() {
477                                self.edges.push(Edge {
478                                    src: id.clone(),
479                                    dst: nested_id,
480                                    kind: EdgeKind::Contains,
481                                });
482                            }
483                        }
484                    }
485                    _ => {}
486                }
487            }
488        }
489    }
490
491    fn maybe_visit_constant(&mut self, node: TsNode<'_>, scope: &[String]) {
492        // Capture module-level bindings as `Constant` nodes. Any
493        // simple-identifier assignment at module scope is an importable symbol
494        // (`from pkg import default_config`, `__version__`, `T = TypeVar(...)`),
495        // not just SCREAMING_SNAKE_CASE — restricting to all-caps dropped the
496        // bulk of a module's public surface. Visibility follows Python's
497        // leading-underscore convention via `make_node`.
498        //
499        // Only plain `identifier = …` / `identifier: T = …` targets are taken;
500        // tuple/attribute/subscript targets (`a, b = …`, `self.x = …`) are
501        // skipped — they aren't module-level named symbols.
502        let mut cursor = node.walk();
503        for child in node.named_children(&mut cursor) {
504            if child.kind() == "assignment" {
505                if let Some(left) = child.child_by_field_name("left") {
506                    if left.kind() == "identifier" {
507                        let name = self.text(left).to_owned();
508                        if name.is_empty() {
509                            continue;
510                        }
511                        let id = NodeId::new();
512                        let graph_node =
513                            self.make_node(id, NodeKind::Constant, name, scope, node, false);
514                        self.nodes.push(graph_node);
515                    }
516                }
517            }
518        }
519    }
520
521    // ── Pass 3: collect import statements ────────────────────────────────────
522
523    fn collect_imports(&mut self, node: TsNode<'_>) {
524        let mut cursor = node.walk();
525        let children: Vec<TsNode<'_>> = node.named_children(&mut cursor).collect();
526        for child in children {
527            match child.kind() {
528                "import_statement" => {
529                    // `import foo`, `import foo.bar`, `import foo as f`
530                    let mut c = child.walk();
531                    for name_node in child.named_children(&mut c) {
532                        let leaf = match name_node.kind() {
533                            "dotted_name" => {
534                                let text = self.text(name_node);
535                                text.split('.').next_back().map(|s| s.to_owned())
536                            }
537                            "aliased_import" => name_node
538                                .child_by_field_name("name")
539                                .map(|n| self.text(n))
540                                .map(|t| t.split('.').next_back().unwrap_or(t).to_owned()),
541                            _ => None,
542                        };
543                        if let Some(name) = leaf {
544                            self.deferred_imports.push((self.module_id.clone(), name));
545                        }
546                    }
547                }
548                "import_from_statement" => {
549                    // `from foo import bar, baz`
550                    // Named children: first is the source module (dotted_name or
551                    // relative_import), the rest are the imported names.
552                    let mut c = child.walk();
553                    let all_children: Vec<TsNode<'_>> = child.named_children(&mut c).collect();
554                    for name_node in all_children.iter().skip(1) {
555                        let leaf = match name_node.kind() {
556                            "dotted_name" => {
557                                let text = self.text(*name_node);
558                                Some(text.split('.').next_back().unwrap_or(text).to_owned())
559                            }
560                            "aliased_import" => name_node
561                                .child_by_field_name("name")
562                                .map(|n| self.text(n))
563                                .map(|t| t.split('.').next_back().unwrap_or(t).to_owned()),
564                            "wildcard_import" => None,
565                            _ => None,
566                        };
567                        if let Some(name) = leaf {
568                            self.deferred_imports.push((self.module_id.clone(), name));
569                        }
570                    }
571                }
572                _ => {}
573            }
574        }
575    }
576
577    // ── Call collection ───────────────────────────────────────────────────────
578
579    fn collect_calls(&mut self, node: TsNode<'_>, caller_id: &NodeId) {
580        let mut cursor = node.walk();
581        let children: Vec<TsNode<'_>> = node.named_children(&mut cursor).collect();
582        for child in children {
583            if child.kind() == "call" {
584                if let Some(callee) = self.callee_name(child) {
585                    self.record_call(caller_id.clone(), callee);
586                }
587                if let Some(args) = child.child_by_field_name("arguments") {
588                    self.collect_calls(args, caller_id);
589                }
590            } else {
591                self.collect_calls(child, caller_id);
592            }
593        }
594    }
595
596    fn callee_name(&self, call_node: TsNode<'_>) -> Option<String> {
597        let func = call_node.child_by_field_name("function")?;
598        match func.kind() {
599            "identifier" => Some(self.text(func).to_owned()),
600            "attribute" => func
601                .child_by_field_name("attribute")
602                .map(|n| self.text(n).to_owned()),
603            _ => None,
604        }
605    }
606
607    fn record_call(&mut self, caller_id: NodeId, callee_name: String) {
608        if callee_name.is_empty() {
609            return;
610        }
611        if let Some(callee_id) = self.fn_index.get(&callee_name).cloned() {
612            let edge = Edge {
613                src: caller_id,
614                dst: callee_id,
615                kind: EdgeKind::Calls,
616            };
617            if !self.edges.contains(&edge) {
618                self.edges.push(edge);
619            }
620        } else if !self
621            .deferred_calls
622            .iter()
623            .any(|(c, n)| c == &caller_id && n == &callee_name)
624        {
625            self.deferred_calls.push((caller_id, callee_name));
626        }
627    }
628
629    // ── Helpers ───────────────────────────────────────────────────────────────
630
631    /// Returns true if this `function_definition` node is `async def`.
632    fn fn_is_async(node: TsNode<'_>) -> bool {
633        let mut c = node.walk();
634        // `async` appears as an anonymous child (keyword) before `def`
635        let result = node.children(&mut c).any(|n| n.kind() == "async");
636        result
637    }
638
639    /// Returns true if the body subtree contains a `yield` or `yield_from` expression.
640    fn body_has_yield(node: TsNode<'_>) -> bool {
641        if node.kind() == "yield" || node.kind() == "yield_from" {
642            return true;
643        }
644        // Don't descend into nested function definitions — their yields are not
645        // generators of the outer function.
646        if node.kind() == "function_definition" {
647            return false;
648        }
649        let mut c = node.walk();
650        let found = node.named_children(&mut c).any(Self::body_has_yield);
651        found
652    }
653
654    /// Extract decorator names from a `decorated_definition` node.
655    fn collect_decorators(&self, node: TsNode<'_>) -> Vec<String> {
656        let mut c = node.walk();
657        node.named_children(&mut c)
658            .filter(|n| n.kind() == "decorator")
659            .filter_map(|d| self.decorator_name(d))
660            .collect()
661    }
662
663    /// Get the callable name from a `decorator` node.
664    fn decorator_name(&self, decorator: TsNode<'_>) -> Option<String> {
665        let mut c = decorator.walk();
666        let child = decorator.named_children(&mut c).next()?;
667        match child.kind() {
668            "identifier" => Some(self.text(child).to_owned()),
669            "attribute" => child
670                .child_by_field_name("attribute")
671                .map(|n| self.text(n).to_owned()),
672            "call" => child
673                .child_by_field_name("function")
674                .and_then(|f| match f.kind() {
675                    "identifier" => Some(self.text(f).to_owned()),
676                    "attribute" => f
677                        .child_by_field_name("attribute")
678                        .map(|n| self.text(n).to_owned()),
679                    _ => None,
680                }),
681            _ => None,
682        }
683    }
684
685    /// Extract type identifiers from a type-annotation node.
686    /// Records all identifiers found (e.g. `List[MyType]` → ["List", "MyType"]).
687    fn extract_param_types(&mut self, fn_node: TsNode<'_>, fn_id: &NodeId) {
688        let Some(params) = fn_node.child_by_field_name("parameters") else {
689            return;
690        };
691        let mut c = params.walk();
692        for param in params.named_children(&mut c) {
693            let type_node = match param.kind() {
694                "typed_parameter" | "typed_default_parameter" => param.child_by_field_name("type"),
695                _ => None,
696            };
697            if let Some(t) = type_node {
698                for name in self.collect_type_names(t) {
699                    self.deferred_uses.push((fn_id.clone(), name));
700                }
701            }
702        }
703    }
704
705    fn extract_return_type(&mut self, fn_node: TsNode<'_>, fn_id: &NodeId) {
706        if let Some(ret) = fn_node.child_by_field_name("return_type") {
707            for name in self.collect_type_names(ret) {
708                self.deferred_uses.push((fn_id.clone(), name));
709            }
710        }
711    }
712
713    /// Walk a type-annotation subtree and collect all identifier names.
714    fn collect_type_names(&self, node: TsNode<'_>) -> Vec<String> {
715        let mut names = Vec::new();
716        self.walk_type_names(node, &mut names);
717        names
718    }
719
720    fn walk_type_names(&self, node: TsNode<'_>, out: &mut Vec<String>) {
721        match node.kind() {
722            "identifier" => {
723                let name = self.text(node).to_owned();
724                if !is_builtin_type(&name) {
725                    out.push(name);
726                }
727            }
728            _ => {
729                let mut c = node.walk();
730                for child in node.named_children(&mut c) {
731                    self.walk_type_names(child, out);
732                }
733            }
734        }
735    }
736}
737
738/// Returns true for built-in Python types and typing-module generics that do
739/// not correspond to user-defined symbols.
740fn is_builtin_type(name: &str) -> bool {
741    matches!(
742        name,
743        "int"
744            | "str"
745            | "bool"
746            | "float"
747            | "complex"
748            | "bytes"
749            | "bytearray"
750            | "None"
751            | "list"
752            | "dict"
753            | "set"
754            | "frozenset"
755            | "tuple"
756            | "type"
757            | "object"
758            | "Any"
759            | "Optional"
760            | "Union"
761            | "List"
762            | "Dict"
763            | "Set"
764            | "FrozenSet"
765            | "Tuple"
766            | "Callable"
767            | "Type"
768            | "ClassVar"
769            | "Final"
770            | "Literal"
771            | "TypeVar"
772            | "Generic"
773            | "Protocol"
774            | "Sequence"
775            | "Iterable"
776            | "Iterator"
777            | "Generator"
778            | "Coroutine"
779            | "Awaitable"
780            | "AsyncIterator"
781            | "AsyncGenerator"
782            | "NoReturn"
783            | "Never"
784            | "Self"
785            | "Annotated"
786            | "TypeAlias"
787            | "ParamSpec"
788            | "TypeVarTuple"
789            | "overload"
790            | "abstractmethod"
791            | "staticmethod"
792            | "classmethod"
793    )
794}
795
796// ── Tests ─────────────────────────────────────────────────────────────────────
797
798#[cfg(test)]
799mod tests {
800    use super::PythonParser;
801    use crate::parser::LanguageParser;
802    use gitcortex_core::schema::{EdgeKind, NodeKind};
803    use std::path::Path;
804
805    fn parse(
806        src: &str,
807    ) -> (
808        Vec<gitcortex_core::graph::Node>,
809        Vec<gitcortex_core::graph::Edge>,
810    ) {
811        let r = PythonParser::new()
812            .parse(Path::new("test.py"), src)
813            .unwrap();
814        (r.nodes, r.edges)
815    }
816
817    #[allow(clippy::type_complexity)]
818    fn parse_full(
819        src: &str,
820    ) -> (
821        Vec<gitcortex_core::graph::Node>,
822        Vec<gitcortex_core::graph::Edge>,
823        Vec<(gitcortex_core::graph::NodeId, String)>,
824        Vec<(gitcortex_core::graph::NodeId, String)>,
825        Vec<(gitcortex_core::graph::NodeId, String)>,
826        Vec<(gitcortex_core::graph::NodeId, String)>,
827    ) {
828        let r = PythonParser::new()
829            .parse(Path::new("test.py"), src)
830            .unwrap();
831        (
832            r.nodes,
833            r.edges,
834            r.deferred_calls,
835            r.deferred_uses,
836            r.deferred_implements,
837            r.deferred_imports,
838        )
839    }
840
841    #[test]
842    fn parses_free_function() {
843        let (nodes, _) = parse("def greet(name):\n    return name\n");
844        // Module node + Function node
845        assert_eq!(nodes.len(), 2);
846        let fns: Vec<_> = nodes
847            .iter()
848            .filter(|n| n.kind == NodeKind::Function)
849            .collect();
850        assert_eq!(fns.len(), 1);
851        assert_eq!(fns[0].name, "greet");
852    }
853
854    #[test]
855    fn parses_class_and_method() {
856        let src = "class Person:\n    def greet(self):\n        pass\n";
857        let (nodes, edges) = parse(src);
858        let classes: Vec<_> = nodes
859            .iter()
860            .filter(|n| n.kind == NodeKind::Struct)
861            .collect();
862        let methods: Vec<_> = nodes
863            .iter()
864            .filter(|n| n.kind == NodeKind::Method)
865            .collect();
866        assert_eq!(classes.len(), 1);
867        assert_eq!(methods.len(), 1);
868        let contains: Vec<_> = edges
869            .iter()
870            .filter(|e| e.kind == EdgeKind::Contains)
871            .collect();
872        assert!(!contains.is_empty());
873    }
874
875    #[test]
876    fn detects_call_edges() {
877        let src = "def caller():\n    callee()\ndef callee():\n    pass\n";
878        let (_, edges) = parse(src);
879        let calls: Vec<_> = edges.iter().filter(|e| e.kind == EdgeKind::Calls).collect();
880        assert_eq!(calls.len(), 1);
881    }
882
883    #[test]
884    fn detects_base_class_implements() {
885        let src = "class Base:\n    pass\nclass Child(Base):\n    pass\n";
886        let (_, _, _, _, implements, _) = parse_full(src);
887        assert!(
888            implements.iter().any(|(_, name)| name == "Base"),
889            "expected Implements edge to Base, got: {implements:?}"
890        );
891    }
892
893    #[test]
894    fn detects_type_annotation_uses() {
895        let src = "class Foo:\n    pass\ndef bar(x: Foo) -> Foo:\n    pass\n";
896        let (_, _, _, uses, _, _) = parse_full(src);
897        let uses_foo: Vec<_> = uses.iter().filter(|(_, n)| n == "Foo").collect();
898        assert!(
899            uses_foo.len() >= 2,
900            "expected at least 2 Uses edges to Foo, got: {uses_foo:?}"
901        );
902    }
903
904    #[test]
905    fn detects_decorator_uses() {
906        let src = "class Foo:\n    @property\n    def name(self):\n        return self._name\n";
907        let (_, _, _, uses, _, _) = parse_full(src);
908        assert!(
909            uses.iter().any(|(_, n)| n == "property"),
910            "expected Uses edge to 'property' decorator, got: {uses:?}"
911        );
912    }
913
914    #[test]
915    fn detects_import_statement() {
916        let src = "import os\nimport sys\n\ndef main():\n    pass\n";
917        let (_, _, _, _, _, imports) = parse_full(src);
918        assert!(
919            imports.iter().any(|(_, n)| n == "os"),
920            "expected import 'os', got: {imports:?}"
921        );
922        assert!(
923            imports.iter().any(|(_, n)| n == "sys"),
924            "expected import 'sys', got: {imports:?}"
925        );
926    }
927
928    #[test]
929    fn detects_from_import_statement() {
930        let src = "from os.path import join, exists\n\ndef main():\n    pass\n";
931        let (_, _, _, _, _, imports) = parse_full(src);
932        assert!(
933            imports.iter().any(|(_, n)| n == "join"),
934            "expected import 'join', got: {imports:?}"
935        );
936        assert!(
937            imports.iter().any(|(_, n)| n == "exists"),
938            "expected import 'exists', got: {imports:?}"
939        );
940    }
941
942    #[test]
943    fn module_node_is_emitted() {
944        let (nodes, _) = parse("x = 1\n");
945        let modules: Vec<_> = nodes
946            .iter()
947            .filter(|n| n.kind == NodeKind::Module)
948            .collect();
949        assert_eq!(modules.len(), 1);
950        assert_eq!(modules[0].name, "test");
951    }
952
953    #[test]
954    fn async_function_flagged() {
955        let src = "async def fetch():\n    pass\n";
956        let (nodes, _) = parse(src);
957        let fns: Vec<_> = nodes
958            .iter()
959            .filter(|n| n.kind == NodeKind::Function)
960            .collect();
961        assert_eq!(fns.len(), 1);
962        assert!(fns[0].metadata.is_async);
963    }
964
965    // ── Regression: Protocol / Interface ─────────────────────────────────────
966
967    #[test]
968    fn protocol_class_becomes_interface() {
969        let src = "from typing import Protocol\n\nclass MyProto(Protocol):\n    def do_it(self) -> None:\n        ...\n";
970        let (nodes, _) = parse(src);
971        let ifaces: Vec<_> = nodes
972            .iter()
973            .filter(|n| n.kind == NodeKind::Interface)
974            .collect();
975        assert_eq!(ifaces.len(), 1, "expected 1 Interface, got: {ifaces:?}");
976        assert_eq!(ifaces[0].name, "MyProto");
977        assert!(
978            ifaces[0].metadata.is_abstract,
979            "Protocol should be is_abstract"
980        );
981    }
982
983    #[test]
984    fn non_protocol_class_is_struct() {
985        let src = "class Plain:\n    def work(self):\n        pass\n";
986        let (nodes, _) = parse(src);
987        let structs: Vec<_> = nodes
988            .iter()
989            .filter(|n| n.kind == NodeKind::Struct)
990            .collect();
991        assert_eq!(structs.len(), 1);
992        assert_eq!(structs[0].name, "Plain");
993        assert!(!structs[0].metadata.is_abstract);
994    }
995
996    // ── Regression: decorators ────────────────────────────────────────────────
997
998    #[test]
999    fn property_decorator_yields_property_kind() {
1000        let src =
1001            "class Foo:\n    @property\n    def bar(self) -> str:\n        return self._bar\n";
1002        let (nodes, _) = parse(src);
1003        let props: Vec<_> = nodes
1004            .iter()
1005            .filter(|n| n.kind == NodeKind::Property)
1006            .collect();
1007        assert_eq!(props.len(), 1, "expected 1 Property node, got: {props:?}");
1008        assert_eq!(props[0].name, "bar");
1009        assert!(props[0].metadata.is_property);
1010    }
1011
1012    #[test]
1013    fn staticmethod_decorator_sets_is_static() {
1014        let src = "class Foo:\n    @staticmethod\n    def create(x: int) -> 'Foo':\n        return Foo()\n";
1015        let (nodes, _) = parse(src);
1016        let methods: Vec<_> = nodes
1017            .iter()
1018            .filter(|n| n.kind == NodeKind::Method)
1019            .collect();
1020        assert_eq!(methods.len(), 1);
1021        assert_eq!(methods[0].name, "create");
1022        assert!(
1023            methods[0].metadata.is_static,
1024            "staticmethod should set is_static"
1025        );
1026    }
1027
1028    #[test]
1029    fn classmethod_decorator_sets_is_static() {
1030        let src = "class Foo:\n    @classmethod\n    def from_str(cls, s: str) -> 'Foo':\n        return cls()\n";
1031        let (nodes, _) = parse(src);
1032        let methods: Vec<_> = nodes
1033            .iter()
1034            .filter(|n| n.kind == NodeKind::Method)
1035            .collect();
1036        assert_eq!(methods.len(), 1);
1037        assert_eq!(methods[0].name, "from_str");
1038        assert!(
1039            methods[0].metadata.is_static,
1040            "classmethod should set is_static"
1041        );
1042    }
1043
1044    #[test]
1045    fn dataclass_decorator_class_is_struct() {
1046        let src = "from dataclasses import dataclass\n\n@dataclass\nclass Point:\n    x: float\n    y: float\n";
1047        let (nodes, _) = parse(src);
1048        let structs: Vec<_> = nodes
1049            .iter()
1050            .filter(|n| n.kind == NodeKind::Struct)
1051            .collect();
1052        assert_eq!(structs.len(), 1, "expected Struct for @dataclass Point");
1053        assert_eq!(structs[0].name, "Point");
1054    }
1055
1056    // ── Regression: generator / async-generator ──────────────────────────────
1057
1058    #[test]
1059    fn generator_function_sets_is_generator() {
1060        let src = "def numbers():\n    yield 1\n    yield 2\n";
1061        let (nodes, _) = parse(src);
1062        let fns: Vec<_> = nodes
1063            .iter()
1064            .filter(|n| n.kind == NodeKind::Function)
1065            .collect();
1066        assert_eq!(fns.len(), 1);
1067        assert!(
1068            fns[0].metadata.is_generator,
1069            "yield fn should be is_generator"
1070        );
1071        assert!(!fns[0].metadata.is_async);
1072    }
1073
1074    #[test]
1075    fn async_generator_is_both_async_and_generator() {
1076        let src = "async def stream():\n    yield 1\n    yield 2\n";
1077        let (nodes, _) = parse(src);
1078        let fns: Vec<_> = nodes
1079            .iter()
1080            .filter(|n| n.kind == NodeKind::Function)
1081            .collect();
1082        assert_eq!(fns.len(), 1);
1083        assert!(
1084            fns[0].metadata.is_async,
1085            "async generator should be is_async"
1086        );
1087        assert!(
1088            fns[0].metadata.is_generator,
1089            "async generator should be is_generator"
1090        );
1091    }
1092
1093    #[test]
1094    fn nested_yield_does_not_pollute_outer_function() {
1095        // The outer function is NOT a generator; only the inner lambda/nested fn yields.
1096        let src = "def outer():\n    def inner():\n        yield 1\n    return inner()\n";
1097        let (nodes, _) = parse(src);
1098        let fns: Vec<_> = nodes
1099            .iter()
1100            .filter(|n| n.kind == NodeKind::Function)
1101            .collect();
1102        let outer = fns
1103            .iter()
1104            .find(|n| n.name == "outer")
1105            .expect("outer not found");
1106        assert!(
1107            !outer.metadata.is_generator,
1108            "outer should NOT be generator — yield is in nested fn"
1109        );
1110    }
1111
1112    // ── Regression: constants ─────────────────────────────────────────────────
1113
1114    #[test]
1115    fn module_level_bindings_detected() {
1116        // Module-level assignments are importable symbols regardless of case:
1117        // ALL_CAPS constants, lowercase config objects, dunders, and TypeVars
1118        // all count. Locals inside functions must NOT be captured here.
1119        let src = "MAX_SIZE = 100\nDEFAULT_NAME = 'anon'\ndefault_config = {}\n__version__ = '1.0'\nT = 1\n\ndef f():\n    local = 1\n    return local\n";
1120        let (nodes, _) = parse(src);
1121        let names: Vec<&str> = nodes
1122            .iter()
1123            .filter(|n| n.kind == NodeKind::Constant)
1124            .map(|n| n.name.as_str())
1125            .collect();
1126        for expected in [
1127            "MAX_SIZE",
1128            "DEFAULT_NAME",
1129            "default_config",
1130            "__version__",
1131            "T",
1132        ] {
1133            assert!(
1134                names.contains(&expected),
1135                "expected module-level binding `{expected}` as Constant; got {names:?}"
1136            );
1137        }
1138        assert!(
1139            !names.contains(&"local"),
1140            "function-local assignment must not be captured as a module Constant"
1141        );
1142    }
1143
1144    // ── Regression: nested classes ────────────────────────────────────────────
1145
1146    #[test]
1147    fn nested_class_emits_contains_edge_from_parent() {
1148        let src = "class Outer:\n    class Inner:\n        pass\n";
1149        let (nodes, edges) = parse(src);
1150        let outer = nodes
1151            .iter()
1152            .find(|n| n.name == "Outer")
1153            .expect("Outer not found");
1154        let inner = nodes
1155            .iter()
1156            .find(|n| n.name == "Inner")
1157            .expect("Inner not found");
1158        let contains: Vec<_> = edges
1159            .iter()
1160            .filter(|e| e.kind == EdgeKind::Contains)
1161            .collect();
1162        assert!(
1163            contains
1164                .iter()
1165                .any(|e| e.src == outer.id && e.dst == inner.id),
1166            "expected Contains Outer → Inner"
1167        );
1168    }
1169
1170    // ── Regression: type annotation Uses deferred entries ────────────────────
1171
1172    #[test]
1173    fn multiple_type_annotations_produce_uses_entries() {
1174        let src =
1175            "class Req:\n    pass\nclass Resp:\n    pass\ndef handler(r: Req) -> Resp:\n    pass\n";
1176        let (_, _, _, uses, _, _) = parse_full(src);
1177        let uses_req: Vec<_> = uses.iter().filter(|(_, n)| n == "Req").collect();
1178        let uses_resp: Vec<_> = uses.iter().filter(|(_, n)| n == "Resp").collect();
1179        assert!(!uses_req.is_empty(), "expected Uses edge to Req");
1180        assert!(!uses_resp.is_empty(), "expected Uses edge to Resp");
1181    }
1182
1183    // ── Regression: visibility ────────────────────────────────────────────────
1184
1185    #[test]
1186    fn private_method_has_private_visibility() {
1187        use gitcortex_core::schema::Visibility;
1188        let src = "class Foo:\n    def _internal(self):\n        pass\n    def public(self):\n        pass\n";
1189        let (nodes, _) = parse(src);
1190        let internal = nodes
1191            .iter()
1192            .find(|n| n.name == "_internal")
1193            .expect("_internal not found");
1194        let public = nodes
1195            .iter()
1196            .find(|n| n.name == "public")
1197            .expect("public not found");
1198        assert_eq!(internal.metadata.visibility, Visibility::Private);
1199        assert_eq!(public.metadata.visibility, Visibility::Pub);
1200    }
1201
1202    // ── Regression: call detection ────────────────────────────────────────────
1203
1204    #[test]
1205    fn calls_edge_between_two_functions() {
1206        let src = "def helper():\n    pass\n\ndef main():\n    helper()\n";
1207        let (_, edges) = parse(src);
1208        let calls: Vec<_> = edges.iter().filter(|e| e.kind == EdgeKind::Calls).collect();
1209        assert_eq!(calls.len(), 1, "expected exactly 1 Calls edge");
1210    }
1211
1212    #[test]
1213    fn method_call_via_self_creates_calls_edge() {
1214        let src = "class Svc:\n    def run(self):\n        self.process()\n    def process(self):\n        pass\n";
1215        let (nodes, edges) = parse(src);
1216        let run = nodes
1217            .iter()
1218            .find(|n| n.name == "run")
1219            .expect("run not found");
1220        let process = nodes
1221            .iter()
1222            .find(|n| n.name == "process")
1223            .expect("process not found");
1224        let calls: Vec<_> = edges.iter().filter(|e| e.kind == EdgeKind::Calls).collect();
1225        // self.process() resolves immediately because "process" is pre-indexed in pass 1
1226        assert!(
1227            calls.iter().any(|e| e.src == run.id && e.dst == process.id),
1228            "expected Calls edge run → process, got: {calls:?}"
1229        );
1230    }
1231
1232    // ── Regression: import edge collection ───────────────────────────────────
1233
1234    #[test]
1235    fn aliased_import_uses_alias_name() {
1236        let src = "import numpy as np\nimport pandas as pd\n";
1237        let (_, _, _, _, _, imports) = parse_full(src);
1238        // For aliased imports, the leaf of the original module name is recorded
1239        let names: Vec<&str> = imports.iter().map(|(_, n)| n.as_str()).collect();
1240        assert!(
1241            names.contains(&"numpy"),
1242            "expected import 'numpy', got: {names:?}"
1243        );
1244        assert!(
1245            names.contains(&"pandas"),
1246            "expected import 'pandas', got: {names:?}"
1247        );
1248    }
1249
1250    #[test]
1251    fn dotted_import_records_leaf_module() {
1252        let src = "import os.path\n";
1253        let (_, _, _, _, _, imports) = parse_full(src);
1254        let names: Vec<&str> = imports.iter().map(|(_, n)| n.as_str()).collect();
1255        assert!(
1256            names.contains(&"path"),
1257            "expected leaf 'path' from 'import os.path', got: {names:?}"
1258        );
1259    }
1260}