From 424a2b25562d8e404a42d44300f106d8d98bb330 Mon Sep 17 00:00:00 2001 From: Eric Sethna <14333569+esethna@users.noreply.github.com> Date: Wed, 9 Sep 2026 14:15:16 -0700 Subject: [PATCH 1/2] Cap progress request bodies and stored module ID lists (MM-70622). Authenticated users could send unbounded completedModuleIds arrays that the handler decoded fully and merged into a KV record that never shrinks. Co-authored-by: Cursor --- server/plugin.go | 1 + server/progress/http.go | 23 +++++++++++++++- server/progress/http_test.go | 52 +++++++++++++++++++++++++++++++++++ server/progress/ids.go | 38 +++++++++++++++++++++++++ server/progress/ids_test.go | 29 +++++++++++++++++++ server/progress/store.go | 7 +++++ server/progress/store_test.go | 40 +++++++++++++++++++++++++++ server/router_test.go | 16 +++++++++++ 8 files changed, 205 insertions(+), 1 deletion(-) diff --git a/server/plugin.go b/server/plugin.go index 594a8b2..1104194 100644 --- a/server/plugin.go +++ b/server/plugin.go @@ -106,6 +106,7 @@ func (p *Plugin) ServeHTTP(_ *plugin.Context, w http.ResponseWriter, r *http.Req http.NotFound(w, r) return } + r.Body = http.MaxBytesReader(w, r.Body, progress.MaxRequestBodyBytes) p.router.ServeHTTP(w, r) } diff --git a/server/progress/http.go b/server/progress/http.go index cdf7159..fa37f4e 100644 --- a/server/progress/http.go +++ b/server/progress/http.go @@ -5,6 +5,8 @@ package progress import ( "encoding/json" + "errors" + "io" "net/http" "github.com/mattermost/mattermost/server/public/pluginapi" @@ -89,13 +91,32 @@ func (h *Handler) PutProgress(w http.ResponseWriter, r *http.Request) { return } + r.Body = http.MaxBytesReader(w, r.Body, MaxRequestBodyBytes) + body, err := io.ReadAll(r.Body) + if err != nil { + var maxBytesErr *http.MaxBytesError + if errors.As(err, &maxBytesErr) { + writeError(w, http.StatusRequestEntityTooLarge, "request body too large") + return + } + writeError(w, http.StatusBadRequest, "invalid json") + return + } var req PutRequest - if err := json.NewDecoder(r.Body).Decode(&req); err != nil { + if err := json.Unmarshal(body, &req); err != nil { writeError(w, http.StatusBadRequest, "invalid json") return } + if err := validatePutRequest(req); err != nil { + writeError(w, http.StatusBadRequest, err.Error()) + return + } rec, err := h.store.Put(access.UserFromContext(r.Context()), guideID, req) if err != nil { + if errors.Is(err, errTooManyStoredModuleIDs) { + writeError(w, http.StatusBadRequest, err.Error()) + return + } writeError(w, http.StatusInternalServerError, "failed to save progress") return } diff --git a/server/progress/http_test.go b/server/progress/http_test.go index 36073d8..3c342a9 100644 --- a/server/progress/http_test.go +++ b/server/progress/http_test.go @@ -7,6 +7,7 @@ import ( "encoding/json" "net/http" "net/http/httptest" + "strconv" "strings" "testing" @@ -113,3 +114,54 @@ func TestPutInvalidJSONAndGuideID(t *testing.T) { badID := call(h.PutProgress, http.MethodPut, "/api/v1/progress/Not-Valid", `{}`, "user1", map[string]string{"guideId": "Not-Valid"}) assert.Equal(t, http.StatusBadRequest, badID.Code) } + +func TestPutRejectsOversizedBody(t *testing.T) { + h := newTestHandler(newTestStore(newMemKV()), stubPolicy{guideEnabled: true}) + body := strings.Repeat("a", MaxRequestBodyBytes+1) + + w := call(h.PutProgress, http.MethodPut, "/api/v1/progress/ai-quick-start", body, "user1", map[string]string{"guideId": "ai-quick-start"}) + assert.Equal(t, http.StatusRequestEntityTooLarge, w.Code) + assert.Contains(t, w.Body.String(), "request body too large") +} + +func TestPutRejectsInvalidAndTooManyModuleIDs(t *testing.T) { + h := newTestHandler(newTestStore(newMemKV()), stubPolicy{guideEnabled: true}) + + invalid := call(h.PutProgress, http.MethodPut, "/api/v1/progress/ai-quick-start", `{"completedModuleIds":["../x"],"moduleIds":["chat"]}`, "user1", map[string]string{"guideId": "ai-quick-start"}) + assert.Equal(t, http.StatusBadRequest, invalid.Code) + assert.Contains(t, invalid.Body.String(), "invalid module id") + + tooMany := make([]string, maxRequestModuleIDs+1) + for i := range tooMany { + tooMany[i] = "m" + strconv.Itoa(i) + } + body, err := json.Marshal(PutRequest{CompletedModuleIDs: tooMany, ModuleIDs: []string{"m0"}}) + require.NoError(t, err) + + w := call(h.PutProgress, http.MethodPut, "/api/v1/progress/ai-quick-start", string(body), "user1", map[string]string{"guideId": "ai-quick-start"}) + assert.Equal(t, http.StatusBadRequest, w.Code) + assert.Contains(t, w.Body.String(), "too many module ids") +} + +func TestPutRejectsWhenMergedRecordWouldExceedCap(t *testing.T) { + kv := newMemKV() + s := newTestStore(kv) + ids := make([]string, maxStoredModuleIDs) + for i := range ids { + ids[i] = "m" + strconv.Itoa(i) + } + require.NoError(t, kv.Set(progressKey("user1", "ai-quick-start"), Record{ + V: 1, + GuideID: "ai-quick-start", + CompletedModuleIDs: ids, + })) + + h := newTestHandler(s, stubPolicy{guideEnabled: true}) + w := call(h.PutProgress, http.MethodPut, "/api/v1/progress/ai-quick-start", `{"completedModuleIds":["brand-new"],"moduleIds":["brand-new"]}`, "user1", map[string]string{"guideId": "ai-quick-start"}) + assert.Equal(t, http.StatusBadRequest, w.Code) + assert.Contains(t, w.Body.String(), "too many completed modules") + + rec, err := s.Get("user1", "ai-quick-start") + require.NoError(t, err) + assert.Len(t, rec.CompletedModuleIDs, maxStoredModuleIDs) +} diff --git a/server/progress/ids.go b/server/progress/ids.go index bfdf6f2..dba97cf 100644 --- a/server/progress/ids.go +++ b/server/progress/ids.go @@ -3,6 +3,25 @@ package progress +import ( + "errors" + "strings" +) + +const ( + // MaxRequestBodyBytes is a conservative global HTTP body cap. Academy + // has no file uploads and progress payloads are a few hundred bytes. + MaxRequestBodyBytes = 64 << 10 + maxRequestModuleIDs = 256 + maxStoredModuleIDs = 1024 +) + +var ( + errTooManyModuleIDs = errors.New("too many module ids") + errInvalidModuleID = errors.New("invalid module id") + errTooManyStoredModuleIDs = errors.New("too many completed modules") +) + func validGuideID(id string) bool { if id == "" || len(id) > 128 { return false @@ -28,3 +47,22 @@ func validUserID(id string) bool { } return true } + +func validateModuleIDs(ids []string) error { + if len(ids) > maxRequestModuleIDs { + return errTooManyModuleIDs + } + for _, id := range ids { + if !validGuideID(strings.TrimSpace(id)) { + return errInvalidModuleID + } + } + return nil +} + +func validatePutRequest(req PutRequest) error { + if err := validateModuleIDs(req.CompletedModuleIDs); err != nil { + return err + } + return validateModuleIDs(req.ModuleIDs) +} diff --git a/server/progress/ids_test.go b/server/progress/ids_test.go index d8e1fbd..8323800 100644 --- a/server/progress/ids_test.go +++ b/server/progress/ids_test.go @@ -4,6 +4,7 @@ package progress import ( + "strconv" "testing" "github.com/stretchr/testify/assert" @@ -24,3 +25,31 @@ func TestValidUserID(t *testing.T) { assert.False(t, validUserID("user-id")) assert.False(t, validUserID("../x")) } + +func TestValidatePutRequest(t *testing.T) { + assert.NoError(t, validatePutRequest(PutRequest{ + CompletedModuleIDs: []string{"chat"}, + ModuleIDs: []string{"chat", "search"}, + })) + assert.ErrorIs(t, validatePutRequest(PutRequest{ + CompletedModuleIDs: []string{"../x"}, + ModuleIDs: []string{"chat"}, + }), errInvalidModuleID) + assert.ErrorIs(t, validatePutRequest(PutRequest{ + CompletedModuleIDs: []string{"chat"}, + ModuleIDs: []string{"AI"}, + }), errInvalidModuleID) + assert.ErrorIs(t, validatePutRequest(PutRequest{ + CompletedModuleIDs: []string{""}, + }), errInvalidModuleID) + + tooMany := make([]string, maxRequestModuleIDs+1) + for i := range tooMany { + tooMany[i] = "m" + strconv.Itoa(i) + } + assert.ErrorIs(t, validatePutRequest(PutRequest{CompletedModuleIDs: tooMany}), errTooManyModuleIDs) + assert.ErrorIs(t, validatePutRequest(PutRequest{ModuleIDs: tooMany}), errTooManyModuleIDs) + + atCap := tooMany[:maxRequestModuleIDs] + assert.NoError(t, validatePutRequest(PutRequest{CompletedModuleIDs: atCap, ModuleIDs: atCap})) +} diff --git a/server/progress/store.go b/server/progress/store.go index ae26866..4cd9028 100644 --- a/server/progress/store.go +++ b/server/progress/store.go @@ -90,6 +90,10 @@ func (s *Store) Put(userID, guideID string, req PutRequest) (Record, error) { key := progressKey(userID, guideID) now := time.Now().Unix() + if err := validatePutRequest(req); err != nil { + return Record{}, err + } + completed := normalizeIDs(req.CompletedModuleIDs) curriculum := normalizeIDs(req.ModuleIDs) @@ -105,6 +109,9 @@ func (s *Store) Put(userID, guideID string, req PutRequest) (Record, error) { } merged := normalizeIDs(append(prev.CompletedModuleIDs, completed...)) + if len(merged) > maxStoredModuleIDs { + return nil, errTooManyStoredModuleIDs + } next = Record{ V: 1, GuideID: guideID, diff --git a/server/progress/store_test.go b/server/progress/store_test.go index 683fa43..b84a97a 100644 --- a/server/progress/store_test.go +++ b/server/progress/store_test.go @@ -4,6 +4,7 @@ package progress import ( + "strconv" "testing" "github.com/stretchr/testify/assert" @@ -30,6 +31,45 @@ func TestPutRequestCompleteness(t *testing.T) { require.True(t, containsAll(have, need)) } +func TestPutRejectsInvalidRequestAndMergedCap(t *testing.T) { + s := newTestStore(newMemKV()) + + _, err := s.Put("user1", "ai-quick-start", PutRequest{ + CompletedModuleIDs: []string{"../x"}, + ModuleIDs: []string{"chat"}, + }) + require.ErrorIs(t, err, errInvalidModuleID) + + kv := newMemKV() + s = newTestStore(kv) + ids := make([]string, maxStoredModuleIDs) + for i := range ids { + ids[i] = "m" + strconv.Itoa(i) + } + require.NoError(t, kv.Set(progressKey("user1", "ai-quick-start"), Record{ + V: 1, + GuideID: "ai-quick-start", + CompletedModuleIDs: ids, + })) + + _, err = s.Put("user1", "ai-quick-start", PutRequest{ + CompletedModuleIDs: []string{ids[0]}, + ModuleIDs: []string{ids[0]}, + }) + require.NoError(t, err) + + _, err = s.Put("user1", "ai-quick-start", PutRequest{ + CompletedModuleIDs: []string{"brand-new"}, + ModuleIDs: []string{"brand-new"}, + }) + require.ErrorIs(t, err, errTooManyStoredModuleIDs) + + rec, err := s.Get("user1", "ai-quick-start") + require.NoError(t, err) + assert.Len(t, rec.CompletedModuleIDs, maxStoredModuleIDs) + assert.NotContains(t, rec.CompletedModuleIDs, "brand-new") +} + func TestGetEmptyRecord(t *testing.T) { s := newTestStore(newMemKV()) rec, err := s.Get("user1", "boards") diff --git a/server/router_test.go b/server/router_test.go index 5fe1759..47c880c 100644 --- a/server/router_test.go +++ b/server/router_test.go @@ -3,6 +3,7 @@ package main import ( "net/http" "net/http/httptest" + "strings" "testing" "github.com/stretchr/testify/assert" @@ -101,3 +102,18 @@ func TestAcademyAllowListEnforcedOnEveryProgressRoute(t *testing.T) { }) } } + +func TestServeHTTPRejectsOversizedBody(t *testing.T) { + p := &Plugin{} + cfg := defaultUserAccessConfig() + p.setConfiguration(&configuration{UserAccessConfig: &cfg}) + p.progressHandler = progress.NewHandler(nil, p, nil) + p.router = p.buildRouter() + + r := httptest.NewRequest(http.MethodPut, "/api/v1/progress/ai-quick-start", strings.NewReader(strings.Repeat("a", progress.MaxRequestBodyBytes+1))) + r.Header.Set("Mattermost-User-Id", "user1") + w := httptest.NewRecorder() + p.ServeHTTP(nil, w, r) + + assert.Equal(t, http.StatusRequestEntityTooLarge, w.Code) +} From ea6216190dc6966a2dbd8b39a4d872a6e0d00deb Mon Sep 17 00:00:00 2001 From: Eric Sethna <14333569+esethna@users.noreply.github.com> Date: Wed, 9 Sep 2026 14:20:40 -0700 Subject: [PATCH 2/2] Reuse err after reading the progress body so govet shadow passes. Co-authored-by: Cursor --- server/progress/http.go | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/server/progress/http.go b/server/progress/http.go index fa37f4e..5f4020f 100644 --- a/server/progress/http.go +++ b/server/progress/http.go @@ -103,11 +103,11 @@ func (h *Handler) PutProgress(w http.ResponseWriter, r *http.Request) { return } var req PutRequest - if err := json.Unmarshal(body, &req); err != nil { + if err = json.Unmarshal(body, &req); err != nil { writeError(w, http.StatusBadRequest, "invalid json") return } - if err := validatePutRequest(req); err != nil { + if err = validatePutRequest(req); err != nil { writeError(w, http.StatusBadRequest, err.Error()) return }