fix(agent): improve WebSocket and SSH fallback handling (#2441)

Make WebSocket reconnect and SSH fallback transitions reliable across asynchronous disconnects, stale callbacks, and overlapping connections. Keep the SSH listener available while disconnected and allow a verified WebSocket connection to take precedence when it recovers.

Co-authored-by: henrygd <hank@henrygd.me>
This commit is contained in:
spatiumstas
2026-10-01 20:33:19 +03:00
committed by henrygd
parent d8c2b1f310
commit 4ffd83677d
10 changed files with 688 additions and 65 deletions

View File

@@ -34,11 +34,26 @@ type ServerOptions struct {
// and begins listening for connections. Returns an error if the server
// is already running or if there's an issue starting the server.
func (a *Agent) StartServer(opts ServerOptions) error {
server, listener, err := a.prepareSSHServer(opts)
if err != nil {
return err
}
return a.serveSSHServer(server, listener)
}
var errSSHServerRunning = errors.New("server already started")
// prepareSSHServer binds the listener before Serve starts so a concurrent stop
// can always close it, including when the WebSocket wins the connection race.
func (a *Agent) prepareSSHServer(opts ServerOptions) (*ssh.Server, net.Listener, error) {
a.serverMu.Lock()
defer a.serverMu.Unlock()
if disableSSH, _ := utils.GetEnv("DISABLE_SSH"); disableSSH == "true" {
return errors.New("SSH disabled")
return nil, nil, errors.New("SSH disabled")
}
if a.server != nil {
return errors.New("server already started")
return nil, nil, errSSHServerRunning
}
slog.Info("Starting SSH server", "addr", opts.Addr, "network", opts.Network)
@@ -46,21 +61,18 @@ func (a *Agent) StartServer(opts ServerOptions) error {
if opts.Network == "unix" {
// remove existing socket file if it exists
if err := os.Remove(opts.Addr); err != nil && !os.IsNotExist(err) {
return err
return nil, nil, err
}
}
// start listening on the address
ln, err := net.Listen(opts.Network, opts.Addr)
if err != nil {
return err
return nil, nil, err
}
defer ln.Close()
// set default handler
ssh.Handle(a.handleSession)
a.server = &ssh.Server{
server := &ssh.Server{
Handler: a.handleSession,
ServerConfigCallback: newSSHServerConfig,
// check public key(s)
PublicKeyHandler: func(ctx ssh.Context, key ssh.PublicKey) bool {
@@ -82,8 +94,20 @@ func (a *Agent) StartServer(opts ServerOptions) error {
IdleTimeout: 70 * time.Second,
}
// Start SSH server on the listener
return a.server.Serve(ln)
a.server = server
a.serverListener = ln
return server, ln, nil
}
func (a *Agent) serveSSHServer(server *ssh.Server, listener net.Listener) error {
err := server.Serve(listener)
a.serverMu.Lock()
if a.server == server {
a.server = nil
a.serverListener = nil
}
a.serverMu.Unlock()
return err
}
// newSSHServerConfig returns a separate config for each connection because
@@ -115,9 +139,8 @@ func (a *Agent) getHubVersion(sessionCtx ssh.Context) semver.Version {
// appropriate encoding format based on hub version, and exits with appropriate
// status codes.
func (a *Agent) handleSession(s ssh.Session) {
a.connectionManager.eventChan <- SSHConnect
sessionCtx := s.Context()
a.connectionManager.sshConnectionOpened(sessionCtx)
hubVersion := a.getHubVersion(sessionCtx)
@@ -164,12 +187,13 @@ func (a *Agent) handleSSHRequest(w io.Writer, req *common.HubRequest[cbor.RawMes
}
ctx := &HandlerContext{
Client: nil,
Agent: a,
Request: req,
RequestID: nil,
HubVerified: true,
SendResponse: sshResponder,
Client: nil,
Agent: a,
Request: req,
RequestID: nil,
HubVerified: true,
ConnectionType: system.ConnectionTypeSSH,
SendResponse: sshResponder,
}
if handler, ok := a.handlerRegistry.GetHandler(req.Action); ok {
@@ -184,7 +208,9 @@ func (a *Agent) handleSSHRequest(w io.Writer, req *common.HubRequest[cbor.RawMes
// handleLegacyStats serves the legacy one-shot stats payload for older hubs
func (a *Agent) handleLegacyStats(w io.Writer, hubVersion semver.Version) error {
stats := a.gatherStats(common.DataRequestOptions{CacheTimeMs: defaultDataCacheTimeMs})
return a.writeToSession(w, stats, hubVersion)
response := *stats
response.Info.ConnectionType = system.ConnectionTypeSSH
return a.writeToSession(w, &response, hubVersion)
}
// writeToSession encodes and writes system statistics to the session.
@@ -261,12 +287,21 @@ func GetNetwork(addr string) string {
// StopServer stops the SSH server if it's running.
// It returns an error if the server is not running or if there's an error stopping it.
func (a *Agent) StopServer() error {
a.serverMu.Lock()
if a.server == nil {
a.serverMu.Unlock()
return errors.New("SSH server not running")
}
server := a.server
listener := a.serverListener
a.server = nil
a.serverListener = nil
a.serverMu.Unlock()
slog.Info("Stopping SSH server")
_ = a.server.Close()
a.server = nil
if listener != nil {
_ = listener.Close()
}
_ = server.Close()
return nil
}