Skip to content
Closed
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
1 change: 1 addition & 0 deletions server/plugin.go
Original file line number Diff line number Diff line change
Expand Up @@ -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)
}

Expand Down
23 changes: 22 additions & 1 deletion server/progress/http.go
Original file line number Diff line number Diff line change
Expand Up @@ -5,6 +5,8 @@ package progress

import (
"encoding/json"
"errors"
"io"
"net/http"

"github.com/mattermost/mattermost/server/public/pluginapi"
Expand Down Expand Up @@ -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
}
Expand Down
52 changes: 52 additions & 0 deletions server/progress/http_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -7,6 +7,7 @@ import (
"encoding/json"
"net/http"
"net/http/httptest"
"strconv"
"strings"
"testing"

Expand Down Expand Up @@ -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)
}
38 changes: 38 additions & 0 deletions server/progress/ids.go
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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)
}
29 changes: 29 additions & 0 deletions server/progress/ids_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -4,6 +4,7 @@
package progress

import (
"strconv"
"testing"

"github.com/stretchr/testify/assert"
Expand All @@ -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}))
}
7 changes: 7 additions & 0 deletions server/progress/store.go
Original file line number Diff line number Diff line change
Expand Up @@ -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)

Expand All @@ -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,
Expand Down
40 changes: 40 additions & 0 deletions server/progress/store_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -4,6 +4,7 @@
package progress

import (
"strconv"
"testing"

"github.com/stretchr/testify/assert"
Expand All @@ -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")
Expand Down
16 changes: 16 additions & 0 deletions server/router_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -3,6 +3,7 @@ package main
import (
"net/http"
"net/http/httptest"
"strings"
"testing"

"github.com/stretchr/testify/assert"
Expand Down Expand Up @@ -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)
}
Loading