Recursive node support, enum expressions, or/and/not

This commit is contained in:
Ganesh Viswanathan 2019-01-18 21:52:29 -06:00
commit cd8a263f85
7 changed files with 75 additions and 38 deletions

View file

@ -4,18 +4,6 @@ import regex
import "."/[getters, globals, grammar, treesitter/runtime] import "."/[getters, globals, grammar, treesitter/runtime]
const gAtoms = @[
"field_identifier",
"identifier",
"shift_expression",
"math_expression",
"number_literal",
"preproc_arg",
"primitive_type",
"sized_type_specifier",
"type_identifier"
].toSet()
proc saveNodeData(node: TSNode): bool = proc saveNodeData(node: TSNode): bool =
let name = $node.tsNodeType() let name = $node.tsNodeType()
if name in gAtoms: if name in gAtoms:
@ -25,10 +13,10 @@ proc saveNodeData(node: TSNode): bool =
if name == "primitive_type" and node.tsNodeParent.tsNodeType() == "sized_type_specifier": if name == "primitive_type" and node.tsNodeParent.tsNodeType() == "sized_type_specifier":
return true return true
if name == "number_literal" and $node.tsNodeParent.tsNodeType() in ["shift_expression", "math_expression"]: if name == "number_literal" and $node.tsNodeParent.tsNodeType() in gExpressions:
return true return true
if name in ["math_expression", "primitive_type", "sized_type_specifier"]: if name in ["primitive_type", "sized_type_specifier"]:
val = val.getType() val = val.getType()
let let
@ -52,6 +40,10 @@ proc saveNodeData(node: TSNode): bool =
ppname == "function_declarator": ppname == "function_declarator":
gStateRT.data.add(("function_declarator", "")) gStateRT.data.add(("function_declarator", ""))
elif name in gExpressions:
if $node.tsNodeParent.tsNodeType() notin gExpressions:
gStateRT.data.add((name, node.getNodeVal()))
elif name in ["abstract_pointer_declarator", "enumerator", "field_declaration", "function_declarator"]: elif name in ["abstract_pointer_declarator", "enumerator", "field_declaration", "function_declarator"]:
gStateRT.data.add((name.replace("abstract_", ""), "")) gStateRT.data.add((name.replace("abstract_", ""), ""))
@ -65,14 +57,19 @@ proc searchAstForNode(ast: ref Ast, node: TSNode): bool =
return return
if ast.children.len != 0: if ast.children.len != 0:
if childNames.contains(ast.regex): if childNames.contains(ast.regex) or
(childNames.len == 0 and ast.recursive):
if node.getTSNodeNamedChildCountSansComments() != 0: if node.getTSNodeNamedChildCountSansComments() != 0:
var flag = true var flag = true
for i in 0 .. node.tsNodeNamedChildCount()-1: for i in 0 .. node.tsNodeNamedChildCount()-1:
if $node.tsNodeNamedChild(i).tsNodeType() != "comment": if $node.tsNodeNamedChild(i).tsNodeType() != "comment":
let let
nodeChild = node.tsNodeNamedChild(i) nodeChild = node.tsNodeNamedChild(i)
astChild = ast.getAstChildByName($nodeChild.tsNodeType()) astChild =
if not ast.recursive:
ast.getAstChildByName($nodeChild.tsNodeType())
else:
ast
if not searchAstForNode(astChild, nodeChild): if not searchAstForNode(astChild, nodeChild):
flag = false flag = false
break break

View file

@ -236,12 +236,16 @@ converter toKind*(kind: string): Kind =
else: else:
exactlyOne exactlyOne
proc getNameKind*(name: string): tuple[name: string, kind: Kind] = proc getNameKind*(name: string): tuple[name: string, kind: Kind, recursive: bool] =
result.name = name if name[0] == '^':
result.recursive = true
result.name = name[1 .. ^1]
else:
result.name = name
result.kind = $name[^1] result.kind = $name[^1]
if result.kind != exactlyOne: if result.kind != exactlyOne:
result.name = name[0 .. ^2] result.name = result.name[0 .. ^2]
proc getTSNodeNamedChildCountSansComments*(node: TSNode): int = proc getTSNodeNamedChildCountSansComments*(node: TSNode): int =
if node.tsNodeNamedChildCount() != 0: if node.tsNodeNamedChildCount() != 0:
@ -261,8 +265,9 @@ proc getTSNodeNamedChildNames*(node: TSNode): seq[string] =
proc getRegexForAstChildren*(ast: ref Ast): string = proc getRegexForAstChildren*(ast: ref Ast): string =
result = "^" result = "^"
for i in 0 .. ast.children.len-1: for i in 0 .. ast.children.len-1:
let kind: string = ast.children[i].kind let
let begin = if result[^1] == '|': "" else: "(?:" kind: string = ast.children[i].kind
begin = if result[^1] == '|': "" else: "(?:"
case kind: case kind:
of "!": of "!":
result &= &"{begin}{ast.children[i].name}|" result &= &"{begin}{ast.children[i].name}|"

View file

