diff --git a/.github/workflows/pr-ci.yml b/.github/workflows/pr-ci.yml index b69972a12..578232f87 100644 --- a/.github/workflows/pr-ci.yml +++ b/.github/workflows/pr-ci.yml @@ -275,6 +275,30 @@ jobs: install-kind: false requires-secret: false + - label: ssh-agent-forward + runner: ubuntu-latest + free-disk-space: false + install-kind: false + requires-secret: false + + - label: ssh-ports-attributes + runner: ubuntu-latest + free-disk-space: false + install-kind: false + requires-secret: false + + - label: ssh-tunnel-mode + runner: ubuntu-latest + free-disk-space: false + install-kind: false + requires-secret: false + + - label: ssh-credentials-server-race + runner: ubuntu-latest + free-disk-space: false + install-kind: false + requires-secret: false + - label: build runner: ubuntu-latest free-disk-space: false diff --git a/cmd/internal/agentcontainer/credentials_server.go b/cmd/internal/agentcontainer/credentials_server.go index 35ae78ac5..b69536114 100644 --- a/cmd/internal/agentcontainer/credentials_server.go +++ b/cmd/internal/agentcontainer/credentials_server.go @@ -4,10 +4,12 @@ import ( "context" "encoding/base64" "encoding/json" + "errors" "fmt" "net" "os" "strconv" + "syscall" "github.com/devsy-org/devsy/cmd/flags" "github.com/devsy-org/devsy/pkg/agent/tunnel" @@ -91,6 +93,16 @@ func (cmd *CredentialsServerCmd) Run(ctx context.Context, port int) error { runCtx, cancel := context.WithCancel(ctx) defer cancel() + ln, err := claimPort(port) + if err != nil { + if errors.Is(err, errPortOwnedByAnotherSession) { + cmd.logPortOwnedByAnotherSession(ctx, port) + return nil + } + return err + } + defer func() { _ = ln.Close() }() + tunnelClient, err := tunnelserver.NewTunnelClient(os.Stdin, os.Stdout, true, ExitCodeIO) if err != nil { return fmt.Errorf("error creating tunnel client: %w", err) @@ -100,12 +112,6 @@ func (cmd *CredentialsServerCmd) Run(ctx context.Context, port int) error { return fmt.Errorf("ping client: %w", err) } - ln, err := claimPort(port) - if err != nil { - return err - } - defer func() { _ = ln.Close() }() - cmd.maybeForwardPorts(runCtx, tunnelClient) if err := cmd.configureDockerHelper(port); err != nil { @@ -126,9 +132,13 @@ func (cmd *CredentialsServerCmd) Run(ctx context.Context, port int) error { cleanupGitSigning := cmd.configureGitSigningKey() defer cleanupGitSigning() - return credentials.RunCredentialsServerWithListener(runCtx, ln, tunnelClient) + return credentials.RunCredentialsServerWithListener(runCtx, ln, tunnelClient, cmd.User) } +// errPortOwnedByAnotherSession marks a bind failure as another session +// already owning the port, not a real error. +var errPortOwnedByAnotherSession = errors.New("credentials server port owned by another session") + // claimPort binds port and returns the listener, holding it exclusively so // no other session can bind the same port until the caller closes it (or // hands it to RunCredentialsServerWithListener). Only one session's @@ -137,13 +147,40 @@ func claimPort(port int) (net.Listener, error) { addr := net.JoinHostPort("localhost", strconv.Itoa(port)) ln, err := net.Listen("tcp", addr) if err != nil { - return nil, fmt.Errorf( - "port %d not available (another session may own the credentials server): %w", + if errors.Is(err, syscall.EADDRINUSE) { + return nil, fmt.Errorf("%w: %w", errPortOwnedByAnotherSession, err) + } + return nil, fmt.Errorf("port %d not available: %w", port, err) + } + return ln, nil +} + +func (cmd *CredentialsServerCmd) logPortOwnedByAnotherSession(ctx context.Context, port int) { + owner, err := credentials.FetchOwner(ctx, port) + switch { + case err != nil: + log.Debugf( + "skipping credentials server for user %s: port %d is taken and its owner could not be determined: %v", + cmd.User, port, err, ) + case owner == "" || owner == cmd.User: + log.Debugf( + "skipping credentials server for user %s: another session already provides it on port %d", + cmd.User, + port, + ) + default: + log.Warnf( + "credentials server for user %s was not started: port %d is already owned by user %s's session; "+ + "git/docker/signing credential helpers for %s will not work until that session ends", + cmd.User, + port, + owner, + cmd.User, + ) } - return ln, nil } func (cmd *CredentialsServerCmd) maybeForwardPorts( diff --git a/cmd/internal/agentcontainer/credentials_server_test.go b/cmd/internal/agentcontainer/credentials_server_test.go index a4aa77b3b..8616f8173 100644 --- a/cmd/internal/agentcontainer/credentials_server_test.go +++ b/cmd/internal/agentcontainer/credentials_server_test.go @@ -1,12 +1,21 @@ package agentcontainer import ( + "context" + "errors" + "fmt" "net" + "strings" "sync" "testing" + "time" + "github.com/devsy-org/devsy/pkg/agent/tunnel" + "github.com/devsy-org/devsy/pkg/credentials" + "github.com/devsy-org/devsy/pkg/log" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" + "google.golang.org/grpc" ) func TestClaimPort_SucceedsWhenPortFree(t *testing.T) { @@ -23,7 +32,7 @@ func TestClaimPort_ErrorsWhenPortHeld(t *testing.T) { _, err = claimPort(port) require.Error(t, err) - assert.Contains(t, err.Error(), "not available") + assert.ErrorIs(t, err, errPortOwnedByAnotherSession) } func TestClaimPort_BecomesClaimableAfterHolderReleases(t *testing.T) { @@ -79,3 +88,100 @@ func TestClaimPort_OnlyOneConcurrentCallerWins(t *testing.T) { _ = winner.Close() } } + +func startFakeCredentialsServer(t *testing.T, owner string) int { + t.Helper() + + ln, err := net.Listen("tcp", "localhost:0") + require.NoError(t, err) + port := ln.Addr().(*net.TCPAddr).Port + + ctx, cancel := context.WithCancel(context.Background()) + t.Cleanup(cancel) + go func() { + _ = credentials.RunCredentialsServerWithListener(ctx, ln, &fakeCredentialsClient{}, owner) + }() + + require.Eventually(t, func() bool { + conn, dialErr := net.Dial("tcp", ln.Addr().String()) + if dialErr != nil { + return false + } + _ = conn.Close() + return true + }, time.Second, 5*time.Millisecond, "fake credentials server must become dialable") + + return port +} + +func TestCredentialsServerCmd_Run_SameOwnerCollisionIsSilentNoOp(t *testing.T) { + port := startFakeCredentialsServer(t, "alice") + + var sink strings.Builder + log.Init(log.Config{Verbosity: 2, Format: "json"}) + remove := log.AddSink(&sink) + defer remove() + + cmd := &CredentialsServerCmd{User: "alice"} + err := cmd.Run(context.Background(), port) + require.NoError(t, err, "losing to the same owner's session must not be an error") + _ = log.Sync() + + assert.NotContains(t, sink.String(), "\"level\":\"warn\"", "same-owner collision must not warn") +} + +func TestCredentialsServerCmd_Run_DifferentOwnerCollisionWarnsButDoesNotError(t *testing.T) { + port := startFakeCredentialsServer(t, "alice") + + var sink strings.Builder + log.Init(log.Config{Verbosity: 2, Format: "json"}) + remove := log.AddSink(&sink) + defer remove() + + cmd := &CredentialsServerCmd{User: "root"} + err := cmd.Run(context.Background(), port) + require.NoError(t, err, "losing the race must still not fail the session") + _ = log.Sync() + + logged := sink.String() + assert.Contains(t, logged, "root", "warning must name the user left without credentials") + assert.Contains(t, logged, "alice", "warning must name the owning session") +} + +type fakeCredentialsClient struct{} + +func (fakeCredentialsClient) GitCredentials( + _ context.Context, _ *tunnel.Message, _ ...grpc.CallOption, +) (*tunnel.Message, error) { + return nil, fmt.Errorf("not implemented") +} + +func (fakeCredentialsClient) DockerCredentials( + _ context.Context, _ *tunnel.Message, _ ...grpc.CallOption, +) (*tunnel.Message, error) { + return nil, fmt.Errorf("not implemented") +} + +func (fakeCredentialsClient) GitSSHSignature( + _ context.Context, _ *tunnel.Message, _ ...grpc.CallOption, +) (*tunnel.Message, error) { + return nil, fmt.Errorf("not implemented") +} + +func (fakeCredentialsClient) GPGPublicKeys( + _ context.Context, _ *tunnel.Message, _ ...grpc.CallOption, +) (*tunnel.Message, error) { + return nil, fmt.Errorf("not implemented") +} + +func (fakeCredentialsClient) DevsyConfig( + _ context.Context, _ *tunnel.Message, _ ...grpc.CallOption, +) (*tunnel.Message, error) { + return nil, fmt.Errorf("not implemented") +} + +func TestClaimPort_WrapsNonAddrInUseErrorsWithoutSentinel(t *testing.T) { + _, err := claimPort(-1) + require.Error(t, err) + assert.False(t, errors.Is(err, errPortOwnedByAnotherSession)) +} diff --git a/e2e/tests/ssh/agent_forward.go b/e2e/tests/ssh/agent_forward.go index 47d2c1958..547390854 100644 --- a/e2e/tests/ssh/agent_forward.go +++ b/e2e/tests/ssh/agent_forward.go @@ -21,8 +21,7 @@ import ( // per-connection socket directory must be cleaned up on disconnect. var _ = ginkgo.Describe( "devsy ssh agent forwarding", - ginkgo.Label("ssh"), - ginkgo.Label("agent-forward"), + ginkgo.Label("ssh-agent-forward"), ginkgo.Ordered, func() { var ( @@ -253,8 +252,6 @@ var _ = ginkgo.Describe( ginkgo.It( "connection without any agent request still cleans up", - ginkgo.Label("ssh"), - ginkgo.Label("agent-forward"), ginkgo.SpecTimeout(framework.TimeoutModerate()), func(ctx ginkgo.SpecContext) { tmpDir, err := os.MkdirTemp("", "devsy-ssh-cm-clean-") @@ -325,7 +322,6 @@ var _ = ginkgo.Describe( ginkgo.It( "parallel sessions on one connection observe the same socket concurrently", - ginkgo.Label("agent-forward"), ginkgo.SpecTimeout(framework.TimeoutModerate()), func(_ ginkgo.SpecContext) { controlPath, closeCM, err := framework.OpenSSHControlMaster( diff --git a/e2e/tests/ssh/credentials_server_race.go b/e2e/tests/ssh/credentials_server_race.go new file mode 100644 index 000000000..721b21aff --- /dev/null +++ b/e2e/tests/ssh/credentials_server_race.go @@ -0,0 +1,85 @@ +package ssh + +import ( + "context" + "os" + "sync" + "time" + + "github.com/devsy-org/devsy/e2e/framework" + "github.com/onsi/ginkgo/v2" + "github.com/onsi/gomega" +) + +const ( + cmdWorkspace = "workspace" + cmdSSH = "ssh" +) + +var _ = ginkgo.Describe( + "devsy ssh credentials server race", + ginkgo.Label("ssh-credentials-server-race"), + ginkgo.Ordered, + func() { + var initialDir string + + ginkgo.BeforeEach(func() { + var err error + initialDir, err = os.Getwd() + framework.ExpectNoError(err) + }) + + ginkgo.It( + "should not surface credentials-server port errors when two ssh sessions race for the same workspace", + ginkgo.SpecTimeout(framework.TimeoutModerate()), + func(ctx context.Context) { + tempDir, err := framework.CopyToTempDir("tests/ssh/testdata/local-test") + framework.ExpectNoError(err) + + f := framework.NewDefaultFramework(initialDir + "/bin") + _ = f.DevsyProviderAdd(ctx, "docker") + err = f.DevsyProviderUse(ctx, "docker") + framework.ExpectNoError(err) + + ginkgo.DeferCleanup(func(cleanupCtx context.Context) { + _ = f.DevsyWorkspaceDelete(cleanupCtx, tempDir) + framework.CleanupTempDir(initialDir, tempDir) + }) + + upDeadline := time.Now().Add(5 * time.Minute) + upCtx, cancelUp := context.WithDeadline(ctx, upDeadline) + defer cancelUp() + err = f.DevsyUp(upCtx, tempDir) + framework.ExpectNoError(err) + + const sessions = 2 + var wg sync.WaitGroup + stderrs := make([]string, sessions) + runErrs := make([]error, sessions) + for i := range sessions { + wg.Add(1) + go func(i int) { + defer wg.Done() + sshCtx, cancelSSH := context.WithDeadline( + ctx, + time.Now().Add(30*time.Second), + ) + defer cancelSSH() + _, stderr, sshErr := f.ExecCommandCapture(sshCtx, []string{ + cmdWorkspace, cmdSSH, tempDir, "--command", "sleep 2", + }) + stderrs[i] = stderr + runErrs[i] = sshErr + }(i) + } + wg.Wait() + + for i := range sessions { + framework.ExpectNoError(runErrs[i], "ssh session %d stderr: %s", i, stderrs[i]) + gomega.Expect(stderrs[i]).NotTo(gomega.ContainSubstring("not available")) + gomega.Expect(stderrs[i]).NotTo(gomega.ContainSubstring("credentials server")) + } + }, + ) + }, +) diff --git a/e2e/tests/ssh/ports_attributes_test.go b/e2e/tests/ssh/ports_attributes.go similarity index 98% rename from e2e/tests/ssh/ports_attributes_test.go rename to e2e/tests/ssh/ports_attributes.go index df579385e..2da537802 100644 --- a/e2e/tests/ssh/ports_attributes_test.go +++ b/e2e/tests/ssh/ports_attributes.go @@ -17,7 +17,7 @@ import ( ) var _ = ginkgo.Describe("devsy portsAttributes e2e", - ginkgo.Label("ssh"), func() { + ginkgo.Label("ssh-ports-attributes"), func() { var initialDir string ginkgo.BeforeEach(func() { @@ -30,7 +30,7 @@ var _ = ginkgo.Describe("devsy portsAttributes e2e", "should forward port with onAutoForward=silent and skip port with onAutoForward=ignore", ginkgo.SpecTimeout(framework.TimeoutShort()), func(ctx context.Context) { - if runtime.GOOS == "windows" { + if runtime.GOOS == osWindows { ginkgo.Skip("skipping on windows") } @@ -132,7 +132,7 @@ var _ = ginkgo.Describe("devsy portsAttributes e2e", "should forward port with notify policy and apply label metadata", ginkgo.SpecTimeout(framework.TimeoutShort()), func(ctx context.Context) { - if runtime.GOOS == "windows" { + if runtime.GOOS == osWindows { ginkgo.Skip("skipping on windows") } @@ -209,7 +209,7 @@ var _ = ginkgo.Describe("devsy portsAttributes e2e", "should skip forwarding when requireLocalPort=true and host port is occupied", ginkgo.SpecTimeout(framework.TimeoutShort()), func(ctx context.Context) { - if runtime.GOOS == "windows" { + if runtime.GOOS == osWindows { ginkgo.Skip("skipping on windows") } diff --git a/e2e/tests/ssh/ssh.go b/e2e/tests/ssh/ssh.go index ae081679e..c85da1067 100644 --- a/e2e/tests/ssh/ssh.go +++ b/e2e/tests/ssh/ssh.go @@ -74,7 +74,6 @@ var _ = ginkgo.Describe("devsy ssh test suite", ginkgo.Label("ssh"), ginkgo.Orde ginkgo.It( "should start workspace with GPG forwarding when host uses SSH signing format", - ginkgo.Label("gpg"), ginkgo.SpecTimeout(framework.TimeoutModerate()), func(ctx ginkgo.SpecContext) { if runtime.GOOS == osWindows { @@ -124,7 +123,6 @@ var _ = ginkgo.Describe("devsy ssh test suite", ginkgo.Label("ssh"), ginkgo.Orde ginkgo.It( "should expose the host GPG secret key in the container via agent forwarding", - ginkgo.Label("gpg"), ginkgo.SpecTimeout(framework.TimeoutModerate()), func(ctx ginkgo.SpecContext) { if runtime.GOOS == osWindows { diff --git a/e2e/tests/ssh/ssh_tunnel_mode_test.go b/e2e/tests/ssh/ssh_tunnel_mode.go similarity index 73% rename from e2e/tests/ssh/ssh_tunnel_mode_test.go rename to e2e/tests/ssh/ssh_tunnel_mode.go index 09df6df4a..7692a7b9f 100644 --- a/e2e/tests/ssh/ssh_tunnel_mode_test.go +++ b/e2e/tests/ssh/ssh_tunnel_mode.go @@ -1,12 +1,17 @@ package ssh import ( + "bytes" "context" + "fmt" "net" "os" + "os/exec" "path/filepath" "runtime" "strings" + "sync" + "syscall" "time" "github.com/devsy-org/devsy/e2e/framework" @@ -15,9 +20,96 @@ import ( "github.com/onsi/gomega" ) +const tunnelActiveMarker = "waiting for shutdown signal" + +const tunnelActiveTimeout = 4 * time.Minute + +type safeBuffer struct { + mu sync.Mutex + buf bytes.Buffer +} + +func (b *safeBuffer) Write(p []byte) (int, error) { + b.mu.Lock() + defer b.mu.Unlock() + return b.buf.Write(p) +} + +func (b *safeBuffer) String() string { + b.mu.Lock() + defer b.mu.Unlock() + return b.buf.String() +} + +type tunnelUpProcess struct { + cmd *exec.Cmd + output *safeBuffer + exited chan struct{} + exitErr error +} + +func startTunnelUp( + f *framework.Framework, workspace string, extraArgs ...string, +) (*tunnelUpProcess, error) { + args := []string{ + cmdWorkspace, "up", + names.Flag(names.Debug), + names.Flag(names.IDE), "none", + names.Flag(names.SSHTunnel), + } + args = append(args, extraArgs...) + args = append(args, workspace) + + // #nosec G204 -- test binary with controlled arguments + cmd := exec.Command(filepath.Join(f.DevsyBinDir, f.DevsyBinName), args...) + out := &safeBuffer{} + cmd.Stdout = out + cmd.Stderr = out + if err := cmd.Start(); err != nil { + return nil, fmt.Errorf("start devsy up --ssh-tunnel: %w", err) + } + + p := &tunnelUpProcess{cmd: cmd, output: out, exited: make(chan struct{})} + go func() { + p.exitErr = cmd.Wait() + close(p.exited) + }() + return p, nil +} + +func (p *tunnelUpProcess) waitUntilActive(ctx context.Context) { + gomega.Eventually(func() (string, error) { + out := p.output.String() + select { + case <-p.exited: + return out, gomega.StopTrying( + "devsy up exited before the tunnel became active", + ).Wrap(p.exitErr) + default: + return out, nil + } + }).WithContext(ctx).WithTimeout(tunnelActiveTimeout).WithPolling(100 * time.Millisecond). + Should(gomega.ContainSubstring(tunnelActiveMarker)) +} + +func (p *tunnelUpProcess) stop() { + select { + case <-p.exited: + return + default: + } + _ = p.cmd.Process.Signal(syscall.SIGINT) + select { + case <-p.exited: + case <-time.After(15 * time.Second): + _ = p.cmd.Process.Kill() + <-p.exited + } +} + var _ = ginkgo.Describe( "devsy ssh tunnel mode", - ginkgo.Label("ssh"), + ginkgo.Label("ssh-tunnel-mode"), ginkgo.Ordered, func() { var initialDir string @@ -48,10 +140,11 @@ var _ = ginkgo.Describe( framework.CleanupTempDir(initialDir, tempDir) }) - devsyUpCtx, cancel := context.WithDeadline(ctx, time.Now().Add(5*time.Minute)) - defer cancel() - err = f.DevsyUp(devsyUpCtx, tempDir, names.Flag(names.SSHTunnel)) + proc, err := startTunnelUp(f, tempDir) framework.ExpectNoError(err) + ginkgo.DeferCleanup(proc.stop) + + proc.waitUntilActive(ctx) devsySSHCtx, cancelSSH := context.WithDeadline(ctx, time.Now().Add(20*time.Second)) defer cancelSSH() @@ -83,16 +176,11 @@ var _ = ginkgo.Describe( framework.CleanupTempDir(initialDir, tempDir) }) - devsyUpCtx, cancel := context.WithDeadline(ctx, time.Now().Add(5*time.Minute)) - defer cancel() - err = f.DevsyUp( - devsyUpCtx, - tempDir, - names.Flag(names.SSHTunnel), - "--ssh-config", - sshConfigPath, - ) + proc, err := startTunnelUp(f, tempDir, "--ssh-config", sshConfigPath) framework.ExpectNoError(err) + ginkgo.DeferCleanup(proc.stop) + + proc.waitUntilActive(ctx) configBytes, err := os.ReadFile(filepath.Clean(sshConfigPath)) framework.ExpectNoError(err) @@ -136,16 +224,11 @@ var _ = ginkgo.Describe( framework.CleanupTempDir(initialDir, tempDir) }) - devsyUpCtx, cancel := context.WithDeadline(ctx, time.Now().Add(5*time.Minute)) - defer cancel() - err = f.DevsyUp( - devsyUpCtx, - tempDir, - names.Flag(names.SSHTunnel), - "--ssh-config", - sshConfigPath, - ) + proc, err := startTunnelUp(f, tempDir, "--ssh-config", sshConfigPath) framework.ExpectNoError(err) + ginkgo.DeferCleanup(proc.stop) + + proc.waitUntilActive(ctx) configBytes, err := os.ReadFile(filepath.Clean(sshConfigPath)) framework.ExpectNoError(err) @@ -190,10 +273,11 @@ var _ = ginkgo.Describe( framework.CleanupTempDir(initialDir, tempDir) }) - devsyUpCtx, cancel := context.WithDeadline(ctx, time.Now().Add(5*time.Minute)) - defer cancel() - err = f.DevsyUp(devsyUpCtx, tempDir, names.Flag(names.SSHTunnel)) + proc, err := startTunnelUp(f, tempDir) framework.ExpectNoError(err) + ginkgo.DeferCleanup(proc.stop) + + proc.waitUntilActive(ctx) for i := range 3 { sshCtx, cancelSSH := context.WithDeadline(ctx, time.Now().Add(20*time.Second)) diff --git a/pkg/credentials/server.go b/pkg/credentials/server.go index 9858e2a86..600e7b350 100644 --- a/pkg/credentials/server.go +++ b/pkg/credentials/server.go @@ -10,6 +10,7 @@ import ( "net/http" "os" "strconv" + "strings" "time" "github.com/devsy-org/devsy/pkg/agent/tunnel" @@ -44,12 +45,13 @@ func RunCredentialsServer( ctx context.Context, port int, client CredentialsClient, + owner string, ) error { ln, err := net.Listen("tcp", net.JoinHostPort("localhost", strconv.Itoa(port))) if err != nil { return fmt.Errorf("listen on port %d: %w", port, err) } - return RunCredentialsServerWithListener(ctx, ln, client) + return RunCredentialsServerWithListener(ctx, ln, client, owner) } // RunCredentialsServerWithListener is like RunCredentialsServer, but takes an @@ -60,9 +62,10 @@ func RunCredentialsServerWithListener( ctx context.Context, ln net.Listener, client CredentialsClient, + owner string, ) error { srv := &http.Server{ - Handler: newCredentialsHandler(ctx, client), + Handler: newCredentialsHandler(ctx, client, owner), ReadHeaderTimeout: 10 * time.Second, ReadTimeout: 30 * time.Second, IdleTimeout: 120 * time.Second, @@ -91,18 +94,31 @@ type credentialsHandlerFunc func( context.Context, http.ResponseWriter, *http.Request, CredentialsClient, ) error +// ownerPath reports the owner a losing session's port claim collided with. +const ownerPath = "/owner" + // newCredentialsHandler returns an http.Handler that routes requests to the // appropriate handler function, which calls the CredentialsClient to get the // credentials and writes them to the response. // // Root is a readiness probe (see waitForServer); it must return 200 so the // server is detected as up. Unknown paths still 404 below. -func newCredentialsHandler(ctx context.Context, client CredentialsClient) http.Handler { +func newCredentialsHandler( + ctx context.Context, + client CredentialsClient, + owner string, +) http.Handler { routes := map[string]credentialsHandlerFunc{ "/": func(_ context.Context, writer http.ResponseWriter, _ *http.Request, _ CredentialsClient) error { writer.WriteHeader(http.StatusOK) return nil }, + ownerPath: func(_ context.Context, writer http.ResponseWriter, _ *http.Request, _ CredentialsClient) error { + writer.Header().Set("Content-Type", "text/plain; charset=utf-8") + writer.WriteHeader(http.StatusOK) + _, err := writer.Write([]byte(owner)) + return err + }, "/git-credentials": handleGitCredentialsRequest, "/docker-credentials": handleDockerCredentialsRequest, "/git-ssh-signature": handleGitSSHSignatureRequest, @@ -127,6 +143,50 @@ func newCredentialsHandler(ctx context.Context, client CredentialsClient) http.H }) } +const fetchOwnerTimeout = 2 * time.Second + +// maxOwnerResponseSize bounds how much of the /owner response FetchOwner +// reads. Port claimPort failed to bind, so whatever is listening there +// isn't necessarily our own credentials server; cap the read instead of +// trusting it to behave. +const maxOwnerResponseSize = 4096 + +// fetchOwnerClient never follows redirects: /owner always answers 200 with +// a plain-text body, so a redirect means the port isn't ours to trust. +var fetchOwnerClient = &http.Client{ + CheckRedirect: func(*http.Request, []*http.Request) error { + return http.ErrUseLastResponse + }, +} + +// FetchOwner returns "" without error if owner is unset or ownerPath is missing. +func FetchOwner(ctx context.Context, port int) (string, error) { + timeoutCtx, cancel := context.WithTimeout(ctx, fetchOwnerTimeout) + defer cancel() + + url := fmt.Sprintf("http://localhost:%d%s", port, ownerPath) + req, err := http.NewRequestWithContext(timeoutCtx, http.MethodGet, url, nil) + if err != nil { + return "", err + } + + resp, err := fetchOwnerClient.Do(req) + if err != nil { + return "", err + } + defer func() { _ = resp.Body.Close() }() + + if resp.StatusCode != http.StatusOK { + return "", nil + } + + body, err := io.ReadAll(io.LimitReader(resp.Body, maxOwnerResponseSize)) + if err != nil { + return "", err + } + return strings.TrimSpace(string(body)), nil +} + func GetPort() (int, error) { strPort := cmp.Or(os.Getenv(config.EnvCredentialsServerPort), DefaultPort) port, err := strconv.Atoi(strPort) diff --git a/pkg/credentials/server_test.go b/pkg/credentials/server_test.go index a5253430b..10d56dac4 100644 --- a/pkg/credentials/server_test.go +++ b/pkg/credentials/server_test.go @@ -4,6 +4,8 @@ import ( "context" "encoding/json" "fmt" + "io" + "net" "net/http" "net/http/httptest" "strings" @@ -14,7 +16,6 @@ import ( "github.com/stretchr/testify/require" ) -// errReader is an io.Reader that always returns an error. type errReader struct{ err error } func (e *errReader) Read([]byte) (int, error) { return 0, e.err } @@ -101,3 +102,85 @@ func TestHandleGitSSHSignature_GRPCSuccess_ReturnsJSON200(t *testing.T) { require.NoError(t, err) assert.Equal(t, "abc123", body["signature"]) } + +func TestOwnerEndpoint_ReturnsConfiguredOwner(t *testing.T) { + mock := &mockCredentialsClient{} + handler := newCredentialsHandler(context.Background(), mock, "alice") + + req := httptest.NewRequest(http.MethodGet, ownerPath, nil) + w := httptest.NewRecorder() + handler.ServeHTTP(w, req) + + resp := w.Result() + assert.Equal(t, http.StatusOK, resp.StatusCode) + body, err := io.ReadAll(resp.Body) + require.NoError(t, err) + assert.Equal(t, "alice", string(body)) +} + +func TestFetchOwner_ReturnsConfiguredOwner(t *testing.T) { + ln, err := net.Listen("tcp", "localhost:0") + require.NoError(t, err) + port := ln.Addr().(*net.TCPAddr).Port + + ctx := t.Context() + go func() { + _ = RunCredentialsServerWithListener(ctx, ln, &mockCredentialsClient{}, "bob") + }() + require.NoError(t, waitForServer(ctx, port)) + + owner, err := FetchOwner(context.Background(), port) + require.NoError(t, err) + assert.Equal(t, "bob", owner) +} + +func TestFetchOwner_EmptyWhenEndpointMissing(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + w.WriteHeader(http.StatusNotFound) + })) + defer server.Close() + + var port int + _, err := fmt.Sscanf(server.URL, "http://127.0.0.1:%d", &port) + require.NoError(t, err) + + owner, err := FetchOwner(context.Background(), port) + require.NoError(t, err) + assert.Empty(t, owner) +} + +func TestFetchOwner_DoesNotFollowRedirects(t *testing.T) { + evil := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + _, _ = w.Write([]byte("evil-owner")) + })) + defer evil.Close() + + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + http.Redirect(w, r, evil.URL, http.StatusFound) + })) + defer server.Close() + + var port int + _, err := fmt.Sscanf(server.URL, "http://127.0.0.1:%d", &port) + require.NoError(t, err) + + owner, err := FetchOwner(context.Background(), port) + require.NoError(t, err) + assert.Empty(t, owner, "a redirect must not be followed to another owner value") +} + +func TestFetchOwner_CapsResponseSize(t *testing.T) { + oversized := strings.Repeat("a", maxOwnerResponseSize*2) + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + _, _ = w.Write([]byte(oversized)) + })) + defer server.Close() + + var port int + _, err := fmt.Sscanf(server.URL, "http://127.0.0.1:%d", &port) + require.NoError(t, err) + + owner, err := FetchOwner(context.Background(), port) + require.NoError(t, err) + assert.LessOrEqual(t, len(owner), maxOwnerResponseSize) +} diff --git a/pkg/credentials/start.go b/pkg/credentials/start.go index db53a9e26..6177718f8 100644 --- a/pkg/credentials/start.go +++ b/pkg/credentials/start.go @@ -23,7 +23,7 @@ func StartCredentialsServer( } go func() { - err := RunCredentialsServer(ctx, port, client) + err := RunCredentialsServer(ctx, port, client, "") if err != nil { log.Errorf("error running git credentials server: error=%v", err) }