fix(server): prevent control replacement lifecycle leaks (#5424)

This commit is contained in:
fatedier
2026-07-21 18:09:09 +08:00
committed by GitHub
parent fe79598ee4
commit d486018885
10 changed files with 2241 additions and 152 deletions

View File

@@ -18,6 +18,7 @@ import (
"bytes"
"context"
"crypto/tls"
"errors"
"fmt"
"io"
"net"
@@ -51,7 +52,6 @@ import (
"github.com/fatedier/frp/pkg/util/xlog"
"github.com/fatedier/frp/server/controller"
"github.com/fatedier/frp/server/group"
"github.com/fatedier/frp/server/metrics"
"github.com/fatedier/frp/server/ports"
"github.com/fatedier/frp/server/proxy"
"github.com/fatedier/frp/server/registry"
@@ -64,6 +64,8 @@ const (
vhostReadWriteTimeout time.Duration = 30 * time.Second
)
var errControlReplaced = errors.New("control was replaced during login")
func init() {
crypto.DefaultSalt = "frp"
// Disable quic-go's receive buffer warning.
@@ -161,9 +163,10 @@ func NewService(cfg *v1.ServerConfig) (*Service, error) {
return nil, err
}
clientRegistry := registry.NewClientRegistry()
svr := &Service{
ctlManager: NewControlManager(),
clientRegistry: registry.NewClientRegistry(),
ctlManager: NewControlManager(clientRegistry),
clientRegistry: clientRegistry,
pxyManager: proxy.NewManager(),
pluginManager: plugin.NewManager(),
rc: &controller.ResourceController{
@@ -469,6 +472,9 @@ func (svr *Service) handleConnection(ctx context.Context, conn net.Conn, interna
if err != nil {
xl.Warnf("register control error: %v", err)
if ctl != nil {
svr.ctlManager.Remove(ctl)
}
if writeErr := writeWithDeadline(conn, connWriteTimeout, func() error {
return acceptedConn.conn.WriteMsg(&msg.LoginResp{
Version: version.Full(),
@@ -477,29 +483,27 @@ func (svr *Service) handleConnection(ctx context.Context, conn net.Conn, interna
}); writeErr != nil {
xl.Warnf("write login error response error: %v", writeErr)
}
conn.Close()
if ctl != nil {
_ = ctl.Close()
} else {
conn.Close()
}
return
}
if err = writeWithDeadline(conn, connWriteTimeout, func() error {
return acceptedConn.conn.WriteMsg(&msg.LoginResp{
Version: version.Full(),
RunID: ctl.runID,
Error: "",
if err = svr.completeControlLogin(ctl, func() error {
return writeWithDeadline(conn, connWriteTimeout, func() error {
return acceptedConn.conn.WriteMsg(&msg.LoginResp{
Version: version.Full(),
RunID: ctl.runID,
Error: "",
})
})
}); err != nil {
xl.Warnf("write login response error: %v", err)
svr.ctlManager.Del(m.RunID, ctl)
svr.clientRegistry.MarkOfflineByRunID(m.RunID)
conn.Close()
xl.Warnf("complete control login error: %v", err)
svr.ctlManager.Remove(ctl)
_ = ctl.Close()
return
}
ctl.Start()
metrics.Server.NewClient()
go func() {
// block until control closed
ctl.WaitClosed()
svr.ctlManager.Del(m.RunID, ctl)
}()
case *msg.NewWorkConn:
if err := svr.RegisterWorkConn(acceptedConn.conn, m); err != nil {
_ = acceptedConn.conn.WriteMsg(&msg.StartWorkConn{
@@ -527,6 +531,17 @@ func (svr *Service) handleConnection(ctx context.Context, conn net.Conn, interna
}
}
func (svr *Service) completeControlLogin(ctl *Control, writeSuccess func() error) error {
committed, err := svr.ctlManager.completeLogin(ctl, writeSuccess)
if err != nil {
return err
}
if !committed {
return errControlReplaced
}
return nil
}
type acceptedConnection struct {
conn *msg.Conn
wireProtocol string
@@ -768,16 +783,15 @@ func (svr *Service) RegisterControl(
}
ctl, err := NewControl(ctx, &SessionContext{
RC: svr.rc,
PxyManager: svr.pxyManager,
PluginManager: svr.pluginManager,
AuthVerifier: authVerifier,
EncryptionKey: svr.auth.EncryptionKey(),
Conn: ctlConn,
LoginMsg: loginMsg,
ServerCfg: svr.cfg,
ClientRegistry: svr.clientRegistry,
WireProtocol: wireProtocol,
RC: svr.rc,
PxyManager: svr.pxyManager,
PluginManager: svr.pluginManager,
AuthVerifier: authVerifier,
EncryptionKey: svr.auth.EncryptionKey(),
Conn: ctlConn,
LoginMsg: loginMsg,
ServerCfg: svr.cfg,
WireProtocol: wireProtocol,
})
if err != nil {
xl.Warnf("create new controller error: %v", err)
@@ -785,18 +799,17 @@ func (svr *Service) RegisterControl(
return nil, fmt.Errorf("unexpected error when creating new controller")
}
if oldCtl := svr.ctlManager.Add(loginMsg.RunID, ctl); oldCtl != nil {
oldCtl.WaitClosed()
if err := svr.ctlManager.Add(ctl); err != nil {
return ctl, err
}
ctl.WaitForHandoff()
remoteAddr := ctlConn.RemoteAddr().String()
if host, _, err := net.SplitHostPort(remoteAddr); err == nil {
remoteAddr = host
active, err := svr.ctlManager.Activate(ctl)
if err != nil {
return ctl, err
}
_, conflict := svr.clientRegistry.Register(loginMsg.User, loginMsg.ClientID, loginMsg.RunID, loginMsg.Hostname, loginMsg.Version, remoteAddr, wireProtocol)
if conflict {
svr.ctlManager.Del(loginMsg.RunID, ctl)
return nil, fmt.Errorf("client_id [%s] for user [%s] is already online", loginMsg.ClientID, loginMsg.User)
if !active {
return ctl, errControlReplaced
}
return ctl, nil
@@ -830,20 +843,25 @@ func (svr *Service) RegisterWorkConn(workConn *msg.Conn, newMsg *msg.NewWorkConn
xl.Warnf("invalid NewWorkConn with run id [%s]", newMsg.RunID)
return err
}
return ctl.RegisterWorkConn(proxy.NewWorkConn(workConn))
return svr.ctlManager.RegisterWorkConn(ctl, proxy.NewWorkConn(workConn))
}
func (svr *Service) RegisterVisitorConn(visitorConn net.Conn, newMsg *msg.NewVisitorConn, wireProtocol string) error {
visitorUser := ""
admit := func(visitorUser string) error {
return svr.rc.VisitorManager.NewConn(newMsg.ProxyName, visitorConn, newMsg.Timestamp, newMsg.SignKey,
newMsg.UseEncryption, newMsg.UseCompression, visitorUser, wireProtocol)
}
// TODO(deprecation): Compatible with old versions, can be without runID, user is empty. In later versions, it will be mandatory to include runID.
// If runID is required, it is not compatible with versions prior to v0.50.0.
if newMsg.RunID != "" {
ctl, exist := svr.ctlManager.GetByID(newMsg.RunID)
if !exist {
admitted, err := svr.ctlManager.admitVisitorByRunID(newMsg.RunID, admit)
if err != nil {
return err
}
if !admitted {
return fmt.Errorf("no client control found for run id [%s]", newMsg.RunID)
}
visitorUser = ctl.sessionCtx.LoginMsg.User
return nil
}
return svr.rc.VisitorManager.NewConn(newMsg.ProxyName, visitorConn, newMsg.Timestamp, newMsg.SignKey,
newMsg.UseEncryption, newMsg.UseCompression, visitorUser, wireProtocol)
return admit("")
}