154 lines
4.1 KiB
Go
154 lines
4.1 KiB
Go
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, ",")
|
|
}
|