gitproto: add SSH transport via per-RPC ssh exec · Entire

gitproto: add SSH transport via per-RPC ssh exec

406a5ab→main·

Sessions

9dc3927cec43View transcript

Changes

5

package gitproto

import (
    "bytes"
    "context"
    "errors"
    "fmt"
    "io"
    "net/url"
    "os/exec"
    "strings"
    "sync"
)

// SSHLookPath is replaceable in tests.
var SSHLookPath = exec.LookPath

// SSHConn represents a Git transport over the local ssh binary.
type SSHConn struct {
    Label       string
    EndpointURL *url.URL
    sshPath     string
    progressOut io.Writer
}

// NewSSHConn creates a new SSH transport connection backed by the local ssh
// binary.
func NewSSHConn(ep *url.URL, label string) (*SSHConn, error) {
    sshPath, err := SSHLookPath("ssh")
    if err != nil {
        return nil, fmt.Errorf("locate ssh binary: %w", err)
    }
    normalizeEndpointPath(ep)
    return &SSHConn{
        Label:       label,
        EndpointURL: ep,
        sshPath:     sshPath,
    }, nil
}

func (c *SSHConn) Endpoint() *url.URL { return c.EndpointURL }

func (c *SSHConn) ProgressWriter() io.Writer { return c.progressOut }

func (c *SSHConn) SetProgressWriter(w io.Writer) { c.progressOut = w }

func (c *SSHConn) Close() error { return nil }

func (c *SSHConn) RequestInfoRefs(ctx context.Context, service string, gitProtocol string) ([]byte, error) {
    cmd, stderr, err := c.startRPC(ctx, service, gitProtocol)
    if err != nil {
        return nil, err
    }
    if err := cmd.Stdin.Close(); err != nil {
        return nil, fmt.Errorf("close ssh stdin for %s: %w", service, err)
    }
    data, readErr := io.ReadAll(cmd.Stdout)
    waitErr := cmd.wait()
    if ctx.Err() != nil {
        return nil, errors.Join(ctx.Err(), readErr, stderr.wrap(waitErr))
    }
    if readErr != nil {
        return nil, fmt.Errorf("%s info-refs: %w", service, readErr)
    }
    if waitErr != nil {
        return nil, fmt.Errorf("%s info-refs: %w", service, stderr.wrap(waitErr))
    }
    return data, nil
}

// Additional code omitted for brevity...

func (e *sshCommandError) wrap(err error) error {
    if err == nil {
        return nil
    }
    if msg := e.String(); msg != "" {
        return fmt.Errorf("%w: %s", err, msg)
    }
    return err
}

package gitproto

import (
    "context"
    "errors"
    "io"
    "os"
    "path/filepath"
    "strconv"
    "strings"
    "testing"
    "time"
)

func TestNewSSHConnRequiresBinary(t *testing.T) {
    orig := SSHLookPath
    t.Cleanup(func() { SSHLookPath = orig })
    SSHLookPath = func(string) (string, error) {
        return "", errors.New("not found")
    }

ep, err := transport.ParseURL("ssh://example.com/repo.git")
    if err != nil {
        t.Fatalf("parse url: %v", err)
    }
    _, err = NewSSHConn(ep, "source")
    if err == nil || !strings.Contains(err.Error(), "locate ssh binary") {
        t.Fatalf("NewSSHConn error = %v, want locate ssh binary failure", err)
    }
}

// Additional test cases omitted for brevity...