scryer-engine 0.2.1

Tree-sitter and stack-graphs AST indexing engine for Scryer code intelligence
use std::path::Path;
use tree_sitter::{Node, Parser};

use crate::payload::{ParsedFilePayload, RawScope, RawSymbol};

/// Tree-sitter AST parser and entity extractor for Python source code.
pub struct PythonAstParser;

impl PythonAstParser {
    /// Parse Python source code and extract all declarations, signatures, docstrings, and lexical scopes.
    pub fn parse(
        relative_path: &Path,
        content: &[u8],
        content_hash: &str,
    ) -> anyhow::Result<ParsedFilePayload> {
        let source_str = std::str::from_utf8(content)
            .map_err(|e| anyhow::anyhow!("Invalid UTF-8 in {}: {e}", relative_path.display()))?;

        let mut parser = Parser::new();
        let language = tree_sitter_python::LANGUAGE;
        parser
            .set_language(&language.into())
            .map_err(|e| anyhow::anyhow!("Failed to set Python tree-sitter language: {e}"))?;

        let tree = parser.parse(content, None).ok_or_else(|| {
            anyhow::anyhow!("Tree-sitter failed to parse {}", relative_path.display())
        })?;

        let base_module = compute_base_module(relative_path);
        let mut extractor = PythonAstExtractor::new(source_str, base_module);

        // Add root module scope
        let root_node = tree.root_node();
        let root_scope = extractor.add_scope("module", root_node, None);

        extractor.traverse(root_node, Some(root_scope), &[]);

        let line_count = source_str.lines().count() as u32;

        Ok(ParsedFilePayload {
            relative_path: relative_path.to_path_buf(),
            content_hash: content_hash.to_string(),
            language: "python".to_string(),
            line_count: if line_count == 0 { 1 } else { line_count },
            byte_size: content.len(),
            scopes: extractor.scopes,
            symbols: extractor.symbols,
            references: Vec::new(),
            edges: Vec::new(),
        })
    }
}

/// Compute the base Python module path from the relative file path.
/// E.g.: `pkg/utils.py` -> "pkg.utils", `pkg/__init__.py` -> "pkg"
fn compute_base_module(rel_path: &Path) -> String {
    let mut components = Vec::new();

    for part in rel_path.iter() {
        let part_str = part.to_string_lossy();
        if let Some(stem) = part_str.strip_suffix(".py") {
            if stem != "__init__" && stem != "__main__" {
                components.push(stem.to_string());
            }
        } else {
            components.push(part_str.to_string());
        }
    }

    if components.is_empty() {
        "__main__".to_string()
    } else {
        components.join(".")
    }
}

struct PythonAstExtractor<'a> {
    source: &'a str,
    base_module: String,
    scopes: Vec<RawScope>,
    symbols: Vec<RawSymbol>,
}

impl<'a> PythonAstExtractor<'a> {
    fn new(source: &'a str, base_module: String) -> Self {
        Self {
            source,
            base_module,
            scopes: Vec::new(),
            symbols: Vec::new(),
        }
    }

    fn add_scope(&mut self, kind: &str, node: Node, parent_scope: Option<usize>) -> usize {
        let local_id = self.scopes.len();
        self.scopes.push(RawScope {
            local_id,
            parent_local_id: parent_scope,
            scope_kind: kind.to_string(),
            start_byte: node.start_byte(),
            end_byte: node.end_byte(),
            start_line: (node.start_position().row + 1) as u32,
            end_line: (node.end_position().row + 1) as u32,
        });
        local_id
    }

