diff --git a/go.mod b/go.mod index 1ebfcf2c1..f91993d21 100644 --- a/go.mod +++ b/go.mod @@ -27,6 +27,7 @@ require ( github.com/jstemmer/go-junit-report/v2 v2.1.0 github.com/karrick/godirwalk v1.17.0 github.com/manifoldco/promptui v0.9.0 + github.com/modelcontextprotocol/go-sdk v1.7.0 github.com/peterebden/go-cli-init/v5 v5.2.1 github.com/peterebden/go-deferred-regex v1.1.0 github.com/peterebden/go-sri v1.1.1 @@ -77,6 +78,7 @@ require ( github.com/go-ole/go-ole v1.3.0 // indirect github.com/golang/glog v1.2.5 // indirect github.com/google/go-containerregistry v0.21.7 // indirect + github.com/google/jsonschema-go v0.4.3 // indirect github.com/google/s2a-go v0.1.9 // indirect github.com/googleapis/enterprise-certificate-proxy v0.3.18 // indirect github.com/googleapis/gax-go/v2 v2.23.0 // indirect @@ -93,10 +95,13 @@ require ( github.com/prometheus/client_model v0.6.2 // indirect github.com/prometheus/procfs v0.21.1 // indirect github.com/secure-systems-lab/go-securesystemslib v0.11.0 // indirect + github.com/segmentio/asm v1.1.3 // indirect + github.com/segmentio/encoding v0.5.4 // indirect github.com/shoenig/go-m1cpu v0.2.2 // indirect github.com/sigstore/protobuf-specs v0.5.1 // indirect github.com/tklauser/go-sysconf v0.4.0 // indirect github.com/tklauser/numcpus v0.12.0 // indirect + github.com/yosida95/uritemplate/v3 v3.0.2 // indirect github.com/youmark/pkcs8 v0.0.0-20240726163527-a2c0da244d78 // indirect github.com/yusufpapurcu/wmi v1.2.4 // indirect go.opentelemetry.io/auto/sdk v1.2.1 // indirect diff --git a/go.sum b/go.sum index 873a473c6..73b93c149 100644 --- a/go.sum +++ b/go.sum @@ -82,6 +82,8 @@ github.com/go-ole/go-ole v1.3.0 h1:Dt6ye7+vXGIKZ7Xtk4s6/xVdGDQynvom7xCFEdWr6uE= github.com/go-ole/go-ole v1.3.0/go.mod h1:5LS6F96DhAwUc7C+1HLexzMXY1xGRSryjyPPKW6zv78= github.com/go-stack/stack v1.8.0/go.mod h1:v0f6uXyyMGvRgIKkXu+yp6POWl0qKG85gN/melR3HDY= github.com/gogo/protobuf v1.3.2/go.mod h1:P1XiOD3dCwIKUDQYPy72D8LYyHL2YPYrpS2s69NZV8Q= +github.com/golang-jwt/jwt/v5 v5.3.1 h1:kYf81DTWFe7t+1VvL7eS+jKFVWaUnK9cB1qbwn63YCY= +github.com/golang-jwt/jwt/v5 v5.3.1/go.mod h1:fxCRLWMO43lRc8nhHWY6LGqRcf+1gQWArsqaEUEa5bE= github.com/golang/glog v0.0.0-20160126235308-23def4e6c14b/go.mod h1:SBH7ygxi8pfUlaOkMMuAQtPIUF8ecWP5IEl/CR7VP2Q= github.com/golang/glog v1.2.5 h1:DrW6hGnjIhtvhOIiAKT6Psh/Kd/ldepEa81DKeiRJ5I= github.com/golang/glog v1.2.5/go.mod h1:6AhwSGph0fcJtXVM/PEHPqZlFeoLxhs7/t5UDAwmO+w= @@ -109,6 +111,8 @@ github.com/google/go-cmp v0.7.0 h1:wk8382ETsv4JYUZwIsn6YpYiWiBsYLSJiTsyBybVuN8= github.com/google/go-cmp v0.7.0/go.mod h1:pXiqmnSA92OHEEa9HXL2W4E7lf9JzCmGVUdgjX3N/iU= github.com/google/go-containerregistry v0.21.7 h1:/vPFuVXDjtFREsVArW+0h1CIl5urnOhzei4X2DMW9IU= github.com/google/go-containerregistry v0.21.7/go.mod h1:kjSbt7/zMsKLWfnHrIvKvhXHUw91jbe9DNjPPJ32gXE= +github.com/google/jsonschema-go v0.4.3 h1:/DBOLZTfDow7pe2GmaJNhltueGTtDKICi8V8p+DQPd0= +github.com/google/jsonschema-go v0.4.3/go.mod h1:r5quNTdLOYEz95Ru18zA0ydNbBuYoo9tgaYcxEYhJVE= github.com/google/s2a-go v0.1.9 h1:LGD7gtMgezd8a/Xak7mEWL0PjoTQFvpRudN895yqKW0= github.com/google/s2a-go v0.1.9/go.mod h1:YA0Ei2ZQL3acow2O62kdp9UlnvMmU7kA6Eutn0dXayM= github.com/google/shlex v0.0.0-20191202100458-e7afc7fbc510 h1:El6M4kTTCOh6aBiKaUGG7oYTSPP8MxqL4YI3kZKwcP4= @@ -164,6 +168,8 @@ github.com/mattn/go-colorable v0.1.13 h1:fFA4WZxdEF4tXPZVKMLwD8oUnCTTo08duU7wxec github.com/mattn/go-colorable v0.1.13/go.mod h1:7S9/ev0klgBDR4GtXTXX8a3vIGJpMovkB8vQcUbaXHg= github.com/mattn/go-isatty v0.0.20 h1:xfD0iDuEKnDkl03q4limB+vH+GxLEtL/jb4xVJSWWEY= github.com/mattn/go-isatty v0.0.20/go.mod h1:W+V8PltTTMOvKvAeJH7IuucS94S2C6jfK/D7dTCTo3Y= +github.com/modelcontextprotocol/go-sdk v1.7.0 h1:yqjY2dsbKAC0LSuWZVBMrHgiG8ukXv6NRo0JiALay44= +github.com/modelcontextprotocol/go-sdk v1.7.0/go.mod h1:dL7u98E/zjJTGzEq+j30jQ8K2k1mb6LeAH4inEcSGts= github.com/munnerz/goautoneg v0.0.0-20191010083416-a7dc8b61c822 h1:C3w9PqII01/Oq1c1nUAm88MOHcQC9l5mIlSMApZMrHA= github.com/munnerz/goautoneg v0.0.0-20191010083416-a7dc8b61c822/go.mod h1:+n7T8mK8HuQTcFwEeznm/DIxMOiR9yIdICNftLE1DvQ= github.com/opencontainers/go-digest v1.0.0 h1:apOUWs51W5PlhuyGyz9FCeeBIOUDA/6nW8Oi/yOhh5U= @@ -206,6 +212,10 @@ github.com/rogpeppe/go-internal v1.14.1 h1:UQB4HGPB6osV0SQTLymcB4TgvyWu6ZyliaW0t github.com/rogpeppe/go-internal v1.14.1/go.mod h1:MaRKkUm5W0goXpeCfT7UZI6fk/L7L7so1lCWt35ZSgc= github.com/secure-systems-lab/go-securesystemslib v0.11.0 h1:iuCR9kcMFD4QurdKrGvPLoKZLv9YvwPYVr0473BdtFs= github.com/secure-systems-lab/go-securesystemslib v0.11.0/go.mod h1:+PMOTjUGwHj2vcZ+TFKlb1tXRbrdWE1LYDT5i9JC80Q= +github.com/segmentio/asm v1.1.3 h1:WM03sfUOENvvKexOLp+pCqgb/WDjsi7EK8gIsICtzhc= +github.com/segmentio/asm v1.1.3/go.mod h1:Ld3L4ZXGNcSLRg4JBsZ3//1+f/TjYl0Mzen/DQy1EJg= +github.com/segmentio/encoding v0.5.4 h1:OW1VRern8Nw6ITAtwSZ7Idrl3MXCFwXHPgqESYfvNt0= +github.com/segmentio/encoding v0.5.4/go.mod h1:HS1ZKa3kSN32ZHVZ7ZLPLXWvOVIiZtyJnO1gPH1sKt0= github.com/shirou/gopsutil/v3 v3.24.5 h1:i0t8kL+kQTvpAYToeuiVk3TgDeKOFioZO3Ztz/iZ9pI= github.com/shirou/gopsutil/v3 v3.24.5/go.mod h1:bsoOS1aStSs9ErQ1WWfxllSeS1K5D+U30r2NfcubMVk= github.com/shoenig/go-m1cpu v0.2.2 h1:4nc55oVv7nygGnfI9bhLCLzUEs4794y0Bkqx4q2zy7Y= @@ -243,6 +253,8 @@ github.com/tklauser/numcpus v0.12.0 h1:NR85qdvHA9pFse3x3weVZ0r0ST8R6l5RHbZrlRaqo github.com/tklauser/numcpus v0.12.0/go.mod h1:ABHeXzJnr/qqwguhClkZKT1/8VABcYrsyUiUGobwWJg= github.com/ulikunitz/xz v0.5.15 h1:9DNdB5s+SgV3bQ2ApL10xRc35ck0DuIX/isZvIk+ubY= github.com/ulikunitz/xz v0.5.15/go.mod h1:nbz6k7qbPmH4IRqmfOplQw/tblSgqTqBwxkY0oWt/14= +github.com/yosida95/uritemplate/v3 v3.0.2 h1:Ed3Oyj9yrmi9087+NczuL5BwkIc4wvTb5zIM+UJPGz4= +github.com/yosida95/uritemplate/v3 v3.0.2/go.mod h1:ILOh0sOhIJR3+L/8afwt/kE++YT040gmv5BQTMR2HP4= github.com/youmark/pkcs8 v0.0.0-20240726163527-a2c0da244d78 h1:ilQV1hzziu+LLM3zUTJ0trRztfwgjqKnBWNtSRkbmwM= github.com/youmark/pkcs8 v0.0.0-20240726163527-a2c0da244d78/go.mod h1:aL8wCCfTfSfmXjznFBSZNN13rSJjlIOI1fUNAtF7rmI= github.com/yuin/goldmark v1.1.27/go.mod h1:3hX8gzYuyVAZsxl0MRgGTJEmQBFcNTphYh9decYSb74= diff --git a/src/BUILD.plz b/src/BUILD.plz index 4ef5ea6fd..caa400f0e 100644 --- a/src/BUILD.plz +++ b/src/BUILD.plz @@ -27,6 +27,7 @@ go_binary( "//src/generate", "//src/hashes", "//src/help", + "//src/mcp", "//src/metrics", "//src/output", "//src/plz", diff --git a/src/mcp/BUILD b/src/mcp/BUILD new file mode 100644 index 000000000..d4715d0bf --- /dev/null +++ b/src/mcp/BUILD @@ -0,0 +1,32 @@ +go_library( + name = "mcp", + srcs = [ + "mcp.go", + "tools.go", + ], + pgo_file = "//:pgo", + visibility = ["PUBLIC"], + deps = [ + "///third_party/go/github.com_modelcontextprotocol_go-sdk//mcp", + "//src/cli/logging", + "//src/core", + "//src/parse", + "//src/plz", + "//src/query", + "//src/version", + ], +) + +go_test( + name = "mcp_test", + srcs = ["mcp_test.go"], + external = True, + deps = [ + ":mcp", + "///third_party/go/github.com_modelcontextprotocol_go-sdk//mcp", + "///third_party/go/github.com_stretchr_testify//assert", + "///third_party/go/github.com_stretchr_testify//require", + "//src/core", + "//src/fs", + ], +) diff --git a/src/mcp/mcp.go b/src/mcp/mcp.go new file mode 100644 index 000000000..52119dc90 --- /dev/null +++ b/src/mcp/mcp.go @@ -0,0 +1,205 @@ +// Package mcp implements a Model Context Protocol server that exposes plz query +// functionality over stdio. The build graph is parsed once at startup and kept +// in memory between queries, so clients don't pay the graph construction cost +// on every query as they would invoking plz directly. +package mcp + +import ( + "context" + "fmt" + "io" + "os" + "strings" + "sync" + + sdk "github.com/modelcontextprotocol/go-sdk/mcp" + + "github.com/thought-machine/please/src/cli/logging" + "github.com/thought-machine/please/src/core" + "github.com/thought-machine/please/src/parse" + "github.com/thought-machine/please/src/plz" + "github.com/thought-machine/please/src/version" +) + +var log = logging.Log + +// A Server holds the cached build state that queries are served from. +type Server struct { + // mu serialises queries and reloads; reloads swap out the state, and some + // queries (filter) temporarily mutate it. + mu sync.Mutex + state *core.BuildState + graph *core.BuildGraph + config *core.Configuration + + transport sdk.Transport +} + +// Option provides a mechanism to set options on the MCP server. +type Option func(*Server) + +// WithTransport overrides the transport that the server receives requests on and +// sends responses over. The default is stdio. Tests can pass one half of +// sdk.NewInMemoryTransports() and drive the server with a real MCP client. +func WithTransport(t sdk.Transport) Option { + return func(s *Server) { + s.transport = t + } +} + +// WithState supplies a pre-parsed build state to serve queries from, in place of +// parsing the build graph at startup. +func WithState(state *core.BuildState) Option { + return func(s *Server) { + s.state = state + } +} + +// NewServer instantiates a new MCP server. +func NewServer(config *core.Configuration, opts ...Option) *Server { + s := &Server{ + config: config, + graph: core.NewGraph(), + } + + for _, opt := range opts { + opt(s) + } + + return s +} + +// Serve runs an MCP server until the client disconnects or ctx is cancelled. +// By default, the server uses lazy loading to parse the build graph on demand. +func (s *Server) Serve(ctx context.Context) error { + if s.transport == nil { + s.transport = stdioTransport() + } + if s.state == nil { + log.Notice("Serving MCP with lazy-loaded build graph...") + } else { + log.Notice("Serving MCP for %d targets", len(s.state.Graph.AllTargets())) + } + + srv := sdk.NewServer(&sdk.Implementation{ + Name: "please", + Title: "Please build system", + Version: version.PleaseVersion, + }, nil) + s.registerTools(srv) + return srv.Run(ctx, s.transport) +} + +// stdioTransport returns a transport communicating over stdin and stdout. +// The MCP protocol runs over stdout, so anything else that writes there would +// corrupt the framing. Point os.Stdout at stderr for the life of the process +// and hand the real stdout to the transport; queries capture os.Stdout per-call. +// Stdout is wrapped so that the transport doesn't close it when the session ends. +func stdioTransport() sdk.Transport { + out := os.Stdout + os.Stdout = os.Stderr + return &sdk.IOTransport{Reader: os.Stdin, Writer: nopCloserWriter{out}} +} + +// nopCloserWriter is an io.WriteCloser with a trivial Close method. +type nopCloserWriter struct { + io.Writer +} + +func (nopCloserWriter) Close() error { return nil } + +// parseGraph parses the entire build graph into a fresh build state. +// On success the new state replaces the current one; on failure the old state is kept. +// Callers must hold s.mu (except before the server has started). +func (s *Server) parseGraph() error { + state := core.NewBuildState(s.config) + state.NeedBuild = false + parse.InitParser(state) + plz.RunHost(core.WholeGraph, state) + if failed, _, _ := state.Failures(); failed { + return fmt.Errorf("failed to parse the build graph; see server logs for details") + } + s.state = state + return nil +} + +// withState runs f against a build state under the server lock, converting panics +// into errors so a misbehaving query can't kill the server. +// Under lazy-loading, it parses the given targets on demand into the persistent graph. +func (s *Server) withState(targets []string, f func(state *core.BuildState) error) (err error) { + s.mu.Lock() + defer s.mu.Unlock() + defer func() { + if p := recover(); p != nil { + err = fmt.Errorf("query failed: %v", p) + } + }() + + if s.state != nil { + return f(s.state) + } + + state := core.NewBuildState(s.config) + state.NeedBuild = false + state.Graph = s.graph + + if len(targets) == 0 { + return f(state) + } + + labels := make([]core.BuildLabel, 0, len(targets)) + for _, t := range targets { + l, err := core.TryParseBuildLabel(t, "", "") + if err != nil { + return fmt.Errorf("invalid build label %q: %w", t, err) + } + labels = append(labels, l) + } + plz.RunHost(labels, state) + if failed, _, _ := state.Failures(); failed { + var errs []string + for r := range state.Results() { + if r.Status.IsFailure() { + errs = append(errs, fmt.Sprintf("%s (%s): %s", r.Label, r.Status, r.Err)) + } + } + return fmt.Errorf("failed to parse the build graph: %s", strings.Join(errs, "; ")) + } + + return f(state) +} + +// resolveLabels parses a set of label strings, expands pseudo-targets (:all and /...) +// against the graph and verifies that every resulting target exists. +// Verification matters: the query functions call TargetOrDie on the labels they're +// given, which would kill the server on an unknown target. +func resolveLabels(state *core.BuildState, in []string) ([]core.BuildLabel, error) { + labels := make([]core.BuildLabel, 0, len(in)) + for _, l := range in { + label, err := core.TryParseBuildLabel(l, "", "") + if err != nil { + return nil, fmt.Errorf("invalid build label %q: %w", l, err) + } + labels = append(labels, label) + } + expanded := state.ExpandLabels(labels) + for _, l := range expanded { + if state.Graph.Target(l) == nil { + return nil, fmt.Errorf("target %s not found in the build graph", l) + } + } + if len(expanded) == 0 { + return nil, fmt.Errorf("no targets matched the given labels") + } + return expanded, nil +} + +// textResult wraps a string as an MCP tool result. +func textResult(text string) *sdk.CallToolResult { + if text == "" { + text = "(no output)" + } + return &sdk.CallToolResult{ + Content: []sdk.Content{&sdk.TextContent{Text: text}}, + } +} diff --git a/src/mcp/mcp_test.go b/src/mcp/mcp_test.go new file mode 100644 index 000000000..457b99f6e --- /dev/null +++ b/src/mcp/mcp_test.go @@ -0,0 +1,389 @@ +package mcp_test + +import ( + "context" + "encoding/json" + "os" + "testing" + "time" + + sdk "github.com/modelcontextprotocol/go-sdk/mcp" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + "github.com/thought-machine/please/src/core" + "github.com/thought-machine/please/src/fs" + "github.com/thought-machine/please/src/mcp" +) + +// targetsResult mirrors the structured result the target-returning tools produce. +type targetsResult struct { + Targets map[string]map[string]any `json:"targets"` +} + +// filesResult mirrors the structured result the file-mapping tools produce. +type filesResult struct { + Files map[string][]string `json:"files"` + Targets map[string]map[string]any `json:"targets"` +} + +func TestListTools(t *testing.T) { + a := assert.New(t) + r := require.New(t) + + session := newTestSession(t, testState()) + res, err := session.ListTools(t.Context(), nil) + r.NoError(err) + + names := make([]string, len(res.Tools)) + for i, tool := range res.Tools { + names[i] = tool.Name + } + a.ElementsMatch([]string{ + "deps", "revdeps", "somepath", "print", "alltargets", "filter", + "whatinputs", "whatoutputs", "inputs", "outputs", "reload_graph", + }, names) +} + +func TestTargetQueries(t *testing.T) { + tests := []struct { + name string + tool string + args map[string]any + expectedTargets []string + }{ + { + name: "print", + tool: "print", + args: map[string]any{"targets": []string{"//package1:target1"}}, + expectedTargets: []string{"//package1:target1"}, + }, + { + // deps reports the dependencies of the queried targets, not the targets themselves. + name: "deps", + tool: "deps", + args: map[string]any{"targets": []string{"//package1:target1"}}, + expectedTargets: []string{"//package2:target2"}, + }, + { + name: "revdeps", + tool: "revdeps", + args: map[string]any{"targets": []string{"//package2:target2"}}, + expectedTargets: []string{"//package1:target1"}, + }, + { + name: "alltargets", + tool: "alltargets", + args: map[string]any{}, + expectedTargets: []string{"//package1:target1", "//package2:target2"}, + }, + { + name: "pseudo-target expansion", + tool: "print", + args: map[string]any{"targets": []string{"//package1:all"}}, + expectedTargets: []string{"//package1:target1"}, + }, + } + + for _, test := range tests { + t.Run(test.name, func(t *testing.T) { + a := assert.New(t) + r := require.New(t) + + session := newTestSession(t, testState()) + res, err := session.CallTool(t.Context(), &sdk.CallToolParams{ + Name: test.tool, + Arguments: test.args, + }) + r.NoError(err) + r.False(res.IsError, "tool returned an error: %s", contentText(res)) + + var out targetsResult + decodeStructured(t, res, &out) + a.ElementsMatch(test.expectedTargets, keys(out.Targets)) + }) + } +} + +func TestFieldsAreRestricted(t *testing.T) { + a := assert.New(t) + r := require.New(t) + + session := newTestSession(t, testState()) + res, err := session.CallTool(t.Context(), &sdk.CallToolParams{ + Name: "print", + Arguments: map[string]any{ + "targets": []string{"//package1:target1"}, + "fields": []string{"deps"}, + }, + }) + r.NoError(err) + r.False(res.IsError, "tool returned an error: %s", contentText(res)) + + var out targetsResult + decodeStructured(t, res, &out) + r.Contains(out.Targets, "//package1:target1") + a.Equal([]string{"deps"}, keys(out.Targets["//package1:target1"])) +} + +func TestWhatInputs(t *testing.T) { + a := assert.New(t) + r := require.New(t) + + session := newTestSession(t, testState()) + res, err := session.CallTool(t.Context(), &sdk.CallToolParams{ + Name: "whatinputs", + Arguments: map[string]any{"files": []string{"package1/file1.txt"}}, + }) + r.NoError(err) + r.False(res.IsError, "tool returned an error: %s", contentText(res)) + + var out filesResult + decodeStructured(t, res, &out) + a.Equal(map[string][]string{ + "package1/file1.txt": {"//package1:target1"}, + }, out.Files) + a.ElementsMatch([]string{"//package1:target1"}, keys(out.Targets)) +} + +// The text-returning tools write through an io.Writer rather than stdout, so their +// output arrives as text content. +func TestTextQueries(t *testing.T) { + tests := []struct { + name string + tool string + args map[string]any + expectedText string + }{ + { + name: "inputs", + tool: "inputs", + args: map[string]any{"targets": []string{"//package1:target1"}}, + expectedText: "package1/file1.txt\n", + }, + { + name: "outputs", + tool: "outputs", + args: map[string]any{"targets": []string{"//package1:target1"}}, + expectedText: "plz-out/gen/package1/out1.txt\n", + }, + { + name: "outputs as JSON", + tool: "outputs", + args: map[string]any{ + "targets": []string{"//package1:target1"}, + "json": true, + }, + expectedText: "{\n \"//package1:target1\": [\n \"plz-out/gen/package1/out1.txt\"\n ]\n}\n", + }, + } + + for _, test := range tests { + t.Run(test.name, func(t *testing.T) { + a := assert.New(t) + r := require.New(t) + + session := newTestSession(t, testState()) + res, err := session.CallTool(t.Context(), &sdk.CallToolParams{ + Name: test.tool, + Arguments: test.args, + }) + r.NoError(err) + r.False(res.IsError, "tool returned an error: %s", contentText(res)) + a.Equal(test.expectedText, contentText(res)) + }) + } +} + +// Unknown targets must come back as tool errors; the underlying query functions +// call TargetOrDie, which would otherwise take the server down with it. +func TestErrorsAreReportedNotFatal(t *testing.T) { + tests := []struct { + name string + tool string + args map[string]any + expectedMessage string + }{ + { + name: "unknown target", + tool: "print", + args: map[string]any{"targets": []string{"//package1:nope"}}, + expectedMessage: "not found in the build graph", + }, + { + name: "invalid label", + tool: "print", + args: map[string]any{"targets": []string{"not a label"}}, + expectedMessage: "invalid build label", + }, + { + name: "unknown field", + tool: "print", + args: map[string]any{ + "targets": []string{"//package1:target1"}, + "fields": []string{"nonsense"}, + }, + expectedMessage: "unknown field nonsense", + }, + } + + for _, test := range tests { + t.Run(test.name, func(t *testing.T) { + a := assert.New(t) + r := require.New(t) + + session := newTestSession(t, testState()) + res, err := session.CallTool(t.Context(), &sdk.CallToolParams{ + Name: test.tool, + Arguments: test.args, + }) + r.NoError(err) + a.True(res.IsError) + a.Contains(contentText(res), test.expectedMessage) + + // The session is still usable afterwards. + _, err = session.ListTools(t.Context(), nil) + a.NoError(err) + }) + } +} + +func TestLazyLoading(t *testing.T) { + a := assert.New(t) + r := require.New(t) + + ctx, cancel := context.WithTimeout(t.Context(), 5*time.Second) + t.Cleanup(cancel) + + // Create a temporary mock package directory and BUILD file to test parsing on-demand + // without requiring any external plugins/compilers in the test sandbox. + err := os.MkdirAll("mockpkg", 0755) + r.NoError(err) + t.Cleanup(func() { os.RemoveAll("mockpkg") }) + + err = os.WriteFile("mockpkg/BUILD", []byte(`filegroup(name = "foo", srcs = ["bar.txt"])`), 0644) + r.NoError(err) + err = os.WriteFile("mockpkg/bar.txt", []byte("hello"), 0644) + r.NoError(err) + + config, err := core.ReadConfigFiles(fs.HostFS, []string{".plzconfig"}, nil) + r.NoError(err) + + serverTransport, clientTransport := sdk.NewInMemoryTransports() + srv := mcp.NewServer( + config, + mcp.WithTransport(serverTransport), + ) + done := make(chan struct{}) + go func() { + defer close(done) + srv.Serve(ctx) + }() + t.Cleanup(func() { + cancel() + <-done + }) + + client := sdk.NewClient(&sdk.Implementation{Name: "test", Version: "1"}, nil) + session, err := client.Connect(ctx, clientTransport, nil) + r.NoError(err) + t.Cleanup(func() { session.Close() }) + + // Since we are lazy-loading, calling "print" with an actual target (like "//mockpkg:foo") + // should parse the target on the fly and succeed. + res, err := session.CallTool(ctx, &sdk.CallToolParams{ + Name: "print", + Arguments: map[string]any{ + "targets": []string{"//mockpkg:foo"}, + }, + }) + r.NoError(err) + r.False(res.IsError, "tool returned an error: %s", contentText(res)) + + var out targetsResult + decodeStructured(t, res, &out) + a.Contains(keys(out.Targets), "//mockpkg:foo") +} + +// newTestSession starts a server serving the given state over an in-memory +// transport and returns a client session connected to it. +func newTestSession(t *testing.T, state *core.BuildState) *sdk.ClientSession { + t.Helper() + r := require.New(t) + + ctx, cancel := context.WithCancel(t.Context()) + t.Cleanup(cancel) + + serverTransport, clientTransport := sdk.NewInMemoryTransports() + srv := mcp.NewServer( + state.Config, + mcp.WithTransport(serverTransport), + mcp.WithState(state), + ) + done := make(chan struct{}) + go func() { + defer close(done) + srv.Serve(ctx) + }() + t.Cleanup(func() { + cancel() + <-done + }) + + client := sdk.NewClient(&sdk.Implementation{Name: "test", Version: "1"}, nil) + session, err := client.Connect(ctx, clientTransport, nil) + r.NoError(err) + t.Cleanup(func() { session.Close() }) + return session +} + +// testState returns a build state with a small hand-built graph in it. +func testState() *core.BuildState { + state := core.NewDefaultBuildState() + pkg1 := core.NewPackage("package1") + pkg2 := core.NewPackage("package2") + + target2 := core.NewBuildTarget(core.NewBuildLabel(pkg2.Name, "target2")) + target1 := core.NewBuildTarget(core.NewBuildLabel(pkg1.Name, "target1")) + target1.AddSource(core.FileLabel{File: "file1.txt", Package: pkg1.Name}) + target1.AddDependency(target2.Label) + target1.AddOutput("out1.txt") + + for _, target := range []*core.BuildTarget{target1, target2} { + state.Graph.AddTarget(target) + } + pkg1.AddTarget(target1) + pkg2.AddTarget(target2) + state.Graph.AddPackage(pkg1) + state.Graph.AddPackage(pkg2) + return state +} + +// decodeStructured decodes the structured content of a tool result into out. +func decodeStructured(t *testing.T, res *sdk.CallToolResult, out any) { + t.Helper() + r := require.New(t) + + b, err := json.Marshal(res.StructuredContent) + r.NoError(err) + r.NoError(json.Unmarshal(b, out)) +} + +// contentText returns the concatenated text content of a tool result. +func contentText(res *sdk.CallToolResult) string { + text := "" + for _, content := range res.Content { + if tc, ok := content.(*sdk.TextContent); ok { + text += tc.Text + } + } + return text +} + +func keys[V any](m map[string]V) []string { + ret := make([]string, 0, len(m)) + for k := range m { + ret = append(ret, k) + } + return ret +} diff --git a/src/mcp/tools.go b/src/mcp/tools.go new file mode 100644 index 000000000..23b83915d --- /dev/null +++ b/src/mcp/tools.go @@ -0,0 +1,478 @@ +package mcp + +import ( + "bytes" + "context" + "fmt" + "reflect" + "sort" + "strings" + + sdk "github.com/modelcontextprotocol/go-sdk/mcp" + + "github.com/thought-machine/please/src/core" + "github.com/thought-machine/please/src/parse" + "github.com/thought-machine/please/src/query" +) + +// A targetsResult carries structured definitions of a set of targets, in the same +// format as plz query print --json. +type targetsResult struct { + Targets map[string]map[string]any `json:"targets" jsonschema:"Target definitions keyed by build label, in plz query print --json format."` +} + +// A pathResult is a targetsResult with an ordered path through the graph. +type pathResult struct { + Path []string `json:"path" jsonschema:"The dependency path, in order from the first target to the second."` + Targets map[string]map[string]any `json:"targets" jsonschema:"Definitions of the targets on the path, keyed by build label."` +} + +// A filesResult is a targetsResult with a mapping from query files to matching labels. +type filesResult struct { + Files map[string][]string `json:"files" jsonschema:"Build labels matched for each queried file. Files matching no target map to an empty list."` + Targets map[string]map[string]any `json:"targets" jsonschema:"Definitions of the matched targets, keyed by build label."` +} + +// addTool registers a tool whose handler returns plain text. +func addTool[In any]( + srv *sdk.Server, + name, description string, + h func(ctx context.Context, in In) (string, error), +) { + tool := &sdk.Tool{Name: name, Description: description} + sdk.AddTool(srv, tool, func( + ctx context.Context, + req *sdk.CallToolRequest, + in In, + ) (*sdk.CallToolResult, any, error) { + out, err := h(ctx, in) + if err != nil { + return nil, nil, err + } + return textResult(out), nil, nil + }) +} + +// addStructuredTool registers a tool whose handler returns a structured result; +// the SDK serialises it into the result's structured content. +func addStructuredTool[In, Out any]( + srv *sdk.Server, + name, description string, + h func(ctx context.Context, in In) (Out, error), +) { + tool := &sdk.Tool{Name: name, Description: description} + sdk.AddTool(srv, tool, func( + ctx context.Context, + req *sdk.CallToolRequest, + in In, + ) (*sdk.CallToolResult, Out, error) { + out, err := h(ctx, in) + return nil, out, err + }) +} + +type depsArgs struct { + Targets []string `json:"targets" jsonschema:"Build labels to query, e.g. //src/core:core. Pseudo-targets like //src/... and //src/core:all are supported."` + Hidden bool `json:"hidden,omitempty" jsonschema:"Include hidden targets (names beginning with an underscore)."` + Level int `json:"level,omitempty" jsonschema:"Maximum depth to traverse; omit or 0 for unlimited."` + Fields []string `json:"fields,omitempty" jsonschema:"Restrict the returned target definitions to these fields (e.g. srcs, deps, outs). Omit for all fields."` +} + +type revdepsArgs struct { + Targets []string `json:"targets" jsonschema:"Build labels to query, e.g. //src/core:core. Pseudo-targets like //src/... and //src/core:all are supported."` + Hidden bool `json:"hidden,omitempty" jsonschema:"Include hidden targets (names beginning with an underscore)."` + Level int `json:"level,omitempty" jsonschema:"Levels of reverse dependencies to include; -1 for the full transitive set. Omitting it or 0 means 1 level, like plz query revdeps."` + Fields []string `json:"fields,omitempty" jsonschema:"Restrict the returned target definitions to these fields (e.g. srcs, deps, outs). Omit for all fields."` +} + +type somepathArgs struct { + From string `json:"from" jsonschema:"Build label to start from."` + To string `json:"to" jsonschema:"Build label to find a path to."` + Except []string `json:"except,omitempty" jsonschema:"Build labels to exclude from the path."` + Hidden bool `json:"hidden,omitempty" jsonschema:"Include hidden targets in the path."` + Fields []string `json:"fields,omitempty" jsonschema:"Restrict the returned target definitions to these fields (e.g. srcs, deps, outs). Omit for all fields."` +} + +type printArgs struct { + Targets []string `json:"targets" jsonschema:"Build labels to print, e.g. //src/core:core."` + Fields []string `json:"fields,omitempty" jsonschema:"Restrict the returned target definitions to these fields (e.g. srcs, deps, outs). Omit for all fields."` +} + +type alltargetsArgs struct { + Targets []string `json:"targets,omitempty" jsonschema:"Packages to list targets in, e.g. //src/... . Omit to list the entire graph."` + Hidden bool `json:"hidden,omitempty" jsonschema:"Include hidden targets (names beginning with an underscore)."` + Fields []string `json:"fields,omitempty" jsonschema:"Restrict the returned target definitions to these fields (e.g. srcs, deps, outs). Omit for all fields."` +} + +type filterArgs struct { + Targets []string `json:"targets,omitempty" jsonschema:"Build labels to filter. Omit to filter the entire graph."` + Include []string `json:"include,omitempty" jsonschema:"Only include targets with at least one of these labels/tags."` + Exclude []string `json:"exclude,omitempty" jsonschema:"Exclude targets with any of these labels/tags."` + Hidden bool `json:"hidden,omitempty" jsonschema:"Include hidden targets (names beginning with an underscore)."` + Fields []string `json:"fields,omitempty" jsonschema:"Restrict the returned target definitions to these fields (e.g. srcs, deps, outs). Omit for all fields."` +} + +type labelsArgs struct { + Targets []string `json:"targets" jsonschema:"Build labels to query, e.g. //src/core:core."` +} + +type outputsArgs struct { + Targets []string `json:"targets" jsonschema:"Build labels to query, e.g. //src/core:core."` + JSON bool `json:"json,omitempty" jsonschema:"Print the outputs as JSON."` +} + +type whatinputsArgs struct { + Files []string `json:"files" jsonschema:"File paths relative to the repo root. Files that aren't an input to any target map to an empty list."` + Hidden bool `json:"hidden,omitempty" jsonschema:"Report hidden targets rather than their parent."` + Fields []string `json:"fields,omitempty" jsonschema:"Restrict the returned target definitions to these fields (e.g. srcs, deps, outs). Omit for all fields."` +} + +type whatoutputsArgs struct { + Files []string `json:"files" jsonschema:"Output file paths relative to the repo root (within plz-out)."` + Fields []string `json:"fields,omitempty" jsonschema:"Restrict the returned target definitions to these fields (e.g. srcs, deps, outs). Omit for all fields."` +} + +// registerTools registers all the query tools on the given MCP server. +func (s *Server) registerTools(srv *sdk.Server) { + s.registerGraphTools(srv) + s.registerTargetTools(srv) + s.registerFileTools(srv) + s.registerAdminTools(srv) +} + +// registerGraphTools registers the tools that walk the dependency graph. +func (s *Server) registerGraphTools(srv *sdk.Server) { + addStructuredTool(srv, "deps", + "Returns the transitive dependencies of a set of build targets, with each target's definition.", + func(ctx context.Context, in depsArgs) (targetsResult, error) { + level := in.Level + if level <= 0 { + level = -1 + } + return s.targetsQuery(in.Targets, in.Fields, func(state *core.BuildState) (core.BuildLabels, error) { + labels, err := resolveLabels(state, in.Targets) + if err != nil { + return nil, err + } + return query.DepsLabels(state, labels, in.Hidden, level), nil + }) + }) + + addStructuredTool(srv, "revdeps", + "Returns the targets that depend on a set of build targets (reverse dependencies), with each target's definition.", + func(ctx context.Context, in revdepsArgs) (targetsResult, error) { + level := in.Level + if level == 0 { + level = 1 + } + targets := append(in.Targets, "//...") + return s.targetsQuery(targets, in.Fields, func(state *core.BuildState) (core.BuildLabels, error) { + labels, err := resolveLabels(state, in.Targets) + if err != nil { + return nil, err + } + return query.ReverseDepsLabels(state, labels, level, in.Hidden), nil + }) + }) + + addStructuredTool(srv, "somepath", + "Finds a dependency path between two build targets, returning the path in order and each target's definition.", + func(ctx context.Context, in somepathArgs) (pathResult, error) { + ret := pathResult{} + if err := validatePrintFields(in.Fields); err != nil { + return ret, err + } + targets := append([]string{in.From, in.To}, in.Except...) + err := s.withState(targets, func(state *core.BuildState) error { + path, err := s.somePath(state, in) + if err != nil { + return err + } + ret.Path = labelStrings(path) + ret.Targets = targetMaps(state, path, in.Fields) + return nil + }) + return ret, err + }) +} + +// somePath resolves the labels in the given args and finds a path between them. +func (s *Server) somePath(state *core.BuildState, in somepathArgs) (core.BuildLabels, error) { + from, err := resolveLabels(state, []string{in.From}) + if err != nil { + return nil, err + } + to, err := resolveLabels(state, []string{in.To}) + if err != nil { + return nil, err + } + except := []core.BuildLabel{} + if len(in.Except) > 0 { + if except, err = resolveLabels(state, in.Except); err != nil { + return nil, err + } + } + return query.SomePathLabels(state.Graph, from, to, except, in.Hidden) +} + +// registerTargetTools registers the tools that inspect individual targets and target sets. +func (s *Server) registerTargetTools(srv *sdk.Server) { + addStructuredTool(srv, "print", + "Returns the definition of build targets as they exist in the build graph, in plz query print --json format.", + func(ctx context.Context, in printArgs) (targetsResult, error) { + return s.targetsQuery(in.Targets, in.Fields, func(state *core.BuildState) (core.BuildLabels, error) { + return resolveLabels(state, in.Targets) + }) + }) + + addStructuredTool(srv, "alltargets", + "Returns all the build targets in the graph, optionally filtered to a set of packages, with each target's definition.", + func(ctx context.Context, in alltargetsArgs) (targetsResult, error) { + targets := in.Targets + if len(targets) == 0 { + targets = []string{"//..."} + } + return s.targetsQuery(targets, in.Fields, func(state *core.BuildState) (core.BuildLabels, error) { + labels, err := resolveWholeGraphLabels(state, in.Targets) + if err != nil { + return nil, err + } + return visibleLabels(labels, in.Hidden), nil + }) + }) + + addStructuredTool(srv, "filter", + "Filters a set of targets by include/exclude labels (as passed to plz --include / --exclude), returning each match's definition.", + func(ctx context.Context, in filterArgs) (targetsResult, error) { + targets := in.Targets + if len(targets) == 0 { + targets = []string{"//..."} + } + return s.targetsQuery(targets, in.Fields, func(state *core.BuildState) (core.BuildLabels, error) { + labels, err := resolveWholeGraphLabels(state, in.Targets) + if err != nil { + return nil, err + } + state.SetIncludeAndExclude(in.Include, in.Exclude) + defer state.SetIncludeAndExclude(nil, nil) + ret := core.BuildLabels{} + for _, l := range visibleLabels(labels, in.Hidden) { + if state.ShouldInclude(state.Graph.TargetOrDie(l)) { + ret = append(ret, l) + } + } + return ret, nil + }) + }) +} + +// registerFileTools registers the tools that map between files and targets. +func (s *Server) registerFileTools(srv *sdk.Server) { + addStructuredTool(srv, "whatinputs", + "Finds the build targets that the given files are inputs (sources) to, with each target's definition.", + func(ctx context.Context, in whatinputsArgs) (filesResult, error) { + targets := make([]string, 0, len(in.Files)) + if s.state == nil { + // If we haven't generated the build graph yet, identify the + // targets associated with the files in the request. + state := core.NewBuildState(s.config) + for _, file := range in.Files { + targets = append(targets, core.FindOwningPackage(state, file).String()) + } + } + return s.filesQuery(targets, in.Fields, func(state *core.BuildState) map[string]core.BuildLabels { + return query.WhatInputsLabels(state.Graph, in.Files, in.Hidden) + }) + }) + + addStructuredTool(srv, "whatoutputs", + "Finds the build targets that produce the given output files, with each target's definition.", + func(ctx context.Context, in whatoutputsArgs) (filesResult, error) { + return s.filesQuery([]string{"//..."}, in.Fields, func(state *core.BuildState) map[string]core.BuildLabels { + return query.WhatOutputsLabels(state.Graph, in.Files) + }) + }) + + addTool(srv, "inputs", + "Lists all the input files (sources) of a set of build targets.", + func(ctx context.Context, in labelsArgs) (string, error) { + var buf bytes.Buffer + err := s.withState(in.Targets, func(state *core.BuildState) error { + labels, err := resolveLabels(state, in.Targets) + if err != nil { + return err + } + query.TargetInputs(&buf, state.Graph, labels) + return nil + }) + return buf.String(), err + }) + + addTool(srv, "outputs", + "Lists the output files of a set of build targets.", + func(ctx context.Context, in outputsArgs) (string, error) { + var buf bytes.Buffer + err := s.withState(in.Targets, func(state *core.BuildState) error { + labels, err := resolveLabels(state, in.Targets) + if err != nil { + return err + } + query.TargetOutputs(&buf, state.Graph, labels, in.JSON) + return nil + }) + return buf.String(), err + }) +} + +// registerAdminTools registers the tools that manage the server itself. +func (s *Server) registerAdminTools(srv *sdk.Server) { + addTool(srv, "reload_graph", + "Re-parses the build graph from the BUILD files on disk. Use this after BUILD files have changed; other queries are answered from a cached graph. Configuration (.plzconfig) changes require a server restart.", + func(ctx context.Context, in struct{}) (string, error) { + s.mu.Lock() + defer s.mu.Unlock() + if s.state != nil { + if err := s.parseGraph(); err != nil { + return "", err + } + return fmt.Sprintf("Build graph reloaded: %d targets in %d packages.", + len(s.state.Graph.AllTargets()), len(s.state.Graph.PackageMap())), nil + } + s.graph = core.NewGraph() + return "Build graph cache cleared. Targets will be parsed on demand.", nil + }) +} + +// targetsQuery runs f to obtain a set of labels and returns their structured definitions. +func (s *Server) targetsQuery( + targets []string, + fields []string, + f func(state *core.BuildState) (core.BuildLabels, error), +) (targetsResult, error) { + ret := targetsResult{} + if err := validatePrintFields(fields); err != nil { + return ret, err + } + err := s.withState(targets, func(state *core.BuildState) error { + labels, err := f(state) + if err != nil { + return err + } + ret.Targets = targetMaps(state, labels, fields) + return nil + }) + return ret, err +} + +// filesQuery runs f to obtain a file-to-labels mapping and returns it along with the +// structured definitions of all matched targets. +func (s *Server) filesQuery( + targets []string, + fields []string, + f func(state *core.BuildState) map[string]core.BuildLabels, +) (filesResult, error) { + ret := filesResult{} + if err := validatePrintFields(fields); err != nil { + return ret, err + } + err := s.withState(targets, func(state *core.BuildState) error { + ret.Files = map[string][]string{} + all := core.BuildLabels{} + for file, labels := range f(state) { + ret.Files[file] = labelStrings(labels) + all = append(all, labels...) + } + ret.Targets = targetMaps(state, all, fields) + return nil + }) + return ret, err +} + +// targetMaps builds the print-style JSON representation of each of the given targets. +func targetMaps( + state *core.BuildState, + labels core.BuildLabels, + fields []string, +) map[string]map[string]any { + order := parse.BuildRuleArgOrder(state) + ret := make(map[string]map[string]any, len(labels)) + for _, l := range labels { + if target := state.Graph.Target(l); target != nil { + ret[l.String()] = query.TargetToMap(order, fields, target) + } + } + return ret +} + +// labelStrings converts a set of build labels to their string forms. +func labelStrings(labels core.BuildLabels) []string { + ret := make([]string, len(labels)) + for i, l := range labels { + ret[i] = l.String() + } + return ret +} + +// visibleLabels filters out hidden targets unless hidden is set. +func visibleLabels(labels core.BuildLabels, hidden bool) core.BuildLabels { + if hidden { + return labels + } + ret := make(core.BuildLabels, 0, len(labels)) + for _, l := range labels { + if !strings.HasPrefix(l.Name, "_") { + ret = append(ret, l) + } + } + return ret +} + +// resolveWholeGraphLabels is like resolveLabels but treats an empty input as the whole graph. +func resolveWholeGraphLabels(state *core.BuildState, in []string) ([]core.BuildLabel, error) { + if len(in) == 0 { + return state.ExpandLabels(core.WholeGraph), nil + } + return resolveLabels(state, in) +} + +// validatePrintFields checks that the given field names exist on a build target; +// query.Print dies on unknown fields, which would take the server with it. +func validatePrintFields(fields []string) error { + if len(fields) == 0 { + return nil + } + valid := validPrintFieldNames() + for _, f := range fields { + if _, present := valid[f]; !present { + names := make([]string, 0, len(valid)) + for name := range valid { + names = append(names, name) + } + sort.Strings(names) + return fmt.Errorf("unknown field %s; known fields are: %s", f, strings.Join(names, ", ")) + } + } + return nil +} + +// validPrintFieldNames returns the set of field names query.Print understands, +// mirroring its name resolution (the 'name' struct tag, else the lowercased field name). +func validPrintFieldNames() map[string]struct{} { + valid := map[string]struct{}{} + add := func(t reflect.Type) { + for i := 0; i < t.NumField(); i++ { + f := t.Field(i) + if name := f.Tag.Get("name"); name != "" { + valid[name] = struct{}{} + } else { + valid[strings.ToLower(f.Name)] = struct{}{} + } + } + } + targetType := reflect.TypeOf(core.BuildTarget{}) + add(targetType) + if testField, ok := targetType.FieldByName("Test"); ok { + add(testField.Type.Elem()) + } + return valid +} diff --git a/src/please.go b/src/please.go index 9757c5024..8e434b6c8 100644 --- a/src/please.go +++ b/src/please.go @@ -35,6 +35,7 @@ import ( "github.com/thought-machine/please/src/generate" "github.com/thought-machine/please/src/hashes" "github.com/thought-machine/please/src/help" + "github.com/thought-machine/please/src/mcp" "github.com/thought-machine/please/src/metrics" "github.com/thought-machine/please/src/output" "github.com/thought-machine/please/src/plz" @@ -267,6 +268,9 @@ var opts struct { } `positional-args:"true" required:"true"` } `command:"watch" description:"Watches sources of targets for changes and rebuilds them"` + Mcp struct { + } `command:"mcp" description:"Runs a Model Context Protocol server over stdio, answering queries about the build graph from an in-memory cache"` + Update struct { Force bool `long:"force" description:"Forces a re-download of the new version."` NoVerify bool `long:"noverify" description:"Skips signature and hash verification of downloaded version"` @@ -830,12 +834,12 @@ var buildFunctions = map[string]func() int{ }, "query.input": func() int { return runQuery(true, opts.Query.Input.Args.Targets, func(state *core.BuildState) { - query.TargetInputs(state.Graph, state.ExpandOriginalLabels()) + query.TargetInputs(os.Stdout, state.Graph, state.ExpandOriginalLabels()) }) }, "query.output": func() int { return runQuery(true, opts.Query.Output.Args.Targets, func(state *core.BuildState) { - query.TargetOutputs(state.Graph, state.ExpandOriginalLabels(), opts.Query.Output.JSON) + query.TargetOutputs(os.Stdout, state.Graph, state.ExpandOriginalLabels(), opts.Query.Output.JSON) }) }, "query.completions": func() int { @@ -1012,6 +1016,13 @@ var buildFunctions = map[string]func() int{ watch.Watch(state, state.ExpandOriginalLabels(), args, opts.Watch.NoTest, runPlease) return toExitCode(success, state) }, + "mcp": func() int { + if err := mcp.NewServer(config).Serve(context.Background()); err != nil { + log.Error("%s", err) + return 1 + } + return 0 + }, "generate": func() int { opts.BuildFlags.Include = append(opts.BuildFlags.Include, "codegen") diff --git a/src/query/deps.go b/src/query/deps.go index a67d884d6..3c08b0c73 100644 --- a/src/query/deps.go +++ b/src/query/deps.go @@ -17,17 +17,45 @@ func Deps(out io.Writer, state *core.BuildState, labels []core.BuildLabel, hidde fmt.Fprintf(out, " edge [fontname=\"Helvetica,Arial,sans-serif\"]\n") fmt.Fprintf(out, " rankdir=\"LR\"\n") } + visit := func(dep, parent *core.BuildTarget, level int) { + if formatdot { + printTargetDot(out, dep, parent) + } else { + printTarget(out, dep, level) + } + } done := map[core.BuildLabel]bool{} for _, label := range labels { - deps(out, state, state.Graph.TargetOrDie(label), done, targetLevel, 0, hidden, formatdot) + walkDeps(state, state.Graph.TargetOrDie(label), done, targetLevel, 0, hidden, visit) } if formatdot { fmt.Fprintf(out, "}\n") } } -// deps looks at all the deps of the given target & recurses into them, printing as appropriate. -func deps(out io.Writer, state *core.BuildState, target *core.BuildTarget, done map[core.BuildLabel]bool, targetLevel, currentLevel int, hidden, formatdot bool) { +// DepsLabels returns all transitive dependencies of a set of targets in traversal order. +func DepsLabels(state *core.BuildState, labels []core.BuildLabel, hidden bool, targetLevel int) core.BuildLabels { + ret := core.BuildLabels{} + visit := func(dep, parent *core.BuildTarget, level int) { + ret = append(ret, dep.Label) + } + done := map[core.BuildLabel]bool{} + for _, label := range labels { + walkDeps(state, state.Graph.TargetOrDie(label), done, targetLevel, 0, hidden, visit) + } + return ret +} + +// walkDeps looks at all the deps of the given target & recurses into them, +// calling visit for each dependency as it's discovered. +func walkDeps( + state *core.BuildState, + target *core.BuildTarget, + done map[core.BuildLabel]bool, + targetLevel, currentLevel int, + hidden bool, + visit func(dep, parent *core.BuildTarget, level int), +) { if currentLevel == targetLevel { return } @@ -39,18 +67,14 @@ func deps(out io.Writer, state *core.BuildState, target *core.BuildTarget, done } done[l] = true if dep := state.Graph.TargetOrDie(l); hidden || !dep.HasParent() { - // dep is to be printed; either we're printing hidden deps or it has no parent (i.e. is not hidden) - if formatdot { - printTargetDot(out, dep, target) - } else { - printTarget(out, dep, currentLevel) - } - deps(out, state, dep, done, targetLevel, currentLevel+1, hidden, formatdot) + // dep is to be visited; either we're including hidden deps or it has no parent (i.e. is not hidden) + visit(dep, target, currentLevel) + walkDeps(state, dep, done, targetLevel, currentLevel+1, hidden, visit) } else if dep.Label.Parent() == target.Label.Parent() { // This is a hidden dependency of the current target, recurse without increasing depth - deps(out, state, dep, done, targetLevel, currentLevel, hidden, formatdot) + walkDeps(state, dep, done, targetLevel, currentLevel, hidden, visit) } else { - deps(out, state, dep, done, targetLevel, currentLevel+1, hidden, formatdot) + walkDeps(state, dep, done, targetLevel, currentLevel+1, hidden, visit) } } } diff --git a/src/query/inputs.go b/src/query/inputs.go index 1cd416742..8df98da55 100644 --- a/src/query/inputs.go +++ b/src/query/inputs.go @@ -2,6 +2,7 @@ package query import ( "fmt" + "io" "sort" "golang.org/x/exp/maps" @@ -10,7 +11,7 @@ import ( ) // TargetInputs prints all inputs for a single target. -func TargetInputs(graph *core.BuildGraph, labels []core.BuildLabel) { +func TargetInputs(out io.Writer, graph *core.BuildGraph, labels []core.BuildLabel) { inputPaths := map[string]bool{} for _, label := range labels { for sourcePath := range core.IterInputPaths(graph, graph.TargetOrDie(label)) { @@ -21,6 +22,6 @@ func TargetInputs(graph *core.BuildGraph, labels []core.BuildLabel) { keys := maps.Keys(inputPaths) sort.Strings(keys) for _, path := range keys { - fmt.Printf("%s\n", path) + fmt.Fprintf(out, "%s\n", path) } } diff --git a/src/query/outputs.go b/src/query/outputs.go index 3ab38034c..bad481ea7 100644 --- a/src/query/outputs.go +++ b/src/query/outputs.go @@ -3,39 +3,39 @@ package query import ( "encoding/json" "fmt" - "os" + "io" "path/filepath" "github.com/thought-machine/please/src/core" ) // TargetOutputs prints all output files for a set of targets. -func TargetOutputs(graph *core.BuildGraph, labels []core.BuildLabel, useJSON bool) { +func TargetOutputs(out io.Writer, graph *core.BuildGraph, labels []core.BuildLabel, useJSON bool) { if useJSON { - targetOutputsJSON(graph, labels) + targetOutputsJSON(out, graph, labels) } else { - targetOutputsFlat(graph, labels) + targetOutputsFlat(out, graph, labels) } } -func targetOutputsFlat(graph *core.BuildGraph, labels []core.BuildLabel) { +func targetOutputsFlat(out io.Writer, graph *core.BuildGraph, labels []core.BuildLabel) { for _, label := range labels { target := graph.TargetOrDie(label) - for _, out := range target.Outputs() { - fmt.Printf("%s\n", filepath.Join(target.OutDir(), out)) + for _, o := range target.Outputs() { + fmt.Fprintf(out, "%s\n", filepath.Join(target.OutDir(), o)) } } } -func targetOutputsJSON(graph *core.BuildGraph, labels []core.BuildLabel) { +func targetOutputsJSON(out io.Writer, graph *core.BuildGraph, labels []core.BuildLabel) { data := map[string][]string{} for _, label := range labels { target := graph.TargetOrDie(label) - for _, out := range target.Outputs() { - data[label.String()] = append(data[label.String()], filepath.Join(target.OutDir(), out)) + for _, o := range target.Outputs() { + data[label.String()] = append(data[label.String()], filepath.Join(target.OutDir(), o)) } } - encoder := json.NewEncoder(os.Stdout) + encoder := json.NewEncoder(out) encoder.SetIndent("", " ") if err := encoder.Encode(data); err != nil { log.Fatalf("failed to write JSON: %v", err) diff --git a/src/query/print.go b/src/query/print.go index 6ecd7e63b..e5edf1bd5 100644 --- a/src/query/print.go +++ b/src/query/print.go @@ -72,6 +72,17 @@ func handleSpecialFields(specials specialFieldsMap, target *core.BuildTarget, na return reflect.ValueOf(fun(target)), true } +// TargetToMap converts a build target into a map of its fields keyed by BUILD rule argument +// name, the same representation as `plz query print --json`. A non-empty fields list +// restricts the result to those fields. order is as returned by parse.BuildRuleArgOrder. +func TargetToMap( + order map[string]int, + fields []string, + target *core.BuildTarget, +) map[string]interface{} { + return targetToValueMap(order, fields, target) +} + // targetToValueMap creates a map of fields on BuildTarget keyed by the name tag on the struct field annotation. It // handles converting fields like named fields, or complex fields so that this can be serialised to json. func targetToValueMap(order map[string]int, fieldsToInclude []string, target *core.BuildTarget) map[string]interface{} { diff --git a/src/query/reverse_deps.go b/src/query/reverse_deps.go index cca1ea2f4..c84392628 100644 --- a/src/query/reverse_deps.go +++ b/src/query/reverse_deps.go @@ -10,19 +10,28 @@ import ( // ReverseDeps finds all transitive targets that depend on the set of input labels. func ReverseDeps(state *core.BuildState, labels []core.BuildLabel, level int, hidden bool) { + for _, l := range ReverseDepsLabels(state, labels, level, hidden) { + fmt.Println(l.String()) + } +} + +// ReverseDepsLabels returns the sorted labels of all transitive targets that depend on the +// set of input labels. +func ReverseDepsLabels( + state *core.BuildState, + labels []core.BuildLabel, + level int, + hidden bool, +) core.BuildLabels { targets := FindRevdeps(state, labels, hidden, true, true, level) ls := make(core.BuildLabels, 0, len(targets)) - for target := range targets { if state.ShouldInclude(target) { ls = append(ls, target.Label) } } sort.Sort(ls) - - for _, l := range ls { - fmt.Println(l.String()) - } + return ls } // node represents a node in the build graph and the depth we visited it at. diff --git a/src/query/somepath.go b/src/query/somepath.go index 75b91bdc1..166a351cc 100644 --- a/src/query/somepath.go +++ b/src/query/somepath.go @@ -7,9 +7,27 @@ import ( "github.com/thought-machine/please/src/core" ) -// SomePath finds and returns a path between two targets, or between one and a set of targets. +// SomePath finds and prints a path between two targets, or between one and a set of targets. // Useful for a "why on earth do I depend on this thing" type query. func SomePath(graph *core.BuildGraph, from, to, except []core.BuildLabel, showHidden bool) error { + path, err := SomePathLabels(graph, from, to, except, showHidden) + if err != nil { + return err + } + fmt.Println("Found path:") + for _, l := range path { + fmt.Printf(" %s\n", l) + } + return nil +} + +// SomePathLabels finds and returns a path between two targets, or between one and a set of +// targets, or an error if no path exists between them. +func SomePathLabels( + graph *core.BuildGraph, + from, to, except []core.BuildLabel, + showHidden bool, +) (core.BuildLabels, error) { s := somepath{ graph: graph, except: make(map[core.BuildLabel]struct{}, len(except)), @@ -21,7 +39,6 @@ func SomePath(graph *core.BuildGraph, from, to, except []core.BuildLabel, showHi for _, l1 := range expandAllTargets(graph, from) { for _, l2 := range expandAllTargets(graph, to) { if path := s.SomePath(l1, l2); len(path) != 0 { - fmt.Println("Found path:") if !showHidden { // Filter path to just non-hidden targets for i, x := range path { @@ -29,17 +46,14 @@ func SomePath(graph *core.BuildGraph, from, to, except []core.BuildLabel, showHi } path = slices.Compact(path) } - for _, l := range path { - fmt.Printf(" %s\n", l) - } - return nil + return path, nil } } } if len(from) == 1 && len(to) == 1 { - return fmt.Errorf("Couldn't find any dependency path between %s and %s", from[0], to[0]) + return nil, fmt.Errorf("Couldn't find any dependency path between %s and %s", from[0], to[0]) } - return fmt.Errorf("Couldn't find any dependency path between those targets") + return nil, fmt.Errorf("Couldn't find any dependency path between those targets") } // expandAllTargets expands any :all labels in the given set. diff --git a/src/query/whatinputs.go b/src/query/whatinputs.go index 1fcfb33bd..8d52689aa 100644 --- a/src/query/whatinputs.go +++ b/src/query/whatinputs.go @@ -29,6 +29,16 @@ func WhatInputs(graph *core.BuildGraph, files []string, hidden, printFiles, igno } } +// WhatInputsLabels returns the targets that have each of the given files as sources, +// keyed by file. Files that are not a source to any target map to an empty list. +func WhatInputsLabels( + graph *core.BuildGraph, + files []string, + hidden bool, +) map[string]core.BuildLabels { + return whatInputs(graph.AllTargets(), files, hidden) +} + func whatInputs(targets []*core.BuildTarget, files []string, hidden bool) map[string]core.BuildLabels { filesMap := make(map[string]map[core.BuildLabel]struct{}, len(files)) for _, file := range files { diff --git a/src/query/whatoutputs.go b/src/query/whatoutputs.go index c55517949..0afa34117 100644 --- a/src/query/whatoutputs.go +++ b/src/query/whatoutputs.go @@ -29,6 +29,17 @@ func WhatOutputs(graph *core.BuildGraph, files []string, printFiles bool) { } } +// WhatOutputsLabels returns the targets responsible for producing each of the given files, +// keyed by file. Files that are not an output of any target map to an empty list. +func WhatOutputsLabels(graph *core.BuildGraph, files []string) map[string]core.BuildLabels { + targets := graph.AllTargets() + ret := make(map[string]core.BuildLabels, len(files)) + for _, f := range files { + ret[f] = whatOutputs(targets, f) + } + return ret +} + func whatOutputs(targets []*core.BuildTarget, file string) []core.BuildLabel { ret := []core.BuildLabel{} for _, t := range targets { diff --git a/third_party/go/BUILD b/third_party/go/BUILD index 155ab7ae0..94d4e3ab7 100644 --- a/third_party/go/BUILD +++ b/third_party/go/BUILD @@ -729,3 +729,39 @@ go_repo( module = "go.opentelemetry.io/auto/sdk", version = "v1.2.1", ) + +go_repo( + licences = ["MIT"], + module = "github.com/modelcontextprotocol/go-sdk", + version = "v1.7.0", +) + +go_repo( + licences = ["MIT"], + module = "github.com/google/jsonschema-go", + version = "v0.4.3", +) + +go_repo( + licences = ["MIT"], + module = "github.com/segmentio/encoding", + version = "v0.5.4", +) + +go_repo( + licences = ["MIT"], + module = "github.com/segmentio/asm", + version = "v1.1.3", +) + +go_repo( + licences = ["BSD-3-Clause"], + module = "github.com/yosida95/uritemplate/v3", + version = "v3.0.2", +) + +go_repo( + licences = ["MIT"], + module = "github.com/golang-jwt/jwt/v5", + version = "v5.3.1", +)