From 347420e77fc56f93280e1c05aeaf435dc2dfce86 Mon Sep 17 00:00:00 2001 From: Tonis Tiigi Date: Thu, 6 Sep 2018 15:25:45 -0700 Subject: [PATCH] sshprovider: allow keys from local files Signed-off-by: Tonis Tiigi --- client/client_test.go | 52 +++++++++- cmd/buildctl/build.go | 4 +- session/sshforward/ssh.go | 2 +- .../sshforward/sshprovider/agentprovider.go | 99 ++++++++++++++----- 4 files changed, 126 insertions(+), 31 deletions(-) diff --git a/client/client_test.go b/client/client_test.go index f70d8d40a..0e4053211 100644 --- a/client/client_test.go +++ b/client/client_test.go @@ -5,7 +5,9 @@ import ( "context" "crypto/rand" "crypto/rsa" + "crypto/x509" "encoding/json" + "encoding/pem" "fmt" "io" "io/ioutil" @@ -104,7 +106,7 @@ func testSSHMount(t *testing.T, sb integration.Sandbox) { defer clean() ssh, err := sshprovider.NewSSHAgentProvider([]sshprovider.AgentConfig{{ - Socket: sockPath, + Paths: []string{sockPath}, }}) require.NoError(t, err) @@ -190,6 +192,54 @@ func testSSHMount(t *testing.T, sb integration.Sandbox) { dt, err = ioutil.ReadFile(filepath.Join(destDir, "out")) require.NoError(t, err) require.Contains(t, string(dt), "agent refused operation") + + // valid socket from key on disk + st = llb.Image("alpine:latest"). + Run(llb.Shlex(`apk add --no-cache openssh`)). + Run(llb.Shlex(`sh -c 'ssh-add -l > /out/out'`), + llb.AddSSHSocket()) + + out = st.AddMount("/out", llb.Scratch()) + def, err = out.Marshal() + require.NoError(t, err) + + k, err = rsa.GenerateKey(rand.Reader, 1024) + require.NoError(t, err) + + dt = pem.EncodeToMemory( + &pem.Block{ + Type: "RSA PRIVATE KEY", + Bytes: x509.MarshalPKCS1PrivateKey(k), + }, + ) + + tmpDir, err := ioutil.TempDir("", "buildkit") + require.NoError(t, err) + defer os.RemoveAll(tmpDir) + + err = ioutil.WriteFile(filepath.Join(tmpDir, "key"), dt, 0600) + require.NoError(t, err) + + ssh, err = sshprovider.NewSSHAgentProvider([]sshprovider.AgentConfig{{ + Paths: []string{filepath.Join(tmpDir, "key")}, + }}) + require.NoError(t, err) + + destDir, err = ioutil.TempDir("", "buildkit") + require.NoError(t, err) + defer os.RemoveAll(destDir) + + _, err = c.Solve(context.TODO(), def, SolveOpt{ + Exporter: ExporterLocal, + ExporterOutputDir: destDir, + Session: []session.Attachable{ssh}, + }, nil) + require.NoError(t, err) + + dt, err = ioutil.ReadFile(filepath.Join(destDir, "out")) + require.NoError(t, err) + require.Contains(t, string(dt), "1024") + require.Contains(t, string(dt), "(RSA)") } func testExtraHosts(t *testing.T, sb integration.Sandbox) { diff --git a/cmd/buildctl/build.go b/cmd/buildctl/build.go index bf4f56d11..a82a5b810 100644 --- a/cmd/buildctl/build.go +++ b/cmd/buildctl/build.go @@ -85,7 +85,7 @@ var buildCommand = cli.Command{ }, cli.StringSliceFlag{ Name: "ssh", - Usage: "Allow forwarding SSH agent to the builder. Format default|[=]", + Usage: "Allow forwarding SSH agent to the builder. Format default|[=|[,]]", }, }, } @@ -387,7 +387,7 @@ func parseSSHSpecs(inp []string) ([]sshprovider.AgentConfig, error) { ID: parts[0], } if len(parts) > 1 { - cfg.Socket = parts[1] + cfg.Paths = strings.Split(parts[1], ",") } configs = append(configs, cfg) } diff --git a/session/sshforward/ssh.go b/session/sshforward/ssh.go index 1a41e9a4c..a4effef60 100644 --- a/session/sshforward/ssh.go +++ b/session/sshforward/ssh.go @@ -97,7 +97,7 @@ func MountSSHSocket(ctx context.Context, c session.Caller, opt SocketOpt) (sockP id = DefaultID } - go s.run(ctx, l, id) + go s.run(ctx, l, id) // erroring per connection allowed return sockPath, func() error { err := l.Close() diff --git a/session/sshforward/sshprovider/agentprovider.go b/session/sshforward/sshprovider/agentprovider.go index 718009ecc..b27688542 100644 --- a/session/sshforward/sshprovider/agentprovider.go +++ b/session/sshforward/sshprovider/agentprovider.go @@ -2,8 +2,8 @@ package sshprovider import ( "context" - "fmt" "io" + "io/ioutil" "net" "os" "time" @@ -19,22 +19,23 @@ import ( ) type AgentConfig struct { - ID string - Socket string + ID string + Paths []string } func NewSSHAgentProvider(confs []AgentConfig) (session.Attachable, error) { - m := map[string]string{} + m := map[string]source{} for _, conf := range confs { - if conf.Socket == "" { - conf.Socket = os.Getenv("SSH_AUTH_SOCK") + if len(conf.Paths) == 0 || len(conf.Paths) == 1 && conf.Paths[0] == "" { + conf.Paths = []string{os.Getenv("SSH_AUTH_SOCK")} } - if conf.Socket == "" { + if conf.Paths[0] == "" { return nil, errors.Errorf("invalid empty ssh agent socket, make sure SSH_AUTH_SOCK is set") } - if err := validateSSHAgentSocket(conf.Socket); err != nil { + src, err := toAgentSource(conf.Paths) + if err != nil { return nil, err } if conf.ID == "" { @@ -43,14 +44,19 @@ func NewSSHAgentProvider(confs []AgentConfig) (session.Attachable, error) { if _, ok := m[conf.ID]; ok { return nil, errors.Errorf("invalid duplicate ID %s", conf.ID) } - m[conf.ID] = conf.Socket + m[conf.ID] = src } return &socketProvider{m: m}, nil } +type source struct { + agent agent.Agent + socket string +} + type socketProvider struct { - m map[string]string + m map[string]source } func (sp *socketProvider) Register(server *grpc.Server) { @@ -77,20 +83,26 @@ func (sp *socketProvider) ForwardAgent(stream sshforward.SSH_ForwardAgentServer) id = v[0] } - socket, ok := sp.m[id] + src, ok := sp.m[id] if !ok { - fmt.Printf("unset11 %s\n", id) return errors.Errorf("unset ssh forward key %s", id) } - conn, err := net.DialTimeout("unix", socket, time.Second) - if err != nil { - return errors.Wrapf(err, "failed to connect to %s", socket) - } - s1, s2 := sockPair() - a := &readOnlyAgent{agent.NewClient(conn)} + var a agent.Agent - defer conn.Close() + if src.socket != "" { + conn, err := net.DialTimeout("unix", src.socket, time.Second) + if err != nil { + return errors.Wrapf(err, "failed to connect to %s", src.socket) + } + + a = &readOnlyAgent{agent.NewClient(conn)} + defer conn.Close() + } else { + a = src.agent + } + + s1, s2 := sockPair() eg, ctx := errgroup.WithContext(context.TODO()) @@ -106,16 +118,49 @@ func (sp *socketProvider) ForwardAgent(stream sshforward.SSH_ForwardAgentServer) return eg.Wait() } -func validateSSHAgentSocket(socket string) error { - conn, err := net.DialTimeout("unix", socket, time.Second) - if err != nil { - return errors.Wrapf(err, "failed to connect to %s", socket) +func toAgentSource(paths []string) (source, error) { + var keys bool + var socket string + a := agent.NewKeyring() + for _, p := range paths { + if socket != "" { + return source{}, errors.New("only single socket allowed") + } + fi, err := os.Stat(p) + if err != nil { + return source{}, errors.WithStack(err) + } + if fi.Mode()&os.ModeSocket > 0 { + if keys { + return source{}, errors.Errorf("invalid combination of keys and sockets") + } + socket = p + continue + } + keys = true + f, err := os.Open(p) + if err != nil { + return source{}, errors.Wrapf(err, "failed to open %s", p) + } + dt, err := ioutil.ReadAll(&io.LimitedReader{R: f, N: 100 * 1024}) + if err != nil { + return source{}, errors.Wrapf(err, "failed to read %s", p) + } + + k, err := ssh.ParseRawPrivateKey(dt) + if err != nil { + return source{}, errors.Wrapf(err, "failed to parse %s", p) // TODO: prompt passphrase? + } + if err := a.Add(agent.AddedKey{PrivateKey: k}); err != nil { + return source{}, errors.Wrapf(err, "failed to add %s to agent", p) + } } - defer conn.Close() - if _, err := agent.NewClient(conn).List(); err != nil { - return errors.Wrapf(err, "failed to verify %s as ssh agent socket", socket) + + if socket != "" { + return source{socket: socket}, nil } - return nil + + return source{agent: a}, nil } func sockPair() (io.ReadWriteCloser, io.ReadWriteCloser) {