mirror of
https://github.com/Rain-kl/OpenFlare.git
synced 2026-10-08 16:46:37 +08:00
feat(option): seed and validate origin error page options
This commit is contained in:
@@ -4,14 +4,18 @@
|
||||
package option
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"regexp"
|
||||
"strconv"
|
||||
"strings"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/model"
|
||||
openrestyrender "github.com/Rain-kl/Wavelet/pkg/render/openresty"
|
||||
)
|
||||
|
||||
const maxOriginErrorPageHTMLBytes = 256 << 10 // 256 KiB
|
||||
|
||||
var openRestyOptionValidators = map[string]func(key, value string) error{
|
||||
model.ConfigKeyOpenRestyDefaultServerReturnStatus: validateOpenRestyDefaultServerReturnStatus,
|
||||
model.ConfigKeyOpenRestyWorkerProcesses: validateOpenRestyWorkerProcesses,
|
||||
@@ -54,11 +58,18 @@ var openRestyOptionValidators = map[string]func(key, value string) error{
|
||||
model.ConfigKeyOpenRestyDefaultLimitConnPerIP: validateNonNegativeIntegerOption,
|
||||
model.ConfigKeyOpenRestyDefaultLimitRate: validateOpenRestyDefaultLimitRate,
|
||||
model.ConfigKeyOpenRestyDefaultLimitReqPerIP: validateOpenRestyDefaultLimitReqPerIP,
|
||||
model.ConfigKeyOriginErrorPageEnabled: validateBooleanOption,
|
||||
model.ConfigKeyOriginErrorPageStatusCodes: validateOriginErrorPageStatusCodes,
|
||||
model.ConfigKeyOriginErrorPageHTML: validateOriginErrorPageHTML,
|
||||
}
|
||||
|
||||
var openRestyDefaultLimitRatePattern = regexp.MustCompile(`^\d+[kKmM]?$`)
|
||||
|
||||
func validateOpenRestyOption(key, value string) error {
|
||||
// HTML 按原始字节长度校验,避免 TrimSpace 影响上限判断
|
||||
if key == model.ConfigKeyOriginErrorPageHTML {
|
||||
return validateOriginErrorPageHTML(key, value)
|
||||
}
|
||||
trimmed := strings.TrimSpace(value)
|
||||
if validator, ok := openRestyOptionValidators[key]; ok {
|
||||
return validator(key, trimmed)
|
||||
@@ -207,3 +218,31 @@ func validateOpenRestyDefaultLimitReqPerIP(key, trimmed string) error {
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func validateOriginErrorPageStatusCodes(key, trimmed string) error {
|
||||
if trimmed == "" {
|
||||
return fmt.Errorf("%s 不能为空", key)
|
||||
}
|
||||
var tags []string
|
||||
if err := json.Unmarshal([]byte(trimmed), &tags); err != nil {
|
||||
return fmt.Errorf("%s 必须为 JSON 字符串数组", key)
|
||||
}
|
||||
if len(tags) == 0 {
|
||||
return fmt.Errorf("%s 至少包含一个状态码标签", key)
|
||||
}
|
||||
codes, err := openrestyrender.ExpandStatusCodeTags(tags)
|
||||
if err != nil {
|
||||
return fmt.Errorf("%s: %v", key, err)
|
||||
}
|
||||
if len(codes) == 0 {
|
||||
return fmt.Errorf("%s 展开后不能为空", key)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func validateOriginErrorPageHTML(key, value string) error {
|
||||
if len(value) > maxOriginErrorPageHTMLBytes {
|
||||
return fmt.Errorf("%s 长度不能超过 %d 字节(256 KiB)", key, maxOriginErrorPageHTMLBytes)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
@@ -0,0 +1,96 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package option
|
||||
|
||||
import (
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/model"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
func TestValidateOriginErrorPageStatusCodes(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
value string
|
||||
wantErr string
|
||||
}{
|
||||
{
|
||||
name: "合法单码与区间",
|
||||
value: `["522","500-502"]`,
|
||||
},
|
||||
{
|
||||
name: "默认区间",
|
||||
value: `["500-599"]`,
|
||||
},
|
||||
{
|
||||
name: "非法标签",
|
||||
value: `["abc"]`,
|
||||
wantErr: "无效状态码",
|
||||
},
|
||||
{
|
||||
name: "非 JSON 数组",
|
||||
value: `500-599`,
|
||||
wantErr: "必须为 JSON 字符串数组",
|
||||
},
|
||||
{
|
||||
name: "空数组",
|
||||
value: `[]`,
|
||||
wantErr: "至少包含一个状态码标签",
|
||||
},
|
||||
{
|
||||
name: "越界状态码",
|
||||
value: `["399"]`,
|
||||
wantErr: "状态码须在",
|
||||
},
|
||||
{
|
||||
name: "空字符串",
|
||||
value: "",
|
||||
wantErr: "不能为空",
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
t.Parallel()
|
||||
err := validateOpenRestyOption(model.ConfigKeyOriginErrorPageStatusCodes, tt.value)
|
||||
if tt.wantErr == "" {
|
||||
require.NoError(t, err)
|
||||
return
|
||||
}
|
||||
require.Error(t, err)
|
||||
assert.Contains(t, err.Error(), tt.wantErr)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestValidateOriginErrorPageHTML(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
require.NoError(t, validateOpenRestyOption(model.ConfigKeyOriginErrorPageHTML, ""))
|
||||
require.NoError(t, validateOpenRestyOption(model.ConfigKeyOriginErrorPageHTML, "<html>ok</html>"))
|
||||
|
||||
oversized := strings.Repeat("a", maxOriginErrorPageHTMLBytes+1)
|
||||
err := validateOpenRestyOption(model.ConfigKeyOriginErrorPageHTML, oversized)
|
||||
require.Error(t, err)
|
||||
assert.Contains(t, err.Error(), "长度不能超过")
|
||||
|
||||
// 恰好上限应通过
|
||||
atLimit := strings.Repeat("b", maxOriginErrorPageHTMLBytes)
|
||||
require.NoError(t, validateOpenRestyOption(model.ConfigKeyOriginErrorPageHTML, atLimit))
|
||||
}
|
||||
|
||||
func TestValidateOriginErrorPageEnabled(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
require.NoError(t, validateOpenRestyOption(model.ConfigKeyOriginErrorPageEnabled, "true"))
|
||||
require.NoError(t, validateOpenRestyOption(model.ConfigKeyOriginErrorPageEnabled, "false"))
|
||||
err := validateOpenRestyOption(model.ConfigKeyOriginErrorPageEnabled, "yes")
|
||||
require.Error(t, err)
|
||||
assert.Contains(t, err.Error(), "true 或 false")
|
||||
}
|
||||
Reference in New Issue
Block a user