diff --git a/internal/lsp/code_lens.go b/internal/lsp/code_lens.go index 392dde3..aa1ea40 100644 --- a/internal/lsp/code_lens.go +++ b/internal/lsp/code_lens.go @@ -3,6 +3,7 @@ package lsp import ( "context" "fmt" + "path/filepath" "go.lsp.dev/protocol" @@ -22,50 +23,102 @@ func (s *Server) CodeLens(ctx context.Context, params *protocol.CodeLensParams) if !ok { return []protocol.CodeLens{}, nil } - - semanticDoc := newSemanticDocument(doc.Path, doc.Content, s.parsePathForNavigation(doc.Path)) + // Check cancellation before the (potentially expensive) initial parse so a + // cancelled request never surfaces as a misleading empty success. + if err := ctx.Err(); err != nil { + return nil, err + } + parse := s.newRequestParse(ctx) + rootResult := parse(doc.Path) + if err := ctx.Err(); err != nil { + return nil, err + } + semanticDoc := newSemanticDocument(doc.Path, doc.Content, rootResult) if !semanticDoc.valid() { return []protocol.CodeLens{}, nil } - return semanticDoc.codeLenses(), nil -} - -func (s *Server) CodeLensResolve(ctx context.Context, params *protocol.CodeLens) (*protocol.CodeLens, error) { - if params == nil || params.Command != nil { - return params, nil + // Resolve every lens up front in one workspace pass: the client then never + // calls CodeLensResolve, so there is no per-version command memo to go stale + // when a cross-file reference changes. + candidatePaths := s.referenceCandidatePaths() + resolveType := func(path string, result *ridl.ParseResult, name string) *definitionMatch { + return resolveTypeDefinitionWith(parse, path, result, name) } - - data, ok := decodeCodeLensData(params.Data) - if !ok { - return params, nil + resolveError := func(path string, result *ridl.ParseResult, name string) *definitionMatch { + return resolveErrorDefinitionWith(parse, path, result, name) } - doc, ok := s.docs.Get(data.URI) - if !ok { - return params, nil + base := semanticDoc.codeLenses() + lenses := make([]protocol.CodeLens, 0, len(base)) + for i := range base { + data, ok := decodeCodeLensData(base[i].Data) + if !ok { + continue + } + target := semanticDoc.referenceTargetAt( + protocol.Position{Line: data.Line, Character: data.Character}, + resolveType, resolveError, + ) + if target == nil || target.definition == nil { + continue + } + locations := collectReferenceLocationsWith(parse, candidatePaths, s.contentForPath, target, false, resolveType, resolveError) + // Check after the (potentially expensive) cross-file collection so a + // cancelled request — including on the last/only symbol — never returns + // a partial under-counted result to the client. + if err := ctx.Err(); err != nil { + return nil, err + } + lens := base[i] + lens.Command = &protocol.Command{ + Title: referenceCountTitle(len(locations)), + Command: showReferencesCommand, + Arguments: []any{protocol.DocumentURI(data.URI), base[i].Range.Start, locations}, + } + lens.Data = nil + lenses = append(lenses, lens) } - semanticDoc := newSemanticDocument(doc.Path, doc.Content, s.parsePathForNavigation(doc.Path)) - if !semanticDoc.valid() { - return params, nil - } + return lenses, nil +} - target := semanticDoc.referenceTargetAt(protocol.Position{ - Line: data.Line, - Character: data.Character, - }, s.resolveTypeDefinition, s.resolveErrorDefinition) - if target == nil || target.definition == nil { - return params, nil - } +// CodeLensResolve is a passthrough: CodeLens returns fully-resolved lenses, so +// resolution never needs to recompute. +func (s *Server) CodeLensResolve(ctx context.Context, params *protocol.CodeLens) (*protocol.CodeLens, error) { + return params, nil +} - locations := s.collectReferenceLocations(target, false) - params.Command = &protocol.Command{ - Title: referenceCountTitle(len(locations)), - Command: showReferencesCommand, - Arguments: []any{protocol.DocumentURI(data.URI), params.Range.Start, locations}, +// newRequestParse builds a request-scoped parser that AST-parses each file at +// most once: an open buffer's already-built result is reused; everything else is +// parsed AST-only (no schema build, no import recursion). Keyed by cleaned path +// so each file is parsed once per request. +// +// The overlay map is snapshotted once here; s.docs.FindByPath and s.contentForPath +// read live document state on each miss. Handlers run concurrently, so a request +// is eventually-consistent-by-design: a mid-request DidChange self-heals on the +// client's next codeLens request. +func (s *Server) newRequestParse(ctx context.Context) parseFn { + overlays := s.overlayContents() + memo := map[string]*ridl.ParseResult{} + return func(path string) *ridl.ParseResult { + // Honour request cancellation before each (possibly expensive) AST parse. + if ctx.Err() != nil { + return nil + } + key := filepath.Clean(path) + if result, ok := memo[key]; ok { + return result + } + var result *ridl.ParseResult + if doc, ok := s.docs.FindByPath(path); ok && doc.Result != nil && doc.Result.Root != nil { + result = doc.Result + } else { + result, _ = s.parser.ParseAST(ctx, s.workspace.Root(), path, overlays) + } + memo[key] = result + return result } - return params, nil } func (d *semanticDocument) codeLenses() []protocol.CodeLens { diff --git a/internal/lsp/code_lens_perf_test.go b/internal/lsp/code_lens_perf_test.go new file mode 100644 index 0000000..4f266be --- /dev/null +++ b/internal/lsp/code_lens_perf_test.go @@ -0,0 +1,149 @@ +package lsp + +import ( + "context" + "os" + "path/filepath" + "strconv" + "testing" + + "go.lsp.dev/protocol" + "go.uber.org/zap" +) + +const baseRIDL = `webrpc = v1 + +name = test +version = v0.0.1 + +struct User + - id: uint64 +` + +// TestNewRequestParseMemoizesPerPath asserts that calling the request-scoped +// parse function twice with the same path returns the identical *ParseResult +// pointer, proving the file is parsed at most once per request. +func TestNewRequestParseMemoizesPerPath(t *testing.T) { + srv, _, dir := setupServer(t) + + basePath := filepath.Join(dir, "base.ridl") + if err := os.WriteFile(basePath, []byte(baseRIDL), 0o644); err != nil { + t.Fatal(err) + } + + parse := srv.newRequestParse(context.Background()) + + r1 := parse(basePath) + if r1 == nil { + t.Fatal("expected non-nil ParseResult, got nil — test would be vacuous") + } + if r1.Root == nil { + t.Fatal("expected ParseResult.Root != nil, got nil — test would be vacuous") + } + + r2 := parse(basePath) + if r1 != r2 { + t.Errorf("parse returned different pointers for the same path — file was parsed more than once") + } +} + +// TestNewRequestParseKeyedByCleanPath asserts that path variants that resolve +// to the same cleaned path (e.g. "dir/./file.ridl" vs "dir/file.ridl") share +// the same memo slot and produce the identical *ParseResult pointer. +func TestNewRequestParseKeyedByCleanPath(t *testing.T) { + srv, _, dir := setupServer(t) + + basePath := filepath.Join(dir, "base.ridl") + if err := os.WriteFile(basePath, []byte(baseRIDL), 0o644); err != nil { + t.Fatal(err) + } + + parse := srv.newRequestParse(context.Background()) + + // Two string paths that differ only by a redundant "." segment. + p1 := dir + "/base.ridl" + p2 := dir + "/./base.ridl" // deliberately un-cleaned via string concat + + r1 := parse(p1) + if r1 == nil { + t.Fatal("expected non-nil ParseResult for p1, got nil — test would be vacuous") + } + + r2 := parse(p2) + if r2 == nil { + t.Fatal("expected non-nil ParseResult for p2, got nil") + } + if r1 != r2 { + t.Errorf("different un-cleaned paths that resolve to the same file returned different ParseResult pointers — clean-path keying is broken") + } +} + +// BenchmarkCodeLensEager measures the cost of a single CodeLens request +// against a workspace of 40 importer files all referencing two symbols in +// a shared base file. Run with -bench BenchmarkCodeLensEager -run x. +func BenchmarkCodeLensEager(b *testing.B) { + // Replicate setupServer body for *testing.B (setupServer only accepts *testing.T). + dir := b.TempDir() + client := newMockClient() + srv := NewServer(zap.NewNop()) + srv.SetClient(client) + srv.workspace.SetRoot(dir) + + baseContent := `webrpc = v1 + +name = base +version = v0.0.1 + +struct User + - id: uint64 + +struct Account + - userId: uint64 +` + basePath := filepath.Join(dir, "base.ridl") + if err := os.WriteFile(basePath, []byte(baseContent), 0o644); err != nil { + b.Fatal(err) + } + + // Write 40 importer files each referencing User and Account from base. + for i := 0; i < 40; i++ { + content := `webrpc = v1 + +name = f` + strconv.Itoa(i) + ` +version = v0.0.1 + +import + - base.ridl + +struct S` + strconv.Itoa(i) + ` + - u: User + - a: Account +` + p := filepath.Join(dir, "f"+strconv.Itoa(i)+".ridl") + if err := os.WriteFile(p, []byte(content), 0o644); err != nil { + b.Fatal(err) + } + } + + // DidOpen base.ridl so it is the target document. + baseURI := string(PathToURI(basePath)) + ctx := context.Background() + if err := srv.DidOpen(ctx, &protocol.DidOpenTextDocumentParams{ + TextDocument: protocol.TextDocumentItem{ + URI: protocol.DocumentURI(baseURI), + Text: baseContent, + Version: 1, + }, + }); err != nil { + b.Fatal(err) + } + + b.ResetTimer() + for i := 0; i < b.N; i++ { + if _, err := srv.CodeLens(ctx, &protocol.CodeLensParams{ + TextDocument: protocol.TextDocumentIdentifier{URI: protocol.DocumentURI(baseURI)}, + }); err != nil { + b.Fatal(err) + } + } +} diff --git a/internal/lsp/code_lens_test.go b/internal/lsp/code_lens_test.go index cf95f68..6f914d0 100644 --- a/internal/lsp/code_lens_test.go +++ b/internal/lsp/code_lens_test.go @@ -55,8 +55,14 @@ service TestService if userLens == nil { t.Fatalf("missing code lens for User declaration in %#v", lenses) } - if userLens.Command != nil { - t.Fatalf("expected unresolved code lens command, got %#v", userLens.Command) + if userLens.Command == nil { + t.Fatalf("expected eager resolved code lens command, got nil") + } + if userLens.Command.Command != showReferencesCommand { + t.Fatalf("unexpected code lens command %q", userLens.Command.Command) + } + if userLens.Data != nil { + t.Fatalf("expected Data cleared on resolved lens, got %#v", userLens.Data) } notFoundLens := findCodeLensAtPosition(lenses, positionAt(t, content, "UserNotFound")) @@ -65,7 +71,7 @@ service TestService } } -func TestCodeLensResolveBuildsShowReferencesCommand(t *testing.T) { +func TestCodeLensBuildsShowReferencesCommand(t *testing.T) { srv, _, dir := setupServer(t) ctx := context.Background() @@ -106,30 +112,208 @@ service TestService if lens == nil { t.Fatalf("missing code lens for User declaration in %#v", lenses) } + if lens.Command == nil { + t.Fatalf("expected resolved code lens command, got nil") + } + if lens.Command.Command != showReferencesCommand { + t.Fatalf("unexpected code lens command %q", lens.Command.Command) + } + if lens.Command.Title != "2 references" { + t.Fatalf("unexpected code lens title %q", lens.Command.Title) + } + if len(lens.Command.Arguments) != 3 { + t.Fatalf("expected show-references arguments, got %#v", lens.Command.Arguments) + } + + locations, ok := lens.Command.Arguments[2].([]protocol.Location) + if !ok { + t.Fatalf("expected reference locations in code lens args, got %#v", lens.Command.Arguments[2]) + } + if len(locations) != 2 { + t.Fatalf("expected 2 reference locations, got %#v", locations) + } +} - resolved, err := srv.CodeLensResolve(ctx, lens) +func TestCodeLensResolveIsPassthrough(t *testing.T) { + srv, _, _ := setupServer(t) + ctx := context.Background() + + in := &protocol.CodeLens{ + Range: protocol.Range{Start: protocol.Position{Line: 4}}, + Command: &protocol.Command{Title: "1 reference", Command: showReferencesCommand}, + } + out, err := srv.CodeLensResolve(ctx, in) if err != nil { + t.Fatalf("CodeLensResolve error: %v", err) + } + if out != in { + t.Fatal("expected CodeLensResolve to return the same lens unchanged") + } +} + +// TestCodeLensImportedTypeRefCount asserts that the CodeLens on a type +// declaration counts references from other files that import and use it. +// base.ridl defines struct Base; user.ridl imports it and uses Base in a +// field. The lens on Base must show "1 reference" (the field in user.ridl). +// includeDeclaration is false in CodeLens, so the declaration itself does not +// add to the count. +func TestCodeLensImportedTypeRefCount(t *testing.T) { + srv, _, dir := setupServer(t) + ctx := context.Background() + + baseContent := `webrpc = v1 + +name = base +version = v0.0.1 + +struct Base + - id: uint64 +` + basePath := filepath.Join(dir, "base.ridl") + if err := os.WriteFile(basePath, []byte(baseContent), 0o644); err != nil { t.Fatal(err) } - if resolved == nil || resolved.Command == nil { - t.Fatalf("expected resolved code lens command, got %#v", resolved) + + userContent := `webrpc = v1 + +name = user +version = v0.0.1 + +import + - base.ridl + +struct User + - base: Base +` + userPath := filepath.Join(dir, "user.ridl") + if err := os.WriteFile(userPath, []byte(userContent), 0o644); err != nil { + t.Fatal(err) } - if resolved.Command.Command != showReferencesCommand { - t.Fatalf("unexpected code lens command %q", resolved.Command.Command) + + // Open base.ridl — CodeLens is requested on it; user.ridl is picked up via + // the workspace walk inside referenceCandidatePaths. + baseURI := fileURI(basePath) + _ = srv.DidOpen(ctx, &protocol.DidOpenTextDocumentParams{ + TextDocument: protocol.TextDocumentItem{ + URI: protocol.DocumentURI(baseURI), + Text: baseContent, + Version: 1, + }, + }) + + lenses, err := srv.CodeLens(ctx, &protocol.CodeLensParams{ + TextDocument: protocol.TextDocumentIdentifier{URI: protocol.DocumentURI(baseURI)}, + }) + if err != nil { + t.Fatal(err) } - if resolved.Command.Title != "2 references" { - t.Fatalf("unexpected code lens title %q", resolved.Command.Title) + + baseLens := findCodeLensAtPosition(lenses, positionAt(t, baseContent, "Base\n")) + if baseLens == nil { + t.Fatalf("missing code lens for Base declaration in %#v", lenses) + } + if baseLens.Command == nil { + t.Fatal("expected resolved code lens command, got nil") } - if len(resolved.Command.Arguments) != 3 { - t.Fatalf("expected show-references arguments, got %#v", resolved.Command.Arguments) + if baseLens.Command.Title != "1 reference" { + t.Fatalf("expected \"1 reference\" for Base (used once in user.ridl field), got %q", baseLens.Command.Title) } +} - locations, ok := resolved.Command.Arguments[2].([]protocol.Location) - if !ok { - t.Fatalf("expected reference locations in code lens args, got %#v", resolved.Command.Arguments[2]) +// TestCodeLensErrorRefCount asserts that the CodeLens on an error declaration +// counts the single method that lists it in its errors clause. +func TestCodeLensErrorRefCount(t *testing.T) { + srv, _, dir := setupServer(t) + ctx := context.Background() + + content := `webrpc = v1 + +name = errtest +version = v0.0.1 + +error 1 Foo "msg" HTTP 404 + +service FooService + - M() => (ok: bool) errors Foo +` + path := filepath.Join(dir, "errref.ridl") + if err := os.WriteFile(path, []byte(content), 0o644); err != nil { + t.Fatal(err) } - if len(locations) != 2 { - t.Fatalf("expected 2 reference locations, got %#v", locations) + + uri := fileURI(path) + _ = srv.DidOpen(ctx, &protocol.DidOpenTextDocumentParams{ + TextDocument: protocol.TextDocumentItem{ + URI: protocol.DocumentURI(uri), + Text: content, + Version: 1, + }, + }) + + lenses, err := srv.CodeLens(ctx, &protocol.CodeLensParams{ + TextDocument: protocol.TextDocumentIdentifier{URI: protocol.DocumentURI(uri)}, + }) + if err != nil { + t.Fatal(err) + } + + fooLens := findCodeLensAtPosition(lenses, positionAt(t, content, "Foo ")) + if fooLens == nil { + t.Fatalf("missing code lens for Foo error declaration in %#v", lenses) + } + if fooLens.Command == nil { + t.Fatal("expected resolved code lens command, got nil") + } + if fooLens.Command.Title != "1 reference" { + t.Fatalf("expected \"1 reference\" for Foo error (used in M() errors clause), got %q", fooLens.Command.Title) + } +} + +// TestCodeLensReturnsErrorOnCanceledContext verifies that CodeLens propagates +// ctx.Err() rather than returning a (misleading) nil-error empty/partial result +// when the request context is already cancelled. +// +// DidOpen uses a live context so s.docs.Get succeeds and the test exercises the +// ctx check inside CodeLens itself — if DidOpen also used the cancelled ctx the +// doc would never be registered and the early-exit `!ok` branch would mask the +// real assertion. +func TestCodeLensReturnsErrorOnCanceledContext(t *testing.T) { + srv, _, dir := setupServer(t) + + content := `webrpc = v1 + +name = canceltest +version = v0.0.1 + +struct Foo + - id: uint64 +` + path := filepath.Join(dir, "cancel.ridl") + if err := os.WriteFile(path, []byte(content), 0o644); err != nil { + t.Fatal(err) + } + + uri := fileURI(path) + // Register the document with a live context so s.docs.Get succeeds. + _ = srv.DidOpen(context.Background(), &protocol.DidOpenTextDocumentParams{ + TextDocument: protocol.TextDocumentItem{ + URI: protocol.DocumentURI(uri), + Text: content, + Version: 1, + }, + }) + + ctx, cancel := context.WithCancel(context.Background()) + cancel() // cancel before the CodeLens call + + lenses, err := srv.CodeLens(ctx, &protocol.CodeLensParams{ + TextDocument: protocol.TextDocumentIdentifier{URI: protocol.DocumentURI(uri)}, + }) + if err == nil { + t.Fatalf("expected non-nil error for cancelled context, got nil (lenses: %#v)", lenses) + } + if lenses != nil { + t.Fatalf("expected nil lenses for cancelled context, got %#v", lenses) } } diff --git a/internal/lsp/definition.go b/internal/lsp/definition.go index fe16a6f..74c2c5c 100644 --- a/internal/lsp/definition.go +++ b/internal/lsp/definition.go @@ -68,23 +68,35 @@ func (s *Server) parsePath(ctx context.Context, path string) *ridl.ParseResult { return result } +type parseFn func(path string) *ridl.ParseResult + +type defResolver func(path string, result *ridl.ParseResult, name string) *definitionMatch + func (s *Server) resolveTypeDefinition(path string, result *ridl.ParseResult, name string) *definitionMatch { + return resolveTypeDefinitionWith(s.parsePathForNavigation, path, result, name) +} + +func (s *Server) resolveErrorDefinition(path string, result *ridl.ParseResult, name string) *definitionMatch { + return resolveErrorDefinitionWith(s.parsePathForNavigation, path, result, name) +} + +func resolveTypeDefinitionWith(parse parseFn, path string, result *ridl.ParseResult, name string) *definitionMatch { if name == "" || isBuiltInRIDLType(name) { return nil } - return s.resolveNamedDefinition(path, result, name, findTypeDefinitionToken) + return resolveNamedDefinitionWith(parse, path, result, name, findTypeDefinitionToken) } -func (s *Server) resolveErrorDefinition(path string, result *ridl.ParseResult, name string) *definitionMatch { +func resolveErrorDefinitionWith(parse parseFn, path string, result *ridl.ParseResult, name string) *definitionMatch { if name == "" { return nil } - return s.resolveNamedDefinition(path, result, name, findErrorDefinitionToken) + return resolveNamedDefinitionWith(parse, path, result, name, findErrorDefinitionToken) } type definitionFinder func(root *ridl.RootNode, name string) *ridl.TokenNode -func (s *Server) resolveNamedDefinition(path string, result *ridl.ParseResult, name string, finder definitionFinder) *definitionMatch { +func resolveNamedDefinitionWith(parse parseFn, path string, result *ridl.ParseResult, name string, finder definitionFinder) *definitionMatch { if result == nil || result.Root == nil { return nil } @@ -94,10 +106,10 @@ func (s *Server) resolveNamedDefinition(path string, result *ridl.ParseResult, n } visited := map[string]struct{}{path: {}} - return s.resolveImportedDefinition(path, result.Root, name, visited, finder) + return resolveImportedDefinitionWith(parse, path, result.Root, name, visited, finder) } -func (s *Server) resolveImportedDefinition(path string, root *ridl.RootNode, name string, visited map[string]struct{}, finder definitionFinder) *definitionMatch { +func resolveImportedDefinitionWith(parse parseFn, path string, root *ridl.RootNode, name string, visited map[string]struct{}, finder definitionFinder) *definitionMatch { if root == nil { return nil } @@ -113,7 +125,7 @@ func (s *Server) resolveImportedDefinition(path string, root *ridl.RootNode, nam } visited[importPath] = struct{}{} - importResult := s.parsePathForNavigation(importPath) + importResult := parse(importPath) if importResult == nil || importResult.Root == nil { continue } @@ -122,7 +134,7 @@ func (s *Server) resolveImportedDefinition(path string, root *ridl.RootNode, nam return definitionForToken(importPath, token) } - if match := s.resolveImportedDefinition(importPath, importResult.Root, name, visited, finder); match != nil { + if match := resolveImportedDefinitionWith(parse, importPath, importResult.Root, name, visited, finder); match != nil { return match } } diff --git a/internal/lsp/references.go b/internal/lsp/references.go index 976ff26..759e96d 100644 --- a/internal/lsp/references.go +++ b/internal/lsp/references.go @@ -374,6 +374,25 @@ func skipWorkspaceDir(name string) bool { } func (s *Server) collectReferenceLocations(target *referenceTarget, includeDeclaration bool) []protocol.Location { + return collectReferenceLocationsWith( + s.parsePathForNavigation, + s.referenceCandidatePaths(), + s.contentForPath, + target, + includeDeclaration, + s.resolveTypeDefinition, + s.resolveErrorDefinition, + ) +} + +func collectReferenceLocationsWith( + parse parseFn, + candidatePaths []string, + contentFor func(path string) (string, bool), + target *referenceTarget, + includeDeclaration bool, + resolveType, resolveError defResolver, +) []protocol.Location { if target == nil || target.definition == nil { return nil } @@ -381,19 +400,19 @@ func (s *Server) collectReferenceLocations(target *referenceTarget, includeDecla locations := make([]protocol.Location, 0, 8) seen := map[string]struct{}{} - for _, path := range s.referenceCandidatePaths() { - result := s.parsePathForNavigation(path) + for _, path := range candidatePaths { + result := parse(path) if result == nil || result.Root == nil { continue } - content, ok := s.contentForPath(path) + content, ok := contentFor(path) if !ok { continue } doc := newSemanticDocument(path, content, result) - for _, location := range doc.referenceLocations(target, s.resolveTypeDefinition, s.resolveErrorDefinition, includeDeclaration) { + for _, location := range doc.referenceLocations(target, resolveType, resolveError, includeDeclaration) { key := referenceLocationKey(location) if _, ok := seen[key]; ok { continue diff --git a/internal/lsp/server.go b/internal/lsp/server.go index 3f662bd..49afd66 100644 --- a/internal/lsp/server.go +++ b/internal/lsp/server.go @@ -61,8 +61,10 @@ func (s *Server) Initialize(ctx context.Context, params *protocol.InitializePara DocumentOnTypeFormattingProvider: &protocol.DocumentOnTypeFormattingOptions{ FirstTriggerCharacter: onTypeFormattingTrigger, }, + // ResolveProvider is false: CodeLens returns fully-resolved lenses, so + // clients must not call codeLens/resolve. CodeLensProvider: &protocol.CodeLensOptions{ - ResolveProvider: true, + ResolveProvider: false, }, ColorProvider: true, DocumentFormattingProvider: true, diff --git a/internal/ridl/parser.go b/internal/ridl/parser.go index e3843ed..1cd8dd8 100644 --- a/internal/ridl/parser.go +++ b/internal/ridl/parser.go @@ -226,6 +226,51 @@ func (p *Parser) parse(ctx context.Context, workspace, path string, overlays map return result, nil } +// ParseAST parses only the RIDL AST at path (overlay-aware), skipping schema +// construction and recursive import resolution. Root is identical to Parse; +// Schema is left empty. +func (p *Parser) ParseAST(ctx context.Context, workspace, path string, overlays map[string]string) (*ParseResult, error) { + if err := ctx.Err(); err != nil { + return nil, err + } + + fsys, _, relPath, err := parserFS(workspace, path, overlays) + if err != nil { + return nil, err + } + + src, err := fs.ReadFile(fsys, relPath) + if err != nil { + return nil, err + } + + // Non-nil empty schema so semanticDocument.valid() passes for AST-only + // readers (reference counting reads Root, never Schema). + result := &ParseResult{Schema: &schema.WebRPCSchema{ + Types: []*schema.Type{}, + Errors: []*schema.Error{}, + Services: []*schema.Service{}, + }} + + astParser, err := newUpstreamParser(relPath, src) + if err != nil { + result.Errors = []error{err} + return result, nil + } + + if err := runUpstreamParser(astParser); err != nil { + // Best-effort: keep the partial root so mid-edit features still work. + rootNode := astParser.root + result.Root = &rootNode + result.Errors = []error{err} + return result, nil + } + + rootNode := astParser.root + result.Root = &rootNode + return result, nil +} + // These substrings couple us to upstream webrpc error wording — there is no typed // error to match on. TestVersionRequiredErrorFormat pins them so an upstream bump // that changes the wording fails CI instead of silently breaking import handling. diff --git a/internal/ridl/parser_ast_test.go b/internal/ridl/parser_ast_test.go new file mode 100644 index 0000000..17bc63f --- /dev/null +++ b/internal/ridl/parser_ast_test.go @@ -0,0 +1,129 @@ +package ridl + +import ( + "context" + "path/filepath" + "testing" +) + +// multiDeclRIDL is a representative file that contains structs, an enum, and an +// error — enough surface area to catch Root divergence between Parse and ParseAST. +const multiDeclRIDL = `webrpc = v1 + +name = multidecl +version = v0.0.1 + +struct User + - id: uint64 + - name: string + +struct Account + - userId: uint64 + +enum Status: uint32 + - Active = 1 + - Inactive = 2 + +error 1 NotFound "not found" HTTP 404 +` + +// TestParseASTRootEquivalentToParse asserts that ParseAST and Parse produce +// structurally identical Root nodes for a valid multi-declaration file: same +// struct/enum/error counts and matching name strings. Schema divergence is +// expected and intentional — only Root is compared. +func TestParseASTRootEquivalentToParse(t *testing.T) { + dir := t.TempDir() + p := filepath.Join(dir, "multi.ridl") + writeTestFile(t, p, multiDeclRIDL) + + ctx := context.Background() + full, err := NewParser().Parse(ctx, dir, p, nil) + if err != nil { + t.Fatalf("Parse error: %v", err) + } + if full.Root == nil { + t.Fatal("Parse returned nil Root") + } + + ast, err := NewParser().ParseAST(ctx, dir, p, nil) + if err != nil { + t.Fatalf("ParseAST error: %v", err) + } + if ast.Root == nil { + t.Fatal("ParseAST returned nil Root") + } + + if got, want := len(ast.Root.Structs()), len(full.Root.Structs()); got != want { + t.Errorf("Structs count: ParseAST=%d, Parse=%d", got, want) + } + if got, want := len(ast.Root.Enums()), len(full.Root.Enums()); got != want { + t.Errorf("Enums count: ParseAST=%d, Parse=%d", got, want) + } + if got, want := len(ast.Root.Errors()), len(full.Root.Errors()); got != want { + t.Errorf("Errors count: ParseAST=%d, Parse=%d", got, want) + } + + for i, s := range full.Root.Structs() { + if i >= len(ast.Root.Structs()) { + break + } + if got, want := ast.Root.Structs()[i].Name().String(), s.Name().String(); got != want { + t.Errorf("Struct[%d] name: ParseAST=%q, Parse=%q", i, got, want) + } + } + for i, e := range full.Root.Enums() { + if i >= len(ast.Root.Enums()) { + break + } + if got, want := ast.Root.Enums()[i].Name().String(), e.Name().String(); got != want { + t.Errorf("Enum[%d] name: ParseAST=%q, Parse=%q", i, got, want) + } + } + for i, e := range full.Root.Errors() { + if i >= len(ast.Root.Errors()) { + break + } + if got, want := ast.Root.Errors()[i].Name().String(), e.Name().String(); got != want { + t.Errorf("Error[%d] name: ParseAST=%q, Parse=%q", i, got, want) + } + } +} + +func TestParseASTReturnsRootAndEmptySchema(t *testing.T) { + dir := t.TempDir() + p := filepath.Join(dir, "a.ridl") + writeTestFile(t, p, "webrpc = v1\n\nname = test\nversion = v0.0.1\n\nstruct User\n - id: uint64\n") + + result, err := NewParser().ParseAST(context.Background(), dir, p, nil) + if err != nil { + t.Fatalf("ParseAST error: %v", err) + } + if result == nil || result.Root == nil { + t.Fatal("expected non-nil Root") + } + if result.Schema == nil { + t.Fatal("expected non-nil empty Schema so semanticDocument.valid() passes") + } + if len(result.Root.Structs()) != 1 { + t.Fatalf("expected 1 struct, got %d", len(result.Root.Structs())) + } +} + +func TestParseASTSkipsImportResolution(t *testing.T) { + dir := t.TempDir() + // Imports a file that does NOT exist on disk. Full Parse would chase it; + // ParseAST must not, and must still return this file's Root. + p := filepath.Join(dir, "a.ridl") + writeTestFile(t, p, "webrpc = v1\n\nname = test\nversion = v0.0.1\n\nimport\n - missing.ridl\n\nstruct User\n - id: uint64\n") + + result, err := NewParser().ParseAST(context.Background(), dir, p, nil) + if err != nil { + t.Fatalf("ParseAST error: %v", err) + } + if result == nil || result.Root == nil || len(result.Root.Structs()) != 1 { + t.Fatal("expected this file's Root regardless of missing import") + } + if result.Schema == nil { + t.Fatal("expected non-nil Schema") + } +}