diff --git a/lsp/cache/parse.go b/lsp/cache/parse.go index 25251ac..1d0143c 100644 --- a/lsp/cache/parse.go +++ b/lsp/cache/parse.go @@ -216,6 +216,14 @@ type ParsedFile struct { // tokens is the identifier set of ast, computed lazily once per parse. tokens map[string]struct{} + + // defs and enumValues index the file's definitions, computed lazily + // once per parse. A re-parse replaces the whole ParsedFile, so the + // caches never go stale. + defsOnce sync.Once + defs map[string]syntax.Node + enumOnce sync.Once + enumValues map[string]*syntax.Identifier } func (p *ParsedFile) Mapper() *mapper.Mapper { @@ -248,6 +256,55 @@ func (p *ParsedFile) Tokens() map[string]struct{} { return tokens } +// Definitions returns the file's top-level definitions indexed by name: +// structs, unions, exceptions, enums, services, consts, and typedefs. The +// node's concrete type identifies the definition kind. +func (p *ParsedFile) Definitions() map[string]syntax.Node { + p.defsOnce.Do(func() { + p.defs = map[string]syntax.Node{} + + if p.ast == nil { + return + } + + for _, n := range p.ast.Nodes { + switch v := n.(type) { + case *syntax.Struct: + p.defs[v.Name.Text] = v + case *syntax.Enum: + p.defs[v.Name.Text] = v + case *syntax.Service: + p.defs[v.Name.Text] = v + case *syntax.Const: + p.defs[v.Name.Text] = v + case *syntax.Typedef: + p.defs[v.Name.Text] = v + } + } + }) + + return p.defs +} + +// EnumValues returns the file's enum value names indexed by name. +func (p *ParsedFile) EnumValues() map[string]*syntax.Identifier { + p.enumOnce.Do(func() { + p.enumValues = map[string]*syntax.Identifier{} + + if p.ast == nil { + return + } + + for _, enum := range p.ast.Enums() { + for _, value := range enum.Values { + p.enumValues[value.Name.Text] = value.Name + } + } + }) + + return p.enumValues +} + func (p *ParsedFile) AggregatedError() error { if len(p.errs) == 0 { return nil diff --git a/lsp/cache/parse_test.go b/lsp/cache/parse_test.go index 01fa109..a63be6c 100644 --- a/lsp/cache/parse_test.go +++ b/lsp/cache/parse_test.go @@ -4,6 +4,7 @@ import ( "testing" "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" ) func TestParse(t *testing.T) { @@ -68,3 +69,49 @@ struct Xtruct3 }) } } + +// TestParsedFileDefinitions pins the definition and enum-value indexes: +// every top-level definition is reachable by name, enum values by name. +func TestParsedFileDefinitions(t *testing.T) { + ss := BuildSnapshotForTest([]*FileChange{ + {URI: "file:///tmp/test.thrift", Version: 0, Content: []byte(`struct S { + 1: required string Name, +} + +union U { + 1: string x, +} + +exception X { + 1: string m, +} + +enum Color { + RED, + GREEN, +} + +service Fed { + void go(), +} + +const i32 LIMIT = 1, +typedef string PilotName`), From: FileChangeTypeDidOpen}, + }) + + pf, err := ss.Parse(t.Context(), "file:///tmp/test.thrift") + require.NoError(t, err) + + defs := pf.Definitions() + require.Len(t, defs, 7) + + for _, name := range []string{"S", "U", "X", "Color", "Fed", "LIMIT", "PilotName"} { + _, ok := defs[name] + assert.True(t, ok, "definition %q missing", name) + } + + values := pf.EnumValues() + require.Len(t, values, 2) + assert.NotNil(t, values["RED"]) + assert.NotNil(t, values["GREEN"]) +} diff --git a/lsp/codejump/definition.go b/lsp/codejump/definition.go index 0dc67da..da98613 100644 --- a/lsp/codejump/definition.go +++ b/lsp/codejump/definition.go @@ -44,29 +44,25 @@ func FindTypeDefinition(ctx context.Context, ss *cache.Snapshot, file uri.URI, a _, identifier := lsputils.ParseIdent(file, ast.Includes(), name) for _, astFile := range definitionFiles(ctx, ss, file, ast, name) { - dstAst, err := parseDefinitionFile(ctx, ss, astFile) + dstPf, err := parseDefinitionFile(ctx, ss, astFile) if err != nil { return astFile, nil, DefinitionNone, err } - if dstException := GetExceptionNode(dstAst, identifier); dstException != nil { - return astFile, dstException.Name, DefinitionException, nil - } - - if dstStruct := GetStructNode(dstAst, identifier); dstStruct != nil { - return astFile, dstStruct.Name, DefinitionStruct, nil - } - - if dstEnum := GetEnumNode(dstAst, identifier); dstEnum != nil { - return astFile, dstEnum.Name, DefinitionEnum, nil - } - - if dstUnion := GetUnionNode(dstAst, identifier); dstUnion != nil { - return astFile, dstUnion.Name, DefinitionUnion, nil - } - - if dstTypedef := GetTypedefNode(dstAst, identifier); dstTypedef != nil { - return astFile, dstTypedef.Name, DefinitionTypedef, nil + switch v := dstPf.Definitions()[identifier].(type) { + case *syntax.Struct: + switch v.Kind { + case syntax.UnionDecl: + return astFile, v.Name, DefinitionUnion, nil + case syntax.ExceptionDecl: + return astFile, v.Name, DefinitionException, nil + } + + return astFile, v.Name, DefinitionStruct, nil + case *syntax.Enum: + return astFile, v.Name, DefinitionEnum, nil + case *syntax.Typedef: + return astFile, v.Name, DefinitionTypedef, nil } } @@ -86,18 +82,20 @@ func FindConstValueDefinition(ctx context.Context, ss *cache.Snapshot, file uri. } _, identifier := lsputils.ParseIdent(file, ast.Includes(), name) + identifier = bareName(identifier) + for _, astFile := range definitionFiles(ctx, ss, file, ast, name) { - dstAst, err := parseDefinitionFile(ctx, ss, astFile) + dstPf, err := parseDefinitionFile(ctx, ss, astFile) if err != nil { return astFile, nil, err } - if dstEnumValue := GetEnumValueIdentifierNode(dstAst, identifier); dstEnumValue != nil { - return astFile, dstEnumValue, nil + if id := dstPf.EnumValues()[identifier]; id != nil { + return astFile, id, nil } - if constIdentifier := GetConstIdentifierNode(dstAst, identifier); constIdentifier != nil { - return astFile, constIdentifier, nil + if cst, ok := dstPf.Definitions()[identifier].(*syntax.Const); ok { + return astFile, cst.Name, nil } } @@ -113,13 +111,13 @@ func FindServiceDefinition(ctx context.Context, ss *cache.Snapshot, file uri.URI _, identifier := lsputils.ParseIdent(file, ast.Includes(), ident.Text) for _, astFile := range definitionFiles(ctx, ss, file, ast, ident.Text) { - dstAst, err := parseDefinitionFile(ctx, ss, astFile) + dstPf, err := parseDefinitionFile(ctx, ss, astFile) if err != nil { return astFile, nil, err } - if dstService := GetServiceNode(dstAst, identifier); dstService != nil { - return astFile, dstService.Name, nil + if svc, ok := dstPf.Definitions()[identifier].(*syntax.Service); ok { + return astFile, svc.Name, nil } } @@ -128,8 +126,8 @@ func FindServiceDefinition(ctx context.Context, ss *cache.Snapshot, file uri.URI // parseDefinitionFile parses the definition file, tolerating parse errors // in the target file (the definitions may still be found in the partial -// AST). -func parseDefinitionFile(ctx context.Context, ss *cache.Snapshot, file uri.URI) (*syntax.Document, error) { +// AST). It returns the parsed file so callers can use its indexes. +func parseDefinitionFile(ctx context.Context, ss *cache.Snapshot, file uri.URI) (*cache.ParsedFile, error) { pf, err := ss.Parse(ctx, file) if err != nil { return nil, err @@ -143,7 +141,7 @@ func parseDefinitionFile(ctx context.Context, ss *cache.Snapshot, file uri.URI) return nil, errNoAST } - return pf.AST(), nil + return pf, nil } func typeNameDefinition(ctx context.Context, ss *cache.Snapshot, file uri.URI, pf *cache.ParsedFile, target *target) ([]protocol.Location, error) { diff --git a/lsp/codejump/hover.go b/lsp/codejump/hover.go index 6075f8a..46e37ac 100644 --- a/lsp/codejump/hover.go +++ b/lsp/codejump/hover.go @@ -42,17 +42,17 @@ func hoverService(ctx context.Context, ss *cache.Snapshot, file uri.URI, pf *cac return "", err } - dstAst, err := parseDefinitionFile(ctx, ss, astFile) + dstPf, err := parseDefinitionFile(ctx, ss, astFile) if err != nil { return "", err } - svc := GetServiceNode(dstAst, id.Text) + svc, _ := dstPf.Definitions()[id.Text].(*syntax.Service) if svc == nil { return "", nil } - return formatNode(dstAst, svc) + return formatNode(dstPf.AST(), svc) } func hoverDefinition(ctx context.Context, ss *cache.Snapshot, file uri.URI, pf *cache.ParsedFile, target *target) (string, error) { @@ -63,31 +63,17 @@ func hoverDefinition(ctx context.Context, ss *cache.Snapshot, file uri.URI, pf * return "", err } - dstAst, err := parseDefinitionFile(ctx, ss, astFile) + dstPf, err := parseDefinitionFile(ctx, ss, astFile) if err != nil { return "", err } - var node syntax.Node - - switch kind { - case DefinitionException: - node = GetExceptionNode(dstAst, id.Text) - case DefinitionStruct: - node = GetStructNode(dstAst, id.Text) - case DefinitionEnum: - node = GetEnumNode(dstAst, id.Text) - case DefinitionUnion: - node = GetUnionNode(dstAst, id.Text) - case DefinitionTypedef: - node = GetTypedefNode(dstAst, id.Text) - } - - if node == nil { + node, ok := dstPf.Definitions()[id.Text] + if !ok || !definitionMatches(node, kind) { return "", nil } - return formatNode(dstAst, node) + return formatNode(dstPf.AST(), node) } func hoverConstValue(ctx context.Context, ss *cache.Snapshot, file uri.URI, pf *cache.ParsedFile, target *target) (string, error) { @@ -96,17 +82,17 @@ func hoverConstValue(ctx context.Context, ss *cache.Snapshot, file uri.URI, pf * return "", err } - dstAst, err := parseDefinitionFile(ctx, ss, astFile) + dstPf, err := parseDefinitionFile(ctx, ss, astFile) if err != nil { return "", err } - if dstEnum := GetEnumNodeByEnumValue(dstAst, id.Text); dstEnum != nil { - return formatNode(dstAst, dstEnum) + if dstEnum := enumOfValue(dstPf, id.Text); dstEnum != nil { + return formatNode(dstPf.AST(), dstEnum) } - if dstConst := GetConstNode(dstAst, id.Text); dstConst != nil { - return formatNode(dstAst, dstConst) + if dstConst, ok := dstPf.Definitions()[id.Text].(*syntax.Const); ok { + return formatNode(dstPf.AST(), dstConst) } return "", nil diff --git a/lsp/codejump/utils.go b/lsp/codejump/utils.go index 500a0a6..f2a07b6 100644 --- a/lsp/codejump/utils.go +++ b/lsp/codejump/utils.go @@ -81,244 +81,109 @@ func definitionFiles(ctx context.Context, ss *cache.Snapshot, file uri.URI, ast return files } -// GetExceptionNode finds an exception declaration by name. -func GetExceptionNode(ast *syntax.Document, name string) *syntax.Struct { - if ast == nil { - return nil - } - - for _, excep := range ast.Exceptions() { - if excep.Name != nil && excep.Name.Text == name { - return excep - } - } - - return nil -} - -// GetStructNode finds a struct declaration by name. -func GetStructNode(ast *syntax.Document, name string) *syntax.Struct { - if ast == nil { - return nil - } - - for _, st := range ast.Structs() { - if st.Name != nil && st.Name.Text == name { - return st - } - } - - return nil -} - -// GetUnionNode finds a union declaration by name. -func GetUnionNode(ast *syntax.Document, name string) *syntax.Struct { - if ast == nil { - return nil - } - - for _, st := range ast.Unions() { - if st.Name != nil && st.Name.Text == name { - return st - } - } - - return nil -} - -// GetEnumNode finds an enum declaration by name. -func GetEnumNode(ast *syntax.Document, name string) *syntax.Enum { - if ast == nil { - return nil - } - - for _, st := range ast.Enums() { - if st.Name != nil && st.Name.Text == name { - return st - } - } - - return nil -} - -// GetEnumNodeByEnumValue finds the enum declaring an enum value reference -// like "EnumName.VALUE" or a bare value name. -func GetEnumNodeByEnumValue(ast *syntax.Document, enumValueName string) *syntax.Enum { - if ast == nil { - return nil - } - - enumName, _, found := strings.Cut(enumValueName, ".") - if found { - return GetEnumNode(ast, enumName) - } - - for _, enum := range ast.Enums() { - for _, value := range enum.Values { - if value.Name != nil && value.Name.Text == enumValueName { - return enum - } - } - } - - return nil -} - -// GetEnumValueIdentifierNode returns the identifier of an enum value -// referenced as "EnumName.VALUE", or as a bare name. -func GetEnumValueIdentifierNode(ast *syntax.Document, name string) *syntax.Identifier { - if ast == nil { - return nil - } - - enumName, identifier, found := strings.Cut(name, ".") - if !found { - // Bare name: search all enum values. - for _, enum := range ast.Enums() { - for _, enumValue := range enum.Values { - if enumValue.Name != nil && enumValue.Name.Text == name { - return enumValue.Name - } - } - } - - return nil - } - - for _, enum := range ast.Enums() { - if enum.Name == nil || enum.Name.Text != enumName { - continue - } - - for _, enumValue := range enum.Values { - if enumValue.Name != nil && enumValue.Name.Text == identifier { - return enumValue.Name - } - } - } - - return nil -} - -// GetConstNode finds a const declaration by name. -func GetConstNode(ast *syntax.Document, name string) *syntax.Const { - if ast == nil { - return nil - } - - for _, cst := range ast.Consts() { - if cst.Name != nil && cst.Name.Text == name { - return cst - } - } +// IsBasicType reports whether t is a built-in base type. +func IsBasicType(t string) bool { + _, ok := basicType[t] - return nil + return ok } -// GetConstIdentifierNode returns the name identifier of a const -// declaration. -func GetConstIdentifierNode(ast *syntax.Document, name string) *syntax.Identifier { - if ast == nil { - return nil - } - - for _, cst := range ast.Consts() { - if cst.Name != nil && cst.Name.Text == name { - return cst.Name - } - } +// IsContainerType reports whether t is a container keyword. +func IsContainerType(t string) bool { + _, ok := containerType[t] - return nil + return ok } -// GetTypedefNode finds a typedef declaration by name. -func GetTypedefNode(ast *syntax.Document, name string) *syntax.Typedef { - if ast == nil { - return nil +// typeReferenceName returns the referenced type name of a FieldType, or "" +// for base types and containers. +func typeReferenceName(ft *syntax.FieldType) string { + if ft == nil { + return "" } - for _, td := range ast.Typedefs() { - if td.Name != nil && td.Name.Text == name { - return td - } + if ft.Kind == syntax.TypeIdent && ft.Ident != nil { + return ft.Ident.Text } - return nil + return "" } -// GetServiceNode finds a service declaration by name. -func GetServiceNode(ast *syntax.Document, name string) *syntax.Service { - if ast == nil { - return nil - } - - for _, svc := range ast.Services() { - if svc.Name != nil && svc.Name.Text == name { - return svc - } +// bareName strips the include qualifier from a name: "base.User" becomes +// "User". References in files that include the definition file use the bare +// name, so qualified literals must match against it too. +func bareName(name string) string { + if i := strings.LastIndexByte(name, '.'); i >= 0 { + return name[i+1:] } - return nil + return name } var basicType = map[string]struct{}{ - "map": {}, - "set": {}, - "list": {}, - "string": {}, + "bool": {}, + "byte": {}, + "i8": {}, "i16": {}, "i32": {}, "i64": {}, - "i8": {}, "double": {}, - "bool": {}, - "byte": {}, + "string": {}, "binary": {}, - "uuid": {}, "slist": {}, + "uuid": {}, } var containerType = map[string]struct{}{ + "list": {}, "map": {}, "set": {}, - "list": {}, } -// IsBasicType reports whether t is a built-in base type. -func IsBasicType(t string) bool { - _, ok := basicType[t] - - return ok -} +// definitionMatches reports whether the node has the expected definition +// kind. +func definitionMatches(n syntax.Node, kind DefinitionKind) bool { + switch v := n.(type) { + case *syntax.Struct: + switch v.Kind { + case syntax.UnionDecl: + return kind == DefinitionUnion + case syntax.ExceptionDecl: + return kind == DefinitionException + } -// IsContainerType reports whether t is a container keyword. -func IsContainerType(t string) bool { - _, ok := containerType[t] + return kind == DefinitionStruct + case *syntax.Enum: + return kind == DefinitionEnum + case *syntax.Typedef: + return kind == DefinitionTypedef + } - return ok + return false } -// typeReferenceName returns the referenced type name of a FieldType, or "" -// for base types and containers. -func typeReferenceName(ft *syntax.FieldType) string { - if ft == nil { - return "" - } - - if ft.Kind == syntax.TypeIdent && ft.Ident != nil { - return ft.Ident.Text +// enumOfValue returns the enum declaring the value with the given name, or +// nil. +func enumOfValue(pf *cache.ParsedFile, name string) *syntax.Enum { + id := pf.EnumValues()[name] + if id == nil { + return nil } - return "" -} + // The value identifier's parent enum is reachable through the node + // path; walk the document to find the enum containing the value. + for _, n := range pf.AST().Nodes { + enum, ok := n.(*syntax.Enum) + if !ok { + continue + } -// bareName strips the include qualifier from a name: "base.User" becomes -// "User". References in files that include the definition file use the bare -// name, so qualified literals must match against it too. -func bareName(name string) string { - if i := strings.LastIndexByte(name, '.'); i >= 0 { - return name[i+1:] + for _, value := range enum.Values { + if value.Name == id { + return enum + } + } } - return name + return nil }