diff --git a/internal/documents/store.go b/internal/documents/store.go index 58d291e..7836fa1 100644 --- a/internal/documents/store.go +++ b/internal/documents/store.go @@ -47,6 +47,24 @@ func (s *Store) Delete(uri string) { delete(s.docs, uri) } +// SetResult attaches a parse result to the stored document via copy-on-write, +// leaving any snapshot a caller already holds unmutated. The update is dropped if +// the document is gone or its version has moved on (a newer change superseded the +// content this result was parsed from), so a stale parse is never cached. +func (s *Store) SetResult(uri string, version int32, result *ridl.ParseResult) { + s.mu.Lock() + defer s.mu.Unlock() + + doc, ok := s.docs[uri] + if !ok || doc.Version != version { + return + } + + updated := *doc + updated.Result = result + s.docs[uri] = &updated +} + func (s *Store) All() []*Document { s.mu.RLock() defer s.mu.RUnlock() diff --git a/internal/documents/store_test.go b/internal/documents/store_test.go new file mode 100644 index 0000000..af5c473 --- /dev/null +++ b/internal/documents/store_test.go @@ -0,0 +1,36 @@ +package documents + +import ( + "testing" + + ridl "github.com/webrpc/ridl-lsp/internal/ridl" +) + +// TestSetResultDropsStaleVersion is the whole reason SetResult takes a version: +// a parse of older content (a slower request) must not clobber the result of a +// newer change that already superseded it. +func TestSetResultDropsStaleVersion(t *testing.T) { + s := NewStore() + s.Set(&Document{URI: "u", Version: 2}) + + s.SetResult("u", 1, &ridl.ParseResult{}) // parsed from superseded content + if doc, _ := s.Get("u"); doc.Result != nil { + t.Fatal("stale-version result must be dropped") + } + + current := &ridl.ParseResult{} + s.SetResult("u", 2, current) + if doc, _ := s.Get("u"); doc.Result != current { + t.Fatal("matching-version result must be applied") + } +} + +// TestSetResultMissingDocIsNoop guards the gone-document branch: a result +// arriving after the document closed must neither panic nor resurrect it. +func TestSetResultMissingDocIsNoop(t *testing.T) { + s := NewStore() + s.SetResult("missing", 1, &ridl.ParseResult{}) + if _, ok := s.Get("missing"); ok { + t.Fatal("SetResult must not create a document") + } +} diff --git a/internal/lsp/cancellation_test.go b/internal/lsp/cancellation_test.go new file mode 100644 index 0000000..016bd48 --- /dev/null +++ b/internal/lsp/cancellation_test.go @@ -0,0 +1,70 @@ +package lsp + +import ( + "context" + "os" + "path/filepath" + "testing" + + "go.lsp.dev/protocol" +) + +// TestParseDocumentSkipsCanceledContext: a cancelled request must not surface +// ctx.Err() as a "context canceled" diagnostic, nor clear the cached parse +// result. Parse returns the context error once ctx is done, and parseDocument +// must treat that as "stop", not "the document is broken" (regression guard for +// the I6 cancellation plumbing). +func TestParseDocumentSkipsCanceledContext(t *testing.T) { + srv, _, dir := setupServer(t) + + path := filepath.Join(dir, "doc.ridl") + if err := os.WriteFile(path, []byte(validRIDL), 0644); err != nil { + t.Fatal(err) + } + uri := fileURI(path) + + _ = srv.DidOpen(context.Background(), &protocol.DidOpenTextDocumentParams{ + TextDocument: protocol.TextDocumentItem{URI: protocol.DocumentURI(uri), Text: validRIDL, Version: 1}, + }) + + doc, ok := srv.docs.Get(uri) + if !ok { + t.Fatal("document missing after DidOpen") + } + if doc.Result == nil { + t.Fatal("expected a cached result for the valid document") + } + + ctx, cancel := context.WithCancel(context.Background()) + cancel() + + if diags := srv.parseDocument(ctx, doc); diags != nil { + t.Fatalf("cancelled parse must yield no diagnostics, got %v", diags) + } + + after, _ := srv.docs.Get(uri) + if after.Result == nil { + t.Fatal("cancelled parse must not clear the cached result") + } +} + +// TestParsePathHonorsCanceledContext: the ctx-aware parse used by the diagnostics +// path (e.g. transitive re-import checks) must stop on a cancelled request rather +// than parsing imports off a dead request. +func TestParsePathHonorsCanceledContext(t *testing.T) { + srv, _, dir := setupServer(t) + + other := filepath.Join(dir, "other.ridl") + if err := os.WriteFile(other, []byte(validRIDL), 0644); err != nil { + t.Fatal(err) + } + + ctx, cancel := context.WithCancel(context.Background()) + cancel() + + // other.ridl is not open, so there is no cached result to short-circuit on: + // parsePath must parse, and that parse must bail on the cancelled context. + if got := srv.parsePath(ctx, other); got != nil { + t.Fatalf("expected nil from parsePath on cancelled ctx, got %v", got) + } +} diff --git a/internal/lsp/code_action.go b/internal/lsp/code_action.go index bd6bf75..4a80d0f 100644 --- a/internal/lsp/code_action.go +++ b/internal/lsp/code_action.go @@ -216,7 +216,9 @@ func (s *Server) missingImportCandidates(doc *semanticDocument) []missingImportC candidates := make([]missingImportCandidate, 0, 4) seen := map[string]struct{}{} for _, unresolved := range doc.unresolvedSymbols(s.resolveTypeDefinition, s.resolveErrorDefinition) { - targetPath, ok := s.uniqueImportCandidatePath(doc.path, unresolved.kind, unresolved.name) + // Interactive code action: uses the background parser, consistent with the + // other navigation paths (see parsePathForNavigation). + targetPath, ok := s.uniqueImportCandidatePath(context.Background(), doc.path, unresolved.kind, unresolved.name) if !ok || docHasImportedPath(doc, targetPath) { continue } @@ -522,16 +524,21 @@ func unresolvedTypeNames(expr string) []string { return names } -func (s *Server) uniqueImportCandidatePath(docPath string, kind referenceKind, name string) (string, bool) { +func (s *Server) uniqueImportCandidatePath(ctx context.Context, docPath string, kind referenceKind, name string) (string, bool) { var definers, reExporters []string seen := map[string]struct{}{} for _, path := range s.referenceCandidatePaths() { + // This scans and parses every candidate in the workspace; bail promptly + // when the driving request (e.g. a diagnostics refresh) is cancelled. + if ctx.Err() != nil { + return "", false + } if path == "" || path == docPath { continue } - result := s.parsePathForNavigation(path) + result := s.parsePath(ctx, path) if result == nil || result.Root == nil { continue } @@ -1066,9 +1073,9 @@ func (s *Server) transitiveReImportCodeActions(doc *documents.Document, diagnost continue } - originalPath, ok := s.uniqueImportCandidatePath(doc.Path, referenceKindType, name) + originalPath, ok := s.uniqueImportCandidatePath(context.Background(), doc.Path, referenceKindType, name) if !ok { - originalPath, ok = s.uniqueImportCandidatePath(doc.Path, referenceKindError, name) + originalPath, ok = s.uniqueImportCandidatePath(context.Background(), doc.Path, referenceKindError, name) } if !ok { continue diff --git a/internal/lsp/definition.go b/internal/lsp/definition.go index 7d9f226..fe16a6f 100644 --- a/internal/lsp/definition.go +++ b/internal/lsp/definition.go @@ -48,12 +48,20 @@ func definitionForToken(path string, token *ridl.TokenNode) *definitionMatch { return &definitionMatch{path: path, token: token} } +// parsePathForNavigation deliberately uses context.Background(): interactive +// navigation is a fast single-file parse, and threading the request ctx through +// the whole resolution-callback layer would be high-churn for negligible benefit. +// The cancellable callers (diagnostics) use parsePath directly. func (s *Server) parsePathForNavigation(path string) *ridl.ParseResult { + return s.parsePath(context.Background(), path) +} + +func (s *Server) parsePath(ctx context.Context, path string) *ridl.ParseResult { if doc, ok := s.docs.FindByPath(path); ok && doc.Result != nil && doc.Result.Root != nil { return doc.Result } - result, err := s.parser.Parse(s.workspace.Root(), path, s.overlayContents()) + result, err := s.parser.Parse(ctx, s.workspace.Root(), path, s.overlayContents()) if err != nil { return nil } diff --git a/internal/lsp/diagnostics.go b/internal/lsp/diagnostics.go index 81304c9..6a5e53f 100644 --- a/internal/lsp/diagnostics.go +++ b/internal/lsp/diagnostics.go @@ -21,8 +21,10 @@ var ( ) func (s *Server) parseAndPublishDiagnostics(ctx context.Context, doc *documents.Document) { - diagnostics := s.parseDocument(doc) - if s.client == nil { + diagnostics := s.parseDocument(ctx, doc) + // A cancelled/superseded request must not overwrite the client's diagnostics + // with a half-computed (or empty) set. + if ctx.Err() != nil || s.client == nil { return } @@ -34,12 +36,19 @@ func (s *Server) parseAndPublishDiagnostics(ctx context.Context, doc *documents. } } -func (s *Server) parseDocument(doc *documents.Document) []protocol.Diagnostic { +func (s *Server) parseDocument(ctx context.Context, doc *documents.Document) []protocol.Diagnostic { overlays := s.overlayContents() - result, err := s.parser.Parse(s.workspace.Root(), doc.Path, overlays) + result, err := s.parser.Parse(ctx, s.workspace.Root(), doc.Path, overlays) + // Bail before touching anything on cancellation. Import recursion swallows + // ctx errors as skipped imports, so Parse can return err == nil with an + // incomplete result after a cancel — caching that would poison the document's + // state, and surfacing ctx.Err() as a diagnostic would flash a bogus error. + if ctx.Err() != nil { + return nil + } if err != nil { - doc.Result = nil + s.docs.SetResult(doc.URI, doc.Version, nil) return []protocol.Diagnostic{ { Range: lineRange(1), @@ -51,11 +60,14 @@ func (s *Server) parseDocument(doc *documents.Document) []protocol.Diagnostic { } if len(result.Errors) == 0 { - doc.Result = result - return s.importDiagnostics(doc) + s.docs.SetResult(doc.URI, doc.Version, result) + return s.importDiagnostics(ctx, doc, result) } - doc.Result = nil + // Cache the partial result even with parse errors: the parser still produces + // a best-effort AST (Root is populated), and reusing it lets navigation work + // — and avoids a re-parse per request — while the user is mid-edit. + s.docs.SetResult(doc.URI, doc.Version, result) diagnostics := make([]protocol.Diagnostic, 0, len(result.Errors)) for _, e := range result.Errors { @@ -151,16 +163,16 @@ func severityWarning() protocol.DiagnosticSeverity { return protocol.DiagnosticSeverityWarning } -func (s *Server) importDiagnostics(doc *documents.Document) []protocol.Diagnostic { - if doc == nil || doc.Result == nil || doc.Result.Root == nil { +func (s *Server) importDiagnostics(ctx context.Context, doc *documents.Document, result *ridl.ParseResult) []protocol.Diagnostic { + if doc == nil || result == nil || result.Root == nil { return nil } - semanticDoc := newSemanticDocument(doc.Path, doc.Content, doc.Result) + semanticDoc := newSemanticDocument(doc.Path, doc.Content, result) referenced := semanticDoc.referencedNames() var diagnostics []protocol.Diagnostic - for _, importNode := range doc.Result.Root.Imports() { + for _, importNode := range result.Root.Imports() { if importNode == nil || importNode.Path() == nil { continue } @@ -172,7 +184,7 @@ func (s *Server) importDiagnostics(doc *documents.Document) []protocol.Diagnosti importPath := importNode.Path().String() resolvedPath := workspace.ResolveImportPath(doc.Path, importPath) - importResult, err := s.parser.Parse(s.workspace.Root(), resolvedPath, s.overlayContents()) + importResult, err := s.parser.Parse(ctx, s.workspace.Root(), resolvedPath, s.overlayContents()) if err != nil || importResult == nil || importResult.Root == nil { continue } @@ -224,13 +236,13 @@ func (s *Server) importDiagnostics(doc *documents.Document) []protocol.Diagnosti } // Check selective imports for transitive re-imports. - for _, importNode := range doc.Result.Root.Imports() { + for _, importNode := range result.Root.Imports() { if importNode == nil || importNode.Path() == nil || len(importNode.Members()) == 0 { continue } importPath := workspace.ResolveImportPath(doc.Path, importNode.Path().String()) - importResult := s.parsePathForNavigation(importPath) + importResult := s.parsePath(ctx, importPath) if importResult == nil || importResult.Root == nil { continue } @@ -247,9 +259,9 @@ func (s *Server) importDiagnostics(doc *documents.Document) []protocol.Diagnosti continue } - originalPath, ok := s.uniqueImportCandidatePath(doc.Path, referenceKindType, name) + originalPath, ok := s.uniqueImportCandidatePath(ctx, doc.Path, referenceKindType, name) if !ok { - originalPath, ok = s.uniqueImportCandidatePath(doc.Path, referenceKindError, name) + originalPath, ok = s.uniqueImportCandidatePath(ctx, doc.Path, referenceKindError, name) } if !ok { continue diff --git a/internal/lsp/diagnostics_test.go b/internal/lsp/diagnostics_test.go index 8a8dcd2..b501a60 100644 --- a/internal/lsp/diagnostics_test.go +++ b/internal/lsp/diagnostics_test.go @@ -268,7 +268,7 @@ func TestInvalidToValidClearsDiagnostics(t *testing.T) { } } -func TestValidToInvalidClearsCachedParseResult(t *testing.T) { +func TestValidToInvalidRetainsPartialParseResult(t *testing.T) { srv, client, dir := setupServer(t) ctx := context.Background() @@ -316,8 +316,11 @@ func TestValidToInvalidClearsCachedParseResult(t *testing.T) { if !ok { t.Fatal("expected changed document to remain tracked") } - if doc.Result != nil { - t.Fatal("expected cached parse result to be cleared after parse failure") + // The parser produces a best-effort AST even on failure; retaining it lets + // navigation reuse the partial result mid-edit instead of re-parsing the same + // invalid buffer on every request (audit I3). + if doc.Result == nil || doc.Result.Root == nil { + t.Fatal("expected a partial parse result to be retained after parse failure") } } diff --git a/internal/lsp/document_cow_test.go b/internal/lsp/document_cow_test.go new file mode 100644 index 0000000..4de5947 --- /dev/null +++ b/internal/lsp/document_cow_test.go @@ -0,0 +1,62 @@ +package lsp + +import ( + "context" + "os" + "path/filepath" + "testing" + + "go.lsp.dev/protocol" +) + +// TestDidChangeDoesNotMutatePriorSnapshot pins the copy-on-write invariant: a +// *Document handed out by the store must never be mutated in place. A concurrent +// reader (handlers run in their own goroutines via jsonrpc2.AsyncHandler) holding +// an earlier snapshot must keep seeing that snapshot's content after DidChange. +func TestDidChangeDoesNotMutatePriorSnapshot(t *testing.T) { + srv, _, dir := setupServer(t) + ctx := context.Background() + + path := filepath.Join(dir, "doc.ridl") + if err := os.WriteFile(path, []byte(validRIDL), 0644); err != nil { + t.Fatal(err) + } + uri := fileURI(path) + + _ = srv.DidOpen(ctx, &protocol.DidOpenTextDocumentParams{ + TextDocument: protocol.TextDocumentItem{URI: protocol.DocumentURI(uri), Text: validRIDL, Version: 1}, + }) + + snapshot, ok := srv.docs.Get(uri) + if !ok { + t.Fatal("document not in store after DidOpen") + } + originalContent := snapshot.Content + + const changed = "webrpc = v1\n\nname = changed\nversion = v0.2.0\n" + _ = srv.DidChange(ctx, &protocol.DidChangeTextDocumentParams{ + TextDocument: protocol.VersionedTextDocumentIdentifier{ + TextDocumentIdentifier: protocol.TextDocumentIdentifier{URI: protocol.DocumentURI(uri)}, + Version: 2, + }, + ContentChanges: []protocol.TextDocumentContentChangeEvent{{Text: changed}}, + }) + + if snapshot.Content != originalContent { + t.Fatalf("prior snapshot was mutated in place: content is now %q, want %q", snapshot.Content, originalContent) + } + if snapshot.Version != 1 { + t.Fatalf("prior snapshot version mutated: got %d, want 1", snapshot.Version) + } + + updated, ok := srv.docs.Get(uri) + if !ok { + t.Fatal("document missing after DidChange") + } + if updated.Content != changed { + t.Fatalf("store did not reflect new content: got %q", updated.Content) + } + if updated.Version != 2 { + t.Fatalf("store version: got %d, want 2", updated.Version) + } +} diff --git a/internal/lsp/references.go b/internal/lsp/references.go index e176497..976ff26 100644 --- a/internal/lsp/references.go +++ b/internal/lsp/references.go @@ -9,6 +9,7 @@ import ( "strings" "go.lsp.dev/protocol" + "go.uber.org/zap" ridl "github.com/webrpc/ridl-lsp/internal/ridl" ) @@ -330,8 +331,23 @@ func (s *Server) referenceCandidatePaths() []string { return paths } - _ = filepath.WalkDir(root, func(path string, entry os.DirEntry, err error) error { - if err != nil || entry == nil || entry.IsDir() || filepath.Ext(path) != ".ridl" { + walkErr := filepath.WalkDir(root, func(path string, entry os.DirEntry, err error) error { + if err != nil { + // Degrade to partial results instead of silently truncating: a + // swallowed error reads as "searched the whole workspace" when it didn't. + s.logger.Warn("workspace scan: skipping unreadable entry", zap.String("path", path), zap.Error(err)) + return nil + } + if entry.IsDir() { + // Pruning keeps a monorepo scan bounded — these trees never hold + // project schemas but can be enormous, and the scan runs synchronously + // on the request. + if path != root && skipWorkspaceDir(entry.Name()) { + return filepath.SkipDir + } + return nil + } + if filepath.Ext(path) != ".ridl" { return nil } if _, ok := seen[path]; ok { @@ -341,11 +357,22 @@ func (s *Server) referenceCandidatePaths() []string { paths = append(paths, path) return nil }) + if walkErr != nil { + s.logger.Warn("workspace scan did not complete", zap.String("root", root), zap.Error(walkErr)) + } sort.Strings(paths) return paths } +func skipWorkspaceDir(name string) bool { + switch name { + case ".git", "node_modules", "vendor": + return true + } + return strings.HasPrefix(name, ".") +} + func (s *Server) collectReferenceLocations(target *referenceTarget, includeDeclaration bool) []protocol.Location { if target == nil || target.definition == nil { return nil diff --git a/internal/lsp/server.go b/internal/lsp/server.go index b6ed27a..ce3a950 100644 --- a/internal/lsp/server.go +++ b/internal/lsp/server.go @@ -131,9 +131,14 @@ func (s *Server) DidChange(ctx context.Context, params *protocol.DidChangeTextDo } if len(params.ContentChanges) > 0 { - doc.Content = params.ContentChanges[len(params.ContentChanges)-1].Text - doc.Version = params.TextDocument.Version - s.docs.Set(doc) + // Copy-on-write: never mutate the stored *Document in place, since other + // handlers may hold the prior snapshot. The new content invalidates the + // cached parse result. + updated := *doc + updated.Content = params.ContentChanges[len(params.ContentChanges)-1].Text + updated.Version = params.TextDocument.Version + updated.Result = nil + s.docs.Set(&updated) s.refreshOpenDocuments(ctx) } diff --git a/internal/lsp/workspace_walk_test.go b/internal/lsp/workspace_walk_test.go new file mode 100644 index 0000000..a35f1c9 --- /dev/null +++ b/internal/lsp/workspace_walk_test.go @@ -0,0 +1,49 @@ +package lsp + +import ( + "os" + "path/filepath" + "slices" + "testing" +) + +// TestReferenceCandidatePathsPrunesHeavyDirs: the workspace scan that backs +// find-references / workspace-symbols / missing-import quick-fixes must not +// descend into .git, node_modules, vendor, or hidden directories. They never +// hold project schemas and can dominate a monorepo scan that runs synchronously +// on the request (audit I7). +func TestReferenceCandidatePathsPrunesHeavyDirs(t *testing.T) { + srv, _, dir := setupServer(t) + + write := func(rel string) string { + p := filepath.Join(dir, filepath.FromSlash(rel)) + if err := os.MkdirAll(filepath.Dir(p), 0o755); err != nil { + t.Fatal(err) + } + if err := os.WriteFile(p, []byte(validRIDL), 0o644); err != nil { + t.Fatal(err) + } + return p + } + + wanted := write("api.ridl") + write("node_modules/dep/schema.ridl") + write("vendor/lib/schema.ridl") + write(".git/hooks/schema.ridl") + write(".cache/schema.ridl") + + paths := srv.referenceCandidatePaths() + + if !slices.Contains(paths, wanted) { + t.Fatalf("expected top-level api.ridl in candidates, got %v", paths) + } + for _, p := range paths { + rel, _ := filepath.Rel(dir, p) + rel = filepath.ToSlash(rel) + for _, pruned := range []string{"node_modules/", "vendor/", ".git/", ".cache/"} { + if len(rel) >= len(pruned) && rel[:len(pruned)] == pruned { + t.Fatalf("candidate from pruned dir leaked: %q", rel) + } + } + } +} diff --git a/internal/ridl/layout_canary_test.go b/internal/ridl/layout_canary_test.go index 8925852..6d467ae 100644 --- a/internal/ridl/layout_canary_test.go +++ b/internal/ridl/layout_canary_test.go @@ -1,6 +1,7 @@ package ridl import ( + "context" "path/filepath" "testing" ) @@ -64,7 +65,7 @@ func TestUpstreamLayoutCanary(t *testing.T) { writeTestFile(t, path, content) - result, err := NewParser().Parse(dir, path, nil) + result, err := NewParser().Parse(context.Background(), dir, path, nil) if err != nil { t.Fatalf("Parse returned error: %v", err) } diff --git a/internal/ridl/parser.go b/internal/ridl/parser.go index 98e63be..4ee944d 100644 --- a/internal/ridl/parser.go +++ b/internal/ridl/parser.go @@ -1,6 +1,7 @@ package ridl import ( + "context" "io/fs" "os" "path/filepath" @@ -170,11 +171,17 @@ func NewParser() *Parser { // Parse parses the RIDL file at path using workspace as the preferred fs root. // overlays maps document paths to in-memory content (open editor buffers). -func (p *Parser) Parse(workspace, path string, overlays map[string]string) (*ParseResult, error) { - return p.parse(workspace, path, overlays, map[string]struct{}{}) +func (p *Parser) Parse(ctx context.Context, workspace, path string, overlays map[string]string) (*ParseResult, error) { + return p.parse(ctx, workspace, path, overlays, map[string]struct{}{}) } -func (p *Parser) parse(workspace, path string, overlays map[string]string, visited map[string]struct{}) (*ParseResult, error) { +func (p *Parser) parse(ctx context.Context, workspace, path string, overlays map[string]string, visited map[string]struct{}) (*ParseResult, error) { + // Bail before each (possibly recursive, import-chasing) parse so a superseded + // or client-cancelled request stops promptly instead of walking the import graph. + if err := ctx.Err(); err != nil { + return nil, err + } + fsys, root, relPath, err := parserFS(workspace, path, overlays) if err != nil { return nil, err @@ -196,7 +203,7 @@ func (p *Parser) parse(workspace, path string, overlays map[string]string, visit if err := runUpstreamParser(astParser); err != nil { rootNode := astParser.root result.Root = &rootNode - result.Schema = p.buildPartialSchema(workspace, path, result.Root, overlays, visited) + result.Schema = p.buildPartialSchema(ctx, workspace, path, result.Root, overlays, visited) result.Errors = []error{err} return result, nil } @@ -204,7 +211,7 @@ func (p *Parser) parse(workspace, path string, overlays map[string]string, visit rootNode := astParser.root result.Root = &rootNode imported := len(visited) > 0 - result.Schema = p.buildPartialSchema(workspace, path, result.Root, overlays, visited) + result.Schema = p.buildPartialSchema(ctx, workspace, path, result.Root, overlays, visited) schemaDoc, err := ridl.NewParser(fsys, root, relPath).Parse() if err != nil { @@ -231,7 +238,7 @@ func isVersionOptionalSchemaError(err error, imported bool) bool { return imported || strings.Contains(err.Error(), "stack trace:") } -func (p *Parser) buildPartialSchema(workspace, path string, root *RootNode, overlays map[string]string, visited map[string]struct{}) *schema.WebRPCSchema { +func (p *Parser) buildPartialSchema(ctx context.Context, workspace, path string, root *RootNode, overlays map[string]string, visited map[string]struct{}) *schema.WebRPCSchema { doc := &schema.WebRPCSchema{ Types: []*schema.Type{}, Errors: []*schema.Error{}, @@ -273,7 +280,7 @@ func (p *Parser) buildPartialSchema(workspace, path string, root *RootNode, over continue } - importResult, err := p.parse(workspace, importPath, overlays, cloneVisited(visited)) + importResult, err := p.parse(ctx, workspace, importPath, overlays, cloneVisited(visited)) if err != nil || importResult == nil || importResult.Schema == nil { continue } diff --git a/internal/ridl/parser_test.go b/internal/ridl/parser_test.go index 0ecf033..2700824 100644 --- a/internal/ridl/parser_test.go +++ b/internal/ridl/parser_test.go @@ -1,6 +1,7 @@ package ridl import ( + "context" "os" "path/filepath" "strings" @@ -33,7 +34,7 @@ service SharedService writeTestFile(t, mainPath, mainContent) writeTestFile(t, sharedPath, sharedContent) - result, err := NewParser().Parse(dir, mainPath, nil) + result, err := NewParser().Parse(context.Background(), dir, mainPath, nil) if err != nil { t.Fatalf("Parse returned error: %v", err) } @@ -62,7 +63,7 @@ service MainService writeTestFile(t, path, content) - result, err := NewParser().Parse(dir, path, nil) + result, err := NewParser().Parse(context.Background(), dir, path, nil) if err != nil { t.Fatalf("Parse returned error: %v", err) } @@ -74,6 +75,19 @@ service MainService } } +func TestParseHonorsCanceledContext(t *testing.T) { + dir := t.TempDir() + path := filepath.Join(dir, "x.ridl") + writeTestFile(t, path, "webrpc = v1\n\nname = x\nversion = v1.0.0\n") + + ctx, cancel := context.WithCancel(context.Background()) + cancel() + + if _, err := NewParser().Parse(ctx, dir, path, nil); err == nil { + t.Fatal("expected Parse to return the context error when the context is already cancelled") + } +} + func writeTestFile(t *testing.T, path, content string) { t.Helper()