mirror of
https://github.com/henrygd/beszel.git
synced 2026-09-30 21:37:47 +02:00
The updater and on-demand requests kept separate copies of the SSH client (sys.client and the transport's client) and synced them after each request. That allowed a closed client to be reinstalled over a newer one and leaked connections that were replaced without being closed. The updater now dials, opens sessions and tears down timed-out connections through the transport. The dial keeps the TCP keepalive and handshake deadline, and an OnConnect callback handles the per-connection resets. The transport is created under a lock and agentVersion is now atomic.
293 lines
8.0 KiB
Go
293 lines
8.0 KiB
Go
package transport
|
|
|
|
import (
|
|
"context"
|
|
"errors"
|
|
"fmt"
|
|
"io"
|
|
"net"
|
|
"strings"
|
|
"sync"
|
|
"time"
|
|
|
|
"github.com/blang/semver"
|
|
"github.com/fxamacker/cbor/v2"
|
|
"github.com/henrygd/beszel/internal/common"
|
|
"golang.org/x/crypto/ssh"
|
|
)
|
|
|
|
// sshKeepAliveInterval is the TCP keep-alive idle interval for SSH connections
|
|
// to agents. Enabling OS-level keep-alives lets the hub eventually detect a
|
|
// dead peer on an otherwise idle connection instead of trusting it forever.
|
|
// This is a backstop for genuine network death; an application-level wedge
|
|
// (agent process hung while its kernel keeps ACKing) is caught by per-operation
|
|
// timeouts instead (see issue #2041).
|
|
const sshKeepAliveInterval = 30 * time.Second
|
|
|
|
// sshHandshakeTimeout bounds the SSH handshake after the TCP connection is
|
|
// established. ssh.ClientConfig.Timeout only covers the TCP connect, so a peer
|
|
// that accepts the connection but never sends an SSH banner would otherwise
|
|
// block the caller forever (GHSA-h9jh-29rh-w464).
|
|
var sshHandshakeTimeout = 10 * time.Second
|
|
|
|
// SSHTransport implements Transport over SSH connections. It owns the single
|
|
// SSH connection to an agent, which is shared by all requests and sessions.
|
|
type SSHTransport struct {
|
|
mu sync.Mutex
|
|
client *ssh.Client
|
|
config *ssh.ClientConfig
|
|
host string
|
|
port string
|
|
timeout time.Duration
|
|
onConnect func(agentVersion semver.Version)
|
|
}
|
|
|
|
// SSHTransportConfig holds configuration for creating an SSH transport.
|
|
type SSHTransportConfig struct {
|
|
Host string
|
|
Port string
|
|
Config *ssh.ClientConfig
|
|
Timeout time.Duration
|
|
// OnConnect is called when a new connection is established, before it is
|
|
// available to other callers, with the agent version from its SSH server
|
|
// version string. It runs under the transport lock and must not call back
|
|
// into the transport.
|
|
OnConnect func(agentVersion semver.Version)
|
|
}
|
|
|
|
// NewSSHTransport creates a new SSH transport with the given configuration.
|
|
func NewSSHTransport(cfg SSHTransportConfig) *SSHTransport {
|
|
timeout := cfg.Timeout
|
|
if timeout == 0 {
|
|
timeout = 4 * time.Second
|
|
}
|
|
return &SSHTransport{
|
|
config: cfg.Config,
|
|
host: cfg.Host,
|
|
port: cfg.Port,
|
|
timeout: timeout,
|
|
onConnect: cfg.OnConnect,
|
|
}
|
|
}
|
|
|
|
// GetClient returns the current SSH client, or nil if not connected.
|
|
func (t *SSHTransport) GetClient() *ssh.Client {
|
|
t.mu.Lock()
|
|
defer t.mu.Unlock()
|
|
return t.client
|
|
}
|
|
|
|
// Request sends a request to the agent via SSH and unmarshals the response.
|
|
func (t *SSHTransport) Request(ctx context.Context, action common.WebSocketAction, req any, dest any) (err error) {
|
|
if err := ctx.Err(); err != nil {
|
|
return err
|
|
}
|
|
client, err := t.Connect(ctx)
|
|
if err != nil {
|
|
return err
|
|
}
|
|
|
|
// Closing only the session still depends on the peer processing SSH packets.
|
|
// Close the captured connection to release every blocked read/write, including
|
|
// concurrent sessions; subsequent requests can reconnect.
|
|
stop := closeOnCancellation(ctx, func() { t.CloseClient(client) })
|
|
defer func() {
|
|
stop()
|
|
if err != nil && ctx.Err() != nil {
|
|
err = ctx.Err()
|
|
}
|
|
if isConnectionError(err) {
|
|
t.CloseClient(client)
|
|
}
|
|
}()
|
|
|
|
session, err := t.NewSession(ctx, client)
|
|
if err != nil {
|
|
return err
|
|
}
|
|
defer session.Close()
|
|
|
|
stdout, err := session.StdoutPipe()
|
|
if err != nil {
|
|
return err
|
|
}
|
|
stdin, err := session.StdinPipe()
|
|
if err != nil {
|
|
return err
|
|
}
|
|
if err := session.Shell(); err != nil {
|
|
return err
|
|
}
|
|
|
|
// Send request
|
|
hubReq := common.HubRequest[any]{Action: action, Data: req}
|
|
if err := cbor.NewEncoder(stdin).Encode(hubReq); err != nil {
|
|
return fmt.Errorf("failed to encode request: %w", err)
|
|
}
|
|
stdin.Close()
|
|
|
|
// Read response
|
|
var resp common.AgentResponse
|
|
if err := cbor.NewDecoder(stdout).Decode(&resp); err != nil {
|
|
return fmt.Errorf("failed to decode response: %w", err)
|
|
}
|
|
|
|
if resp.Error != "" {
|
|
return errors.New(resp.Error)
|
|
}
|
|
|
|
if err := session.Wait(); err != nil {
|
|
return err
|
|
}
|
|
|
|
return UnmarshalResponse(resp, action, dest)
|
|
}
|
|
|
|
// IsConnected returns true if the SSH connection is active.
|
|
func (t *SSHTransport) IsConnected() bool {
|
|
return t.GetClient() != nil
|
|
}
|
|
|
|
// Close terminates the SSH connection.
|
|
func (t *SSHTransport) Close() {
|
|
t.CloseClient(t.GetClient())
|
|
}
|
|
|
|
// CloseClient closes client and clears it if it is still the current
|
|
// connection, so a late close never discards a replacement connection.
|
|
func (t *SSHTransport) CloseClient(client *ssh.Client) {
|
|
t.mu.Lock()
|
|
if t.client == client {
|
|
t.client = nil
|
|
}
|
|
t.mu.Unlock()
|
|
if client != nil {
|
|
client.Close()
|
|
}
|
|
}
|
|
|
|
// closeOnCancellation stops I/O when ctx is cancelled. The returned function
|
|
// waits for any in-progress close so it cannot outlive the operation.
|
|
func closeOnCancellation(ctx context.Context, closeConn func()) func() {
|
|
done := make(chan struct{})
|
|
stop := context.AfterFunc(ctx, func() {
|
|
closeConn()
|
|
close(done)
|
|
})
|
|
return func() {
|
|
if !stop() {
|
|
<-done
|
|
}
|
|
}
|
|
}
|
|
|
|
// Connect returns the current client or establishes a new cancellable SSH
|
|
// connection, calling OnConnect when a new connection is stored.
|
|
func (t *SSHTransport) Connect(ctx context.Context) (*ssh.Client, error) {
|
|
if client := t.GetClient(); client != nil {
|
|
return client, nil
|
|
}
|
|
if t.config == nil {
|
|
return nil, errors.New("SSH config not set")
|
|
}
|
|
|
|
network := "tcp"
|
|
host := t.host
|
|
if strings.HasPrefix(host, "/") {
|
|
network = "unix"
|
|
} else {
|
|
host = net.JoinHostPort(host, t.port)
|
|
}
|
|
|
|
dialer := net.Dialer{Timeout: t.config.Timeout, KeepAlive: sshKeepAliveInterval}
|
|
conn, err := dialer.DialContext(ctx, network, host)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
_ = conn.SetDeadline(time.Now().Add(sshHandshakeTimeout))
|
|
stop := closeOnCancellation(ctx, func() { conn.Close() })
|
|
sshConn, chans, reqs, err := ssh.NewClientConn(conn, host, t.config)
|
|
stop()
|
|
if ctx.Err() != nil {
|
|
conn.Close()
|
|
return nil, ctx.Err()
|
|
}
|
|
if err != nil {
|
|
conn.Close()
|
|
return nil, err
|
|
}
|
|
// clear the handshake deadline so it doesn't apply to the long-lived connection
|
|
_ = conn.SetDeadline(time.Time{})
|
|
client := ssh.NewClient(sshConn, chans, reqs)
|
|
|
|
t.mu.Lock()
|
|
if existing := t.client; existing != nil {
|
|
t.mu.Unlock()
|
|
client.Close()
|
|
return existing, nil
|
|
}
|
|
// Initialize per-connection state (e.g. the agent version, which selects the
|
|
// protocol) before other callers can reuse the client.
|
|
if t.onConnect != nil {
|
|
agentVersion, _ := extractAgentVersion(string(client.Conn.ServerVersion()))
|
|
t.onConnect(agentVersion)
|
|
}
|
|
t.client = client
|
|
t.mu.Unlock()
|
|
return client, nil
|
|
}
|
|
|
|
// NewSession opens a session on client, bounded by the transport timeout
|
|
// independently of ctx. The connection is closed if session creation stalls.
|
|
func (t *SSHTransport) NewSession(ctx context.Context, client *ssh.Client) (*ssh.Session, error) {
|
|
ctx, cancel := context.WithTimeout(ctx, t.timeout)
|
|
defer cancel()
|
|
stop := closeOnCancellation(ctx, func() { t.CloseClient(client) })
|
|
session, err := client.NewSession()
|
|
stop()
|
|
if ctx.Err() != nil {
|
|
if session != nil {
|
|
session.Close()
|
|
}
|
|
return nil, ctx.Err()
|
|
}
|
|
return session, err
|
|
}
|
|
|
|
// extractAgentVersion extracts the beszel version from SSH server version string.
|
|
func extractAgentVersion(versionString string) (semver.Version, error) {
|
|
_, after, _ := strings.Cut(versionString, "_")
|
|
return semver.Parse(after)
|
|
}
|
|
|
|
// RequestWithRetry sends a request with automatic retry on connection failures.
|
|
func (t *SSHTransport) RequestWithRetry(ctx context.Context, action common.WebSocketAction, req any, dest any, retries int) error {
|
|
var lastErr error
|
|
for attempt := 0; attempt <= retries; attempt++ {
|
|
err := t.Request(ctx, action, req, dest)
|
|
if err == nil {
|
|
return nil
|
|
}
|
|
lastErr = err
|
|
|
|
// Check if it's a connection error that warrants a retry
|
|
if isConnectionError(err) && attempt < retries {
|
|
continue
|
|
}
|
|
return err
|
|
}
|
|
return lastErr
|
|
}
|
|
|
|
// isConnectionError checks if an error indicates a connection problem.
|
|
func isConnectionError(err error) bool {
|
|
if err == nil {
|
|
return false
|
|
}
|
|
errStr := err.Error()
|
|
return strings.Contains(errStr, "connection") ||
|
|
strings.Contains(errStr, "EOF") ||
|
|
strings.Contains(errStr, "closed") ||
|
|
errors.Is(err, io.EOF)
|
|
}
|