diff --git a/internal/mcp/deps.go b/internal/mcp/deps.go index d31806e..3c4a9d6 100644 --- a/internal/mcp/deps.go +++ b/internal/mcp/deps.go @@ -33,29 +33,19 @@ type Parser interface { ParseWithContext(ctx context.Context, filePath string, content []byte) ([]model.Node, []model.Edge, error) } -// ImpactAnalyzer defines the blast-radius analysis contract for graph nodes. -// @intent Injects an analyzer that calculates the impact radius of node changes into the server handler. +// ImpactAnalyzer defines the bounded blast-radius analysis contract for graph nodes. +// @intent inject a node/depth-capped blast-radius analyzer so a single MCP request cannot +// expand into an unbounded graph walk. // @see mcp.handlers.getImpactRadius type ImpactAnalyzer interface { - ImpactRadius(ctx context.Context, nodeID uint, depth int) ([]model.Node, error) -} - -// BoundedImpactAnalyzer extends ImpactAnalyzer with node and depth caps to prevent runaway traversal on large graphs. -// @intent expose bounded blast-radius analysis for handlers that must protect shared MCP requests from unbounded graph walks. -type BoundedImpactAnalyzer interface { ImpactRadiusBounded(ctx context.Context, nodeID uint, depth int, opts impactpkg.RadiusOptions) (*impactpkg.RadiusResult, error) } -// FlowTracer defines the call-flow tracing contract for graph nodes. -// @intent Connects an analyzer to the server that reconstructs call flows starting from a given node. +// FlowTracer defines the bounded call-flow tracing contract for graph nodes. +// @intent inject a node-capped call-flow tracer so a deep call chain cannot expand into an +// unbounded traversal. // @see mcp.handlers.traceFlow type FlowTracer interface { - TraceFlow(ctx context.Context, startNodeID uint) (*model.Flow, error) -} - -// BoundedFlowTracer extends FlowTracer with a node cap to prevent runaway traversal on deeply nested call chains. -// @intent let MCP handlers trace deep call chains without letting one request expand into an unbounded traversal. -type BoundedFlowTracer interface { TraceFlowBounded(ctx context.Context, startNodeID uint, opts flowspkg.TraceOptions) (*flowspkg.TraceResult, error) } @@ -193,4 +183,8 @@ type Deps struct { MaxFileBytes int64 MaxTotalParsedBytes int64 + + // RefreshSearchDocuments overrides the search-document refresh used after a build; nil uses + // service.RefreshSearchDocuments. Kept as an injectable field so tests need no package globals. + RefreshSearchDocuments func(ctx context.Context, db *gorm.DB) (int, error) } diff --git a/internal/mcp/handler_analysis.go b/internal/mcp/handler_analysis.go index 71cfed4..320a486 100644 --- a/internal/mcp/handler_analysis.go +++ b/internal/mcp/handler_analysis.go @@ -161,32 +161,13 @@ func (h *handlers) getImpactRadius(ctx context.Context, request mcp.CallToolRequ return "", nodeNotFoundErr(qn) } - var nodes []model.Node - truncated := false - if bounded, ok := h.deps.ImpactAnalyzer.(BoundedImpactAnalyzer); ok { - res, err := bounded.ImpactRadiusBounded(ctx, node.ID, depth, impactpkg.RadiusOptions{MaxDepth: maxDepth, MaxNodes: maxNodes}) - if err != nil { - log.ErrorContext(ctx, "impact analysis error", append(obs.TraceLogArgs(ctx), "node_id", node.ID, trace.SlogError(err))...) - return "", trace.Wrap(err, "impact analysis error") - } - nodes = res.Nodes - truncated = res.Truncated - } else { - if maxDepth > 0 && depth > maxDepth { - depth = maxDepth - truncated = true - } - var err error - nodes, err = h.deps.ImpactAnalyzer.ImpactRadius(ctx, node.ID, depth) - if err != nil { - log.ErrorContext(ctx, "impact analysis error", append(obs.TraceLogArgs(ctx), "node_id", node.ID, trace.SlogError(err))...) - return "", trace.Wrap(err, "impact analysis error") - } - if maxNodes > 0 && len(nodes) > maxNodes { - nodes = nodes[:maxNodes] - truncated = true - } + res, err := h.deps.ImpactAnalyzer.ImpactRadiusBounded(ctx, node.ID, depth, impactpkg.RadiusOptions{MaxDepth: maxDepth, MaxNodes: maxNodes}) + if err != nil { + log.ErrorContext(ctx, "impact analysis error", append(obs.TraceLogArgs(ctx), "node_id", node.ID, trace.SlogError(err))...) + return "", trace.Wrap(err, "impact analysis error") } + nodes := res.Nodes + truncated := res.Truncated log.InfoContext(ctx, "get_impact_radius completed", append(obs.TraceLogArgs(ctx), "qualified_name", qn, "result_count", len(nodes))...) @@ -249,32 +230,15 @@ func (h *handlers) traceFlow(ctx context.Context, request mcp.CallToolRequest) ( return "", nodeNotFoundErr(qn) } - var flow *model.Flow - truncated := false - containsFallbackCalls := false - fallbackEdgesCount := 0 - if bounded, ok := h.deps.FlowTracer.(BoundedFlowTracer); ok { - res, err := bounded.TraceFlowBounded(ctx, node.ID, flowspkg.TraceOptions{MaxNodes: maxNodes, IncludeFallbackCalls: &includeFallbackCalls}) - if err != nil { - log.ErrorContext(ctx, "trace error", append(obs.TraceLogArgs(ctx), "node_id", node.ID, trace.SlogError(err))...) - return "", trace.Wrap(err, "trace error") - } - flow = res.Flow - truncated = res.Truncated - containsFallbackCalls = res.ContainsFallbackCalls - fallbackEdgesCount = res.FallbackEdgesCount - } else { - var err error - flow, err = h.deps.FlowTracer.TraceFlow(ctx, node.ID) - if err != nil { - log.ErrorContext(ctx, "trace error", append(obs.TraceLogArgs(ctx), "node_id", node.ID, trace.SlogError(err))...) - return "", trace.Wrap(err, "trace error") - } - if maxNodes > 0 && len(flow.Members) > maxNodes { - flow.Members = flow.Members[:maxNodes] - truncated = true - } + res, err := h.deps.FlowTracer.TraceFlowBounded(ctx, node.ID, flowspkg.TraceOptions{MaxNodes: maxNodes, IncludeFallbackCalls: &includeFallbackCalls}) + if err != nil { + log.ErrorContext(ctx, "trace error", append(obs.TraceLogArgs(ctx), "node_id", node.ID, trace.SlogError(err))...) + return "", trace.Wrap(err, "trace error") } + flow := res.Flow + truncated := res.Truncated + containsFallbackCalls := res.ContainsFallbackCalls + fallbackEdgesCount := res.FallbackEdgesCount log.InfoContext(ctx, "trace_flow completed", append(obs.TraceLogArgs(ctx), "qualified_name", qn, "members", len(flow.Members))...) diff --git a/internal/mcp/handler_parse.go b/internal/mcp/handler_parse.go index 7edc214..9a3fe46 100644 --- a/internal/mcp/handler_parse.go +++ b/internal/mcp/handler_parse.go @@ -18,7 +18,13 @@ import ( "github.com/tae2089/code-context-graph/internal/service" ) -var refreshSearchDocuments = service.RefreshSearchDocuments +// @intent refresh search documents through the injected override, defaulting to the service impl. +func (h *handlers) refreshSearchDocuments(ctx context.Context) (int, error) { + if h.deps.RefreshSearchDocuments != nil { + return h.deps.RefreshSearchDocuments(ctx, h.deps.DB) + } + return service.RefreshSearchDocuments(ctx, h.deps.DB) +} // @intent serialize build_or_update_graph results with a fixed JSON schema without changing the wire format. type buildOrUpdateGraphResponse struct { @@ -250,7 +256,7 @@ func (h *handlers) buildOrUpdateGraph(ctx context.Context, request mcp.CallToolR } // search rebuild if h.deps.SearchBackend != nil && h.deps.DB != nil { - if _, err := refreshSearchDocuments(ctx, h.deps.DB); err != nil { + if _, err := h.refreshSearchDocuments(ctx); err != nil { if failClosed { failClosedErr = err failedSteps = append(failedSteps, "search_documents") @@ -274,7 +280,7 @@ func (h *handlers) buildOrUpdateGraph(ctx context.Context, request mcp.CallToolR skippedSteps = appendUniqueStrings(skippedSteps, "communities", "flows") // search only rebuild if h.deps.SearchBackend != nil && h.deps.DB != nil { - if _, err := refreshSearchDocuments(ctx, h.deps.DB); err != nil { + if _, err := h.refreshSearchDocuments(ctx); err != nil { if failClosed { failClosedErr = err failedSteps = append(failedSteps, "search_documents") @@ -433,7 +439,7 @@ func (h *handlers) runPostprocess(ctx context.Context, request mcp.CallToolReque if doFTS { if h.deps.SearchBackend != nil && h.deps.DB != nil { - if _, err := refreshSearchDocuments(ctx, h.deps.DB); err != nil { + if _, err := h.refreshSearchDocuments(ctx); err != nil { if failClosed { failClosedErr = err failedSteps = append(failedSteps, "search_documents") diff --git a/internal/mcp/handlers_test.go b/internal/mcp/handlers_test.go index 423be8d..6dcf5e1 100644 --- a/internal/mcp/handlers_test.go +++ b/internal/mcp/handlers_test.go @@ -1404,9 +1404,7 @@ func TestBuildOrUpdateGraph_DegradedOnSearchDocumentRefreshFailure(t *testing.T) deps := setupTestDeps(t) backend := &failSearchBackend{} deps.SearchBackend = backend - origRefresh := refreshSearchDocuments - defer func() { refreshSearchDocuments = origRefresh }() - refreshSearchDocuments = func(ctx context.Context, db *gorm.DB) (int, error) { + deps.RefreshSearchDocuments = func(ctx context.Context, db *gorm.DB) (int, error) { return 0, errors.New("search document refresh boom") } @@ -1451,9 +1449,7 @@ func TestBuildOrUpdateGraph_MinimalSkipsFTSRebuildOnSearchDocumentRefreshFailure deps := setupTestDeps(t) backend := &failSearchBackend{} deps.SearchBackend = backend - origRefresh := refreshSearchDocuments - defer func() { refreshSearchDocuments = origRefresh }() - refreshSearchDocuments = func(ctx context.Context, db *gorm.DB) (int, error) { + deps.RefreshSearchDocuments = func(ctx context.Context, db *gorm.DB) (int, error) { return 0, errors.New("search document refresh boom") } diff --git a/internal/mcp/testmocks_test.go b/internal/mcp/testmocks_test.go index 999e5ef..29bc336 100644 --- a/internal/mcp/testmocks_test.go +++ b/internal/mcp/testmocks_test.go @@ -365,23 +365,6 @@ type mockFlowTracer struct { remainderErr error } -func (m *mockFlowTracer) TraceFlow(ctx context.Context, startNodeID uint) (*model.Flow, error) { - m.calls++ - flow := m.returnFlow - if flow == nil { - flow = &model.Flow{ - Namespace: ctxns.FromContext(ctx), - Name: "flow_from_mock", - Members: []model.FlowMembership{{ - NodeID: startNodeID, - Ordinal: 0, - Namespace: ctxns.FromContext(ctx), - }}, - } - } - return flow, m.traceErr -} - func (m *mockFlowTracer) TraceFlowBounded(ctx context.Context, startNodeID uint, opts flows.TraceOptions) (*flows.TraceResult, error) { m.calls++ m.opts = append(m.opts, opts) diff --git a/internal/service/build.go b/internal/service/build.go index 8a45ae0..7e71b6f 100644 --- a/internal/service/build.go +++ b/internal/service/build.go @@ -100,8 +100,8 @@ func (s *GraphService) bindAndReleaseNodeBatch(ctx context.Context, txStore stor parsed.tsComments = nil parsed.sourceLines = nil - if testBuildBatchReleaseHook != nil { - testBuildBatchReleaseHook(batches, idx) + if s.onBatchRelease != nil { + s.onBatchRelease(batches, idx) } return nil } @@ -450,7 +450,7 @@ func (s *GraphService) flushBuildEdges(ctx context.Context, txStore store.GraphS return err } end := min(start+buildEdgeResolveChunkSize, len(implementsEdges)) - resolved, err := resolveBuildEdges(ctx, txStore, implementsEdges[start:end], resolveOptions) + resolved, err := s.edgeResolver()(ctx, txStore, implementsEdges[start:end], resolveOptions) if err != nil { s.logger().ErrorContext(ctx, "resolve deferred implements edges failed", append(obs.TraceLogArgs(ctx), "start", start, "end", end, "error", err)...) return trace.Wrap(err, "resolve deferred implements edges") @@ -474,7 +474,7 @@ func (s *GraphService) flushBuildEdges(ctx context.Context, txStore store.GraphS end := min(start+buildEdgeResolveChunkSize, len(parsed.edges)) chunk := parsed.edges[start:end] resolveInput := chunkWithImportWarmup(chunk, importsByPath[parsed.relPath]) - resolved, err := resolveBuildEdges(ctx, txStore, resolveInput, resolveOptions) + resolved, err := s.edgeResolver()(ctx, txStore, resolveInput, resolveOptions) if err != nil { s.logger().ErrorContext(ctx, "resolve deferred edges failed", append(obs.TraceLogArgs(ctx), "file", parsed.relPath, "error", err)...) return trace.Wrap(err, "resolve deferred edges for "+parsed.relPath) diff --git a/internal/service/indexer.go b/internal/service/indexer.go index ec3a12d..86809e6 100644 --- a/internal/service/indexer.go +++ b/internal/service/indexer.go @@ -18,12 +18,7 @@ import ( type languagePackageInfo = treesitter.PackageInfo -var ( - testBuildBatchReleaseHook func([]parsedBuildNodeBatch, int) - resolveBuildEdges resolveBuildEdgesFn = edgeresolve.ResolveWithOptions -) - -// @intent abstract build-time edge resolution so tests and build paths can swap resolver behavior without rewiring callers. +// @intent abstract build-time edge resolution so tests can inject resolver behavior per GraphService. type resolveBuildEdgesFn func(ctx context.Context, lookup edgeresolve.NodeLookup, edges []model.Edge, options edgeresolve.ResolveOptions) ([]model.Edge, error) const ( @@ -89,6 +84,20 @@ type GraphService struct { Walkers map[string]*treesitter.Walker Parsers map[string]Parser Logger *slog.Logger + + // resolveEdges overrides build-time edge resolution; nil uses edgeresolve.ResolveWithOptions. + // onBatchRelease, when set, is notified after each node batch is persisted and its buffers + // are released. Both are test seams kept per-instance so the server has no mutable globals. + resolveEdges resolveBuildEdgesFn + onBatchRelease func([]parsedBuildNodeBatch, int) +} + +// @intent resolve build edges through the injected resolver, defaulting to the production resolver. +func (s *GraphService) edgeResolver() resolveBuildEdgesFn { + if s.resolveEdges != nil { + return s.resolveEdges + } + return edgeresolve.ResolveWithOptions } // BuildOptions configures one graph build run. diff --git a/internal/service/indexer_test.go b/internal/service/indexer_test.go index ce85c8b..269da81 100644 --- a/internal/service/indexer_test.go +++ b/internal/service/indexer_test.go @@ -498,22 +498,20 @@ func TestFlushBuildEdges_ResolvesAndUpsertsBoundedBatches(t *testing.T) { } var resolveSizes []int - oldResolve := resolveBuildEdges - resolveBuildEdges = func(ctx context.Context, lookup edgeresolve.NodeLookup, edges []model.Edge, options edgeresolve.ResolveOptions) ([]model.Edge, error) { + resolver := func(ctx context.Context, lookup edgeresolve.NodeLookup, edges []model.Edge, options edgeresolve.ResolveOptions) ([]model.Edge, error) { resolveSizes = append(resolveSizes, len(edges)) if len(edges) > buildEdgeResolveChunkSize { t.Fatalf("resolve batch exceeded limit: got %d want <= %d", len(edges), buildEdgeResolveChunkSize) } - return oldResolve(ctx, lookup, edges, options) + return edgeresolve.ResolveWithOptions(ctx, lookup, edges, options) } - t.Cleanup(func() { resolveBuildEdges = oldResolve }) batches := []parsedBuildEdgeBatch{ {relPath: "cmd/main.go", edges: []model.Edge{{Kind: model.EdgeKindImportsFrom, FilePath: "cmd/main.go", Line: 1, Fingerprint: "imports_from:cmd/main.go:github.com/example/project/mcp:1"}, {Kind: model.EdgeKindCalls, FilePath: "cmd/main.go", Line: 2, Fingerprint: "calls:cmd/main.go:h.deps.FlowTracer.TraceFlow:2"}}}, {relPath: "flows/tracer.go", edges: []model.Edge{{Kind: model.EdgeKindImplements, FilePath: "flows/tracer.go", Line: 7, Fingerprint: "implements:flows/tracer.go:flows.Tracer:mcp.FlowTracer"}}}, } - svc := &GraphService{} + svc := &GraphService{resolveEdges: resolver} if err := svc.flushBuildEdges(ctx, st, batches, nil, edgeresolve.ResolveOptions{}); err != nil { t.Fatalf("flushBuildEdges: %v", err) } @@ -543,8 +541,7 @@ func TestFlushBuildEdges_ResolvesImplementsOnlyOnce(t *testing.T) { st.nodesByFP["flows/tracer.go"] = []model.Node{{ID: 30, QualifiedName: "flows/tracer.go", Name: "flows/tracer.go", Kind: model.NodeKindFile, FilePath: "flows/tracer.go", Language: "go"}, {ID: 3, QualifiedName: "flows.Tracer", Name: "Tracer", Kind: model.NodeKindClass, FilePath: "flows/tracer.go", StartLine: 7, EndLine: 7, Language: "go"}, {ID: 4, QualifiedName: "flows.Tracer.TraceFlow", Name: "TraceFlow", Kind: model.NodeKindFunction, FilePath: "flows/tracer.go", StartLine: 9, EndLine: 11, Language: "go"}} var implementsSeen []int - oldResolve := resolveBuildEdges - resolveBuildEdges = func(ctx context.Context, lookup edgeresolve.NodeLookup, edges []model.Edge, options edgeresolve.ResolveOptions) ([]model.Edge, error) { + resolver := func(ctx context.Context, lookup edgeresolve.NodeLookup, edges []model.Edge, options edgeresolve.ResolveOptions) ([]model.Edge, error) { count := 0 for _, edge := range edges { if edge.Kind == model.EdgeKindImplements { @@ -552,9 +549,8 @@ func TestFlushBuildEdges_ResolvesImplementsOnlyOnce(t *testing.T) { } } implementsSeen = append(implementsSeen, count) - return oldResolve(ctx, lookup, edges, options) + return edgeresolve.ResolveWithOptions(ctx, lookup, edges, options) } - t.Cleanup(func() { resolveBuildEdges = oldResolve }) batches := []parsedBuildEdgeBatch{ {relPath: "flows/tracer.go", edges: []model.Edge{{Kind: model.EdgeKindImplements, FilePath: "flows/tracer.go", Line: 7, Fingerprint: "implements:flows/tracer.go:flows.Tracer:mcp.FlowTracer"}}}, @@ -562,7 +558,7 @@ func TestFlushBuildEdges_ResolvesImplementsOnlyOnce(t *testing.T) { {relPath: "cmd/main.go", edges: []model.Edge{{Kind: model.EdgeKindContains, FilePath: "cmd/main.go", Line: 1, Fingerprint: "contains:cmd/main.go:main.Run"}}}, } - svc := &GraphService{} + svc := &GraphService{resolveEdges: resolver} if err := svc.flushBuildEdges(ctx, st, batches, nil, edgeresolve.ResolveOptions{}); err != nil { t.Fatalf("flushBuildEdges: %v", err) } @@ -2100,8 +2096,7 @@ func TestBuild_ReleasesBatchCommentStateAfterBinding(t *testing.T) { tsCommentsNil bool sourceNil bool } - prevHook := testBuildBatchReleaseHook - testBuildBatchReleaseHook = func(batches []parsedBuildNodeBatch, idx int) { + recordRelease := func(batches []parsedBuildNodeBatch, idx int) { snapshots = append(snapshots, struct { batch int tsCommentsNil bool @@ -2112,7 +2107,6 @@ func TestBuild_ReleasesBatchCommentStateAfterBinding(t *testing.T) { sourceNil: batches[idx].sourceLines == nil, }) } - defer func() { testBuildBatchReleaseHook = prevHook }() fakeStore := newRecordingGraphStore(t) svc := &GraphService{ @@ -2120,7 +2114,8 @@ func TestBuild_ReleasesBatchCommentStateAfterBinding(t *testing.T) { Walkers: map[string]*treesitter.Walker{ ".go": treesitter.NewWalker(treesitter.GoSpec), }, - Logger: slog.Default(), + Logger: slog.Default(), + onBatchRelease: recordRelease, } dir := t.TempDir()