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
3 changes: 3 additions & 0 deletions openssh/AGENTS.md
Original file line number Diff line number Diff line change
Expand Up @@ -5,6 +5,9 @@
with a Go-native SSH protocol implementation.
- Keep process launching injectable. Callers own environment and login-shell
policy; the package owns argument construction and connection lifecycle.
- Bind an explicitly supplied runner to the ControlMaster generation it starts
or adopts. Use that runner for every later probe and teardown of the same
generation; never switch execution policy underneath a live master.
- Treat programmatic `Target` values as untrusted until `ValidateTarget`
succeeds before every OpenSSH invocation. Unusual account names belong in
trusted `ssh_config` behind a safe host alias; do not weaken the explicit
Expand Down
6 changes: 6 additions & 0 deletions openssh/THREAT_MODEL.md
Original file line number Diff line number Diff line change
Expand Up @@ -72,6 +72,12 @@ full cleanup window, including a final socket ownership/type inspection at
expiry, so a detached master that binds its socket late is still terminated and
drained before the manager releases ownership.

When a caller supplies a runner for one connection, the manager retains that
runner with the connection generation. Readiness checks, later liveness probes,
replacement teardown, explicit disconnect, and idle cleanup use the same
runner. The caller is responsible for keeping the runner and its captured
execution policy valid for the lifetime of the generation.

### Network and remote peers

