Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
26 changes: 10 additions & 16 deletions internal/mcp/deps.go
Original file line number Diff line number Diff line change
Expand Up @@ -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)
}

Expand Down Expand Up @@ -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)
}
64 changes: 14 additions & 50 deletions internal/mcp/handler_analysis.go
Original file line number Diff line number Diff line change
Expand Up @@ -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))...)

Expand Down Expand Up @@ -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))...)

Expand Down
14 changes: 10 additions & 4 deletions internal/mcp/handler_parse.go
Original file line number Diff line number Diff line change
Expand Up @@ -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 {
Expand Down Expand Up @@ -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")
Expand All @@ -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")
Expand Down Expand Up @@ -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")
Expand Down
8 changes: 2 additions & 6 deletions internal/mcp/handlers_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -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")
}

Expand Down Expand Up @@ -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")
}

Expand Down
17 changes: 0 additions & 17 deletions internal/mcp/testmocks_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down
8 changes: 4 additions & 4 deletions internal/service/build.go
Original file line number Diff line number Diff line change
Expand Up @@ -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
}
Expand Down Expand Up @@ -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")
Expand All @@ -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)
Expand Down
21 changes: 15 additions & 6 deletions internal/service/indexer.go
Original file line number Diff line number Diff line change
Expand Up @@ -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 (
Expand Down Expand Up @@ -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.
Expand Down
23 changes: 9 additions & 14 deletions internal/service/indexer_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -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)
}
Expand Down Expand Up @@ -543,26 +541,24 @@ 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 {
count++
}
}
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"}}},
{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: "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)
}
Expand Down Expand Up @@ -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
Expand All @@ -2112,15 +2107,15 @@ func TestBuild_ReleasesBatchCommentStateAfterBinding(t *testing.T) {
sourceNil: batches[idx].sourceLines == nil,
})
}
defer func() { testBuildBatchReleaseHook = prevHook }()

fakeStore := newRecordingGraphStore(t)
svc := &GraphService{
Store: fakeStore,
Walkers: map[string]*treesitter.Walker{
".go": treesitter.NewWalker(treesitter.GoSpec),
},
Logger: slog.Default(),
Logger: slog.Default(),
onBatchRelease: recordRelease,
}

dir := t.TempDir()
Expand Down
Loading