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
18 changes: 15 additions & 3 deletions server.go
Original file line number Diff line number Diff line change
Expand Up @@ -98,6 +98,7 @@ type Server struct {
AgentForwardingCallback AgentForwardingCallback // callback for allowing agent forwarding, denies all if nil

ConnectionFailedCallback ConnectionFailedCallback // callback to report connection failures
DisconnectCallback DisconnectCallback // callback after an established SSH connection ends

// Timeout fields use their Default* value when nil. A configured duration
// less than or equal to zero disables that timeout.
Expand Down Expand Up @@ -156,6 +157,7 @@ type connectionSettings struct {
logger log.Logger
connCallback ConnCallback
connectionFailedCallback ConnectionFailedCallback
disconnectCallback DisconnectCallback
handler Handler
ptyCallback PtyCallback
sessionRequestCallback SessionRequestCallback
Expand Down Expand Up @@ -201,6 +203,7 @@ func (srv *Server) connectionSettings() *connectionSettings {
logger: srv.Logger,
connCallback: srv.ConnCallback,
connectionFailedCallback: srv.ConnectionFailedCallback,
disconnectCallback: srv.DisconnectCallback,
handler: handler,
ptyCallback: srv.PtyCallback,
sessionRequestCallback: srv.SessionRequestCallback,
Expand Down Expand Up @@ -908,6 +911,18 @@ func (srv *Server) handleConn(newConn net.Conn, active *activeConn, settings ...
return
}
srv.releaseStartup(active)
ctx.SetValue(ContextKeyConn, sshConn)
applyConnMetadata(ctx, sshConn)
publishAuthPermissions(ctx, sshConn.Permissions)
if connectionSettings.disconnectCallback != nil {
defer func() {
cancel()
closeQuietly(sshConn)
srv.untrackActiveConn(active)
tracked = false
connectionSettings.disconnectCallback(ctx, conn)
}()
}
connectionLimiter := resourceLimiter{
limit: int64(connectionSettings.maxConnections),
active: connectionSettings.authenticatedConnections,
Expand All @@ -923,9 +938,6 @@ func (srv *Server) handleConn(newConn net.Conn, active *activeConn, settings ...
srv.trackConn(sshConn, true)
defer srv.trackConn(sshConn, false)

ctx.SetValue(ContextKeyConn, sshConn)
applyConnMetadata(ctx, sshConn)
publishAuthPermissions(ctx, sshConn.Permissions)
maxSessions := connectionSettings.maxSessionsPerConnection
maxChannels := connectionSettings.maxChannelsPerConnection
globalChannelLimiter := &resourceLimiter{
Expand Down
191 changes: 190 additions & 1 deletion server_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -86,6 +86,120 @@ func TestHandleConnInitializesHostSigner(t *testing.T) {
srv.mu.RUnlock()
}

func TestDisconnectCallbackWaitsForConnectionWorkers(t *testing.T) {
type observation struct {
ctxErr error
user string
sessionID string
serverConn *gossh.ServerConn
connectionErr error
}

handlerStarted := make(chan struct{})
handlerCanceled := make(chan struct{})
releaseHandler := make(chan struct{})
var releaseOnce sync.Once
defer releaseOnce.Do(func() { close(releaseHandler) })
disconnected := make(chan observation, 1)
var calls atomic.Int32
srv := &Server{
Handler: func(session Session) {
close(handlerStarted)
<-session.Context().Done()
close(handlerCanceled)
<-releaseHandler
},
DisconnectCallback: func(ctx Context, conn net.Conn) {
calls.Add(1)
_, connectionErr := conn.Write([]byte("closed"))
serverConn, _ := ctx.Value(ContextKeyConn).(*gossh.ServerConn)
disconnected <- observation{
ctxErr: ctx.Err(),
user: ctx.User(),
sessionID: ctx.SessionID(),
serverConn: serverConn,
connectionErr: connectionErr,
}
},
}
session, client, cleanup := newTestSession(t, srv, nil)
defer cleanup()
require.NoError(t, session.Start(""))
<-handlerStarted
require.NoError(t, session.Close())
require.Never(t, func() bool { return calls.Load() != 0 }, 20*time.Millisecond, time.Millisecond)

require.NoError(t, client.Close())
<-handlerCanceled
require.Zero(t, calls.Load(), "callback ran before the connection worker stopped")
releaseOnce.Do(func() { close(releaseHandler) })

select {
case got := <-disconnected:
require.ErrorIs(t, got.ctxErr, context.Canceled)
require.Equal(t, "testuser", got.user)
require.NotEmpty(t, got.sessionID)
require.NotNil(t, got.serverConn)
require.Error(t, got.connectionErr)
case <-time.After(time.Second):
t.Fatal("disconnect callback was not called")
}
require.Never(t, func() bool { return calls.Load() > 1 }, 20*time.Millisecond, time.Millisecond)
}

func TestDisconnectCallbackRunsAfterServerClose(t *testing.T) {
disconnected := make(chan error, 1)
var calls atomic.Int32
srv := &Server{DisconnectCallback: func(ctx Context, _ net.Conn) {
calls.Add(1)
disconnected <- ctx.Err()
}}
session, client, cleanup := newTestSession(t, srv, nil)
defer cleanup()
defer closeQuietly(session)
defer closeQuietly(client)

require.NoError(t, srv.Close())
select {
case err := <-disconnected:
require.ErrorIs(t, err, context.Canceled)
case <-time.After(time.Second):
t.Fatal("disconnect callback was not called after server close")
}
require.NoError(t, srv.Close())
require.Never(t, func() bool { return calls.Load() > 1 }, 20*time.Millisecond, time.Millisecond)
}

func TestDisconnectCallbackCanShutdownServer(t *testing.T) {
shutdownResult := make(chan error, 1)
var srv *Server
srv = &Server{DisconnectCallback: func(Context, net.Conn) {
shutdownResult <- srv.Shutdown(context.Background())
}}
l := newLocalListener()
serveDone := make(chan error, 1)
go func() { serveDone <- srv.Serve(l) }()
t.Cleanup(func() {
_ = srv.Close()
closeQuietly(l)
})
client, err := gossh.Dial("tcp", l.Addr().String(), &gossh.ClientConfig{
User: "user",
HostKeyCallback: gossh.InsecureIgnoreHostKey(),
})
require.NoError(t, err)
defer closeQuietly(client)
require.NoError(t, client.Close())

select {
case err := <-shutdownResult:
require.NoError(t, err)
case <-time.After(time.Second):
t.Fatal("disconnect callback deadlocked while shutting down the server")
}
require.ErrorIs(t, <-serveDone, ErrServerClosed)
}

func TestHandleConnReportsRequiredHostSigner(t *testing.T) {
callbackResult := make(chan error, 1)
srv := &Server{RequireHostSigners: true}
Expand Down Expand Up @@ -287,8 +401,10 @@ func TestConnectionSettingsSnapshot(t *testing.T) {
timeout := time.Second
oldHandlerCalled := false
oldChannelHandlerCalled := false
oldDisconnectCallbackCalled := false
srv := &Server{
Handler: func(Session) { oldHandlerCalled = true },
Handler: func(Session) { oldHandlerCalled = true },
DisconnectCallback: func(Context, net.Conn) { oldDisconnectCallbackCalled = true },
ChannelHandlers: map[string]ChannelHandler{
"test": func(*Server, *gossh.ServerConn, gossh.NewChannel, Context) {
oldChannelHandlerCalled = true
Expand All @@ -299,14 +415,17 @@ func TestConnectionSettingsSnapshot(t *testing.T) {

settings := srv.connectionSettings()
srv.Handle(func(Session) { t.Fatal("snapshot used updated handler") })
srv.DisconnectCallback = func(Context, net.Conn) { t.Fatal("snapshot used updated disconnect callback") }
srv.ChannelHandlers["test"] = func(*Server, *gossh.ServerConn, gossh.NewChannel, Context) {
t.Fatal("snapshot used updated channel handler")
}
timeout = 2 * time.Second

settings.handler(nil)
settings.disconnectCallback(nil, nil)
settings.channelHandlers["test"](nil, nil, nil, nil)
require.True(t, oldHandlerCalled)
require.True(t, oldDisconnectCallbackCalled)
require.True(t, oldChannelHandlerCalled)
require.Equal(t, time.Second, settings.handshakeTimeout)
}
Expand Down Expand Up @@ -1118,11 +1237,79 @@ func TestMaxStartupsRejectsConnectionsAtFullLimit(t *testing.T) {
require.ErrorIs(t, <-serveDone, ErrServerClosed)
}

func TestDisconnectCallbackRunsForAuthenticatedConnectionLimitRejection(t *testing.T) {
type observation struct {
user string
hasServerConn bool
}

maxConnections := 1
disconnected := make(chan observation, 2)
s := &Server{
MaxConnections: &maxConnections,
DisconnectCallback: func(ctx Context, _ net.Conn) {
disconnected <- observation{
user: ctx.User(),
hasServerConn: ctx.Value(ContextKeyConn) != nil,
}
},
}
l := newLocalListener()
serveDone := make(chan error, 1)
go func() { serveDone <- s.Serve(l) }()
var clients []*gossh.Client
t.Cleanup(func() {
for _, client := range clients {
closeQuietly(client)
}
_ = s.Close()
closeQuietly(l)
select {
case <-serveDone:
case <-time.After(time.Second):
t.Error("server did not stop during test cleanup")
}
})
clientConfig := func(user string) *gossh.ClientConfig {
return &gossh.ClientConfig{User: user, HostKeyCallback: gossh.InsecureIgnoreHostKey()}
}

first, err := gossh.Dial("tcp", l.Addr().String(), clientConfig("first"))
require.NoError(t, err)
clients = append(clients, first)
require.Eventually(t, func() bool {
return s.authenticatedConnections.Load() == 1
}, time.Second, time.Millisecond)

second, _ := gossh.Dial("tcp", l.Addr().String(), clientConfig("second"))
if second != nil {
clients = append(clients, second)
closeQuietly(second)
}
select {
case got := <-disconnected:
require.Equal(t, "second", got.user)
require.True(t, got.hasServerConn)
case <-time.After(time.Second):
t.Fatal("connection-limit rejection did not invoke disconnect callback")
}

require.NoError(t, first.Close())
select {
case got := <-disconnected:
require.Equal(t, "first", got.user)
require.True(t, got.hasServerConn)
case <-time.After(time.Second):
t.Fatal("accepted connection did not invoke disconnect callback")
}
}

func TestFailedHandshakeReleasesStartupBeforeCallback(t *testing.T) {
l := newLocalListener()
callbackEntered := make(chan struct{})
releaseCallback := make(chan struct{})
var callbackOnce sync.Once
var disconnectCalls atomic.Int32
s := &Server{
MaxStartups: &MaxStartupsConfig{Start: 1, Full: 1},
ConnectionFailedCallback: func(net.Conn, error) {
Expand All @@ -1131,6 +1318,7 @@ func TestFailedHandshakeReleasesStartupBeforeCallback(t *testing.T) {
<-releaseCallback
})
},
DisconnectCallback: func(Context, net.Conn) { disconnectCalls.Add(1) },
}
serveDone := make(chan error, 1)
go func() { serveDone <- s.Serve(l) }()
Expand All @@ -1145,6 +1333,7 @@ func TestFailedHandshakeReleasesStartupBeforeCallback(t *testing.T) {
case <-time.After(time.Second):
t.Fatal("connection failure callback was not called")
}
require.Zero(t, disconnectCalls.Load())

second, err := net.Dial("tcp", l.Addr().String())
require.NoError(t, err)
Expand Down
6 changes: 6 additions & 0 deletions ssh.go
Original file line number Diff line number Diff line change
Expand Up @@ -91,6 +91,12 @@ type ServerConfigCallback func(ctx Context, config *gossh.ServerConfig)
// Please note: the net.Conn is likely to be closed at this point
type ConnectionFailedCallback func(conn net.Conn, err error)

// DisconnectCallback is called exactly once after a successfully established
// SSH connection ends. The Context is canceled, the connection is closed, and
// connection workers have stopped before the callback runs. Implementations must
// return promptly. Panics from the callback are not recovered.
type DisconnectCallback func(ctx Context, conn net.Conn)

// Window represents the size of a PTY window.
type Window struct {
Width int
Expand Down