[sk] Embed the MCP server This avoids having to run an internal MCP server instance behind the corp proxy and jump through hoops to connect `sk` to it. The downside is that the user needs to personally have read access to all of the DBs, as opposed to just the MCP server itself. Bug: b/491418947 Change-Id: I6772af19ef568c67e8ab33659338bee268f7c509 Reviewed-on: https://skia-review.googlesource.com/c/buildbot/+/1350856 Reviewed-by: Brett Somocurcio <brettos@google.com> Commit-Queue: Brett Somocurcio <brettos@google.com> Auto-Submit: Eric Boren <borenet@google.com>
diff --git a/sk/go/agent/tool/BUILD.bazel b/sk/go/agent/tool/BUILD.bazel index eb2e635..b37e34e 100644 --- a/sk/go/agent/tool/BUILD.bazel +++ b/sk/go/agent/tool/BUILD.bazel
@@ -9,14 +9,14 @@ importpath = "go.skia.org/infra/sk/go/agent/tool", visibility = ["//visibility:public"], deps = [ - "//go/auth", - "//go/httputils", + "//autogardener/go/mcp", + "//go/cleanup", "//go/skerr", + "//go/sklog", "//go/util", - "@com_github_mark3labs_mcp_go//client", - "@com_github_mark3labs_mcp_go//client/transport", + "//mcp/services/skia", "@com_github_mark3labs_mcp_go//mcp", "@com_github_urfave_cli_v2//:cli", - "@org_golang_x_oauth2//google", + "@org_golang_google_genai//:genai", ], )
diff --git a/sk/go/agent/tool/mcp.go b/sk/go/agent/tool/mcp.go index 31b8cf8..fb52dc7 100644 --- a/sk/go/agent/tool/mcp.go +++ b/sk/go/agent/tool/mcp.go
@@ -4,46 +4,53 @@ "context" "fmt" "io" - "os" "strings" - "github.com/mark3labs/mcp-go/client" - "github.com/mark3labs/mcp-go/client/transport" "github.com/mark3labs/mcp-go/mcp" "github.com/urfave/cli/v2" - "go.skia.org/infra/go/auth" - "go.skia.org/infra/go/httputils" + sk_mcp "go.skia.org/infra/autogardener/go/mcp" + "go.skia.org/infra/go/cleanup" "go.skia.org/infra/go/skerr" + "go.skia.org/infra/go/sklog" "go.skia.org/infra/go/util" - "golang.org/x/oauth2/google" + "go.skia.org/infra/mcp/services/skia" + "google.golang.org/genai" ) const ( - mcpServerURL = "https://mcp-skia.luci.app/sse" - mcpServerOverrideEnvVar = "SK_MCP_SERVER_OVERRIDE" + publicFirestoreInstance = "production" + publicTdBtProject = "skia-public" + publicTdBtInstance = "staging" + publicSwarmingServer = "chromium-swarm.appspot.com" + internalFirestoreInstance = "internal" + internalTdBtProject = "google.com:skia-corp" + internalTdBtInstance = "internal" + internalSwarmingServer = "chrome-swarming.appspot.com" ) func createCommandsForMCPTools(ctx context.Context) ([]*cli.Command, error) { - mcpClient, err := initMCP(ctx) + mcpClient, err := initMCP(ctx, false /* Tools are the same for public/internal */) if err != nil { return nil, skerr.Wrap(err) } - tools, err := mcpClient.ListTools(ctx, mcp.ListToolsRequest{}) - if err != nil { - return nil, skerr.Wrap(err) + internalFlag := &cli.BoolFlag{ + Name: "internal", + Usage: "If set, use internal data. The user must have read permissions.", } - var commands []*cli.Command - for _, tool := range tools.Tools { - toolName := tool.Name + for _, tool := range mcpClient.Tools() { + decl := tool.FunctionDeclarations[0] + toolName := decl.Name + flags := append(getFlagsFromSchema(decl.Parameters), defaultFlags...) + flags = append(flags, internalFlag) cmd := &cli.Command{ Name: toolName, - Usage: tool.Description, - Description: tool.Description, - Flags: append(getFlagsFromSchema(tool.InputSchema), defaultFlags...), + Usage: decl.Description, + Description: decl.Description, + Flags: flags, Action: func(c *cli.Context) error { - return callMCPTool(c, toolName) + return callMCPTool(c, toolName, c.Bool(internalFlag.Name)) }, } commands = append(commands, cmd) @@ -59,35 +66,24 @@ }, } -func getFlagsFromSchema(schema mcp.ToolInputSchema) []cli.Flag { +func getFlagsFromSchema(schema *genai.Schema) []cli.Flag { + requiredMap := make(map[string]bool, len(schema.Required)) + for _, required := range schema.Required { + requiredMap[required] = true + } + var flags []cli.Flag for name, prop := range schema.Properties { - propMap, ok := prop.(map[string]interface{}) - if !ok { - continue - } + description := prop.Description + required := requiredMap[name] - description := "" - if d, ok := propMap["description"]; ok { - description = fmt.Sprintf("%v", d) - } - - required := false - for _, req := range schema.Required { - if req == name { - required = true - break - } - } - - propType, _ := propMap["type"].(string) - switch propType { - case "boolean": + switch prop.Type { + case genai.TypeBoolean: flags = append(flags, &cli.BoolFlag{ Name: name, Usage: description, }) - case "number", "integer": + case genai.TypeNumber, genai.TypeInteger: flags = append(flags, &cli.IntFlag{ Name: name, Usage: description, @@ -105,8 +101,8 @@ return flags } -func callMCPTool(ctx *cli.Context, toolName string) error { - mcpClient, err := initMCP(ctx.Context) +func callMCPTool(ctx *cli.Context, toolName string, internal bool) error { + mcpClient, err := initMCP(ctx.Context, internal) if err != nil { return skerr.Wrap(err) } @@ -126,12 +122,7 @@ } } - res, err := mcpClient.CallTool(ctx.Context, mcp.CallToolRequest{ - Params: mcp.CallToolParams{ - Name: toolName, - Arguments: args, - }, - }) + res, err := mcpClient.CallTool(ctx.Context, toolName, args) if err != nil { return skerr.Wrap(err) } @@ -153,32 +144,33 @@ return nil } -func initMCP(ctx context.Context) (*client.Client, error) { - ts, err := google.DefaultTokenSource(ctx, auth.ScopeUserinfoEmail) - if err != nil { - return nil, skerr.Wrap(err) - } - c := httputils.DefaultClientConfig().WithTokenSource(ts).WithoutRetries().Client() - - mcpURL := os.Getenv(mcpServerOverrideEnvVar) - if mcpURL == "" { - mcpURL = mcpServerURL - } - mcpClient, err := client.NewSSEMCPClient(mcpURL, transport.WithHTTPClient(c)) - if err != nil { - return nil, skerr.Wrap(err) - } - - if err := mcpClient.Start(ctx); err != nil { - return nil, skerr.Wrap(err) - } - _, err = mcpClient.Initialize(ctx, mcp.InitializeRequest{ - Params: mcp.InitializeParams{ - ProtocolVersion: mcp.LATEST_PROTOCOL_VERSION, - }, +func initMCP(ctx context.Context, internal bool) (sk_mcp.MCPClient, error) { + srv := &skia.SkiaService{} + cleanup.AtExit(func() { + if err := srv.Shutdown(); err != nil { + sklog.Errorf("Error performing shutdown for service: %v", err) + } }) - if err != nil { + firestoreInstance := publicFirestoreInstance + tdBtProject := publicTdBtProject + tdBtInstance := publicTdBtInstance + swarmingServer := publicSwarmingServer + if internal { + firestoreInstance = internalFirestoreInstance + tdBtProject = internalTdBtProject + tdBtInstance = internalTdBtInstance + swarmingServer = internalSwarmingServer + } + mcpArgs := fmt.Sprintf( + "--firestore_instance=%s --bigtable_project=%s --bigtable_instance=%s --swarming_server=%s", + firestoreInstance, + tdBtProject, + tdBtInstance, + swarmingServer, + ) + if err := srv.Init(mcpArgs); err != nil { return nil, skerr.Wrap(err) } + mcpClient := sk_mcp.NewEmbeddedService(srv) return mcpClient, nil }