@ -1,10 +1,33 @@
import sets, tables import sequtils, sets, tables
import regex import regex
when not declared(CIMPORT): when not declared(CIMPORT):
import "."/treesitter/runtime import "."/treesitter/runtime
const
gAtoms* = @[
"field_identifier",
"identifier",
"number_literal",
"preproc_arg",
"primitive_type",
"sized_type_specifier",
"type_identifier"
].toSet()
gExpressions* = @[
"parenthesized_expression",
"bitwise_expression",
"shift_expression",
"math_expression"
].toSet()
gEnumVals* = @[
"identifier",
"number_literal"
].concat(toSeq(gExpressions.items))
type type
Kind* = enum Kind* = enum
exactlyOne exactlyOne
@ -16,6 +39,7 @@ type
Ast* = object Ast* = object
name*: string name*: string
kind*: Kind kind*: Kind
recursive*: bool
children*: seq[ref Ast] children*: seq[ref Ast]
when not declared(CIMPORT): when not declared(CIMPORT):
tonim*: proc (ast: ref Ast, node: TSNode) tonim*: proc (ast: ref Ast, node: TSNode)

View file

@ -325,14 +325,15 @@ proc initGrammar() =
if gStateRT.consts.addNewIdentifer(fname): if gStateRT.consts.addNewIdentifer(fname):
if i+1 < gStateRT.data.len-fend and if i+1 < gStateRT.data.len-fend and
gStateRT.data[i+1].name in ["identifier", "shift_expression", "math_expression", "number_literal"]: gStateRT.data[i+1].name in gEnumVals:
if " " in gStateRT.data[i+1].val:
gStateRT.data[i+1].val = "(" & gStateRT.data[i+1].val.replace(" ", "") & ")"
gStateRT.data[i+1].val = gStateRT.data[i+1].val.multiReplace([ gStateRT.data[i+1].val = gStateRT.data[i+1].val.multiReplace([
("<<", " shl "), (">>", " shr ") (" ", ""),
("<<", " shl "), (">>", " shr "),
("^", " xor "), ("&", " and "), ("|", " or "),
("~", " not ")
]) ])
gStateRT.constStr &= &" {fname}* = {gStateRT.data[i+1].val}.{nname}\n" gStateRT.constStr &= &" {fname}* = ({gStateRT.data[i+1].val}).{nname}\n"
try: try:
count = gStateRT.data[i+1].val.parseInt() + 1 count = gStateRT.data[i+1].val.parseInt() + 1
except: except:
@ -349,15 +350,12 @@ proc initGrammar() =
(type_identifier?) (type_identifier?)
(enumerator_list (enumerator_list
(enumerator+ (enumerator+
(identifier+) (identifier?)
(number_literal?) (^$1+)
(shift_expression|math_expression?
(number_literal+)
)
) )
) )
) )
""", """ % gEnumVals.join("|"),
proc (ast: ref Ast, node: TSNode) = proc (ast: ref Ast, node: TSNode) =
var var
name = "" name = ""
@ -439,8 +437,9 @@ proc initGrammar() =
proc initRegex(ast: ref Ast) = proc initRegex(ast: ref Ast) =
if ast.children.len != 0: if ast.children.len != 0:
for child in ast.children: if not ast.recursive:
child.initRegex() for child in ast.children:
child.initRegex()
var var
reg: string reg: string

View file

@ -34,8 +34,10 @@ proc readFromTokens(): ref Ast =
quit(1) quit(1)
if gTokens[idx+1] != "comment": if gTokens[idx+1] != "comment":
result = new(Ast) result = new(Ast)
(result.name, result.kind) = gTokens[idx+1].getNameKind() (result.name, result.kind, result.recursive) = gTokens[idx+1].getNameKind()
result.children = @[] result.children = @[]
if result.recursive:
result.children.add(result)
idx += 2 idx += 2
while gTokens[idx] != ")": while gTokens[idx] != ")":
var res = readFromTokens() var res = readFromTokens()
@ -48,9 +50,9 @@ proc readFromTokens(): ref Ast =
idx += 1 idx += 1
proc printAst*(node: ref Ast, offset=""): string = proc printAst*(node: ref Ast, offset=""): string =
result = offset & "(" & node.name & node.kind.toString() result = offset & "(" & (if node.recursive: "^" else: "") & node.name & node.kind.toString()
if node.children.len != 0: if node.children.len != 0 and not node.recursive:
result &= "\n" result &= "\n"
for child in node.children: for child in node.children:
result &= printAst(child, offset & " ") result &= printAst(child, offset & " ")

View file

@ -60,6 +60,12 @@ typedef enum ENUM4 {
enum12 enum12
} ENUM4; } ENUM4;
enum ENUM5 {
enum13 = (1 << 2),
enum14 = ((1 << 3) | 1),
enum15 = (1 << (1 & 1))
};
typedef void * VOIDPTR; typedef void * VOIDPTR;
typedef int * INTPTR; typedef int * INTPTR;

View file

@ -97,6 +97,10 @@ else:
check e3 == enum7 check e3 == enum7
check e4 == enum11 check e4 == enum11
check enum13 == 4
check enum14 == 9
check enum15 == 2
cAddStdDir() cAddStdDir()
## failing tests ## failing tests