Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
19 changes: 12 additions & 7 deletions pkg/agent/delivery/factory.go
Original file line number Diff line number Diff line change
Expand Up @@ -21,6 +21,7 @@ type FactoryOptions struct {
IsRemoteDocker bool
ContainerID string
ExecFunc inject.ExecFunc //nolint:staticcheck // legacy delivery strategies require this type
PodExec PodExecFunc
}

func NewAgentDelivery(opts FactoryOptions) AgentDelivery {
Expand All @@ -38,14 +39,18 @@ func NewAgentDelivery(opts FactoryOptions) AgentDelivery {
}

case driverType == provider.KubernetesDriver:
log.Debugf("using legacy shell delivery for kubernetes driver")
log.Warnf(
"legacy shell delivery is deprecated; platform-native delivery will replace this in a future release",
)
return &LegacyShellDelivery{
ExecFunc: opts.ExecFunc,
DownloadURL: "",
if opts.PodExec == nil {
log.Debugf("kubernetes pod exec unavailable, using legacy shell delivery")
log.Warnf(
"legacy shell delivery is deprecated; platform-native delivery will replace this in a future release",
)
return &LegacyShellDelivery{
ExecFunc: opts.ExecFunc,
DownloadURL: "",
}
}
log.Debugf("using kubernetes-native delivery (exec stream)")
return &KubernetesDelivery{Exec: opts.PodExec}

case opts.IsRemoteDocker:
log.Debugf("using remote docker delivery (docker cp)")
Expand Down
29 changes: 27 additions & 2 deletions pkg/agent/delivery/factory_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -5,6 +5,7 @@ import (
"io"
"testing"

"github.com/devsy-org/devsy/pkg/driver"
"github.com/devsy-org/devsy/pkg/provider"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
Expand Down Expand Up @@ -73,7 +74,28 @@ func TestNewAgentDelivery_CustomDriver(t *testing.T) {
assert.Equal(t, PhasePostStart, d.Phase())
}

func TestNewAgentDelivery_KubernetesDriver(t *testing.T) {
func TestNewAgentDelivery_KubernetesDriver_Native(t *testing.T) {
podExec := func(_ context.Context, _ []string, _ driver.Streams) error {
return nil
}

opts := FactoryOptions{
WorkspaceConfig: &provider.AgentWorkspaceInfo{
Agent: provider.ProviderAgentConfig{
Driver: provider.KubernetesDriver,
},
},
PodExec: podExec,
}

d := NewAgentDelivery(opts)
native, ok := d.(*KubernetesDelivery)
require.True(t, ok)
assert.NotNil(t, native.Exec)
assert.Equal(t, PhasePostStart, d.Phase())
}

func TestNewAgentDelivery_KubernetesDriver_FallsBackWhenNoPodExec(t *testing.T) {
execFn := func(ctx context.Context, cmd string, stdin io.Reader, stdout io.Writer, stderr io.Writer) error {
return nil
}
Expand All @@ -85,10 +107,13 @@ func TestNewAgentDelivery_KubernetesDriver(t *testing.T) {
},
},
ExecFunc: execFn,
// PodExec intentionally nil → legacy fallback.
}

d := NewAgentDelivery(opts)
assert.IsType(t, &LegacyShellDelivery{}, d)
legacy, ok := d.(*LegacyShellDelivery)
require.True(t, ok)
assert.NotNil(t, legacy.ExecFunc)
assert.Equal(t, PhasePostStart, d.Phase())
}

Expand Down
59 changes: 50 additions & 9 deletions pkg/agent/delivery/kubernetes.go
Original file line number Diff line number Diff line change
@@ -1,18 +1,28 @@
package delivery

import (
"bytes"
"context"
"fmt"
"strings"

pkgconfig "github.com/devsy-org/devsy/pkg/config"
"github.com/devsy-org/devsy/pkg/inject"
"github.com/devsy-org/devsy/pkg/driver"
"github.com/devsy-org/devsy/pkg/log"
"github.com/devsy-org/devsy/pkg/version"
)

var _ AgentDelivery = (*KubernetesDelivery)(nil)

// PodExecFunc runs argv in the workspace pod's dev container with the given streams.
type PodExecFunc func(ctx context.Context, argv []string, streams driver.Streams) error

// KubernetesDelivery streams the agent binary into the pod over the cluster's exec API.
type KubernetesDelivery struct {
ExecFunc inject.ExecFunc //nolint:staticcheck // K8s exec routing requires this type
Exec PodExecFunc

// ExpectedVersion defaults to version.GetVersion() when empty.
ExpectedVersion string
}

func (d *KubernetesDelivery) Phase() DeliveryPhase {
Expand All @@ -27,31 +37,62 @@ func (d *KubernetesDelivery) DeliverPostStart(ctx context.Context, opts PostStar
if opts.BinarySource == nil {
return fmt.Errorf("binary source is required for kubernetes delivery")
}
if d.ExecFunc == nil {
if d.Exec == nil {
return fmt.Errorf("exec function is required for kubernetes delivery")
}

destPath := pkgconfig.ContainerDevsyHelperLocation

// Skip delivery when the in-pod binary already matches.
expected := d.expectedVersion()
if actual := d.detectVersion(ctx, destPath); actual != "" && actual == expected {
log.Debugf("remote agent version matches expected version %s, skipping delivery", expected)
return nil
}

binary, err := opts.BinarySource(ctx, opts.Arch)
if err != nil {
return fmt.Errorf("acquire binary: %w", err)
}
defer func() { _ = binary.Close() }()

destPath := pkgconfig.ContainerDevsyHelperLocation
// Write to a temp file and atomically move it into place so a failed stream
// never leaves an executable stub.
script := fmt.Sprintf(
`set -e; t=$(mktemp %s.XXXXXX); cat > "$t" && chmod 755 "$t" && mv "$t" %s || { rm -f "$t"; exit 1; }`,
destPath,
destPath,
`set -e; d=$(dirname %s); mkdir -p "$d"; `+
`t=$(mktemp %s.XXXXXX); `+
`cat > "$t" && chmod 0755 "$t" && mv -f "$t" %s || { rm -f "$t"; exit 1; }`,
destPath, destPath, destPath,
)

if err := d.ExecFunc(ctx, script, binary, nil, nil); err != nil {
if err := d.Exec(ctx, []string{"sh", "-c", script}, driver.Streams{Stdin: binary}); err != nil {
return fmt.Errorf("write binary to container: %w", err)
}

log.Debugf("delivered agent binary to kubernetes container via exec")
log.Debugf("delivered agent binary to pod via kubernetes exec")
return nil
}

func (d *KubernetesDelivery) Cleanup(_ context.Context, _ string) error {
return nil
}

func (d *KubernetesDelivery) expectedVersion() string {
if d.ExpectedVersion != "" {
return d.ExpectedVersion
}
return version.GetVersion()
}

// detectVersion returns the agent version in the pod, or "" if absent or unprobeable.
func (d *KubernetesDelivery) detectVersion(ctx context.Context, destPath string) string {
script := fmt.Sprintf(`[ -x "%s" ] && "%s" --version 2>/dev/null || true`, destPath, destPath)

var stdout bytes.Buffer
err := d.Exec(ctx, []string{"sh", "-c", script}, driver.Streams{Stdout: &stdout})
if err != nil {
log.Debugf("failed to detect agent version in pod: %v", err)
return ""
}
return strings.TrimSpace(stdout.String())
}
129 changes: 87 additions & 42 deletions pkg/agent/delivery/kubernetes_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -9,10 +9,50 @@ import (
"testing"

pkgconfig "github.com/devsy-org/devsy/pkg/config"
"github.com/devsy-org/devsy/pkg/driver"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
)

const testVersion = "v1.2.3"

// recordingExec records each call, replays stdouts[N] as stdout, and returns
// errs[N] on call N (nil when unset).
type recordingExec struct {
calls []recordedCall
stdouts []string
errs []error
}

type recordedCall struct {
argv []string
stdin string
}

func (r *recordingExec) fn(_ context.Context, argv []string, streams driver.Streams) error {
call := recordedCall{argv: argv}
if streams.Stdin != nil {
var buf bytes.Buffer
_, _ = io.Copy(&buf, streams.Stdin)
call.stdin = buf.String()
}
idx := len(r.calls)
r.calls = append(r.calls, call)
if streams.Stdout != nil && idx < len(r.stdouts) {
_, _ = io.WriteString(streams.Stdout, r.stdouts[idx])
}
if idx < len(r.errs) {
return r.errs[idx]
}
return nil
}

func binarySourceFrom(data string) BinarySourceFunc {
return func(_ context.Context, _ string) (io.ReadCloser, error) {
return io.NopCloser(strings.NewReader(data)), nil
}
}

func TestKubernetesDelivery_Phase(t *testing.T) {
d := &KubernetesDelivery{}
assert.Equal(t, PhasePostStart, d.Phase())
Expand All @@ -26,17 +66,14 @@ func TestKubernetesDelivery_DeliverPreStart_ReturnsError(t *testing.T) {
}

func TestKubernetesDelivery_DeliverPostStart_RequiresBinarySource(t *testing.T) {
d := &KubernetesDelivery{
ExecFunc: func(_ context.Context, _ string, _ io.Reader, _ io.Writer, _ io.Writer) error {
return nil
},
}
exec := &recordingExec{}
d := &KubernetesDelivery{Exec: exec.fn}
err := d.DeliverPostStart(context.Background(), PostStartOptions{})
require.Error(t, err)
assert.Contains(t, err.Error(), "binary source is required")
}

func TestKubernetesDelivery_DeliverPostStart_RequiresExecFunc(t *testing.T) {
func TestKubernetesDelivery_DeliverPostStart_RequiresExec(t *testing.T) {
d := &KubernetesDelivery{}
err := d.DeliverPostStart(context.Background(), PostStartOptions{
BinarySource: fakeBinarySource,
Expand All @@ -47,63 +84,71 @@ func TestKubernetesDelivery_DeliverPostStart_RequiresExecFunc(t *testing.T) {

func TestKubernetesDelivery_DeliverPostStart_WritesBinary(t *testing.T) {
binaryData := "test-binary-content"
var capturedCmd string
var capturedStdin bytes.Buffer

execFn := func(_ context.Context, cmd string, stdin io.Reader, _ io.Writer, _ io.Writer) error {
capturedCmd = cmd
if stdin != nil {
_, _ = io.Copy(&capturedStdin, stdin)
}
return nil
}
// Probe returns nothing → deliver.
exec := &recordingExec{stdouts: []string{""}}

d := &KubernetesDelivery{ExecFunc: execFn}
d := &KubernetesDelivery{Exec: exec.fn, ExpectedVersion: testVersion}
err := d.DeliverPostStart(context.Background(), PostStartOptions{
BinarySource: func(_ context.Context, _ string) (io.ReadCloser, error) {
return io.NopCloser(strings.NewReader(binaryData)), nil
},
Arch: "amd64",
BinarySource: binarySourceFrom(binaryData),
Arch: testArch,
})

require.NoError(t, err)

require.Len(t, exec.calls, 2)
destPath := pkgconfig.ContainerDevsyHelperLocation
assert.Contains(t, capturedCmd, "cat >")
assert.Contains(t, capturedCmd, destPath)
assert.Contains(t, capturedCmd, "chmod 755")

assert.Equal(t, binaryData, capturedStdin.String())
probeScript := strings.Join(exec.calls[0].argv, " ")
assert.Contains(t, probeScript, "--version")
assert.Contains(t, probeScript, destPath)

writeScript := strings.Join(exec.calls[1].argv, " ")
assert.Contains(t, writeScript, "mktemp")
assert.Contains(t, writeScript, "chmod 0755")
assert.Contains(t, writeScript, "mv -f")
assert.Contains(t, writeScript, destPath)
assert.Equal(t, binaryData, exec.calls[1].stdin)
}

func TestKubernetesDelivery_DeliverPostStart_BinarySourceError(t *testing.T) {
execFn := func(_ context.Context, _ string, _ io.Reader, _ io.Writer, _ io.Writer) error {
return nil
}
func TestKubernetesDelivery_DeliverPostStart_SkipsWhenVersionMatches(t *testing.T) {
exec := &recordingExec{stdouts: []string{testVersion + "\n"}}

d := &KubernetesDelivery{ExecFunc: execFn}
d := &KubernetesDelivery{Exec: exec.fn, ExpectedVersion: testVersion}
err := d.DeliverPostStart(context.Background(), PostStartOptions{
BinarySource: func(_ context.Context, _ string) (io.ReadCloser, error) {
return nil, fmt.Errorf("download failed")
},
BinarySource: binarySourceFrom("should-not-be-streamed"),
Arch: testArch,
})
require.NoError(t, err)

require.Error(t, err)
assert.Contains(t, err.Error(), "acquire binary")
require.Len(t, exec.calls, 1, "only the version probe should run")
assert.Empty(t, exec.calls[0].stdin)
}

func TestKubernetesDelivery_DeliverPostStart_ExecError(t *testing.T) {
execFn := func(_ context.Context, _ string, _ io.Reader, _ io.Writer, _ io.Writer) error {
return fmt.Errorf("exec failed")
func TestKubernetesDelivery_DeliverPostStart_DeliversWhenProbeErrors(t *testing.T) {
// A failing probe must not abort delivery; the write still succeeds.
probeErr := &recordingExec{
stdouts: []string{""},
errs: []error{fmt.Errorf("probe boom"), nil},
}
d := &KubernetesDelivery{Exec: probeErr.fn, ExpectedVersion: testVersion}

d := &KubernetesDelivery{ExecFunc: execFn}
err := d.DeliverPostStart(context.Background(), PostStartOptions{
BinarySource: fakeBinarySource,
BinarySource: binarySourceFrom("data"),
Arch: testArch,
})
require.NoError(t, err)
assert.Len(t, probeErr.calls, 2)
Comment thread
coderabbitai[bot] marked this conversation as resolved.
}

func TestKubernetesDelivery_DeliverPostStart_BinarySourceError(t *testing.T) {
exec := &recordingExec{stdouts: []string{""}}
d := &KubernetesDelivery{Exec: exec.fn}
err := d.DeliverPostStart(context.Background(), PostStartOptions{
BinarySource: func(_ context.Context, _ string) (io.ReadCloser, error) {
return nil, fmt.Errorf("download failed")
},
})
require.Error(t, err)
assert.Contains(t, err.Error(), "write binary to container")
assert.Contains(t, err.Error(), "acquire binary")
}

func TestKubernetesDelivery_Cleanup_IsNoOp(t *testing.T) {
Expand Down
Loading
Loading