    fn get_node_text(&self, node: Node) -> &'a str {
        &self.source[node.start_byte()..node.end_byte()]
    }

    /// Extract docstring from a function or class block body.
    /// In Python, docstring is the first statement in the body if it's an expression statement containing a string.
    fn extract_body_docstring(&self, body_node: Node) -> Option<String> {
        let mut cursor = body_node.walk();
        for child in body_node.children(&mut cursor) {
            if child.kind() == "expression_statement" {
                let mut expr_cursor = child.walk();
                for expr_child in child.children(&mut expr_cursor) {
                    if expr_child.kind() == "string" {
                        let text = self.get_node_text(expr_child).trim();
                        return Some(clean_python_docstring(text));
                    }
                }
            } else if child.kind() == "comment" {
                continue;
            } else {
                break;
            }
        }
        None
    }

    fn extract_visibility(&self, name: &str) -> String {
        if name.starts_with('_') && !name.starts_with("__") {
            "private".to_string()
        } else {
            "public".to_string()
        }
    }

    fn extract_signature(&self, node: Node) -> String {
        let raw_header = if let Some(body) = node.child_by_field_name("body") {
            &self.source[node.start_byte()..body.start_byte()]
        } else {
            self.get_node_text(node).lines().next().unwrap_or("")
        };
        raw_header
            .trim()
            .trim_end_matches(':')
            .split_whitespace()
            .collect::<Vec<&str>>()
            .join(" ")
    }

    fn make_qualified_name(&self, name: &str, qualifiers: &[String]) -> String {
        let mut parts = vec![self.base_module.clone()];
        parts.extend_from_slice(qualifiers);
        parts.push(name.to_string());
        parts.join(".")
    }

    fn traverse(&mut self, node: Node, current_scope: Option<usize>, qualifiers: &[String]) {
        match node.kind() {
            "function_definition" => {
                let name = node
                    .child_by_field_name("name")
                    .map(|n| self.get_node_text(n))
                    .unwrap_or("anonymous_fn");

                let visibility = self.extract_visibility(name);
                let signature = self.extract_signature(node);
                let docstring = node
                    .child_by_field_name("body")
                    .and_then(|body| self.extract_body_docstring(body));
                let qualified_name = self.make_qualified_name(name, qualifiers);

                let scope_id = self.add_scope("function", node, current_scope);

                self.symbols.push(RawSymbol {
                    scope_local_id: current_scope,
                    name: name.to_string(),
                    qualified_name: qualified_name.clone(),
                    kind: "fn".to_string(),
                    visibility,
                    signature,
                    docstring,
                    start_byte: node.start_byte(),
                    end_byte: node.end_byte(),
                    start_line: (node.start_position().row + 1) as u32,
                    end_line: (node.end_position().row + 1) as u32,
                });

                if let Some(body) = node.child_by_field_name("body") {
                    let mut child_qualifiers = qualifiers.to_vec();
                    child_qualifiers.push(name.to_string());

                    let mut cursor = body.walk();
                    for child in body.children(&mut cursor) {
                        self.traverse(child, Some(scope_id), &child_qualifiers);
                    }
                }
            }

            "class_definition" => {
                let name = node
                    .child_by_field_name("name")
                    .map(|n| self.get_node_text(n))
                    .unwrap_or("AnonymousClass");

                let visibility = self.extract_visibility(name);
                let signature = self.extract_signature(node);
                let docstring = node
                    .child_by_field_name("body")
                    .and_then(|body| self.extract_body_docstring(body));
                let qualified_name = self.make_qualified_name(name, qualifiers);

                let scope_id = self.add_scope("class", node, current_scope);

                self.symbols.push(RawSymbol {
                    scope_local_id: current_scope,
                    name: name.to_string(),
                    qualified_name: qualified_name.clone(),
                    kind: "class".to_string(),
                    visibility,
                    signature,
                    docstring,
                    start_byte: node.start_byte(),
                    end_byte: node.end_byte(),
                    start_line: (node.start_position().row + 1) as u32,
                    end_line: (node.end_position().row + 1) as u32,
                });

                if let Some(body) = node.child_by_field_name("body") {
                    let mut child_qualifiers = qualifiers.to_vec();
                    child_qualifiers.push(name.to_string());

                    let mut cursor = body.walk();
                    for child in body.children(&mut cursor) {
                        self.traverse(child, Some(scope_id), &child_qualifiers);
                    }
                }
            }

            "decorated_definition" => {
                // In Python, decorated_definition wraps a function or class definition
                let mut cursor = node.walk();
                for child in node.children(&mut cursor) {
                    if child.kind() == "function_definition" || child.kind() == "class_definition" {
                        self.traverse(child, current_scope, qualifiers);
                    }
                }
            }

            "assignment" => {
                // Top-level or class-level constants / variables
                if (current_scope.is_none() || current_scope == Some(0) || qualifiers.len() <= 1)
                    && let Some(left) = node
                        .child_by_field_name("left")
                        .filter(|l| l.kind() == "identifier")
                {
                    let name = self.get_node_text(left);
                    let is_const = name.chars().all(|c| c.is_uppercase() || c == '_');
                    let kind = if is_const { "const" } else { "var" };
                    let visibility = self.extract_visibility(name);
                    let signature = self
                        .get_node_text(node)
                        .lines()
                        .next()
                        .unwrap_or("")
                        .trim()
                        .to_string();
                    let qualified_name = self.make_qualified_name(name, qualifiers);

                    self.symbols.push(RawSymbol {
                        scope_local_id: current_scope,
                        name: name.to_string(),
                        qualified_name,
                        kind: kind.to_string(),
                        visibility,
                        signature,
                        docstring: None,
                        start_byte: node.start_byte(),
                        end_byte: node.end_byte(),
                        start_line: (node.start_position().row + 1) as u32,
                        end_line: (node.end_position().row + 1) as u32,
                    });
                }
            }

            _ => {
                let mut cursor = node.walk();
                for child in node.children(&mut cursor) {
                    self.traverse(child, current_scope, qualifiers);
                }
            }
        }
    }
}

/// Clean triple or single quotes from Python docstrings.
fn clean_python_docstring(raw: &str) -> String {
    let mut s = raw.trim();
    if s.starts_with("r\"\"\"") || s.starts_with("u\"\"\"") || s.starts_with("f\"\"\"") {
        s = &s[4..];
    } else if s.starts_with("\"\"\"") || s.starts_with("'''") {
        s = &s[3..];
    } else if s.starts_with('"') || s.starts_with('\'') {
        s = &s[1..];
    }

    if s.ends_with("\"\"\"") || s.ends_with("'''") {
        s = &s[..s.len() - 3];
    } else if s.ends_with('"') || s.ends_with('\'') {
        s = &s[..s.len() - 1];
    }

    s.trim().to_string()
}