The network is untrusted. Host authentication, transport confidentiality, and
Expand Down
77 changes: 62 additions & 15 deletions openssh/manager.go
Original file line number Diff line number Diff line change
Expand Up @@ -53,6 +53,7 @@ type hostEntry struct {
mu sync.Mutex
state string
target Target
runSSH RunSSH
message string
lastActive time.Time
generation Generation
Expand All @@ -64,6 +65,7 @@ type hostEntry struct {
type connectionSnapshot struct {
state string
target Target
runSSH RunSSH
generation Generation
}

Expand Down Expand Up @@ -161,6 +163,7 @@ func snapshotEntry(entry *hostEntry) connectionSnapshot {
return connectionSnapshot{
state: entry.state,
target: entry.target,
runSSH: entry.runSSH,
generation: entry.generation,
}
}
Expand Down Expand Up @@ -225,6 +228,31 @@ func (m *PersistentManager) Connect(
ctx context.Context,
identity string,
target Target,
) (Generation, error) {
return m.connect(ctx, identity, target, m.config.RunSSH)
}

// ConnectWithRunner establishes or adopts the master using runSSH. The runner
// is retained with the resulting generation and is used for every subsequent
// probe and teardown of that master. If identity already has a live generation,
// its retained runner remains authoritative.
func (m *PersistentManager) ConnectWithRunner(
ctx context.Context,
identity string,
target Target,
runSSH RunSSH,
) (Generation, error) {
if runSSH == nil {
return 0, &ConfigError{Destination: target.String(), Reason: "nil SSH runner"}
}
return m.connect(ctx, identity, target, runSSH)
}

func (m *PersistentManager) connect(
ctx context.Context,
identity string,
target Target,
runSSH RunSSH,
) (Generation, error) {
if err := persistentSupportError(); err != nil {
return 0, err
Expand Down Expand Up @@ -280,7 +308,7 @@ func (m *PersistentManager) Connect(
oldSocketPath, pathErr := m.checkedSocketPath(identity, oldTarget)
teardownErr := pathErr
if teardownErr == nil {
teardownErr = m.stopMaster(ctx, oldSocketPath, oldTarget)
teardownErr = m.stopMaster(ctx, oldSocketPath, oldTarget, snapshot.runSSH)
}
m.finishTeardown(identity, entry, generation, done, teardownErr)
if teardownErr != nil {
Expand All @@ -289,12 +317,18 @@ func (m *PersistentManager) Connect(
continue
}
probeTarget := target
probeRunner := runSSH
if activeState(snapshot.state) && snapshot.target.Hostname != "" {
probeTarget = snapshot.target
}
if snapshot.target.Hostname != "" && snapshot.runSSH != nil {
probeRunner = snapshot.runSSH
}
entry.mu.Unlock()

probeState, probeErr := m.probeControlMaster(ctx, socketPath, probeTarget)
probeState, probeErr := m.probeControlMasterWithRunner(
ctx, socketPath, probeTarget, probeRunner,
)
entry.mu.Lock()
if !entryMatches(entry, snapshot) {
entry.mu.Unlock()
Expand All @@ -321,6 +355,7 @@ func (m *PersistentManager) Connect(
generation, done := beginOperation(entry)
entry.state = StateConnected
entry.target = target
entry.runSSH = probeRunner
entry.message = ""
entry.lastActive = time.Now()
entry.operationDone = nil
Expand All @@ -333,6 +368,7 @@ func (m *PersistentManager) Connect(
generation, done := beginOperation(entry)
entry.state = StateConnecting
entry.target = target
entry.runSSH = runSSH
entry.message = ""
entry.lastActive = time.Now()
entry.mu.Unlock()
Expand All @@ -343,7 +379,7 @@ func (m *PersistentManager) Connect(
return 0, m.finishStart(identity, entry, generation, done, target, removeErr)
}
}
if startErr := m.startMaster(ctx, socketPath, target); startErr != nil {
if startErr := m.startMaster(ctx, socketPath, target, runSSH); startErr != nil {
return 0, m.finishStart(identity, entry, generation, done, target, startErr)
}
if finishErr := m.finishStart(identity, entry, generation, done, target, nil); finishErr != nil {
Expand All @@ -357,6 +393,7 @@ func (m *PersistentManager) startMaster(
ctx context.Context,
socketPath string,
target Target,
runSSH RunSSH,
) error {
arguments, err := MasterArguments(
socketPath,
Expand All @@ -366,7 +403,7 @@ func (m *PersistentManager) startMaster(
if err != nil {
return err
}
exitCode, runErr := m.config.RunSSH(ctx, arguments)
exitCode, runErr := runSSH(ctx, arguments)
if runErr != nil || exitCode != 0 {
primary := &CommandError{
Operation: "master start",
Expand All @@ -377,21 +414,23 @@ func (m *PersistentManager) startMaster(
if ctx.Err() != nil {
primary.Err = errors.Join(ctx.Err(), runErr)
}
return m.cleanupFailedStart(socketPath, target, primary)
return m.cleanupFailedStart(socketPath, target, runSSH, primary)
}

establishCtx, cancel := context.WithTimeout(ctx, m.config.EstablishTimeout)
defer cancel()
ticker := time.NewTicker(m.config.EstablishPollInterval)
defer ticker.Stop()
for {
probeState, probeErr := m.probeControlMaster(establishCtx, socketPath, target)
probeState, probeErr := m.probeControlMasterWithRunner(
establishCtx, socketPath, target, runSSH,
)
if probeErr != nil {
primary := probeErr
if establishCtx.Err() != nil {
primary = establishmentContextError(ctx, probeErr)
}
return m.cleanupFailedStart(socketPath, target, primary)
return m.cleanupFailedStart(socketPath, target, runSSH, primary)
}
if probeState == probeAlive {
return nil
Expand All @@ -401,6 +440,7 @@ func (m *PersistentManager) startMaster(
return m.cleanupFailedStart(
socketPath,
target,
runSSH,
establishmentContextError(ctx, nil),
)
case <-ticker.C:
Expand All @@ -419,14 +459,15 @@ func establishmentContextError(ctx context.Context, probeErr error) error {
func (m *PersistentManager) cleanupFailedStart(
socketPath string,
target Target,
runSSH RunSSH,
primary error,
) error {
cleanupCtx, cancel := context.WithTimeout(
context.Background(),
m.config.CleanupTimeout,
)
defer cancel()
cleanupErr := m.cleanupFailedStartMaster(cleanupCtx, socketPath, target)
cleanupErr := m.cleanupFailedStartMaster(cleanupCtx, socketPath, target, runSSH)
if cleanupErr == nil {
return primary
}
Expand All @@ -440,6 +481,7 @@ func (m *PersistentManager) cleanupFailedStartMaster(
ctx context.Context,
socketPath string,
target Target,
runSSH RunSSH,
) error {
ticker := time.NewTicker(m.config.EstablishPollInterval)
defer ticker.Stop()
Expand All @@ -449,7 +491,7 @@ func (m *PersistentManager) cleanupFailedStartMaster(
return err
}
if dialState != socketAbsent {
return m.stopMaster(ctx, socketPath, target)
return m.stopMaster(ctx, socketPath, target, runSSH)
}

select {
Expand All @@ -463,7 +505,7 @@ func (m *PersistentManager) cleanupFailedStartMaster(
}
// The observation context has expired, so give verified teardown
// its own bounded window to stop and drain the late master.
return m.stopMaster(context.Background(), socketPath, target)
return m.stopMaster(context.Background(), socketPath, target, runSSH)
case <-ticker.C:
}
}
Expand Down Expand Up @@ -534,6 +576,7 @@ func (m *PersistentManager) Disconnect(ctx context.Context, identity string) err
return nil
}
target := entry.target
runSSH := entry.runSSH
stopping := entry.state == StateStopping
generation, done := beginOperation(entry)
entry.mu.Unlock()
Expand All @@ -547,7 +590,7 @@ func (m *PersistentManager) Disconnect(ctx context.Context, identity string) err
if stopping {
teardownErr = m.drainStoppedMaster(ctx, socketPath)
} else {
teardownErr = m.stopMaster(ctx, socketPath, target)
teardownErr = m.stopMaster(ctx, socketPath, target, runSSH)
}
}
m.finishTeardown(identity, entry, generation, done, teardownErr)
Expand All @@ -568,6 +611,7 @@ func (m *PersistentManager) finishTeardown(
if err == nil {
entry.state = StateDisconnected
entry.target = Target{}
entry.runSSH = nil
entry.message = ""
entry.lastActive = time.Now()
} else {
Expand Down Expand Up @@ -595,6 +639,7 @@ func (m *PersistentManager) stopMaster(
ctx context.Context,
socketPath string,
target Target,
runSSH RunSSH,
) error {
stopCtx, cancel := context.WithTimeout(ctx, m.config.CleanupTimeout)
defer cancel()
Expand All @@ -613,7 +658,7 @@ func (m *PersistentManager) stopMaster(
if err != nil {
return err
}
exitCode, runErr := m.config.RunSSH(stopCtx, arguments)
exitCode, runErr := runSSH(stopCtx, arguments)
if runErr != nil || exitCode != 0 {
if stopCtx.Err() != nil {
runErr = errors.Join(stopCtx.Err(), runErr)
Expand Down Expand Up @@ -705,9 +750,10 @@ func (m *PersistentManager) IsAlive(
return false, ErrConnectionChanged
}
target := entry.target
runSSH := entry.runSSH
entry.mu.Unlock()
probeState, probeErr := m.probeControlMaster(
ctx, m.SocketPath(identity, target), target,
probeState, probeErr := m.probeControlMasterWithRunner(
ctx, m.SocketPath(identity, target), target, runSSH,
)

entry.mu.Lock()
Expand Down Expand Up @@ -844,6 +890,7 @@ func (m *PersistentManager) disconnectIdleCandidate(
return err
}
target := entry.target
runSSH := entry.runSSH
generation, done := beginOperation(entry)
entry.mu.Unlock()

Expand All @@ -852,7 +899,7 @@ func (m *PersistentManager) disconnectIdleCandidate(
err = ensurePersistentDirectory(m.socketDir)
}
if err == nil {
err = m.stopMaster(ctx, socketPath, target)
err = m.stopMaster(ctx, socketPath, target, runSSH)
}
m.finishTeardown(identity, entry, generation, done, err)
return err
Expand Down
50 changes: 50 additions & 0 deletions openssh/manager_unix_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -375,6 +375,56 @@ func TestPersistentManagerChangesDestinationByTeardownThenConnect(t *testing.T)
assert.Equal("wes@new", manager.Destination("studio"))
}

func TestPersistentManagerBindsRunnerToConnectionGeneration(t *testing.T) {
assert := assert.New(t)
require := require.New(t)
defaultRunner := newFakeSSH()
boundRunner := newFakeSSH()
manager := newTestManager(t, newSocketDir(t), defaultRunner)
t.Cleanup(defaultRunner.closeAll)
t.Cleanup(boundRunner.closeAll)

generation, err := manager.ConnectWithRunner(
context.Background(), "studio", testTarget("wes@studio"), boundRunner.run,
)
require.NoError(err)
alive, err := manager.IsAlive(context.Background(), "studio", generation)
require.NoError(err)
assert.True(alive)
require.NoError(manager.Disconnect(context.Background(), "studio"))

assert.Empty(defaultRunner.callsSnapshot())
assert.Equal(
[]string{"spawn", "check", "check", "exit"},
boundRunner.operationKinds(),
)
}

func TestPersistentManagerUsesOldRunnerForReplacementTeardown(t *testing.T) {
assert := assert.New(t)
require := require.New(t)
defaultRunner := newFakeSSH()
oldRunner := newFakeSSH()
newRunner := newFakeSSH()
manager := newTestManager(t, newSocketDir(t), defaultRunner)
t.Cleanup(defaultRunner.closeAll)
t.Cleanup(oldRunner.closeAll)
t.Cleanup(newRunner.closeAll)

_, err := manager.ConnectWithRunner(
context.Background(), "studio", testTarget("wes@old"), oldRunner.run,
)
require.NoError(err)
_, err = manager.ConnectWithRunner(
context.Background(), "studio", testTarget("wes@new"), newRunner.run,
)
require.NoError(err)

assert.Empty(defaultRunner.callsSnapshot())
assert.Equal([]string{"spawn", "check", "exit"}, oldRunner.operationKinds())
assert.Equal([]string{"spawn", "check"}, newRunner.operationKinds())
}

func TestPersistentManagerArgumentsRemainBoundToOriginalTarget(t *testing.T) {
assert := assert.New(t)
require := require.New(t)
Expand Down
11 changes: 10 additions & 1 deletion openssh/probe.go
Original file line number Diff line number Diff line change
Expand Up @@ -17,6 +17,15 @@ func (m *PersistentManager) probeControlMaster(
ctx context.Context,
socketPath string,
target Target,
) (masterProbeState, error) {
return m.probeControlMasterWithRunner(ctx, socketPath, target, m.config.RunSSH)
}

func (m *PersistentManager) probeControlMasterWithRunner(
ctx context.Context,
socketPath string,
target Target,
runSSH RunSSH,
) (masterProbeState, error) {
dialState, err := inspectControlSocket(ctx, socketPath)
if err != nil {
Expand All @@ -33,7 +42,7 @@ func (m *PersistentManager) probeControlMaster(
if err != nil {
return probeAbsent, err
}
exitCode, runErr := m.config.RunSSH(ctx, arguments)
exitCode, runErr := runSSH(ctx, arguments)
if runErr == nil && exitCode == 0 {
return probeAlive, nil
}
Expand Down