mirror of
https://github.com/henrygd/beszel.git
synced 2026-10-02 06:17:47 +02:00
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:
@@ -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
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user