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::{EdgeConfidence, 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
73struct FileVisitor<'src> {
76 source: &'src [u8],
77 file: PathBuf,
78 module_id: NodeId,
80 nodes: Vec<Node>,
81 edges: Vec<Edge>,
82 class_index: HashMap<String, NodeId>,
84 fn_index: HashMap<String, NodeId>,
86 deferred_calls: Vec<(NodeId, String, u32)>,
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 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 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 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 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 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 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 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 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(body) = node.child_by_field_name("body") {
332 graph_node.metadata.lld.complexity = Some(super::cyclomatic_complexity(
333 body,
334 &super::complexity::python_decision,
335 ));
336 }
337
338 if let Some(cid) = container_id {
339 self.edges.push(Edge {
340 src: cid,
341 dst: id.clone(),
342 kind: EdgeKind::Contains,
343 line: None,
344 confidence: EdgeConfidence::Extracted,
345 });
346 }
347 self.nodes.push(graph_node);
348
349 self.extract_param_types(node, &id);
351 self.extract_return_type(node, &id);
352
353 for dec in decorators {
356 self.deferred_uses.push((id.clone(), dec.clone()));
357 self.deferred_annotated.push((id.clone(), dec.clone()));
358 }
359
360 if let Some(body) = node.child_by_field_name("body") {
361 self.collect_calls(body, &id);
362 }
363 }
364
365 fn visit_class(&mut self, node: TsNode<'_>, scope: &[String], decorators: &[String]) {
366 let Some(name_node) = node.child_by_field_name("name") else {
367 return;
368 };
369 let name = self.text(name_node).to_owned();
370 let id = self
371 .class_index
372 .get(&name)
373 .cloned()
374 .unwrap_or_else(NodeId::new);
375
376 let mut is_protocol = false;
378 if let Some(bases) = node.child_by_field_name("superclasses") {
379 let mut c = bases.walk();
380 for base in bases.named_children(&mut c) {
381 let base_name = match base.kind() {
382 "identifier" => Some(self.text(base).to_owned()),
383 "attribute" => base
384 .child_by_field_name("attribute")
385 .map(|n| self.text(n).to_owned()),
386 _ => None,
387 };
388 if base_name.as_deref() == Some("Protocol") {
389 is_protocol = true;
390 }
391 }
392 }
393
394 let class_kind = if is_protocol {
395 NodeKind::Interface
396 } else {
397 NodeKind::Struct
398 };
399
400 let mut graph_node =
401 self.make_node(id.clone(), class_kind, name.clone(), scope, node, false);
402 if is_protocol {
403 graph_node.metadata.is_abstract = true;
404 }
405 self.nodes.push(graph_node);
406
407 if let Some(bases) = node.child_by_field_name("superclasses") {
409 let mut c = bases.walk();
410 for base in bases.named_children(&mut c) {
411 let base_name = match base.kind() {
412 "identifier" => Some(self.text(base).to_owned()),
413 "attribute" => base
414 .child_by_field_name("attribute")
415 .map(|n| self.text(n).to_owned()),
416 _ => None,
417 };
418 if let Some(b) = base_name {
419 self.deferred_implements.push((id.clone(), b));
420 }
421 }
422 }
423
424 for dec in decorators {
426 self.deferred_uses.push((id.clone(), dec.clone()));
427 self.deferred_annotated.push((id.clone(), dec.clone()));
428 }
429
430 let mut class_scope = scope.to_vec();
431 class_scope.push(name.clone());
432
433 if let Some(body) = node.child_by_field_name("body") {
434 let mut cursor = body.walk();
435 let children: Vec<TsNode<'_>> = body.named_children(&mut cursor).collect();
436 for child in children {
437 match child.kind() {
438 "function_definition" => {
439 let is_async = Self::fn_is_async(child);
440 self.visit_function(child, &class_scope, Some(id.clone()), is_async, &[]);
441 }
442 "decorated_definition" => {
443 let method_decorators = self.collect_decorators(child);
444 let is_async = child
445 .child_by_field_name("definition")
446 .map(Self::fn_is_async)
447 .unwrap_or(false);
448 if let Some(def) = child.child_by_field_name("definition") {
449 match def.kind() {
450 "function_definition" => {
451 self.visit_function(
452 def,
453 &class_scope,
454 Some(id.clone()),
455 is_async,
456 &method_decorators,
457 );
458 }
459 "class_definition" => {
460 self.visit_class(def, &class_scope, &method_decorators);
461 if let Some(nested_name_node) = def.child_by_field_name("name")
463 {
464 let nested_name = self.text(nested_name_node).to_owned();
465 if let Some(nested_id) =
466 self.class_index.get(&nested_name).cloned()
467 {
468 self.edges.push(Edge {
469 src: id.clone(),
470 dst: nested_id,
471 kind: EdgeKind::Contains,
472 line: None,
473 confidence: EdgeConfidence::Extracted,
474 });
475 }
476 }
477 }
478 _ => {}
479 }
480 }
481 }
482 "class_definition" => {
483 self.visit_class(child, &class_scope, &[]);
484 if let Some(nested_name_node) = child.child_by_field_name("name") {
486 let nested_name = self.text(nested_name_node).to_owned();
487 if let Some(nested_id) = self.class_index.get(&nested_name).cloned() {
488 self.edges.push(Edge {
489 src: id.clone(),
490 dst: nested_id,
491 kind: EdgeKind::Contains,
492 line: None,
493 confidence: EdgeConfidence::Extracted,
494 });
495 }
496 }
497 }
498 _ => {}
499 }
500 }
501 }
502 }
503
504 fn maybe_visit_constant(&mut self, node: TsNode<'_>, scope: &[String]) {
505 let mut cursor = node.walk();
516 for child in node.named_children(&mut cursor) {
517 if child.kind() == "assignment" {
518 if let Some(left) = child.child_by_field_name("left") {
519 if left.kind() == "identifier" {
520 let name = self.text(left).to_owned();
521 if name.is_empty() {
522 continue;
523 }
524 let id = NodeId::new();
525 let graph_node =
526 self.make_node(id, NodeKind::Constant, name, scope, node, false);
527 self.nodes.push(graph_node);
528 }
529 }
530 }
531 }
532 }
533
534 fn collect_imports(&mut self, node: TsNode<'_>) {
537 let mut cursor = node.walk();
538 let children: Vec<TsNode<'_>> = node.named_children(&mut cursor).collect();
539 for child in children {
540 match child.kind() {
541 "import_statement" => {
542 let mut c = child.walk();
544 for name_node in child.named_children(&mut c) {
545 let leaf = match name_node.kind() {
546 "dotted_name" => {
547 let text = self.text(name_node);
548 text.split('.').next_back().map(|s| s.to_owned())
549 }
550 "aliased_import" => name_node
551 .child_by_field_name("name")
552 .map(|n| self.text(n))
553 .map(|t| t.split('.').next_back().unwrap_or(t).to_owned()),
554 _ => None,
555 };
556 if let Some(name) = leaf {
557 self.deferred_imports.push((self.module_id.clone(), name));
558 }
559 }
560 }
561 "import_from_statement" => {
562 let mut c = child.walk();
566 let all_children: Vec<TsNode<'_>> = child.named_children(&mut c).collect();
567 for name_node in all_children.iter().skip(1) {
568 let leaf = match name_node.kind() {
569 "dotted_name" => {
570 let text = self.text(*name_node);
571 Some(text.split('.').next_back().unwrap_or(text).to_owned())
572 }
573 "aliased_import" => name_node
574 .child_by_field_name("name")
575 .map(|n| self.text(n))
576 .map(|t| t.split('.').next_back().unwrap_or(t).to_owned()),
577 "wildcard_import" => None,
578 _ => None,
579 };
580 if let Some(name) = leaf {
581 self.deferred_imports.push((self.module_id.clone(), name));
582 }
583 }
584 }
585 _ => {}
586 }
587 }
588 }
589
590 fn collect_calls(&mut self, node: TsNode<'_>, caller_id: &NodeId) {
593 let mut cursor = node.walk();
594 let children: Vec<TsNode<'_>> = node.named_children(&mut cursor).collect();
595 for child in children {
596 if child.kind() == "call" {
597 if let Some(callee) = self.callee_name(child) {
598 let line = child.start_position().row as u32 + 1;
599 self.record_call(caller_id.clone(), callee, line);
600 }
601 if let Some(args) = child.child_by_field_name("arguments") {
602 self.collect_calls(args, caller_id);
603 }
604 } else {
605 self.collect_calls(child, caller_id);
606 }
607 }
608 }
609
610 fn callee_name(&self, call_node: TsNode<'_>) -> Option<String> {
611 let func = call_node.child_by_field_name("function")?;
612 match func.kind() {
613 "identifier" => Some(self.text(func).to_owned()),
614 "attribute" => func
615 .child_by_field_name("attribute")
616 .map(|n| self.text(n).to_owned()),
617 _ => None,
618 }
619 }
620
621 fn record_call(&mut self, caller_id: NodeId, callee_name: String, line: u32) {
622 if callee_name.is_empty() {
623 return;
624 }
625 if let Some(callee_id) = self.fn_index.get(&callee_name).cloned() {
626 let edge = Edge::call(caller_id, callee_id, line);
627 if !self.edges.contains(&edge) {
628 self.edges.push(edge);
629 }
630 } else if !self
631 .deferred_calls
632 .iter()
633 .any(|(c, n, _)| c == &caller_id && n == &callee_name)
634 {
635 self.deferred_calls.push((caller_id, callee_name, line));
636 }
637 }
638
639 fn fn_is_async(node: TsNode<'_>) -> bool {
643 let mut c = node.walk();
644 let result = node.children(&mut c).any(|n| n.kind() == "async");
646 result
647 }
648
649 fn body_has_yield(node: TsNode<'_>) -> bool {
651 if node.kind() == "yield" || node.kind() == "yield_from" {
652 return true;
653 }
654 if node.kind() == "function_definition" {
657 return false;
658 }
659 let mut c = node.walk();
660 let found = node.named_children(&mut c).any(Self::body_has_yield);
661 found
662 }
663
664 fn collect_decorators(&self, node: TsNode<'_>) -> Vec<String> {
666 let mut c = node.walk();
667 node.named_children(&mut c)
668 .filter(|n| n.kind() == "decorator")
669 .filter_map(|d| self.decorator_name(d))
670 .collect()
671 }
672
673 fn decorator_name(&self, decorator: TsNode<'_>) -> Option<String> {
675 let mut c = decorator.walk();
676 let child = decorator.named_children(&mut c).next()?;
677 match child.kind() {
678 "identifier" => Some(self.text(child).to_owned()),
679 "attribute" => child
680 .child_by_field_name("attribute")
681 .map(|n| self.text(n).to_owned()),
682 "call" => child
683 .child_by_field_name("function")
684 .and_then(|f| match f.kind() {
685 "identifier" => Some(self.text(f).to_owned()),
686 "attribute" => f
687 .child_by_field_name("attribute")
688 .map(|n| self.text(n).to_owned()),
689 _ => None,
690 }),
691 _ => None,
692 }
693 }
694
695 fn extract_param_types(&mut self, fn_node: TsNode<'_>, fn_id: &NodeId) {
698 let Some(params) = fn_node.child_by_field_name("parameters") else {
699 return;
700 };
701 let mut c = params.walk();
702 for param in params.named_children(&mut c) {
703 let type_node = match param.kind() {
704 "typed_parameter" | "typed_default_parameter" => param.child_by_field_name("type"),
705 _ => None,
706 };
707 if let Some(t) = type_node {
708 for name in self.collect_type_names(t) {
709 self.deferred_uses.push((fn_id.clone(), name));
710 }
711 }
712 }
713 }
714
715 fn extract_return_type(&mut self, fn_node: TsNode<'_>, fn_id: &NodeId) {
716 if let Some(ret) = fn_node.child_by_field_name("return_type") {
717 for name in self.collect_type_names(ret) {
718 self.deferred_uses.push((fn_id.clone(), name));
719 }
720 }
721 }
722
723 fn collect_type_names(&self, node: TsNode<'_>) -> Vec<String> {
725 let mut names = Vec::new();
726 self.walk_type_names(node, &mut names);
727 names
728 }
729
730 fn walk_type_names(&self, node: TsNode<'_>, out: &mut Vec<String>) {
731 match node.kind() {
732 "identifier" => {
733 let name = self.text(node).to_owned();
734 if !is_builtin_type(&name) {
735 out.push(name);
736 }
737 }
738 _ => {
739 let mut c = node.walk();
740 for child in node.named_children(&mut c) {
741 self.walk_type_names(child, out);
742 }
743 }
744 }
745 }
746}
747
748fn is_builtin_type(name: &str) -> bool {
751 matches!(
752 name,
753 "int"
754 | "str"
755 | "bool"
756 | "float"
757 | "complex"
758 | "bytes"
759 | "bytearray"
760 | "None"
761 | "list"
762 | "dict"
763 | "set"
764 | "frozenset"
765 | "tuple"
766 | "type"
767 | "object"
768 | "Any"
769 | "Optional"
770 | "Union"
771 | "List"
772 | "Dict"
773 | "Set"
774 | "FrozenSet"
775 | "Tuple"
776 | "Callable"
777 | "Type"
778 | "ClassVar"
779 | "Final"
780 | "Literal"
781 | "TypeVar"
782 | "Generic"
783 | "Protocol"
784 | "Sequence"
785 | "Iterable"
786 | "Iterator"
787 | "Generator"
788 | "Coroutine"
789 | "Awaitable"
790 | "AsyncIterator"
791 | "AsyncGenerator"
792 | "NoReturn"
793 | "Never"
794 | "Self"
795 | "Annotated"
796 | "TypeAlias"
797 | "ParamSpec"
798 | "TypeVarTuple"
799 | "overload"
800 | "abstractmethod"
801 | "staticmethod"
802 | "classmethod"
803 )
804}
805
806#[cfg(test)]
809mod tests {
810 use super::PythonParser;
811 use crate::parser::LanguageParser;
812 use gitcortex_core::schema::{EdgeKind, NodeKind};
813 use std::path::Path;
814
815 fn parse(
816 src: &str,
817 ) -> (
818 Vec<gitcortex_core::graph::Node>,
819 Vec<gitcortex_core::graph::Edge>,
820 ) {
821 let r = PythonParser::new()
822 .parse(Path::new("test.py"), src)
823 .unwrap();
824 (r.nodes, r.edges)
825 }
826
827 #[allow(clippy::type_complexity)]
828 fn parse_full(
829 src: &str,
830 ) -> (
831 Vec<gitcortex_core::graph::Node>,
832 Vec<gitcortex_core::graph::Edge>,
833 Vec<(gitcortex_core::graph::NodeId, String, u32)>,
834 Vec<(gitcortex_core::graph::NodeId, String)>,
835 Vec<(gitcortex_core::graph::NodeId, String)>,
836 Vec<(gitcortex_core::graph::NodeId, String)>,
837 ) {
838 let r = PythonParser::new()
839 .parse(Path::new("test.py"), src)
840 .unwrap();
841 (
842 r.nodes,
843 r.edges,
844 r.deferred_calls,
845 r.deferred_uses,
846 r.deferred_implements,
847 r.deferred_imports,
848 )
849 }
850
851 #[test]
852 fn parses_free_function() {
853 let (nodes, _) = parse("def greet(name):\n return name\n");
854 assert_eq!(nodes.len(), 2);
856 let fns: Vec<_> = nodes
857 .iter()
858 .filter(|n| n.kind == NodeKind::Function)
859 .collect();
860 assert_eq!(fns.len(), 1);
861 assert_eq!(fns[0].name, "greet");
862 }
863
864 #[test]
865 fn parses_class_and_method() {
866 let src = "class Person:\n def greet(self):\n pass\n";
867 let (nodes, edges) = parse(src);
868 let classes: Vec<_> = nodes
869 .iter()
870 .filter(|n| n.kind == NodeKind::Struct)
871 .collect();
872 let methods: Vec<_> = nodes
873 .iter()
874 .filter(|n| n.kind == NodeKind::Method)
875 .collect();
876 assert_eq!(classes.len(), 1);
877 assert_eq!(methods.len(), 1);
878 let contains: Vec<_> = edges
879 .iter()
880 .filter(|e| e.kind == EdgeKind::Contains)
881 .collect();
882 assert!(!contains.is_empty());
883 }
884
885 #[test]
886 fn detects_call_edges() {
887 let src = "def caller():\n callee()\ndef callee():\n pass\n";
888 let (_, edges) = parse(src);
889 let calls: Vec<_> = edges.iter().filter(|e| e.kind == EdgeKind::Calls).collect();
890 assert_eq!(calls.len(), 1);
891 }
892
893 #[test]
894 fn detects_base_class_implements() {
895 let src = "class Base:\n pass\nclass Child(Base):\n pass\n";
896 let (_, _, _, _, implements, _) = parse_full(src);
897 assert!(
898 implements.iter().any(|(_, name)| name == "Base"),
899 "expected Implements edge to Base, got: {implements:?}"
900 );
901 }
902
903 #[test]
904 fn detects_type_annotation_uses() {
905 let src = "class Foo:\n pass\ndef bar(x: Foo) -> Foo:\n pass\n";
906 let (_, _, _, uses, _, _) = parse_full(src);
907 let uses_foo: Vec<_> = uses.iter().filter(|(_, n)| n == "Foo").collect();
908 assert!(
909 uses_foo.len() >= 2,
910 "expected at least 2 Uses edges to Foo, got: {uses_foo:?}"
911 );
912 }
913
914 #[test]
915 fn detects_decorator_uses() {
916 let src = "class Foo:\n @property\n def name(self):\n return self._name\n";
917 let (_, _, _, uses, _, _) = parse_full(src);
918 assert!(
919 uses.iter().any(|(_, n)| n == "property"),
920 "expected Uses edge to 'property' decorator, got: {uses:?}"
921 );
922 }
923
924 #[test]
925 fn detects_import_statement() {
926 let src = "import os\nimport sys\n\ndef main():\n pass\n";
927 let (_, _, _, _, _, imports) = parse_full(src);
928 assert!(
929 imports.iter().any(|(_, n)| n == "os"),
930 "expected import 'os', got: {imports:?}"
931 );
932 assert!(
933 imports.iter().any(|(_, n)| n == "sys"),
934 "expected import 'sys', got: {imports:?}"
935 );
936 }
937
938 #[test]
939 fn detects_from_import_statement() {
940 let src = "from os.path import join, exists\n\ndef main():\n pass\n";
941 let (_, _, _, _, _, imports) = parse_full(src);
942 assert!(
943 imports.iter().any(|(_, n)| n == "join"),
944 "expected import 'join', got: {imports:?}"
945 );
946 assert!(
947 imports.iter().any(|(_, n)| n == "exists"),
948 "expected import 'exists', got: {imports:?}"
949 );
950 }
951
952 #[test]
953 fn module_node_is_emitted() {
954 let (nodes, _) = parse("x = 1\n");
955 let modules: Vec<_> = nodes
956 .iter()
957 .filter(|n| n.kind == NodeKind::Module)
958 .collect();
959 assert_eq!(modules.len(), 1);
960 assert_eq!(modules[0].name, "test");
961 }
962
963 #[test]
964 fn async_function_flagged() {
965 let src = "async def fetch():\n pass\n";
966 let (nodes, _) = parse(src);
967 let fns: Vec<_> = nodes
968 .iter()
969 .filter(|n| n.kind == NodeKind::Function)
970 .collect();
971 assert_eq!(fns.len(), 1);
972 assert!(fns[0].metadata.is_async);
973 }
974
975 #[test]
978 fn protocol_class_becomes_interface() {
979 let src = "from typing import Protocol\n\nclass MyProto(Protocol):\n def do_it(self) -> None:\n ...\n";
980 let (nodes, _) = parse(src);
981 let ifaces: Vec<_> = nodes
982 .iter()
983 .filter(|n| n.kind == NodeKind::Interface)
984 .collect();
985 assert_eq!(ifaces.len(), 1, "expected 1 Interface, got: {ifaces:?}");
986 assert_eq!(ifaces[0].name, "MyProto");
987 assert!(
988 ifaces[0].metadata.is_abstract,
989 "Protocol should be is_abstract"
990 );
991 }
992
993 #[test]
994 fn non_protocol_class_is_struct() {
995 let src = "class Plain:\n def work(self):\n pass\n";
996 let (nodes, _) = parse(src);
997 let structs: Vec<_> = nodes
998 .iter()
999 .filter(|n| n.kind == NodeKind::Struct)
1000 .collect();
1001 assert_eq!(structs.len(), 1);
1002 assert_eq!(structs[0].name, "Plain");
1003 assert!(!structs[0].metadata.is_abstract);
1004 }
1005
1006 #[test]
1009 fn property_decorator_yields_property_kind() {
1010 let src =
1011 "class Foo:\n @property\n def bar(self) -> str:\n return self._bar\n";
1012 let (nodes, _) = parse(src);
1013 let props: Vec<_> = nodes
1014 .iter()
1015 .filter(|n| n.kind == NodeKind::Property)
1016 .collect();
1017 assert_eq!(props.len(), 1, "expected 1 Property node, got: {props:?}");
1018 assert_eq!(props[0].name, "bar");
1019 assert!(props[0].metadata.is_property);
1020 }
1021
1022 #[test]
1023 fn staticmethod_decorator_sets_is_static() {
1024 let src = "class Foo:\n @staticmethod\n def create(x: int) -> 'Foo':\n return Foo()\n";
1025 let (nodes, _) = parse(src);
1026 let methods: Vec<_> = nodes
1027 .iter()
1028 .filter(|n| n.kind == NodeKind::Method)
1029 .collect();
1030 assert_eq!(methods.len(), 1);
1031 assert_eq!(methods[0].name, "create");
1032 assert!(
1033 methods[0].metadata.is_static,
1034 "staticmethod should set is_static"
1035 );
1036 }
1037
1038 #[test]
1039 fn classmethod_decorator_sets_is_static() {
1040 let src = "class Foo:\n @classmethod\n def from_str(cls, s: str) -> 'Foo':\n return cls()\n";
1041 let (nodes, _) = parse(src);
1042 let methods: Vec<_> = nodes
1043 .iter()
1044 .filter(|n| n.kind == NodeKind::Method)
1045 .collect();
1046 assert_eq!(methods.len(), 1);
1047 assert_eq!(methods[0].name, "from_str");
1048 assert!(
1049 methods[0].metadata.is_static,
1050 "classmethod should set is_static"
1051 );
1052 }
1053
1054 #[test]
1055 fn dataclass_decorator_class_is_struct() {
1056 let src = "from dataclasses import dataclass\n\n@dataclass\nclass Point:\n x: float\n y: float\n";
1057 let (nodes, _) = parse(src);
1058 let structs: Vec<_> = nodes
1059 .iter()
1060 .filter(|n| n.kind == NodeKind::Struct)
1061 .collect();
1062 assert_eq!(structs.len(), 1, "expected Struct for @dataclass Point");
1063 assert_eq!(structs[0].name, "Point");
1064 }
1065
1066 #[test]
1069 fn generator_function_sets_is_generator() {
1070 let src = "def numbers():\n yield 1\n yield 2\n";
1071 let (nodes, _) = parse(src);
1072 let fns: Vec<_> = nodes
1073 .iter()
1074 .filter(|n| n.kind == NodeKind::Function)
1075 .collect();
1076 assert_eq!(fns.len(), 1);
1077 assert!(
1078 fns[0].metadata.is_generator,
1079 "yield fn should be is_generator"
1080 );
1081 assert!(!fns[0].metadata.is_async);
1082 }
1083
1084 #[test]
1085 fn async_generator_is_both_async_and_generator() {
1086 let src = "async def stream():\n yield 1\n yield 2\n";
1087 let (nodes, _) = parse(src);
1088 let fns: Vec<_> = nodes
1089 .iter()
1090 .filter(|n| n.kind == NodeKind::Function)
1091 .collect();
1092 assert_eq!(fns.len(), 1);
1093 assert!(
1094 fns[0].metadata.is_async,
1095 "async generator should be is_async"
1096 );
1097 assert!(
1098 fns[0].metadata.is_generator,
1099 "async generator should be is_generator"
1100 );
1101 }
1102
1103 #[test]
1104 fn nested_yield_does_not_pollute_outer_function() {
1105 let src = "def outer():\n def inner():\n yield 1\n return inner()\n";
1107 let (nodes, _) = parse(src);
1108 let fns: Vec<_> = nodes
1109 .iter()
1110 .filter(|n| n.kind == NodeKind::Function)
1111 .collect();
1112 let outer = fns
1113 .iter()
1114 .find(|n| n.name == "outer")
1115 .expect("outer not found");
1116 assert!(
1117 !outer.metadata.is_generator,
1118 "outer should NOT be generator — yield is in nested fn"
1119 );
1120 }
1121
1122 #[test]
1125 fn module_level_bindings_detected() {
1126 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";
1130 let (nodes, _) = parse(src);
1131 let names: Vec<&str> = nodes
1132 .iter()
1133 .filter(|n| n.kind == NodeKind::Constant)
1134 .map(|n| n.name.as_str())
1135 .collect();
1136 for expected in [
1137 "MAX_SIZE",
1138 "DEFAULT_NAME",
1139 "default_config",
1140 "__version__",
1141 "T",
1142 ] {
1143 assert!(
1144 names.contains(&expected),
1145 "expected module-level binding `{expected}` as Constant; got {names:?}"
1146 );
1147 }
1148 assert!(
1149 !names.contains(&"local"),
1150 "function-local assignment must not be captured as a module Constant"
1151 );
1152 }
1153
1154 #[test]
1157 fn nested_class_emits_contains_edge_from_parent() {
1158 let src = "class Outer:\n class Inner:\n pass\n";
1159 let (nodes, edges) = parse(src);
1160 let outer = nodes
1161 .iter()
1162 .find(|n| n.name == "Outer")
1163 .expect("Outer not found");
1164 let inner = nodes
1165 .iter()
1166 .find(|n| n.name == "Inner")
1167 .expect("Inner not found");
1168 let contains: Vec<_> = edges
1169 .iter()
1170 .filter(|e| e.kind == EdgeKind::Contains)
1171 .collect();
1172 assert!(
1173 contains
1174 .iter()
1175 .any(|e| e.src == outer.id && e.dst == inner.id),
1176 "expected Contains Outer → Inner"
1177 );
1178 }
1179
1180 #[test]
1183 fn multiple_type_annotations_produce_uses_entries() {
1184 let src =
1185 "class Req:\n pass\nclass Resp:\n pass\ndef handler(r: Req) -> Resp:\n pass\n";
1186 let (_, _, _, uses, _, _) = parse_full(src);
1187 let uses_req: Vec<_> = uses.iter().filter(|(_, n)| n == "Req").collect();
1188 let uses_resp: Vec<_> = uses.iter().filter(|(_, n)| n == "Resp").collect();
1189 assert!(!uses_req.is_empty(), "expected Uses edge to Req");
1190 assert!(!uses_resp.is_empty(), "expected Uses edge to Resp");
1191 }
1192
1193 #[test]
1196 fn private_method_has_private_visibility() {
1197 use gitcortex_core::schema::Visibility;
1198 let src = "class Foo:\n def _internal(self):\n pass\n def public(self):\n pass\n";
1199 let (nodes, _) = parse(src);
1200 let internal = nodes
1201 .iter()
1202 .find(|n| n.name == "_internal")
1203 .expect("_internal not found");
1204 let public = nodes
1205 .iter()
1206 .find(|n| n.name == "public")
1207 .expect("public not found");
1208 assert_eq!(internal.metadata.visibility, Visibility::Private);
1209 assert_eq!(public.metadata.visibility, Visibility::Pub);
1210 }
1211
1212 #[test]
1215 fn calls_edge_between_two_functions() {
1216 let src = "def helper():\n pass\n\ndef main():\n helper()\n";
1217 let (_, edges) = parse(src);
1218 let calls: Vec<_> = edges.iter().filter(|e| e.kind == EdgeKind::Calls).collect();
1219 assert_eq!(calls.len(), 1, "expected exactly 1 Calls edge");
1220 }
1221
1222 #[test]
1223 fn method_call_via_self_creates_calls_edge() {
1224 let src = "class Svc:\n def run(self):\n self.process()\n def process(self):\n pass\n";
1225 let (nodes, edges) = parse(src);
1226 let run = nodes
1227 .iter()
1228 .find(|n| n.name == "run")
1229 .expect("run not found");
1230 let process = nodes
1231 .iter()
1232 .find(|n| n.name == "process")
1233 .expect("process not found");
1234 let calls: Vec<_> = edges.iter().filter(|e| e.kind == EdgeKind::Calls).collect();
1235 assert!(
1237 calls.iter().any(|e| e.src == run.id && e.dst == process.id),
1238 "expected Calls edge run → process, got: {calls:?}"
1239 );
1240 }
1241
1242 #[test]
1245 fn aliased_import_uses_alias_name() {
1246 let src = "import numpy as np\nimport pandas as pd\n";
1247 let (_, _, _, _, _, imports) = parse_full(src);
1248 let names: Vec<&str> = imports.iter().map(|(_, n)| n.as_str()).collect();
1250 assert!(
1251 names.contains(&"numpy"),
1252 "expected import 'numpy', got: {names:?}"
1253 );
1254 assert!(
1255 names.contains(&"pandas"),
1256 "expected import 'pandas', got: {names:?}"
1257 );
1258 }
1259
1260 #[test]
1261 fn dotted_import_records_leaf_module() {
1262 let src = "import os.path\n";
1263 let (_, _, _, _, _, imports) = parse_full(src);
1264 let names: Vec<&str> = imports.iter().map(|(_, n)| n.as_str()).collect();
1265 assert!(
1266 names.contains(&"path"),
1267 "expected leaf 'path' from 'import os.path', got: {names:?}"
1268 );
1269 }
1270}