diff --git a/lsp/cache/session.go b/lsp/cache/session.go index c74e310..85c21d4 100644 --- a/lsp/cache/session.go +++ b/lsp/cache/session.go @@ -59,6 +59,55 @@ func (s *Session) CreateView(folder uri.URI) { s.views = append(s.views, view) } +// AddView registers a view for the workspace folder, returning the +// existing view when the folder is already tracked. +func (s *Session) AddView(folder uri.URI) *View { + s.viewMu.Lock() + defer s.viewMu.Unlock() + + for _, v := range s.views { + if v.folder == folder { + return v + } + } + + view := NewView(folder.Path(), folder, s.overlayFS, s.cache.IncludePaths) + s.views = append(s.views, view) + + return view +} + +// RemoveView drops the view for the workspace folder and forgets every +// cached file-to-view mapping that pointed at it, so ViewOf re-resolves +// against the remaining folders. +func (s *Session) RemoveView(folder uri.URI) { + s.viewMu.Lock() + defer s.viewMu.Unlock() + + for i, v := range s.views { + if v.folder != folder { + continue + } + + s.views = append(s.views[:i], s.views[i+1:]...) + for file, view := range s.viewMap { + if view == v { + delete(s.viewMap, file) + } + } + + return + } +} + +// Views returns the workspace folders' views. +func (s *Session) Views() []*View { + s.viewMu.Lock() + defer s.viewMu.Unlock() + + return append([]*View(nil), s.views...) +} + func (s *Session) ViewOf(fileURI uri.URI) (*View, error) { s.viewMu.Lock() defer s.viewMu.Unlock() diff --git a/lsp/cache/session_test.go b/lsp/cache/session_test.go new file mode 100644 index 0000000..48ebe4b --- /dev/null +++ b/lsp/cache/session_test.go @@ -0,0 +1,118 @@ +package cache + +import ( + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "go.lsp.dev/uri" +) + +func TestSessionViews(t *testing.T) { + folderA := uri.File("/tmp/a") + folderB := uri.File("/tmp/b") + + tests := []struct { + name string + setup func(s *Session) + folders []uri.URI + }{ + { + name: "no views", + setup: func(s *Session) {}, + folders: nil, + }, + { + name: "one view", + setup: func(s *Session) { + s.AddView(folderA) + }, + folders: []uri.URI{folderA}, + }, + { + name: "views in registration order", + setup: func(s *Session) { + s.AddView(folderB) + s.AddView(folderA) + }, + folders: []uri.URI{folderB, folderA}, + }, + { + name: "removed view disappears", + setup: func(s *Session) { + s.AddView(folderA) + s.AddView(folderB) + s.RemoveView(folderA) + }, + folders: []uri.URI{folderB}, + }, + { + name: "removing an untracked folder is a no-op", + setup: func(s *Session) { + s.AddView(folderA) + s.RemoveView(folderB) + }, + folders: []uri.URI{folderA}, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + s := NewSession(New(nil)) + tt.setup(s) + + views := s.Views() + var got []uri.URI + for _, v := range views { + got = append(got, v.Folder()) + } + + assert.Equal(t, tt.folders, got) + }) + } +} + +func TestSessionAddViewDedups(t *testing.T) { + s := NewSession(New(nil)) + + folder := uri.File("/tmp/a") + first := s.AddView(folder) + second := s.AddView(folder) + + assert.Same(t, first, second) + assert.Len(t, s.Views(), 1) +} + +func TestSessionRemoveViewForgetsMappings(t *testing.T) { + s := NewSession(New(nil)) + + folder := uri.File("/tmp/a") + other := uri.File("/tmp/b") + s.AddView(folder) + s.AddView(other) + + fileA := uri.File("/tmp/a/one.thrift") + fileB := uri.File("/tmp/b/two.thrift") + + // Warm the per-URI view cache. + viewA, err := s.ViewOf(fileA) + require.NoError(t, err) + require.Equal(t, folder, viewA.Folder()) + + viewB, err := s.ViewOf(fileB) + require.NoError(t, err) + require.Equal(t, other, viewB.Folder()) + + // Removing the folder drops the cached mapping, so the file resolves + // to the remaining view. + s.RemoveView(folder) + + view, err := s.ViewOf(fileA) + require.NoError(t, err) + assert.Equal(t, other, view.Folder()) + + // A fresh lookup of the other folder's file still resolves. + view, err = s.ViewOf(fileB) + require.NoError(t, err) + assert.Equal(t, other, view.Folder()) +} diff --git a/lsp/cache/view.go b/lsp/cache/view.go index cc5d3b7..efa3050 100644 --- a/lsp/cache/view.go +++ b/lsp/cache/view.go @@ -83,6 +83,11 @@ func (v *View) ContainsFile(uri uri.URI) bool { return strings.HasPrefix(file, "/") } +// Folder returns the workspace folder the view covers. +func (v *View) Folder() uri.URI { + return v.folder +} + func (v *View) MarkFileKnown(fileURI uri.URI) { v.knownFilesMu.Lock() defer v.knownFilesMu.Unlock() @@ -101,6 +106,22 @@ func (v *View) FileKnown(uri uri.URI) bool { return v.knownFiles[uri] } +// KnownFiles returns the known file URIs of the view, sorted for +// deterministic iteration. +func (v *View) KnownFiles() []uri.URI { + v.knownFilesMu.Lock() + defer v.knownFilesMu.Unlock() + + files := make([]uri.URI, 0, len(v.knownFiles)) + for file := range v.knownFiles { + files = append(files, file) + } + + sort.Slice(files, func(i, j int) bool { return files[i] < files[j] }) + + return files +} + // FileChange applies changes to the view: it swaps in a new snapshot (an // O(1) copy-on-write clone) and re-parses the changed files, then runs the // postFns asynchronously with the affected URIs (changed files plus their diff --git a/lsp/impl_test.go b/lsp/impl_test.go index 15c609c..5de4531 100644 --- a/lsp/impl_test.go +++ b/lsp/impl_test.go @@ -1,9 +1,12 @@ package lsp import ( + "os" + "path/filepath" "testing" "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" "go.lsp.dev/protocol" "go.lsp.dev/uri" @@ -404,3 +407,65 @@ struct Other { }) } } + +func Test_DidChangeWorkspaceFolders(t *testing.T) { + ctx := t.Context() + + dirA := t.TempDir() + dirB := t.TempDir() + + require.NoError(t, os.WriteFile(filepath.Join(dirA, "a.thrift"), []byte("struct FromA {}"), 0o644)) + require.NoError(t, os.WriteFile(filepath.Join(dirB, "b.thrift"), []byte("struct FromB {}"), 0o644)) + + srv := NewServer(cache.New(nil), nil, formatter.Options{}) + + // Adding folders walks them and registers their thrift files. + err := srv.DidChangeWorkspaceFolders(ctx, &protocol.DidChangeWorkspaceFoldersParams{ + Event: protocol.WorkspaceFoldersChangeEvent{ + Added: []protocol.WorkspaceFolder{{URI: uri.File(dirA)}}, + }, + }) + require.NoError(t, err) + + err = srv.DidChangeWorkspaceFolders(ctx, &protocol.DidChangeWorkspaceFoldersParams{ + Event: protocol.WorkspaceFoldersChangeEvent{ + Added: []protocol.WorkspaceFolder{{URI: uri.File(dirB)}}, + }, + }) + require.NoError(t, err) + + assert.Len(t, srv.session.Views(), 2) + + files, err := srv.Symbols(ctx, &protocol.WorkspaceSymbolParams{Query: ""}) + require.NoError(t, err) + + syms, ok := files.(protocol.SymbolInformationSlice) + require.True(t, ok) + assert.Equal(t, []string{"FromA", "FromB"}, symbolNames(syms)) + + // Removing a folder drops its view and its symbols. + err = srv.DidChangeWorkspaceFolders(ctx, &protocol.DidChangeWorkspaceFoldersParams{ + Event: protocol.WorkspaceFoldersChangeEvent{ + Removed: []protocol.WorkspaceFolder{{URI: uri.File(dirA)}}, + }, + }) + require.NoError(t, err) + + assert.Len(t, srv.session.Views(), 1) + + files, err = srv.Symbols(ctx, &protocol.WorkspaceSymbolParams{Query: ""}) + require.NoError(t, err) + + syms, ok = files.(protocol.SymbolInformationSlice) + require.True(t, ok) + assert.Equal(t, []string{"FromB"}, symbolNames(syms)) +} + +func symbolNames(syms protocol.SymbolInformationSlice) []string { + names := make([]string, len(syms)) + for i, s := range syms { + names[i] = s.Name + } + + return names +} diff --git a/lsp/server.go b/lsp/server.go index cf8cbe5..918db18 100644 --- a/lsp/server.go +++ b/lsp/server.go @@ -8,6 +8,7 @@ import ( "github.com/karitham/thrift-ls/formatter" "github.com/karitham/thrift-ls/lsp/cache" + "github.com/karitham/thrift-ls/lsp/symbols" ) type Server struct { @@ -114,6 +115,15 @@ func (s *Server) DidChangeWatchedFiles(ctx context.Context, params *protocol.Did } func (s *Server) DidChangeWorkspaceFolders(ctx context.Context, params *protocol.DidChangeWorkspaceFoldersParams) (err error) { + for _, folder := range params.Event.Removed { + s.session.RemoveView(folder.URI) + } + + for _, folder := range params.Event.Added { + s.session.AddView(folder.URI) + s.walkFoldersThriftFile(folder.URI) + } + return nil } @@ -215,7 +225,7 @@ func (s *Server) SignatureHelp(ctx context.Context, params *protocol.SignatureHe } func (s *Server) Symbols(ctx context.Context, params *protocol.WorkspaceSymbolParams) (result protocol.WorkspaceSymbolResult, err error) { - return protocol.SymbolInformationSlice{}, nil + return protocol.SymbolInformationSlice(symbols.WorkspaceSymbols(ctx, s.session, params.Query, 1000)), nil } func (s *Server) TypeDefinition(ctx context.Context, params *protocol.TypeDefinitionParams) (result protocol.DefinitionResult, err error) { diff --git a/lsp/symbols/workspace.go b/lsp/symbols/workspace.go new file mode 100644 index 0000000..d95ad17 --- /dev/null +++ b/lsp/symbols/workspace.go @@ -0,0 +1,74 @@ +package symbols + +import ( + "context" + "sort" + "strings" + + "go.lsp.dev/protocol" + "go.lsp.dev/uri" + + "github.com/karitham/thrift-ls/lsp/cache" +) + +// WorkspaceSymbols returns the workspace symbols matching query: every +// top-level definition and its members across all workspace folders, +// folders and files ordered by URI, symbols in source order. An empty +// query matches everything; the result is capped at maxResults (0 means +// unlimited). Matching is case-insensitive substring on the symbol name. +func WorkspaceSymbols(ctx context.Context, session *cache.Session, query string, maxResults int) []protocol.SymbolInformation { + res := make([]protocol.SymbolInformation, 0, 64) + q := strings.ToLower(query) + + views := session.Views() + sort.Slice(views, func(i, j int) bool { return views[i].Folder() < views[j].Folder() }) + + for _, view := range views { + snapshot, release := view.Snapshot() + + for _, file := range view.KnownFiles() { + for _, sym := range documentSymbolsFlat(ctx, snapshot, file) { + if q != "" && !strings.Contains(strings.ToLower(sym.Name), q) { + continue + } + + res = append(res, sym) + if maxResults > 0 && len(res) >= maxResults { + release() + + return res + } + } + } + + release() + } + + return res +} + +// documentSymbolsFlat returns the document symbols of a file flattened +// into workspace symbols; the location points at the symbol's name. +func documentSymbolsFlat(ctx context.Context, ss *cache.Snapshot, file uri.URI) []protocol.SymbolInformation { + syms := make([]protocol.SymbolInformation, 0, 16) + + for _, sym := range DocumentSymbols(ctx, ss, file) { + flattenSymbol(sym, file, &syms) + } + + return syms +} + +func flattenSymbol(sym *protocol.DocumentSymbol, file uri.URI, out *[]protocol.SymbolInformation) { + *out = append(*out, protocol.SymbolInformation{ + BaseSymbolInformation: protocol.BaseSymbolInformation{ + Name: sym.Name, + Kind: sym.Kind, + }, + Location: protocol.Location{URI: file, Range: sym.SelectionRange}, + }) + + for i := range sym.Children { + flattenSymbol(&sym.Children[i], file, out) + } +} diff --git a/lsp/symbols/workspace_test.go b/lsp/symbols/workspace_test.go new file mode 100644 index 0000000..848c78a --- /dev/null +++ b/lsp/symbols/workspace_test.go @@ -0,0 +1,258 @@ +package symbols + +import ( + "os" + "path/filepath" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "go.lsp.dev/protocol" + "go.lsp.dev/uri" + + "github.com/karitham/thrift-ls/lsp/cache" +) + +// writeTree writes the file contents under dir and returns the directory. +func writeTree(t *testing.T, files map[string]string) string { + t.Helper() + + dir := t.TempDir() + for name, content := range files { + path := filepath.Join(dir, name) + require.NoError(t, os.MkdirAll(filepath.Dir(path), 0o755)) + require.NoError(t, os.WriteFile(path, []byte(content), 0o644)) + } + + return dir +} + +// openTree registers every file of dir with its view, like the server's +// initialization walk does. When only is non-nil, only those file names +// are registered. +func openTree(t *testing.T, session *cache.Session, dir string, only map[string]bool) { + t.Helper() + + folder := uri.File(dir) + view := session.AddView(folder) + + require.NoError(t, filepath.WalkDir(dir, func(path string, d os.DirEntry, err error) error { + if err != nil || d.IsDir() || filepath.Ext(path) != ".thrift" { + return nil + } + + if only != nil && !only[filepath.Base(path)] { + return nil + } + + content, err := os.ReadFile(path) + require.NoError(t, err) + + view.FileChange(t.Context(), []*cache.FileChange{{ + URI: uri.File(path), + Version: 0, + Content: content, + From: cache.FileChangeTypeInitialize, + }}) + + return nil + })) +} + +func TestWorkspaceSymbols(t *testing.T) { + twoFiles := map[string]string{ + "mobile_suit.thrift": `struct MobileSuit { + 1: required string Name, + 2: optional i32 ModelNumber, +} + +enum ZeonForces { + ZAKU_I = 1, + ZAKU_II, +}`, + "federation.thrift": `service Federation { + void Deploy(1: string suitName), + string Query(), +} + +const i32 DEFAULT_HP = 100, +typedef string PilotName`, + } + + queryTree := map[string]string{ + "a.thrift": `struct MobileSuit { 1: string Name } +struct MobileArmor { 1: string Name } +const i32 ZAKU_HP = 100`, + } + + capTree := map[string]string{ + "a.thrift": `struct A { 1: string x } +struct B { 1: string x } +struct C { 1: string x }`, + } + + tests := []struct { + name string + files map[string]string + query string + max int + only map[string]bool // files to register; nil registers all + nested bool // files live in per-folder subdirectories + want []string + }{ + { + name: "all symbols across files, members flattened", + files: twoFiles, + query: "", + want: []string{ + "Federation", "Deploy", "Query", + "DEFAULT_HP", "PilotName", + "MobileSuit", "Name", "ModelNumber", + "ZeonForces", "ZAKU_I", "ZAKU_II", + }, + }, + { + name: "empty query matches everything", + files: queryTree, + query: "", + want: []string{"MobileSuit", "Name", "MobileArmor", "Name", "ZAKU_HP"}, + }, + { + name: "query filters case-insensitively", + files: queryTree, + query: "MOBILE", + want: []string{"MobileSuit", "MobileArmor"}, + }, + { + name: "query matches members", + files: queryTree, + query: "zaku", + want: []string{"ZAKU_HP"}, + }, + { + name: "query matching nothing returns empty", + files: queryTree, + query: "nothing-matches", + want: []string{}, + }, + { + name: "cap limits the result", + files: capTree, + query: "", + max: 4, + want: []string{"A", "x", "B", "x"}, + }, + { + name: "cap counts only matching symbols", + files: queryTree, + query: "mobile", + max: 1, + want: []string{"MobileSuit"}, + }, + { + name: "multiple workspace folders", + files: map[string]string{"a/a.thrift": "struct FromA {}", "b/b.thrift": "struct FromB {}"}, + query: "", + nested: true, + want: []string{"FromA", "FromB"}, + }, + { + name: "unregistered files are excluded", + files: map[string]string{"known.thrift": "struct Known {}", "unknown.thrift": "struct Unknown {}"}, + query: "", + only: map[string]bool{"known.thrift": true}, + want: []string{"Known"}, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + dir := writeTree(t, tt.files) + + session := cache.NewSession(cache.New(nil)) + if tt.nested { + // Each top-level directory is a workspace folder. + dirs := map[string]bool{} + for name := range tt.files { + dirs[filepath.Dir(name)] = true + } + + for d := range dirs { + openTree(t, session, filepath.Join(dir, d), nil) + } + } else { + openTree(t, session, dir, tt.only) + } + + syms := WorkspaceSymbols(t.Context(), session, tt.query, tt.max) + + names := make([]string, len(syms)) + for i, s := range syms { + names[i] = s.Name + } + + assert.Equal(t, tt.want, names) + }) + } +} + +// TestWorkspaceSymbolsKindAndLocation pins the kind and the name range of +// each symbol kind. +func TestWorkspaceSymbolsKindAndLocation(t *testing.T) { + dir := writeTree(t, map[string]string{ + "shapes.thrift": `struct MobileSuit { + 1: required string Name, +} + +enum ZeonForces { + ZAKU_I = 1, + ZAKU_II, +} + +service Federation { + void Deploy(1: string suitName), +} + +const i32 DEFAULT_HP = 100, +typedef string PilotName`, + }) + + session := cache.NewSession(cache.New(nil)) + openTree(t, session, dir, nil) + + file := uri.File(filepath.Join(dir, "shapes.thrift")) + + tests := []struct { + name string + kind protocol.SymbolKind + line uint32 + col uint32 + }{ + {"MobileSuit", protocol.SymbolKindStruct, 0, 7}, + {"Name", protocol.SymbolKindField, 1, 20}, + {"ZeonForces", protocol.SymbolKindEnum, 4, 5}, + {"ZAKU_II", protocol.SymbolKindEnumMember, 6, 1}, + {"Federation", protocol.SymbolKindInterface, 9, 8}, + {"Deploy", protocol.SymbolKindFunction, 10, 6}, + {"DEFAULT_HP", protocol.SymbolKindConstant, 13, 10}, + {"PilotName", protocol.SymbolKindTypeParameter, 14, 15}, + } + + syms := WorkspaceSymbols(t.Context(), session, "", 0) + byName := make(map[string]protocol.SymbolInformation, len(syms)) + for _, s := range syms { + byName[s.Name] = s + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + sym, ok := byName[tt.name] + require.True(t, ok, "symbol %q missing", tt.name) + + assert.Equal(t, tt.kind, sym.Kind) + assert.Equal(t, file, sym.Location.URI) + assert.Equal(t, tt.line, sym.Location.Range.Start.Line) + assert.Equal(t, tt.col, sym.Location.Range.Start.Character) + }) + } +}