package controlapi import ( "bytes" "encoding/json" "errors" "fmt" "io" "mime" "net/http" "net/url" "regexp" "strings" "unicode/utf8" ) const maximumRequestBody = 1 << 20 var ( logicalIDRegex = regexp.MustCompile(`^[A-Za-z0-9][A-Za-z0-9._:-]{0,63}$`) idempotencyKeyRegex = regexp.MustCompile(`^[A-Za-z0-9._:-]{16,128}$`) operationIDRegex = regexp.MustCompile(`^op_[0-9A-HJKMNP-TV-Z]{26}$`) strongETagRegex = regexp.MustCompile(`^"[A-Za-z0-9_-]{24}"$`) ) func readRequestBody(request *http.Request, expectedMediaType string) ([]byte, error) { mediaType, _, err := mime.ParseMediaType(request.Header.Get("Content-Type")) if err != nil || mediaType != expectedMediaType { return nil, fmt.Errorf("Content-Type must be %s", expectedMediaType) } contents, err := io.ReadAll(io.LimitReader(request.Body, maximumRequestBody+1)) if err != nil { return nil, errors.New("read request body") } if len(contents) == 0 || len(contents) > maximumRequestBody { return nil, errors.New("request body is empty or too large") } if err := rejectDuplicateJSONKeys(contents); err != nil { return nil, err } return contents, nil } func decodeStrictJSON(contents []byte, destination any) error { decoder := json.NewDecoder(bytes.NewReader(contents)) decoder.DisallowUnknownFields() if err := decoder.Decode(destination); err != nil { return errors.New("request body does not match the API schema") } if decoder.Decode(&struct{}{}) != io.EOF { return errors.New("request body contains trailing JSON") } return nil } func rejectTopLevelNulls(contents []byte) error { var fields map[string]json.RawMessage if err := json.Unmarshal(contents, &fields); err != nil || fields == nil { return errors.New("request body must be a JSON object") } for _, value := range fields { if bytes.Equal(bytes.TrimSpace(value), []byte("null")) { return errors.New("request body properties cannot be null") } } return nil } func rejectDuplicateJSONKeys(contents []byte) error { decoder := json.NewDecoder(bytes.NewReader(contents)) decoder.UseNumber() var visit func(int) error visit = func(depth int) error { if depth > 64 { return errors.New("request body nesting is too deep") } token, err := decoder.Token() if err != nil { return errors.New("request body is not valid JSON") } delimiter, ok := token.(json.Delim) if !ok { return nil } switch delimiter { case '{': seen := make(map[string]struct{}) for decoder.More() { keyToken, err := decoder.Token() if err != nil { return errors.New("request body is not valid JSON") } key, ok := keyToken.(string) if !ok { return errors.New("request body is not a JSON object") } if _, exists := seen[key]; exists { return errors.New("request body contains a duplicate property") } seen[key] = struct{}{} if err := visit(depth + 1); err != nil { return err } } end, err := decoder.Token() if err != nil || end != json.Delim('}') { return errors.New("request body is not valid JSON") } case '[': for decoder.More() { if err := visit(depth + 1); err != nil { return err } } end, err := decoder.Token() if err != nil || end != json.Delim(']') { return errors.New("request body is not valid JSON") } default: return errors.New("request body is not valid JSON") } return nil } if err := visit(0); err != nil { return err } if _, err := decoder.Token(); err != io.EOF { return errors.New("request body contains trailing JSON") } return nil } func validLogicalID(value string) bool { return logicalIDRegex.MatchString(value) } func validLength(value string, minimum, maximum int) bool { length := utf8.RuneCountInString(value) return utf8.ValidString(value) && length >= minimum && length <= maximum } func validateEndpoint(value string) bool { if !validLength(value, 1, 2048) { return false } parsed, err := url.Parse(value) return err == nil && parsed.Scheme != "" && parsed.User == nil && !strings.ContainsAny(value, "\r\n") } func validStrongETag(value string) bool { return strongETagRegex.MatchString(value) && value != "*" && !strings.Contains(value, ",") }