From dabd2c922f49251ba1a29f39bbe88e7f015689b8 Mon Sep 17 00:00:00 2001 From: "Michael (Parker) Parker" Date: Mon, 5 Jan 2026 21:19:27 -0500 Subject: [PATCH 1/3] Implement pterodactyl security fixes --- router/router.go | 1 + router/router_server.go | 3 +- router/router_server_ws.go | 104 ++++++++++++++++++++++++++++------ router/router_system.go | 31 ++++++++++ router/tokens/websocket.go | 22 +++++++ router/websocket/limiter.go | 91 +++++++++++++++++++++++++++++ router/websocket/listeners.go | 4 +- router/websocket/message.go | 5 +- router/websocket/websocket.go | 23 +++++++- server/connections.go | 19 +++++++ server/server.go | 16 +++++- server/websockets.go | 7 +++ sftp/handler.go | 8 ++- sftp/server.go | 53 +++++++++-------- system/context_bag.go | 58 +++++++++++++++++++ 15 files changed, 392 insertions(+), 53 deletions(-) create mode 100644 router/websocket/limiter.go create mode 100644 server/connections.go create mode 100644 system/context_bag.go diff --git a/router/router.go b/router/router.go index 86cea0d3..e0d0cd04 100644 --- a/router/router.go +++ b/router/router.go @@ -66,6 +66,7 @@ func Configure(m *wserver.Manager, client remote.Client) *gin.Engine { protected.GET("/api/servers", getAllServers) protected.POST("/api/servers", postCreateServer) protected.DELETE("/api/transfers/:server", deleteTransfer) + protected.POST("/api/deauthorize-user", postDeauthorizeUser) // These are server specific routes, and require that the request be authorized, and // that the server exist on the Daemon. diff --git a/router/router_server.go b/router/router_server.go index 47c4eb46..f186ccfa 100644 --- a/router/router_server.go +++ b/router/router_server.go @@ -75,7 +75,6 @@ func getServerInstallLogs(c *gin.Context) { c.JSON(http.StatusOK, gin.H{"data": output}) } - // Handles a request to control the power state of a server. If the action being passed // through is invalid a 404 is returned. Otherwise, a HTTP/202 Accepted response is returned // and the actual power action is run asynchronously so that we don't have to block the @@ -295,6 +294,8 @@ func deleteServer(c *gin.Context) { // Adds any of the JTIs passed through in the body to the deny list for the websocket // preventing any JWT generated before the current time from being used to connect to // the socket or send along commands. +// +// deprecated: prefer /api/deauthorize-user func postServerDenyWSTokens(c *gin.Context) { var data struct { JTIs []string `json:"jtis"` diff --git a/router/router_server_ws.go b/router/router_server_ws.go index 6835c769..39d0d70d 100644 --- a/router/router_server_ws.go +++ b/router/router_server_ws.go @@ -2,14 +2,17 @@ package router import ( "context" + "encoding/json" + "net/http" "time" + "emperror.dev/errors" "github.com/gin-gonic/gin" - "github.com/goccy/go-json" ws "github.com/gorilla/websocket" - "github.com/pelican-dev/wings/router/middleware" "github.com/pelican-dev/wings/router/websocket" + "github.com/pelican-dev/wings/server" + "golang.org/x/time/rate" ) var expectedCloseCodes = []int{ @@ -25,6 +28,27 @@ func getServerWebsocket(c *gin.Context) { manager := middleware.ExtractManager(c) s, _ := manager.Get(c.Param("server")) + // Limit the total number of websockets that can be opened at any one time for + // a server instance. This applies across all users connected to the server, and + // is not applied on a per-user basis. + // + // todo: it would be great to make this per-user instead, but we need to modify + // how we even request this endpoint in order for that to be possible. Some type + // of signed identifier in the URL that is verified on this end and set by the + // panel using a shared secret is likely the easiest option. The benefit of that + // is that we can both scope things to the user before authentication, and also + // verify that the JWT provided by the panel is assigned to the same user. + if s.Websockets().Len() >= 30 { + c.AbortWithStatusJSON(http.StatusBadRequest, gin.H{ + "error": "Too many open websocket connections.", + }) + + return + } + + c.Header("Content-Security-Policy", "default-src 'self'") + c.Header("X-Frame-Options", "DENY") + // Create a context that can be canceled when the user disconnects from this // socket that will also cancel listeners running in separate threads. If the // connection itself is terminated listeners using this context will also be @@ -37,36 +61,61 @@ func getServerWebsocket(c *gin.Context) { middleware.CaptureAndAbort(c, err) return } - defer handler.Connection.Close() // Track this open connection on the server so that we can close them all programmatically // if the server is deleted. s.Websockets().Push(handler.Uuid(), &cancel) handler.Logger().Debug("opening connection to server websocket") + defer s.Websockets().Remove(handler.Uuid()) - defer func() { - s.Websockets().Remove(handler.Uuid()) - handler.Logger().Debug("closing connection to server websocket") + go func() { + select { + // When the main context is canceled (through disconnect, server deletion, or server + // suspension) close the connection itself. + case <-ctx.Done(): + handler.Logger().Debug("closing connection to server websocket") + if err := handler.Connection.Close(); err != nil { + handler.Logger().WithError(err).Error("failed to close websocket connection") + } + break + } }() - // If the server is deleted we need to send a close message to the connected client - // so that they disconnect since there will be no more events sent along. Listen for - // the request context being closed to break this loop, otherwise this routine will - // be left hanging in the background. go func() { select { case <-ctx.Done(): - break + return + // If the server is deleted we need to send a close message to the connected client + // so that they disconnect since there will be no more events sent along. Listen for + // the request context being closed to break this loop, otherwise this routine will + //be left hanging in the background. case <-s.Context().Done(): - _ = handler.Connection.WriteControl(ws.CloseMessage, ws.FormatCloseMessage(ws.CloseGoingAway, "server deleted"), time.Now().Add(time.Second*5)) + cancel() break } }() - for { - j := websocket.Message{} + // Due to how websockets are handled we need to connect to the socket + // and _then_ abort it if the server is suspended. You cannot capture + // the HTTP response in the websocket client, thus we connect and then + // immediately close with failure. + if s.IsSuspended() { + _ = handler.Connection.WriteMessage(ws.CloseMessage, ws.FormatCloseMessage(4409, "server is suspended")) - _, p, err := handler.Connection.ReadMessage() + return + } + + // There is a separate rate limiter that applies to individual message types + // within the actual websocket logic handler. _This_ rate limiter just exists + // to avoid enormous floods of data through the socket since we need to parse + // JSON each time. This rate limit realistically should never be hit since this + // would require sending 50+ messages a second over the websocket (no more than + // 10 per 200ms). + var throttled bool + rl := rate.NewLimiter(rate.Every(time.Millisecond*200), 10) + + for { + t, p, err := handler.Connection.ReadMessage() if err != nil { if ws.IsUnexpectedCloseError(err, expectedCloseCodes...) { handler.Logger().WithField("error", err).Warn("error handling websocket message for server") @@ -74,16 +123,39 @@ func getServerWebsocket(c *gin.Context) { break } + if !rl.Allow() { + if !throttled { + throttled = true + _ = handler.Connection.WriteJSON(websocket.Message{Event: websocket.ThrottledEvent, Args: []string{"global"}}) + } + continue + } + + throttled = false + + // If the message isn't a format we expect, or the length of the message is far larger + // than we'd ever expect, drop it. The websocket upgrader logic does enforce a maximum + // _compressed_ message size of 4Kb but that could decompress to a much larger amount + // of data. + if t != ws.TextMessage || len(p) > 32_768 { + continue + } + // Discard and JSON parse errors into the void and don't continue processing this // specific socket request. If we did a break here the client would get disconnected // from the socket, which is NOT what we want to do. + var j websocket.Message if err := json.Unmarshal(p, &j); err != nil { continue } go func(msg websocket.Message) { if err := handler.HandleInbound(ctx, msg); err != nil { - _ = handler.SendErrorJson(msg, err) + if errors.Is(err, server.ErrSuspended) { + cancel() + } else { + _ = handler.SendErrorJson(msg, err) + } } }(j) } diff --git a/router/router_system.go b/router/router_system.go index 1cc56c7d..cdc5bd58 100644 --- a/router/router_system.go +++ b/router/router_system.go @@ -14,6 +14,7 @@ import ( "github.com/pelican-dev/wings/config" "github.com/pelican-dev/wings/internal/diagnostics" "github.com/pelican-dev/wings/router/middleware" + "github.com/pelican-dev/wings/router/tokens" "github.com/pelican-dev/wings/server" "github.com/pelican-dev/wings/server/installer" "github.com/pelican-dev/wings/system" @@ -256,3 +257,33 @@ func postUpdateConfiguration(c *gin.Context) { Applied: true, }) } + +func postDeauthorizeUser(c *gin.Context) { + var data struct { + User string `json:"user"` + Servers []string `json:"servers"` + } + + if err := c.BindJSON(&data); err != nil { + return + } + + // todo: disconnect websockets more gracefully + m := middleware.ExtractManager(c) + if len(data.Servers) > 0 { + for _, uuid := range data.Servers { + if s, ok := m.Get(uuid); ok { + s.Websockets().CancelAll() + s.Sftp().Cancel(data.User) + tokens.DenyForServer(s.ID(), data.User) + } + } + } else { + for _, s := range m.All() { + s.Websockets().CancelAll() + s.Sftp().Cancel(data.User) + } + } + + c.Status(http.StatusNoContent) +} diff --git a/router/tokens/websocket.go b/router/tokens/websocket.go index 017f8ab2..2ad3b9df 100644 --- a/router/tokens/websocket.go +++ b/router/tokens/websocket.go @@ -24,16 +24,29 @@ var wingsBootTime = time.Now() // This is used to allow the Panel to revoke tokens en-masse for a given user & server // combination since the JTI for tokens is just MD5(user.id + server.uuid). When a server // is booted this listing is fetched from the panel and the Websocket is dynamically updated. +// +// deprecated: prefer use of userDenylist var denylist sync.Map +var userDenylist sync.Map // Adds a JTI to the denylist by marking any JWTs generated before the current time as // being invalid if they use the same JTI. +// +// deprecated: prefer the use of DenyForServer func DenyJTI(jti string) { log.WithField("jti", jti).Debugf("adding \"%s\" to JTI denylist", jti) denylist.Store(jti, time.Now()) } +// DenyForServer adds a user UUID to the denylist marking any existing JWTs issued +// to the user as being invalid. This is associated with the user. +func DenyForServer(s string, u string) { + log.WithField("user_uuid", u).WithField("server_uuid", s).Debugf("denying all JWTs created at or before current time for user \"%s\"", u) + + userDenylist.Store(strings.Join([]string{s, u}, ":"), time.Now()) +} + // WebsocketPayload defines the JWT payload for a websocket connection. This JWT is passed along to // the websocket after it has been connected to by sending an "auth" event. type WebsocketPayload struct { @@ -79,12 +92,21 @@ func (p *WebsocketPayload) Denylisted() bool { // Finally, if the token was issued before a time that is currently denied for this // token instance, ignore the permissions response. + // + // This list is deprecated, but we maintain the check here so that custom instances + // are able to continue working. We'll remove it in a future release. if t, ok := denylist.Load(p.JWTID); ok { if p.IssuedAt.Time.Before(t.(time.Time)) { return true } } + if t, ok := userDenylist.Load(strings.Join([]string{p.ServerUUID, p.UserUUID}, ":")); ok { + if p.IssuedAt.Time.Before(t.(time.Time)) { + return true + } + } + return false } diff --git a/router/websocket/limiter.go b/router/websocket/limiter.go new file mode 100644 index 00000000..57315a96 --- /dev/null +++ b/router/websocket/limiter.go @@ -0,0 +1,91 @@ +package websocket + +import ( + "sync" + "time" + + "golang.org/x/time/rate" +) + +type LimiterBucket struct { + mu sync.RWMutex + limits map[Event]*rate.Limiter + throttles map[Event]bool +} + +func (h *Handler) IsThrottled(e Event) bool { + l := h.limiter.For(e) + + h.limiter.mu.Lock() + defer h.limiter.mu.Unlock() + + if l.Allow() { + h.limiter.throttles[e] = false + + return false + } + + // If not allowed, track the throttling and send an event over the wire + // if one wasn't already sent in the same throttling period. + if v, ok := h.limiter.throttles[e]; !v || !ok { + h.limiter.throttles[e] = true + h.Logger().WithField("event", e).Debug("throttling websocket due to event volume") + + _ = h.unsafeSendJson(&Message{Event: ThrottledEvent, Args: []string{string(e)}}) + } + + return true +} + +func NewLimiter() *LimiterBucket { + return &LimiterBucket{ + limits: make(map[Event]*rate.Limiter, 4), + throttles: make(map[Event]bool, 4), + } +} + +// For returns the internal rate limiter for the given event type. In most +// cases this is a shared rate limiter for events, but certain "heavy" or low-frequency +// events implement their own limiters. +func (l *LimiterBucket) For(e Event) *rate.Limiter { + name := limiterName(e) + + l.mu.RLock() + if v, ok := l.limits[name]; ok { + l.mu.RUnlock() + return v + } + + l.mu.RUnlock() + l.mu.Lock() + defer l.mu.Unlock() + + limit, burst := limitValuesFor(e) + l.limits[name] = rate.NewLimiter(limit, burst) + + return l.limits[name] +} + +// limitValuesFor returns the underlying limit and burst value for the given event. +func limitValuesFor(e Event) (rate.Limit, int) { + // Twice every five seconds. + if e == AuthenticationEvent || e == SendServerLogsEvent { + return rate.Every(time.Second * 5), 2 + } + + // 10 per second. + if e == SendCommandEvent { + return rate.Every(time.Second), 10 + } + + // 4 per second. + return rate.Every(time.Second), 4 +} + +func limiterName(e Event) Event { + if e == AuthenticationEvent || e == SendServerLogsEvent || e == SendCommandEvent { + return e + } + + return "_default" +} diff --git a/router/websocket/listeners.go b/router/websocket/listeners.go index 06236cc0..5183b3a5 100644 --- a/router/websocket/listeners.go +++ b/router/websocket/listeners.go @@ -129,7 +129,7 @@ func (h *Handler) listenForServerEvents(ctx context.Context) error { continue } var sendErr error - message := Message{Event: e.Topic} + message := Message{Event: Event(e.Topic)} if str, ok := e.Data.(string); ok { message.Args = []string{str} } else if b, ok := e.Data.([]byte); ok { @@ -147,7 +147,7 @@ func (h *Handler) listenForServerEvents(ctx context.Context) error { continue } } - onError(message.Event, sendErr) + onError(string(message.Event), sendErr) } break } diff --git a/router/websocket/message.go b/router/websocket/message.go index 85fb77f3..04f3fe9d 100644 --- a/router/websocket/message.go +++ b/router/websocket/message.go @@ -1,5 +1,7 @@ package websocket +type Event string + const ( AuthenticationSuccessEvent = "auth success" TokenExpiringEvent = "token expiring" @@ -11,11 +13,12 @@ const ( SendStatsEvent = "send stats" ErrorEvent = "daemon error" JwtErrorEvent = "jwt error" + ThrottledEvent = Event("throttled") ) type Message struct { // The event to perform. - Event string `json:"event"` + Event Event `json:"event"` // The data to pass along, only used by power/command currently. Other requests // should either omit the field or pass an empty value as it is ignored. diff --git a/router/websocket/websocket.go b/router/websocket/websocket.go index be12c464..420a15bd 100644 --- a/router/websocket/websocket.go +++ b/router/websocket/websocket.go @@ -8,8 +8,6 @@ import ( "sync" "time" - "github.com/pelican-dev/wings/internal/models" - "emperror.dev/errors" "github.com/apex/log" "github.com/gbrlsnchs/jwt/v3" @@ -23,6 +21,7 @@ import ( "github.com/pelican-dev/wings/config" "github.com/pelican-dev/wings/environment" "github.com/pelican-dev/wings/environment/docker" + "github.com/pelican-dev/wings/internal/models" "github.com/pelican-dev/wings/router/tokens" "github.com/pelican-dev/wings/server" ) @@ -46,6 +45,7 @@ type Handler struct { server *server.Server ra server.RequestActivity uuid uuid.UUID + limiter *LimiterBucket } var ( @@ -84,6 +84,7 @@ func NewTokenPayload(token []byte) (*tokens.WebsocketPayload, error) { // GetHandler returns a new websocket handler using the context provided. func GetHandler(s *server.Server, w http.ResponseWriter, r *http.Request, c *gin.Context) (*Handler, error) { upgrader := websocket.Upgrader{ + EnableCompression: true, // Ensure that the websocket request is originating from the Panel itself, // and not some other location. CheckOrigin: func(r *http.Request) bool { @@ -110,12 +111,16 @@ func GetHandler(s *server.Server, w http.ResponseWriter, r *http.Request, c *gin return nil, err } + conn.SetReadLimit(4096) + _ = conn.SetCompressionLevel(5) + return &Handler{ Connection: conn, jwt: nil, server: s, ra: s.NewRequestActivity("", c.ClientIP()), uuid: u, + limiter: NewLimiter(), }, nil } @@ -150,7 +155,7 @@ func (h *Handler) SendJson(v Message) error { // If the user does not have permission to see backup events, do not emit // them over the socket. - if strings.HasPrefix(v.Event, server.BackupCompletedEvent) { + if strings.HasPrefix(string(v.Event), server.BackupCompletedEvent) { if !j.HasPermission(PermissionReceiveBackups) { return nil } @@ -277,6 +282,14 @@ func (h *Handler) setJwt(token *tokens.WebsocketPayload) { // HandleInbound handles an inbound socket request and route it to the proper action. func (h *Handler) HandleInbound(ctx context.Context, m Message) error { + if h.server.IsSuspended() { + return server.ErrSuspended + } + + if h.IsThrottled(m.Event) { + return nil + } + if m.Event != AuthenticationEvent { if err := h.TokenValid(); err != nil { h.unsafeSendJson(Message{ @@ -287,6 +300,10 @@ func (h *Handler) HandleInbound(ctx context.Context, m Message) error { } } + if h.server.IsSuspended() { + return server.ErrSuspended + } + switch m.Event { case AuthenticationEvent: { diff --git a/server/connections.go b/server/connections.go new file mode 100644 index 00000000..846ae444 --- /dev/null +++ b/server/connections.go @@ -0,0 +1,19 @@ +package server + +import ( + "github.com/pelican-dev/wings/system" +) + +// Sftp returns the SFTP connection bag for the server instance. This bag tracks +// all open SFTP connections by individual user and allows for a single user or +// all users to be disconnected by other processes. +func (s *Server) Sftp() *system.ContextBag { + s.Lock() + defer s.Unlock() + + if s.sftpBag == nil { + s.sftpBag = system.NewContextBag(s.Context()) + } + + return s.sftpBag +} diff --git a/server/server.go b/server/server.go index 8e0b9518..58fa339e 100644 --- a/server/server.go +++ b/server/server.go @@ -69,6 +69,7 @@ type Server struct { // The console throttler instance used to control outputs. throttler *ConsoleThrottle throttleOnce sync.Once + sftpBag *system.ContextBag // Tracks open websocket connections for the server. wsBag *WebsocketBag @@ -166,7 +167,6 @@ func DetermineServerTimezone(envvars map[string]interface{}, defaultTimezone str return defaultTimezone } - // parseInvocation parses the start command in the same way we already do in the entrypoint // We can use this to set the container command with all variables replaced. func parseInvocation(invocation string, envvars map[string]interface{}, memory int64, port int, ip string) (parsed string) { @@ -191,7 +191,7 @@ func parseInvocation(invocation string, envvars map[string]interface{}, memory i invocation = strings.Replace(invocation, segment, tempSegments[i], 1) } - // Replace the placeholders outside of protected segments + // Replace the placeholders outside protected segments invocation = strings.ReplaceAll(invocation, placeholder, fmt.Sprint(varval)) // Restore protected segments @@ -201,6 +201,10 @@ func parseInvocation(invocation string, envvars map[string]interface{}, memory i } // Replace the defaults with their configured values. + // and any connected SFTP clients. We don't need to worry about revoking any JWTs + // here since they'll be blocked from re-connecting to the websocket anyways. This + // just forces the client to disconnect and attempt to reconnect (rather than waiting + // on them to send a message and hit that disconnect logic). invocation = strings.ReplaceAll(invocation, "${SERVER_PORT}", strconv.Itoa(port)) invocation = strings.ReplaceAll(invocation, "${SERVER_MEMORY}", strconv.Itoa(int(memory))) invocation = strings.ReplaceAll(invocation, "${SERVER_IP}", ip) @@ -263,11 +267,17 @@ func (s *Server) Sync() error { s.SyncWithEnvironment() + // If the server is suspended immediately disconnect all open websocket connections. + if s.IsSuspended() { + s.Websockets().CancelAll() + s.Sftp().CancelAll() + } + return nil } // SyncWithConfiguration accepts a configuration object for a server and will -// sync all of the values with the existing server state. This only replaces the +// sync all values with the existing server state. This only replaces the // existing configuration and process configuration for the server. The // underlying environment will not be affected. This is because this function // can be called from scoped where the server may not be fully initialized, diff --git a/server/websockets.go b/server/websockets.go index e86f88cc..6aa03435 100644 --- a/server/websockets.go +++ b/server/websockets.go @@ -25,6 +25,13 @@ func (s *Server) Websockets() *WebsocketBag { return s.wsBag } +func (w *WebsocketBag) Len() int { + w.mu.Lock() + defer w.mu.Unlock() + + return len(w.conns) +} + // Push adds a new websocket connection to the end of the stack. func (w *WebsocketBag) Push(u uuid.UUID, cancel *context.CancelFunc) { w.mu.Lock() diff --git a/sftp/handler.go b/sftp/handler.go index aed12dbc..04cb3ad2 100644 --- a/sftp/handler.go +++ b/sftp/handler.go @@ -107,7 +107,7 @@ func (h *Handler) Filewrite(request *sftp.Request) (io.WriterAt, error) { h.mu.Lock() defer h.mu.Unlock() - + if err := h.fs.IsIgnored(request.Filepath); err != nil { return nil, err } @@ -158,7 +158,7 @@ func (h *Handler) Filecmd(request *sftp.Request) error { if err := h.fs.IsIgnored(request.Filepath); err != nil { return err } - + switch request.Method { // Allows a user to make changes to the permissions of a given file or directory // on their server using their SFTP client. @@ -312,3 +312,7 @@ func (h *Handler) can(permission string) bool { } return false } + +func (h *Handler) User() string { + return h.events.user +} diff --git a/sftp/server.go b/sftp/server.go index d19051ef..92286f8e 100644 --- a/sftp/server.go +++ b/sftp/server.go @@ -126,10 +126,10 @@ func (c *SFTPServer) AcceptInbound(conn net.Conn, config *ssh.ServerConfig) erro go ssh.DiscardRequests(reqs) for ch := range chans { - // If its not a session channel we just move on because its not something we + // If not a session channel we just move on because it's not something we // know how to handle at this point. if ch.ChannelType() != "session" { - ch.Reject(ssh.UnknownChannelType, "unknown channel type") + _ = ch.Reject(ssh.UnknownChannelType, "unknown channel type") continue } @@ -143,37 +143,40 @@ func (c *SFTPServer) AcceptInbound(conn net.Conn, config *ssh.ServerConfig) erro // Channels have a type that is dependent on the protocol. For SFTP // this is "subsystem" with a payload that (should) be "sftp". Discard // anything else we receive ("pty", "shell", etc) - req.Reply(req.Type == "subsystem" && string(req.Payload[4:]) == "sftp", nil) + _ = req.Reply(req.Type == "subsystem" && string(req.Payload[4:]) == "sftp", nil) } }(requests) - // If no UUID has been set on this inbound request then we can assume we - // have screwed up something in the authentication code. This is a sanity - // check, but should never be encountered (ideally...). - // - // This will also attempt to match a specific server out of the global server - // store and return nil if there is no match. - uuid := sconn.Permissions.Extensions["uuid"] - srv := c.manager.Find(func(s *server.Server) bool { - if uuid == "" { - return false + if srv, ok := c.manager.Get(sconn.Permissions.Extensions["uuid"]); ok { + if err := c.Handle(sconn, srv, channel); err != nil { + return err } - return s.ID() == uuid - }) - if srv == nil { - continue } + } + return nil +} - // Spin up a SFTP server instance for the authenticated user's server allowing - // them access to the underlying filesystem. - handler, err := NewHandler(sconn, srv) - if err != nil { - return errors.WithStackIf(err) - } - rs := sftp.NewRequestServer(channel, handler.Handlers()) - if err := rs.Serve(); err == io.EOF { +// Handle spins up a SFTP server instance for the authenticated user's server allowing +// them access to the underlying filesystem. +func (c *SFTPServer) Handle(conn *ssh.ServerConn, srv *server.Server, channel ssh.Channel) error { + handler, err := NewHandler(conn, srv) + if err != nil { + return errors.WithStackIf(err) + } + + ctx := srv.Sftp().Context(handler.User()) + rs := sftp.NewRequestServer(channel, handler.Handlers()) + + go func() { + select { + case <-ctx.Done(): + srv.Log().WithField("user", conn.User()).Warn("sftp: terminating active session") _ = rs.Close() } + }() + + if err := rs.Serve(); err == io.EOF { + _ = rs.Close() } return nil diff --git a/system/context_bag.go b/system/context_bag.go new file mode 100644 index 00000000..55016b41 --- /dev/null +++ b/system/context_bag.go @@ -0,0 +1,58 @@ +package system + +import ( + "context" + "sync" +) + +type ctxHolder struct { + ctx context.Context + cancel context.CancelFunc +} + +type ContextBag struct { + mu sync.Mutex + ctx context.Context + items map[string]ctxHolder +} + +func NewContextBag(ctx context.Context) *ContextBag { + return &ContextBag{ctx: ctx, items: make(map[string]ctxHolder)} +} + +// Context returns a context for the given key. If a value already exists in the +// internal map it is returned, otherwise a new cancelable context is returned. +// This context is shared between all callers until the cancel function is called +// by calling Cancel or CancelAll. +func (cb *ContextBag) Context(key string) context.Context { + cb.mu.Lock() + defer cb.mu.Unlock() + + if _, ok := cb.items[key]; !ok { + ctx, cancel := context.WithCancel(cb.ctx) + cb.items[key] = ctxHolder{ctx, cancel} + } + + return cb.items[key].ctx +} + +func (cb *ContextBag) Cancel(key string) { + cb.mu.Lock() + defer cb.mu.Unlock() + + if v, ok := cb.items[key]; ok { + v.cancel() + delete(cb.items, key) + } +} + +func (cb *ContextBag) CancelAll() { + cb.mu.Lock() + defer cb.mu.Unlock() + + for _, v := range cb.items { + v.cancel() + } + + cb.items = make(map[string]ctxHolder) +} From de3ea74f61df59014c20422459426b92182149a9 Mon Sep 17 00:00:00 2001 From: "Michael (Parker) Parker" Date: Fri, 16 Jan 2026 21:01:10 -0500 Subject: [PATCH 2/3] Implement pterodactyl 292 changes (#158) * Implement pterodactyl 292 changes Add the same change as https://github.com/pterodactyl/wings/pull/292 This adds configuration for the `machone-id` file that is required by hytale Creates and manages machine-id files on a per-server basis Adds code to remove machine-id files when a server is deleted as well. It also adds a group file for use along with the passwd file Updated config for passwd Updated mounts to not set default except for the the correct default. * Update machine-id generation Moved machine-id generation code outside of server create only called during initial server creation Create machine-id file for servers that already exists if the file is missing. Makes sure tempdir is created on wings start * remove append removes the append where not needed --- cmd/root.go | 3 ++ config/config.go | 82 +++++++++++++++++++++++++++++------------ router/router_server.go | 12 +++++- server/mounts.go | 23 ++++++++++-- server/power.go | 7 ++++ server/server.go | 25 ++++++++++++- 6 files changed, 123 insertions(+), 29 deletions(-) diff --git a/cmd/root.go b/cmd/root.go index 8acc73bd..c615b0b4 100644 --- a/cmd/root.go +++ b/cmd/root.go @@ -139,6 +139,9 @@ func rootCmdRun(cmd *cobra.Command, _ []string) { log.WithField("error", err).Fatal("failed to create pelican system user") return } + if err := config.ConfigurePasswd(); err != nil { + log.WithField("error", err).Fatal("failed to create passwd files for pelican") + } log.WithFields(log.Fields{ "username": config.Get().System.Username, "uid": config.Get().System.User.Uid, diff --git a/config/config.go b/config/config.go index 04e8c04e..cf0cd38d 100644 --- a/config/config.go +++ b/config/config.go @@ -17,12 +17,11 @@ import ( "text/template" "time" - "github.com/gbrlsnchs/jwt/v3" - "emperror.dev/errors" "github.com/acobaugh/osrelease" "github.com/apex/log" "github.com/creasty/defaults" + "github.com/gbrlsnchs/jwt/v3" "golang.org/x/sys/unix" "gopkg.in/yaml.v2" @@ -129,9 +128,9 @@ type RemoteQueryConfiguration struct { // be less likely to cause performance issues on the Panel. BootServersPerPage int `default:"50" yaml:"boot_servers_per_page"` - //When using services like Cloudflare Access to manage access to - //a specific system via an external authentication system, - //it is possible to add special headers to bypass authentication. + //When using services like Cloudflare Access to manage access to + //a specific system via an external authentication system, + //it is possible to add special headers to bypass authentication. //The mentioned headers can be appended to queries sent from Wings to the panel. CustomHeaders map[string]string `yaml:"custom_headers"` } @@ -187,11 +186,23 @@ type SystemConfiguration struct { Uid int `yaml:"uid"` Gid int `yaml:"gid"` - // Passwd controls weather a passwd file is mounted in the container - // at /etc/passwd to resolve missing user issues - Passwd bool `json:"mount_passwd" yaml:"mount_passwd" default:"true"` - PasswdFile string `json:"passwd_file" yaml:"passwd_file" default:"/etc/pelican/passwd"` - } `yaml:"user"` + // Passwd controls weather a passwd and group file is mounted in the container + // at /etc/passwd to resolve missing user/group issues inside the container + Passwd struct { + Enable bool `json:"enable" yaml:"enable" default:"true"` + Directory string `json:"directory" yaml:"directory" default:"/etc/pelican"` + } `json:"passwd" yaml:"passwd"` + } `json:"user" yaml:"user"` + + // MachineID manages the mounting of a 'machine-id' file for containers as required for + // some game servers. I.E. Hytale + MachineID struct { + // Enable controls if the machine-id file is generated and mounted into the server container + // This is enabled by default + Enable bool `json:"enable" yaml:"enable" default:"true"` + // FilePath is the full path to the machine-id file that will be generated and mounted + Directory string `json:"directory" yaml:"directory" default:"/etc/pelican/machine-id"` + } `json:"machine_id" yaml:"machine_id"` // The amount of time in seconds that can elapse before a server's disk space calculation is // considered stale and a re-check should occur. DANGER: setting this value too low can seriously @@ -604,19 +615,6 @@ func ConfigureDirectories() error { return err } - log.WithField("filepath", _config.System.User.PasswdFile).Debug("ensuring passwd file exists") - if passwd, err := os.Create(_config.System.User.PasswdFile); err != nil { - return err - } else { - // the WriteFile method returns an error if unsuccessful - err := os.WriteFile(passwd.Name(), []byte(fmt.Sprintf("container:x:%d:%d::/home/container:/usr/sbin/nologin", _config.System.User.Uid, _config.System.User.Gid)), 0644) - // handle this error - if err != nil { - // print it out - fmt.Println(err) - } - } - // There are a non-trivial number of users out there whose data directories are actually a // symlink to another location on the disk. If we do not resolve that final destination at this // point things will appear to work, but endless errors will be encountered when we try to @@ -638,6 +636,11 @@ func ConfigureDirectories() error { return err } + log.WithField("path", _config.System.TmpDirectory).Debug("ensuring temporary data directory exists") + if err := os.MkdirAll(_config.System.TmpDirectory, 0o700); err != nil { + return err + } + log.WithField("path", _config.System.ArchiveDirectory).Debug("ensuring archive data directory exists") if err := os.MkdirAll(_config.System.ArchiveDirectory, 0o700); err != nil { return err @@ -648,9 +651,42 @@ func ConfigureDirectories() error { return err } + log.WithField("path", _config.System.User.Passwd.Directory).Debug("ensuring passwd directory exists") + if err := os.MkdirAll(_config.System.User.Passwd.Directory, 0o700); err != nil { + return err + } + + log.WithField("path", _config.System.MachineID.Directory).Debug("ensuring machine-id directory exists") + if err := os.MkdirAll(_config.System.MachineID.Directory, 0o700); err != nil { + return err + } return nil } +// ConfigurePasswd generates the passwd and group files to be used by +// this looks cleaner than the previous way and is similar to pterodactyl +func ConfigurePasswd() (err error) { + if !_config.System.User.Passwd.Enable { + return + } + log.WithField("filepath", filepath.Join(_config.System.User.Passwd.Directory, "passwd")). + Debug("ensuring passwd file exists") + if err = os.WriteFile(filepath.Join(_config.System.User.Passwd.Directory, "passwd"), + []byte(fmt.Sprintf("container:x:%d:%d::/home/container:/usr/sbin/nologin", + _config.System.User.Uid, _config.System.User.Gid)), 0644); err != nil { + return fmt.Errorf("could not write passwd file: %w", err) + } + + log.WithField("filepath", filepath.Join(_config.System.User.Passwd.Directory, "group")). + Debug("ensuring group file exists") + if err = os.WriteFile(filepath.Join(_config.System.User.Passwd.Directory, "group"), + []byte(fmt.Sprintf("container:x:%d:container", + _config.System.User.Gid)), 0644); err != nil { + return fmt.Errorf("could not write group file: %w", err) + } + return +} + // EnableLogRotation writes a logrotate file for wings to the system logrotate // configuration directory if one exists and a logrotate file is not found. This // allows us to basically automate away the log rotation for most installs, but diff --git a/router/router_server.go b/router/router_server.go index 47c4eb46..e49bbda2 100644 --- a/router/router_server.go +++ b/router/router_server.go @@ -75,7 +75,6 @@ func getServerInstallLogs(c *gin.Context) { c.JSON(http.StatusOK, gin.H{"data": output}) } - // Handles a request to control the power state of a server. If the action being passed // through is invalid a 404 is returned. Otherwise, a HTTP/202 Accepted response is returned // and the actual power action is run asynchronously so that we don't have to block the @@ -281,7 +280,16 @@ func deleteServer(c *gin.Context) { p := fs.Path() _ = fs.UnixFS().Close() if err := os.RemoveAll(p); err != nil { - log.WithFields(log.Fields{"path": p, "error": err}).Warn("failed to remove server files during deletion process") + log.WithFields(log.Fields{"path": p, "error": err}). + Warn("failed to remove server files during deletion process") + } + }(s) + + // remove hanging machine-id file for the server when removing + go func(s *server.Server) { + if err := os.Remove(filepath.Join(config.Get().System.MachineID.Directory, s.ID())); err != nil { + log.WithFields(log.Fields{"server_id": s.ID(), "error": err}). + Warn("failed to remove machine-id file for server") } }(s) diff --git a/server/mounts.go b/server/mounts.go index d9f72e21..a0a53cc7 100644 --- a/server/mounts.go +++ b/server/mounts.go @@ -30,15 +30,32 @@ func (s *Server) Mounts() []environment.Mount { }, } - if config.Get().System.User.Passwd { + if config.Get().System.User.Passwd.Enable { passwdMount := environment.Mount{ - Default: true, Target: "/etc/passwd", - Source: config.Get().System.User.PasswdFile, + Source: filepath.Join(config.Get().System.User.Passwd.Directory, "passwd"), ReadOnly: true, } m = append(m, passwdMount) + + groupMount := environment.Mount{ + Target: "/etc/group", + Source: filepath.Join(config.Get().System.User.Passwd.Directory, "group"), + ReadOnly: true, + } + + m = append(m, groupMount) + } + + if config.Get().System.MachineID.Enable { + machineIDMount := environment.Mount{ + Target: "/etc/machine-id", + Source: filepath.Join(config.Get().System.MachineID.Directory, s.ID()), + ReadOnly: true, + } + + m = append(m, machineIDMount) } // Also include any of this server's custom mounts when returning them. return append(m, s.customMounts()...) diff --git a/server/power.go b/server/power.go index c8a15069..622fd06c 100644 --- a/server/power.go +++ b/server/power.go @@ -214,6 +214,13 @@ func (s *Server) onBeforeStart() error { } } + // create the machine-id file on start in case it's missing + if config.Get().System.MachineID.Enable { + if err := s.CreateMachineID(); err != nil { + return err + } + } + s.Log().Info("completed server preflight, starting boot process...") return nil } diff --git a/server/server.go b/server/server.go index 8e0b9518..88b323f4 100644 --- a/server/server.go +++ b/server/server.go @@ -6,6 +6,7 @@ import ( "net/http" "os" "path" + "path/filepath" "regexp" "strconv" "strings" @@ -166,7 +167,6 @@ func DetermineServerTimezone(envvars map[string]interface{}, defaultTimezone str return defaultTimezone } - // parseInvocation parses the start command in the same way we already do in the entrypoint // We can use this to set the container command with all variables replaced. func parseInvocation(invocation string, envvars map[string]interface{}, memory int64, port int, ip string) (parsed string) { @@ -312,9 +312,32 @@ func (s *Server) CreateEnvironment() error { return err } + // create the machine-id file on install + if config.Get().System.MachineID.Enable { + if err := s.CreateMachineID(); err != nil { + return err + } + } + return s.Environment.Create() } +// CreateMachineID generates the machine-id file for the server +func (s *Server) CreateMachineID() error { + // Hytale wants a machine-id in order to encrypt tokens for the server. So + // write a machine-id file for the server that contains the server's UUID + // without any dashes. + p := filepath.Join(config.Get().System.MachineID.Directory, s.ID()) + s.Log().WithFields(log.Fields{ + "path": p}).Debug("creating machine-id file") + machineID := []byte(strings.ReplaceAll(s.ID(), "-", "")) + if err := os.WriteFile(p, machineID, 0o644); err != nil { + return fmt.Errorf("failed to write machine-id (at '%s') for server '%s': %w", p, s.ID(), err) + } + + return nil +} + // Checks if the server is marked as being suspended or not on the system. func (s *Server) IsSuspended() bool { return s.Config().Suspended From b47ca2d7c56e3e0f42666b54b2bc7c214b072e46 Mon Sep 17 00:00:00 2001 From: "Michael (Parker) Parker" Date: Mon, 5 Jan 2026 21:19:27 -0500 Subject: [PATCH 3/3] Implement pterodactyl security fixes --- router/router.go | 1 + router/router_server.go | 2 + router/router_server_ws.go | 104 ++++++++++++++++++++++++++++------ router/router_system.go | 31 ++++++++++ router/tokens/websocket.go | 22 +++++++ router/websocket/limiter.go | 91 +++++++++++++++++++++++++++++ router/websocket/listeners.go | 4 +- router/websocket/message.go | 5 +- router/websocket/websocket.go | 23 +++++++- server/connections.go | 19 +++++++ server/server.go | 15 ++++- server/websockets.go | 7 +++ sftp/handler.go | 8 ++- sftp/server.go | 53 +++++++++-------- system/context_bag.go | 58 +++++++++++++++++++ 15 files changed, 392 insertions(+), 51 deletions(-) create mode 100644 router/websocket/limiter.go create mode 100644 server/connections.go create mode 100644 system/context_bag.go diff --git a/router/router.go b/router/router.go index 86cea0d3..e0d0cd04 100644 --- a/router/router.go +++ b/router/router.go @@ -66,6 +66,7 @@ func Configure(m *wserver.Manager, client remote.Client) *gin.Engine { protected.GET("/api/servers", getAllServers) protected.POST("/api/servers", postCreateServer) protected.DELETE("/api/transfers/:server", deleteTransfer) + protected.POST("/api/deauthorize-user", postDeauthorizeUser) // These are server specific routes, and require that the request be authorized, and // that the server exist on the Daemon. diff --git a/router/router_server.go b/router/router_server.go index e49bbda2..19a1cb75 100644 --- a/router/router_server.go +++ b/router/router_server.go @@ -303,6 +303,8 @@ func deleteServer(c *gin.Context) { // Adds any of the JTIs passed through in the body to the deny list for the websocket // preventing any JWT generated before the current time from being used to connect to // the socket or send along commands. +// +// deprecated: prefer /api/deauthorize-user func postServerDenyWSTokens(c *gin.Context) { var data struct { JTIs []string `json:"jtis"` diff --git a/router/router_server_ws.go b/router/router_server_ws.go index 6835c769..39d0d70d 100644 --- a/router/router_server_ws.go +++ b/router/router_server_ws.go @@ -2,14 +2,17 @@ package router import ( "context" + "encoding/json" + "net/http" "time" + "emperror.dev/errors" "github.com/gin-gonic/gin" - "github.com/goccy/go-json" ws "github.com/gorilla/websocket" - "github.com/pelican-dev/wings/router/middleware" "github.com/pelican-dev/wings/router/websocket" + "github.com/pelican-dev/wings/server" + "golang.org/x/time/rate" ) var expectedCloseCodes = []int{ @@ -25,6 +28,27 @@ func getServerWebsocket(c *gin.Context) { manager := middleware.ExtractManager(c) s, _ := manager.Get(c.Param("server")) + // Limit the total number of websockets that can be opened at any one time for + // a server instance. This applies across all users connected to the server, and + // is not applied on a per-user basis. + // + // todo: it would be great to make this per-user instead, but we need to modify + // how we even request this endpoint in order for that to be possible. Some type + // of signed identifier in the URL that is verified on this end and set by the + // panel using a shared secret is likely the easiest option. The benefit of that + // is that we can both scope things to the user before authentication, and also + // verify that the JWT provided by the panel is assigned to the same user. + if s.Websockets().Len() >= 30 { + c.AbortWithStatusJSON(http.StatusBadRequest, gin.H{ + "error": "Too many open websocket connections.", + }) + + return + } + + c.Header("Content-Security-Policy", "default-src 'self'") + c.Header("X-Frame-Options", "DENY") + // Create a context that can be canceled when the user disconnects from this // socket that will also cancel listeners running in separate threads. If the // connection itself is terminated listeners using this context will also be @@ -37,36 +61,61 @@ func getServerWebsocket(c *gin.Context) { middleware.CaptureAndAbort(c, err) return } - defer handler.Connection.Close() // Track this open connection on the server so that we can close them all programmatically // if the server is deleted. s.Websockets().Push(handler.Uuid(), &cancel) handler.Logger().Debug("opening connection to server websocket") + defer s.Websockets().Remove(handler.Uuid()) - defer func() { - s.Websockets().Remove(handler.Uuid()) - handler.Logger().Debug("closing connection to server websocket") + go func() { + select { + // When the main context is canceled (through disconnect, server deletion, or server + // suspension) close the connection itself. + case <-ctx.Done(): + handler.Logger().Debug("closing connection to server websocket") + if err := handler.Connection.Close(); err != nil { + handler.Logger().WithError(err).Error("failed to close websocket connection") + } + break + } }() - // If the server is deleted we need to send a close message to the connected client - // so that they disconnect since there will be no more events sent along. Listen for - // the request context being closed to break this loop, otherwise this routine will - // be left hanging in the background. go func() { select { case <-ctx.Done(): - break + return + // If the server is deleted we need to send a close message to the connected client + // so that they disconnect since there will be no more events sent along. Listen for + // the request context being closed to break this loop, otherwise this routine will + //be left hanging in the background. case <-s.Context().Done(): - _ = handler.Connection.WriteControl(ws.CloseMessage, ws.FormatCloseMessage(ws.CloseGoingAway, "server deleted"), time.Now().Add(time.Second*5)) + cancel() break } }() - for { - j := websocket.Message{} + // Due to how websockets are handled we need to connect to the socket + // and _then_ abort it if the server is suspended. You cannot capture + // the HTTP response in the websocket client, thus we connect and then + // immediately close with failure. + if s.IsSuspended() { + _ = handler.Connection.WriteMessage(ws.CloseMessage, ws.FormatCloseMessage(4409, "server is suspended")) - _, p, err := handler.Connection.ReadMessage() + return + } + + // There is a separate rate limiter that applies to individual message types + // within the actual websocket logic handler. _This_ rate limiter just exists + // to avoid enormous floods of data through the socket since we need to parse + // JSON each time. This rate limit realistically should never be hit since this + // would require sending 50+ messages a second over the websocket (no more than + // 10 per 200ms). + var throttled bool + rl := rate.NewLimiter(rate.Every(time.Millisecond*200), 10) + + for { + t, p, err := handler.Connection.ReadMessage() if err != nil { if ws.IsUnexpectedCloseError(err, expectedCloseCodes...) { handler.Logger().WithField("error", err).Warn("error handling websocket message for server") @@ -74,16 +123,39 @@ func getServerWebsocket(c *gin.Context) { break } + if !rl.Allow() { + if !throttled { + throttled = true + _ = handler.Connection.WriteJSON(websocket.Message{Event: websocket.ThrottledEvent, Args: []string{"global"}}) + } + continue + } + + throttled = false + + // If the message isn't a format we expect, or the length of the message is far larger + // than we'd ever expect, drop it. The websocket upgrader logic does enforce a maximum + // _compressed_ message size of 4Kb but that could decompress to a much larger amount + // of data. + if t != ws.TextMessage || len(p) > 32_768 { + continue + } + // Discard and JSON parse errors into the void and don't continue processing this // specific socket request. If we did a break here the client would get disconnected // from the socket, which is NOT what we want to do. + var j websocket.Message if err := json.Unmarshal(p, &j); err != nil { continue } go func(msg websocket.Message) { if err := handler.HandleInbound(ctx, msg); err != nil { - _ = handler.SendErrorJson(msg, err) + if errors.Is(err, server.ErrSuspended) { + cancel() + } else { + _ = handler.SendErrorJson(msg, err) + } } }(j) } diff --git a/router/router_system.go b/router/router_system.go index 1cc56c7d..cdc5bd58 100644 --- a/router/router_system.go +++ b/router/router_system.go @@ -14,6 +14,7 @@ import ( "github.com/pelican-dev/wings/config" "github.com/pelican-dev/wings/internal/diagnostics" "github.com/pelican-dev/wings/router/middleware" + "github.com/pelican-dev/wings/router/tokens" "github.com/pelican-dev/wings/server" "github.com/pelican-dev/wings/server/installer" "github.com/pelican-dev/wings/system" @@ -256,3 +257,33 @@ func postUpdateConfiguration(c *gin.Context) { Applied: true, }) } + +func postDeauthorizeUser(c *gin.Context) { + var data struct { + User string `json:"user"` + Servers []string `json:"servers"` + } + + if err := c.BindJSON(&data); err != nil { + return + } + + // todo: disconnect websockets more gracefully + m := middleware.ExtractManager(c) + if len(data.Servers) > 0 { + for _, uuid := range data.Servers { + if s, ok := m.Get(uuid); ok { + s.Websockets().CancelAll() + s.Sftp().Cancel(data.User) + tokens.DenyForServer(s.ID(), data.User) + } + } + } else { + for _, s := range m.All() { + s.Websockets().CancelAll() + s.Sftp().Cancel(data.User) + } + } + + c.Status(http.StatusNoContent) +} diff --git a/router/tokens/websocket.go b/router/tokens/websocket.go index 017f8ab2..2ad3b9df 100644 --- a/router/tokens/websocket.go +++ b/router/tokens/websocket.go @@ -24,16 +24,29 @@ var wingsBootTime = time.Now() // This is used to allow the Panel to revoke tokens en-masse for a given user & server // combination since the JTI for tokens is just MD5(user.id + server.uuid). When a server // is booted this listing is fetched from the panel and the Websocket is dynamically updated. +// +// deprecated: prefer use of userDenylist var denylist sync.Map +var userDenylist sync.Map // Adds a JTI to the denylist by marking any JWTs generated before the current time as // being invalid if they use the same JTI. +// +// deprecated: prefer the use of DenyForServer func DenyJTI(jti string) { log.WithField("jti", jti).Debugf("adding \"%s\" to JTI denylist", jti) denylist.Store(jti, time.Now()) } +// DenyForServer adds a user UUID to the denylist marking any existing JWTs issued +// to the user as being invalid. This is associated with the user. +func DenyForServer(s string, u string) { + log.WithField("user_uuid", u).WithField("server_uuid", s).Debugf("denying all JWTs created at or before current time for user \"%s\"", u) + + userDenylist.Store(strings.Join([]string{s, u}, ":"), time.Now()) +} + // WebsocketPayload defines the JWT payload for a websocket connection. This JWT is passed along to // the websocket after it has been connected to by sending an "auth" event. type WebsocketPayload struct { @@ -79,12 +92,21 @@ func (p *WebsocketPayload) Denylisted() bool { // Finally, if the token was issued before a time that is currently denied for this // token instance, ignore the permissions response. + // + // This list is deprecated, but we maintain the check here so that custom instances + // are able to continue working. We'll remove it in a future release. if t, ok := denylist.Load(p.JWTID); ok { if p.IssuedAt.Time.Before(t.(time.Time)) { return true } } + if t, ok := userDenylist.Load(strings.Join([]string{p.ServerUUID, p.UserUUID}, ":")); ok { + if p.IssuedAt.Time.Before(t.(time.Time)) { + return true + } + } + return false } diff --git a/router/websocket/limiter.go b/router/websocket/limiter.go new file mode 100644 index 00000000..57315a96 --- /dev/null +++ b/router/websocket/limiter.go @@ -0,0 +1,91 @@ +package websocket + +import ( + "sync" + "time" + + "golang.org/x/time/rate" +) + +type LimiterBucket struct { + mu sync.RWMutex + limits map[Event]*rate.Limiter + throttles map[Event]bool +} + +func (h *Handler) IsThrottled(e Event) bool { + l := h.limiter.For(e) + + h.limiter.mu.Lock() + defer h.limiter.mu.Unlock() + + if l.Allow() { + h.limiter.throttles[e] = false + + return false + } + + // If not allowed, track the throttling and send an event over the wire + // if one wasn't already sent in the same throttling period. + if v, ok := h.limiter.throttles[e]; !v || !ok { + h.limiter.throttles[e] = true + h.Logger().WithField("event", e).Debug("throttling websocket due to event volume") + + _ = h.unsafeSendJson(&Message{Event: ThrottledEvent, Args: []string{string(e)}}) + } + + return true +} + +func NewLimiter() *LimiterBucket { + return &LimiterBucket{ + limits: make(map[Event]*rate.Limiter, 4), + throttles: make(map[Event]bool, 4), + } +} + +// For returns the internal rate limiter for the given event type. In most +// cases this is a shared rate limiter for events, but certain "heavy" or low-frequency +// events implement their own limiters. +func (l *LimiterBucket) For(e Event) *rate.Limiter { + name := limiterName(e) + + l.mu.RLock() + if v, ok := l.limits[name]; ok { + l.mu.RUnlock() + return v + } + + l.mu.RUnlock() + l.mu.Lock() + defer l.mu.Unlock() + + limit, burst := limitValuesFor(e) + l.limits[name] = rate.NewLimiter(limit, burst) + + return l.limits[name] +} + +// limitValuesFor returns the underlying limit and burst value for the given event. +func limitValuesFor(e Event) (rate.Limit, int) { + // Twice every five seconds. + if e == AuthenticationEvent || e == SendServerLogsEvent { + return rate.Every(time.Second * 5), 2 + } + + // 10 per second. + if e == SendCommandEvent { + return rate.Every(time.Second), 10 + } + + // 4 per second. + return rate.Every(time.Second), 4 +} + +func limiterName(e Event) Event { + if e == AuthenticationEvent || e == SendServerLogsEvent || e == SendCommandEvent { + return e + } + + return "_default" +} diff --git a/router/websocket/listeners.go b/router/websocket/listeners.go index 06236cc0..5183b3a5 100644 --- a/router/websocket/listeners.go +++ b/router/websocket/listeners.go @@ -129,7 +129,7 @@ func (h *Handler) listenForServerEvents(ctx context.Context) error { continue } var sendErr error - message := Message{Event: e.Topic} + message := Message{Event: Event(e.Topic)} if str, ok := e.Data.(string); ok { message.Args = []string{str} } else if b, ok := e.Data.([]byte); ok { @@ -147,7 +147,7 @@ func (h *Handler) listenForServerEvents(ctx context.Context) error { continue } } - onError(message.Event, sendErr) + onError(string(message.Event), sendErr) } break } diff --git a/router/websocket/message.go b/router/websocket/message.go index 85fb77f3..04f3fe9d 100644 --- a/router/websocket/message.go +++ b/router/websocket/message.go @@ -1,5 +1,7 @@ package websocket +type Event string + const ( AuthenticationSuccessEvent = "auth success" TokenExpiringEvent = "token expiring" @@ -11,11 +13,12 @@ const ( SendStatsEvent = "send stats" ErrorEvent = "daemon error" JwtErrorEvent = "jwt error" + ThrottledEvent = Event("throttled") ) type Message struct { // The event to perform. - Event string `json:"event"` + Event Event `json:"event"` // The data to pass along, only used by power/command currently. Other requests // should either omit the field or pass an empty value as it is ignored. diff --git a/router/websocket/websocket.go b/router/websocket/websocket.go index be12c464..420a15bd 100644 --- a/router/websocket/websocket.go +++ b/router/websocket/websocket.go @@ -8,8 +8,6 @@ import ( "sync" "time" - "github.com/pelican-dev/wings/internal/models" - "emperror.dev/errors" "github.com/apex/log" "github.com/gbrlsnchs/jwt/v3" @@ -23,6 +21,7 @@ import ( "github.com/pelican-dev/wings/config" "github.com/pelican-dev/wings/environment" "github.com/pelican-dev/wings/environment/docker" + "github.com/pelican-dev/wings/internal/models" "github.com/pelican-dev/wings/router/tokens" "github.com/pelican-dev/wings/server" ) @@ -46,6 +45,7 @@ type Handler struct { server *server.Server ra server.RequestActivity uuid uuid.UUID + limiter *LimiterBucket } var ( @@ -84,6 +84,7 @@ func NewTokenPayload(token []byte) (*tokens.WebsocketPayload, error) { // GetHandler returns a new websocket handler using the context provided. func GetHandler(s *server.Server, w http.ResponseWriter, r *http.Request, c *gin.Context) (*Handler, error) { upgrader := websocket.Upgrader{ + EnableCompression: true, // Ensure that the websocket request is originating from the Panel itself, // and not some other location. CheckOrigin: func(r *http.Request) bool { @@ -110,12 +111,16 @@ func GetHandler(s *server.Server, w http.ResponseWriter, r *http.Request, c *gin return nil, err } + conn.SetReadLimit(4096) + _ = conn.SetCompressionLevel(5) + return &Handler{ Connection: conn, jwt: nil, server: s, ra: s.NewRequestActivity("", c.ClientIP()), uuid: u, + limiter: NewLimiter(), }, nil } @@ -150,7 +155,7 @@ func (h *Handler) SendJson(v Message) error { // If the user does not have permission to see backup events, do not emit // them over the socket. - if strings.HasPrefix(v.Event, server.BackupCompletedEvent) { + if strings.HasPrefix(string(v.Event), server.BackupCompletedEvent) { if !j.HasPermission(PermissionReceiveBackups) { return nil } @@ -277,6 +282,14 @@ func (h *Handler) setJwt(token *tokens.WebsocketPayload) { // HandleInbound handles an inbound socket request and route it to the proper action. func (h *Handler) HandleInbound(ctx context.Context, m Message) error { + if h.server.IsSuspended() { + return server.ErrSuspended + } + + if h.IsThrottled(m.Event) { + return nil + } + if m.Event != AuthenticationEvent { if err := h.TokenValid(); err != nil { h.unsafeSendJson(Message{ @@ -287,6 +300,10 @@ func (h *Handler) HandleInbound(ctx context.Context, m Message) error { } } + if h.server.IsSuspended() { + return server.ErrSuspended + } + switch m.Event { case AuthenticationEvent: { diff --git a/server/connections.go b/server/connections.go new file mode 100644 index 00000000..846ae444 --- /dev/null +++ b/server/connections.go @@ -0,0 +1,19 @@ +package server + +import ( + "github.com/pelican-dev/wings/system" +) + +// Sftp returns the SFTP connection bag for the server instance. This bag tracks +// all open SFTP connections by individual user and allows for a single user or +// all users to be disconnected by other processes. +func (s *Server) Sftp() *system.ContextBag { + s.Lock() + defer s.Unlock() + + if s.sftpBag == nil { + s.sftpBag = system.NewContextBag(s.Context()) + } + + return s.sftpBag +} diff --git a/server/server.go b/server/server.go index 88b323f4..c1e4c3fc 100644 --- a/server/server.go +++ b/server/server.go @@ -70,6 +70,7 @@ type Server struct { // The console throttler instance used to control outputs. throttler *ConsoleThrottle throttleOnce sync.Once + sftpBag *system.ContextBag // Tracks open websocket connections for the server. wsBag *WebsocketBag @@ -191,7 +192,7 @@ func parseInvocation(invocation string, envvars map[string]interface{}, memory i invocation = strings.Replace(invocation, segment, tempSegments[i], 1) } - // Replace the placeholders outside of protected segments + // Replace the placeholders outside protected segments invocation = strings.ReplaceAll(invocation, placeholder, fmt.Sprint(varval)) // Restore protected segments @@ -201,6 +202,10 @@ func parseInvocation(invocation string, envvars map[string]interface{}, memory i } // Replace the defaults with their configured values. + // and any connected SFTP clients. We don't need to worry about revoking any JWTs + // here since they'll be blocked from re-connecting to the websocket anyways. This + // just forces the client to disconnect and attempt to reconnect (rather than waiting + // on them to send a message and hit that disconnect logic). invocation = strings.ReplaceAll(invocation, "${SERVER_PORT}", strconv.Itoa(port)) invocation = strings.ReplaceAll(invocation, "${SERVER_MEMORY}", strconv.Itoa(int(memory))) invocation = strings.ReplaceAll(invocation, "${SERVER_IP}", ip) @@ -263,11 +268,17 @@ func (s *Server) Sync() error { s.SyncWithEnvironment() + // If the server is suspended immediately disconnect all open websocket connections. + if s.IsSuspended() { + s.Websockets().CancelAll() + s.Sftp().CancelAll() + } + return nil } // SyncWithConfiguration accepts a configuration object for a server and will -// sync all of the values with the existing server state. This only replaces the +// sync all values with the existing server state. This only replaces the // existing configuration and process configuration for the server. The // underlying environment will not be affected. This is because this function // can be called from scoped where the server may not be fully initialized, diff --git a/server/websockets.go b/server/websockets.go index e86f88cc..6aa03435 100644 --- a/server/websockets.go +++ b/server/websockets.go @@ -25,6 +25,13 @@ func (s *Server) Websockets() *WebsocketBag { return s.wsBag } +func (w *WebsocketBag) Len() int { + w.mu.Lock() + defer w.mu.Unlock() + + return len(w.conns) +} + // Push adds a new websocket connection to the end of the stack. func (w *WebsocketBag) Push(u uuid.UUID, cancel *context.CancelFunc) { w.mu.Lock() diff --git a/sftp/handler.go b/sftp/handler.go index aed12dbc..04cb3ad2 100644 --- a/sftp/handler.go +++ b/sftp/handler.go @@ -107,7 +107,7 @@ func (h *Handler) Filewrite(request *sftp.Request) (io.WriterAt, error) { h.mu.Lock() defer h.mu.Unlock() - + if err := h.fs.IsIgnored(request.Filepath); err != nil { return nil, err } @@ -158,7 +158,7 @@ func (h *Handler) Filecmd(request *sftp.Request) error { if err := h.fs.IsIgnored(request.Filepath); err != nil { return err } - + switch request.Method { // Allows a user to make changes to the permissions of a given file or directory // on their server using their SFTP client. @@ -312,3 +312,7 @@ func (h *Handler) can(permission string) bool { } return false } + +func (h *Handler) User() string { + return h.events.user +} diff --git a/sftp/server.go b/sftp/server.go index d19051ef..92286f8e 100644 --- a/sftp/server.go +++ b/sftp/server.go @@ -126,10 +126,10 @@ func (c *SFTPServer) AcceptInbound(conn net.Conn, config *ssh.ServerConfig) erro go ssh.DiscardRequests(reqs) for ch := range chans { - // If its not a session channel we just move on because its not something we + // If not a session channel we just move on because it's not something we // know how to handle at this point. if ch.ChannelType() != "session" { - ch.Reject(ssh.UnknownChannelType, "unknown channel type") + _ = ch.Reject(ssh.UnknownChannelType, "unknown channel type") continue } @@ -143,37 +143,40 @@ func (c *SFTPServer) AcceptInbound(conn net.Conn, config *ssh.ServerConfig) erro // Channels have a type that is dependent on the protocol. For SFTP // this is "subsystem" with a payload that (should) be "sftp". Discard // anything else we receive ("pty", "shell", etc) - req.Reply(req.Type == "subsystem" && string(req.Payload[4:]) == "sftp", nil) + _ = req.Reply(req.Type == "subsystem" && string(req.Payload[4:]) == "sftp", nil) } }(requests) - // If no UUID has been set on this inbound request then we can assume we - // have screwed up something in the authentication code. This is a sanity - // check, but should never be encountered (ideally...). - // - // This will also attempt to match a specific server out of the global server - // store and return nil if there is no match. - uuid := sconn.Permissions.Extensions["uuid"] - srv := c.manager.Find(func(s *server.Server) bool { - if uuid == "" { - return false + if srv, ok := c.manager.Get(sconn.Permissions.Extensions["uuid"]); ok { + if err := c.Handle(sconn, srv, channel); err != nil { + return err } - return s.ID() == uuid - }) - if srv == nil { - continue } + } + return nil +} - // Spin up a SFTP server instance for the authenticated user's server allowing - // them access to the underlying filesystem. - handler, err := NewHandler(sconn, srv) - if err != nil { - return errors.WithStackIf(err) - } - rs := sftp.NewRequestServer(channel, handler.Handlers()) - if err := rs.Serve(); err == io.EOF { +// Handle spins up a SFTP server instance for the authenticated user's server allowing +// them access to the underlying filesystem. +func (c *SFTPServer) Handle(conn *ssh.ServerConn, srv *server.Server, channel ssh.Channel) error { + handler, err := NewHandler(conn, srv) + if err != nil { + return errors.WithStackIf(err) + } + + ctx := srv.Sftp().Context(handler.User()) + rs := sftp.NewRequestServer(channel, handler.Handlers()) + + go func() { + select { + case <-ctx.Done(): + srv.Log().WithField("user", conn.User()).Warn("sftp: terminating active session") _ = rs.Close() } + }() + + if err := rs.Serve(); err == io.EOF { + _ = rs.Close() } return nil diff --git a/system/context_bag.go b/system/context_bag.go new file mode 100644 index 00000000..55016b41 --- /dev/null +++ b/system/context_bag.go @@ -0,0 +1,58 @@ +package system + +import ( + "context" + "sync" +) + +type ctxHolder struct { + ctx context.Context + cancel context.CancelFunc +} + +type ContextBag struct { + mu sync.Mutex + ctx context.Context + items map[string]ctxHolder +} + +func NewContextBag(ctx context.Context) *ContextBag { + return &ContextBag{ctx: ctx, items: make(map[string]ctxHolder)} +} + +// Context returns a context for the given key. If a value already exists in the +// internal map it is returned, otherwise a new cancelable context is returned. +// This context is shared between all callers until the cancel function is called +// by calling Cancel or CancelAll. +func (cb *ContextBag) Context(key string) context.Context { + cb.mu.Lock() + defer cb.mu.Unlock() + + if _, ok := cb.items[key]; !ok { + ctx, cancel := context.WithCancel(cb.ctx) + cb.items[key] = ctxHolder{ctx, cancel} + } + + return cb.items[key].ctx +} + +func (cb *ContextBag) Cancel(key string) { + cb.mu.Lock() + defer cb.mu.Unlock() + + if v, ok := cb.items[key]; ok { + v.cancel() + delete(cb.items, key) + } +} + +func (cb *ContextBag) CancelAll() { + cb.mu.Lock() + defer cb.mu.Unlock() + + for _, v := range cb.items { + v.cancel() + } + + cb.items = make(map[string]ctxHolder) +}