use tree_sitter::Node;
use crate::core::{Kind, Symbol};
use crate::lang::{Ctx, LanguagePlugin, extract_with, qualify};
const LANGUAGE: &str = "python";
pub(crate) struct Python;
impl LanguagePlugin for Python {
fn language(&self) -> &'static str {
LANGUAGE
}
fn extensions(&self) -> &[&str] {
&["py"]
}
fn constructor(&self) -> Option<&'static str> {
Some("__init__")
}
fn extract(&self, file: &str, source: &str) -> Vec<Symbol> {
extract_with(
LANGUAGE,
tree_sitter_python::LANGUAGE.into(),
file,
source,
|ctx, root, out| walk(ctx, root, None, Scope::Module, out),
)
}
}
#[derive(Clone, Copy, PartialEq, Eq)]
enum Scope {
Module,
Class,
Enum,
Function,
}
fn walk(ctx: &Ctx, node: Node, parent: Option<&str>, scope: Scope, out: &mut Vec<Symbol>) {
let mut cursor = node.walk();
for child in node.children(&mut cursor) {
match child.kind() {
"class_definition" if scope != Scope::Function => {
if let Some(name) = ctx.field_text(child, "name") {
let (kind, body) = if is_enum(ctx, child) {
(Kind::Enum, Scope::Enum)
} else {
(Kind::Class, Scope::Class)
};
let mut s = ctx.symbol(&name, kind, child, parent);
s.visibility = Some(name_visibility(&name));
out.push(s);
let qualified = qualify(parent, &name, ".");
walk(ctx, child, Some(&qualified), body, out);
}
}
"function_definition" => {
if let Some(name) = ctx.field_text(child, "name") {
let (kind, vis) = match scope {
Scope::Class | Scope::Enum => (Kind::Method, name_visibility(&name)),
Scope::Module => (Kind::Function, name_visibility(&name)),
Scope::Function => (Kind::Function, "local"),
};
let mut s = ctx.symbol(&name, kind, child, parent);
s.visibility = Some(vis);
out.push(s);
let qualified = qualify(parent, &name, ".");
walk(ctx, child, Some(&qualified), Scope::Function, out);
}
}
"class_definition" => {
if let Some(name) = ctx.field_text(child, "name") {
let mut s = ctx.symbol(&name, Kind::Class, child, parent);
s.visibility = Some("local");
out.push(s);
}
}
"assignment" if scope == Scope::Enum => members(ctx, child, parent, out),
"assignment" if scope != Scope::Function => constants(ctx, child, parent, out),
_ => walk(ctx, child, parent, scope, out),
}
}
}
fn is_enum(ctx: &Ctx, class: Node) -> bool {
let Some(bases) = class.child_by_field_name("superclasses") else {
return false;
};
let mut cursor = bases.walk();
bases
.named_children(&mut cursor)
.filter(|b| matches!(b.kind(), "identifier" | "attribute"))
.filter_map(|b| ctx.node_text(b))
.any(|b| {
let last = b.rsplit('.').next().unwrap_or(&b);
["Enum", "Flag", "Choices"]
.iter()
.any(|s| last.ends_with(s))
})
}
fn members(ctx: &Ctx, assign: Node, parent: Option<&str>, out: &mut Vec<Symbol>) {
if assign.child_by_field_name("right").is_none() {
return;
}
bindings(assign, &mut |target| {
if let Some(name) = ctx.node_text(target)
&& !name.starts_with('_')
{
let mut s = ctx.symbol(&name, Kind::Variant, assign, parent);
s.visibility = Some("public");
out.push(s);
}
});
}
fn constants(ctx: &Ctx, assign: Node, parent: Option<&str>, out: &mut Vec<Symbol>) {
bindings(assign, &mut |target| {
if let Some(name) = ctx.node_text(target)
&& is_constant_name(&name)
{
let mut s = ctx.symbol(&name, Kind::Constant, assign, parent);
s.visibility = Some(name_visibility(&name));
out.push(s);
}
});
}
fn bindings(assign: Node, emit: &mut impl FnMut(Node)) {
if let Some(left) = assign.child_by_field_name("left") {
match left.kind() {
"identifier" => emit(left),
"pattern_list" | "tuple_pattern" => {
let mut cursor = left.walk();
left.named_children(&mut cursor)
.filter(|n| n.kind() == "identifier")
.for_each(&mut *emit);
}
_ => {} }
}
if let Some(right) = assign.child_by_field_name("right")
&& right.kind() == "assignment"
{
bindings(right, emit);
}
}
fn is_constant_name(name: &str) -> bool {
name.chars()
.all(|c| c.is_ascii_uppercase() || c.is_ascii_digit() || c == '_')
&& name.chars().filter(char::is_ascii_uppercase).count() >= 2
}
fn name_visibility(name: &str) -> &'static str {
let dunder = name.starts_with("__") && name.ends_with("__");
if name.starts_with('_') && !dunder {
"private"
} else {
"public"
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::lang::testing::find;
fn extract(source: &str) -> Vec<Symbol> {
Python.extract("test.py", source)
}
#[test]
fn extracts_classes_methods_and_functions() {
let src = r#"
class Account:
def deposit(self, amount):
pass
@property
def balance(self):
return 0
def build():
return Account()
"#;
let syms = extract(src);
let account = find(&syms, "Account");
assert_eq!(account.kind, Kind::Class);
assert_eq!(account.parent, None);
let deposit = find(&syms, "deposit");
assert_eq!(deposit.kind, Kind::Method);
assert_eq!(deposit.parent.as_deref(), Some("Account"));
assert_eq!(find(&syms, "balance").kind, Kind::Method);
let build = find(&syms, "build");
assert_eq!(build.kind, Kind::Function);
assert_eq!(build.parent, None);
assert_eq!(account.language, "python");
}
#[test]
fn nested_defs_and_classes_are_local_to_their_enclosing_def() {
let src = r#"
def _multi_decorate(decorators, method):
def _wrapper(self, *args):
def inner():
pass
return inner
class Local:
def hidden(self):
pass
LIMIT = 3
return _wrapper
class Account:
def deposit(self):
@cached
def helper():
pass
"#;
let syms = extract(src);
let wrapper = find(&syms, "_wrapper");
assert_eq!(wrapper.kind, Kind::Function);
assert_eq!(wrapper.parent.as_deref(), Some("_multi_decorate"));
assert_eq!(wrapper.visibility, Some("local"));
let inner = find(&syms, "inner");
assert_eq!(inner.parent.as_deref(), Some("_multi_decorate._wrapper"));
assert_eq!(inner.visibility, Some("local"));
let helper = find(&syms, "helper");
assert_eq!(helper.kind, Kind::Function);
assert_eq!(helper.parent.as_deref(), Some("Account.deposit"));
let local = find(&syms, "Local");
assert_eq!(local.kind, Kind::Class);
assert_eq!(local.parent.as_deref(), Some("_multi_decorate"));
assert_eq!(local.visibility, Some("local"));
for absent in ["hidden", "LIMIT"] {
assert!(!syms.iter().any(|s| s.name == absent), "{absent}: {syms:?}");
}
}
#[test]
fn visible_enum_subclasses_are_enums_of_variants() {
let src = r#"
import enum
from django.db import models
class Color(enum.Enum):
RED = 1
green = auto()
_ignore_ = ["tmp"]
size: int
def describe(self):
return self.name
class Year(models.TextChoices):
FRESHMAN = "FR", _("Freshman")
class Perm(IntFlag, metaclass=Meta):
READ = 4
class Plain(Base):
LIMIT = 3
"#;
let syms = extract(src);
for (name, parent) in [("Color", None), ("Year", None), ("Perm", None)] {
let s = find(&syms, name);
assert_eq!(
(s.kind, s.parent.as_deref()),
(Kind::Enum, parent),
"{name}"
);
}
for (name, parent) in [
("RED", "Color"),
("green", "Color"),
("FRESHMAN", "Year"),
("READ", "Perm"),
] {
let s = find(&syms, name);
assert_eq!(
(s.kind, s.parent.as_deref()),
(Kind::Variant, Some(parent)),
"{name}"
);
}
assert_eq!(find(&syms, "describe").kind, Kind::Method);
for absent in ["_ignore_", "size"] {
assert!(!syms.iter().any(|s| s.name == absent), "{absent}: {syms:?}");
}
assert_eq!(find(&syms, "Plain").kind, Kind::Class);
assert_eq!(find(&syms, "LIMIT").kind, Kind::Constant);
}
#[test]
fn upper_snake_assignments_are_constants() {
let src = r#"
MAX_RETRIES = 3
TIMEOUT: float = 1.5
LOW, HIGH = 1, 9
FIRST = SECOND = 0
_INTERNAL_LIMIT = 2
T = TypeVar("T")
default_widget = None
try:
FAST_PATH = True
except ImportError:
pass
class Account:
DEFAULT_BALANCE = 0
kind = "basic"
def deposit(self, amount):
LOCAL_CAP = 10
self.LIMIT = amount
"#;
let syms = extract(src);
let max = find(&syms, "MAX_RETRIES");
assert_eq!(max.kind, Kind::Constant);
assert_eq!(max.parent, None);
assert_eq!(max.visibility, Some("public"));
for name in ["TIMEOUT", "LOW", "HIGH", "FIRST", "SECOND", "FAST_PATH"] {
assert_eq!(find(&syms, name).kind, Kind::Constant, "{name}");
}
assert_eq!(find(&syms, "_INTERNAL_LIMIT").visibility, Some("private"));
let default = find(&syms, "DEFAULT_BALANCE");
assert_eq!(default.kind, Kind::Constant);
assert_eq!(default.parent.as_deref(), Some("Account"));
for absent in ["T", "default_widget", "kind", "LOCAL_CAP", "LIMIT"] {
assert!(!syms.iter().any(|s| s.name == absent), "{absent}: {syms:?}");
}
}
#[test]
fn empty_and_unparseable_yield_no_symbols() {
assert!(extract("").is_empty());
assert!(extract("# just a comment\n").is_empty());
}
#[test]
fn underscore_names_read_as_private_except_dunders() {
let src = "class Account:\n def _internal(self):\n pass\n def __init__(self):\n pass\n\ndef fetch():\n pass\n";
let syms = extract(src);
assert_eq!(find(&syms, "_internal").visibility, Some("private"));
assert_eq!(find(&syms, "__init__").visibility, Some("public"));
assert_eq!(find(&syms, "fetch").visibility, Some("public"));
}
}