Merge pull request #594 from shadow1ng/dev
发布 / auto-tag (push) Has been cancelled
测试构建 / 代码检查 (push) Has been cancelled
发布 / release (push) Has been cancelled
测试构建 / 单元测试和构建 (push) Has been cancelled
测试构建 / 构建验证 (push) Has been cancelled

Release v2.2.0
This commit is contained in:
ZacharyZcR
2026-07-10 13:57:26 +08:00
committed by GitHub
86 changed files with 9120 additions and 192 deletions
+109 -1
View File
@@ -1,6 +1,114 @@
package lib
import "testing"
import (
"testing"
"gopkg.in/yaml.v2"
)
// =============================================================================
// UnmarshalYAML 测试
// =============================================================================
func TestStrMapUnmarshalYAML(t *testing.T) {
t.Run("正常键值对", func(t *testing.T) {
data := []byte("key1: val1\nkey2: val2\n")
var m StrMap
if err := yaml.Unmarshal(data, &m); err != nil {
t.Fatalf("yaml.Unmarshal error = %v", err)
}
if len(m) != 2 {
t.Fatalf("len = %d, want 2", len(m))
}
if m[0].Key != "key1" || m[0].Value != "val1" {
t.Errorf("m[0] = %+v, want {key1 val1}", m[0])
}
if m[1].Key != "key2" || m[1].Value != "val2" {
t.Errorf("m[1] = %+v, want {key2 val2}", m[1])
}
})
t.Run("单项", func(t *testing.T) {
data := []byte("only: one\n")
var m StrMap
if err := yaml.Unmarshal(data, &m); err != nil {
t.Fatalf("yaml.Unmarshal error = %v", err)
}
if len(m) != 1 || m[0].Key != "only" || m[0].Value != "one" {
t.Fatalf("m = %+v", m)
}
})
t.Run("randomInt 值保留为字符串", func(t *testing.T) {
data := []byte("port: randomInt(1000, 9000)\n")
var m StrMap
if err := yaml.Unmarshal(data, &m); err != nil {
t.Fatalf("yaml.Unmarshal error = %v", err)
}
if len(m) != 1 || m[0].Value != "randomInt(1000, 9000)" {
t.Fatalf("m = %+v", m)
}
})
}
func TestListMapUnmarshalYAML(t *testing.T) {
t.Run("正常列表值", func(t *testing.T) {
data := []byte("users:\n - admin\n - root\npasses:\n - 123\n - 456\n")
var m ListMap
if err := yaml.Unmarshal(data, &m); err != nil {
t.Fatalf("yaml.Unmarshal error = %v", err)
}
if len(m) != 2 {
t.Fatalf("len = %d, want 2", len(m))
}
if m[0].Key != "users" || len(m[0].Value) != 2 || m[0].Value[0] != "admin" || m[0].Value[1] != "root" {
t.Errorf("m[0] = %+v", m[0])
}
if m[1].Key != "passes" || len(m[1].Value) != 2 || m[1].Value[0] != "123" || m[1].Value[1] != "456" {
t.Errorf("m[1] = %+v", m[1])
}
})
t.Run("单个列表", func(t *testing.T) {
data := []byte("cmd:\n - whoami\n")
var m ListMap
if err := yaml.Unmarshal(data, &m); err != nil {
t.Fatalf("yaml.Unmarshal error = %v", err)
}
if len(m) != 1 || m[0].Key != "cmd" || m[0].Value[0] != "whoami" {
t.Fatalf("m = %+v", m)
}
})
t.Run("数字值转字符串", func(t *testing.T) {
data := []byte("ports:\n - 80\n - 443\n")
var m ListMap
if err := yaml.Unmarshal(data, &m); err != nil {
t.Fatalf("yaml.Unmarshal error = %v", err)
}
if m[0].Value[0] != "80" || m[0].Value[1] != "443" {
t.Errorf("数字未转为字符串: %+v", m[0].Value)
}
})
}
func TestStrMapUnmarshalYAML_InvalidValue(t *testing.T) {
// value 是嵌套 map,不是字符串,应报错
data := []byte("key:\n nested: val\n")
var m StrMap
if err := yaml.Unmarshal(data, &m); err == nil {
t.Fatal("期望错误,实际 nil")
}
}
func TestListMapUnmarshalYAML_InvalidValue(t *testing.T) {
// value 是普通字符串而非列表,应报错
data := []byte("key: notalist\n")
var m ListMap
if err := yaml.Unmarshal(data, &m); err == nil {
t.Fatal("期望错误,实际 nil")
}
}
func TestNormalizeHTTPProxyURL(t *testing.T) {
tests := []struct {
+108
View File
@@ -0,0 +1,108 @@
package lib
import (
"testing"
"github.com/google/cel-go/common/types"
)
func TestRegisterCryptoImplementations(t *testing.T) {
overloads := registerCryptoImplementations()
// 建立 operator → index 映射
idx := make(map[string]int, len(overloads))
for i, o := range overloads {
idx[o.Operator] = i
}
t.Run("md5_string", func(t *testing.T) {
i, ok := idx["md5_string"]
if !ok {
t.Fatal("overload md5_string not found")
}
unary := overloads[i].Unary
tests := []struct {
name string
input types.String
want types.String
wantErr bool
}{
{"hello", "hello", "5d41402abc4b2a76b9719d911017c592", false},
{"empty", "", "d41d8cd98f00b204e9800998ecf8427e", false},
{"abc", "abc", "900150983cd24fb0d6963f7d28e17f72", false},
}
for _, tc := range tests {
t.Run(tc.name, func(t *testing.T) {
result := unary(tc.input)
if types.IsError(result) {
t.Fatalf("unexpected error: %v", result)
}
got, ok := result.(types.String)
if !ok {
t.Fatalf("expected types.String, got %T", result)
}
if got != tc.want {
t.Errorf("md5(%q) = %q, want %q", tc.input, got, tc.want)
}
})
}
})
t.Run("md5_string_wrong_type", func(t *testing.T) {
i := idx["md5_string"]
result := overloads[i].Unary(types.Int(42))
if !types.IsError(result) {
t.Errorf("expected error for non-String input, got %v", result)
}
})
t.Run("shiro_key_valid", func(t *testing.T) {
i, ok := idx["shiro_key"]
if !ok {
t.Fatal("overload shiro_key not found")
}
binary := overloads[i].Binary
// kPH+bIxk5D2deZiIxcaaaA== 是常见 shiro 默认 key
result := binary(types.String("kPH+bIxk5D2deZiIxcaaaA=="), types.String("cbc"))
if types.IsError(result) {
t.Fatalf("unexpected error: %v", result)
}
got, ok := result.(types.String)
if !ok {
t.Fatalf("expected types.String, got %T", result)
}
if got == "" {
t.Error("shiro_key returned empty string")
}
})
t.Run("shiro_key_invalid_base64", func(t *testing.T) {
i := idx["shiro_key"]
binary := overloads[i].Binary
// 无效 base64GetShrioCookie 会返回 "",函数返回 NewErr
result := binary(types.String("!!!not_valid_base64!!!"), types.String("cbc"))
if !types.IsError(result) {
t.Errorf("expected error for invalid base64 key, got %v", result)
}
})
t.Run("shiro_key_wrong_key_type", func(t *testing.T) {
i := idx["shiro_key"]
result := overloads[i].Binary(types.Int(1), types.String("cbc"))
if !types.IsError(result) {
t.Error("expected error for non-String key")
}
})
t.Run("shiro_key_wrong_mode_type", func(t *testing.T) {
i := idx["shiro_key"]
result := overloads[i].Binary(types.String("kPH+bIxk5D2deZiIxcaaaA=="), types.Int(0))
if !types.IsError(result) {
t.Error("expected error for non-String mode")
}
})
}
+259
View File
@@ -0,0 +1,259 @@
package lib
import (
"testing"
"github.com/google/cel-go/common/types"
)
func TestRegisterEncodingImplementations(t *testing.T) {
overloads := registerEncodingImplementations()
idx := make(map[string]int, len(overloads))
for i, o := range overloads {
idx[o.Operator] = i
}
t.Run("base64_string", func(t *testing.T) {
unary := overloads[idx["base64_string"]].Unary
tests := []struct {
input types.String
want types.String
}{
{"hello", "aGVsbG8="},
{"", ""},
{"hello world", "aGVsbG8gd29ybGQ="},
}
for _, tc := range tests {
result := unary(tc.input)
if types.IsError(result) {
t.Fatalf("base64_string(%q): unexpected error %v", tc.input, result)
}
got, ok := result.(types.String)
if !ok {
t.Fatalf("expected types.String, got %T", result)
}
if got != tc.want {
t.Errorf("base64_string(%q) = %q, want %q", tc.input, got, tc.want)
}
}
})
t.Run("base64_string_wrong_type", func(t *testing.T) {
result := overloads[idx["base64_string"]].Unary(types.Int(1))
if !types.IsError(result) {
t.Error("expected error for non-String input")
}
})
t.Run("base64_bytes", func(t *testing.T) {
unary := overloads[idx["base64_bytes"]].Unary
tests := []struct {
input types.Bytes
want types.String
}{
{types.Bytes([]byte("hello")), "aGVsbG8="},
{types.Bytes([]byte{}), ""},
}
for _, tc := range tests {
result := unary(tc.input)
if types.IsError(result) {
t.Fatalf("base64_bytes: unexpected error %v", result)
}
got, ok := result.(types.String)
if !ok {
t.Fatalf("expected types.String, got %T", result)
}
if got != tc.want {
t.Errorf("base64_bytes(%v) = %q, want %q", []byte(tc.input), got, tc.want)
}
}
})
t.Run("base64_bytes_wrong_type", func(t *testing.T) {
result := overloads[idx["base64_bytes"]].Unary(types.String("hello"))
if !types.IsError(result) {
t.Error("expected error for non-Bytes input")
}
})
t.Run("base64Decode_string", func(t *testing.T) {
unary := overloads[idx["base64Decode_string"]].Unary
tests := []struct {
input types.String
want types.String
}{
{"aGVsbG8=", "hello"},
{"", ""},
{"aGVsbG8gd29ybGQ=", "hello world"},
}
for _, tc := range tests {
result := unary(tc.input)
if types.IsError(result) {
t.Fatalf("base64Decode_string(%q): unexpected error %v", tc.input, result)
}
got, ok := result.(types.String)
if !ok {
t.Fatalf("expected types.String, got %T", result)
}
if got != tc.want {
t.Errorf("base64Decode_string(%q) = %q, want %q", tc.input, got, tc.want)
}
}
})
t.Run("base64Decode_string_invalid", func(t *testing.T) {
result := overloads[idx["base64Decode_string"]].Unary(types.String("!!!"))
if !types.IsError(result) {
t.Error("expected error for invalid base64 input")
}
})
t.Run("base64Decode_string_wrong_type", func(t *testing.T) {
result := overloads[idx["base64Decode_string"]].Unary(types.Bool(true))
if !types.IsError(result) {
t.Error("expected error for non-String input")
}
})
t.Run("base64Decode_bytes", func(t *testing.T) {
unary := overloads[idx["base64Decode_bytes"]].Unary
result := unary(types.Bytes([]byte("aGVsbG8=")))
if types.IsError(result) {
t.Fatalf("unexpected error: %v", result)
}
got, ok := result.(types.String)
if !ok {
t.Fatalf("expected types.String, got %T", result)
}
if got != "hello" {
t.Errorf("base64Decode_bytes = %q, want %q", got, "hello")
}
})
t.Run("base64Decode_bytes_invalid", func(t *testing.T) {
result := overloads[idx["base64Decode_bytes"]].Unary(types.Bytes([]byte("!!!")))
if !types.IsError(result) {
t.Error("expected error for invalid base64 bytes")
}
})
t.Run("urlencode_string", func(t *testing.T) {
unary := overloads[idx["urlencode_string"]].Unary
// url.QueryEscape: 空格 → "+"
tests := []struct {
input types.String
want types.String
}{
{"hello world", "hello+world"},
{"a=1&b=2", "a%3D1%26b%3D2"},
{"", ""},
}
for _, tc := range tests {
result := unary(tc.input)
if types.IsError(result) {
t.Fatalf("urlencode_string(%q): unexpected error %v", tc.input, result)
}
got, ok := result.(types.String)
if !ok {
t.Fatalf("expected types.String, got %T", result)
}
if got != tc.want {
t.Errorf("urlencode_string(%q) = %q, want %q", tc.input, got, tc.want)
}
}
})
t.Run("urlencode_string_wrong_type", func(t *testing.T) {
result := overloads[idx["urlencode_string"]].Unary(types.Int(0))
if !types.IsError(result) {
t.Error("expected error for non-String input")
}
})
t.Run("urldecode_string", func(t *testing.T) {
unary := overloads[idx["urldecode_string"]].Unary
tests := []struct {
input types.String
want types.String
}{
{"hello%20world", "hello world"},
{"hello+world", "hello world"},
{"", ""},
}
for _, tc := range tests {
result := unary(tc.input)
if types.IsError(result) {
t.Fatalf("urldecode_string(%q): unexpected error %v", tc.input, result)
}
got, ok := result.(types.String)
if !ok {
t.Fatalf("expected types.String, got %T", result)
}
if got != tc.want {
t.Errorf("urldecode_string(%q) = %q, want %q", tc.input, got, tc.want)
}
}
})
t.Run("urldecode_string_invalid", func(t *testing.T) {
// % 后跟非法字符
result := overloads[idx["urldecode_string"]].Unary(types.String("hello%ZZ"))
if !types.IsError(result) {
t.Error("expected error for invalid percent-encoding")
}
})
t.Run("urldecode_string_wrong_type", func(t *testing.T) {
result := overloads[idx["urldecode_string"]].Unary(types.Bool(false))
if !types.IsError(result) {
t.Error("expected error for non-String input")
}
})
t.Run("hexdecode", func(t *testing.T) {
unary := overloads[idx["hexdecode"]].Unary
tests := []struct {
input types.String
want []byte
}{
{"48656c6c6f", []byte("Hello")},
{"", []byte{}},
{"deadbeef", []byte{0xde, 0xad, 0xbe, 0xef}},
}
for _, tc := range tests {
result := unary(tc.input)
if types.IsError(result) {
t.Fatalf("hexdecode(%q): unexpected error %v", tc.input, result)
}
got, ok := result.(types.Bytes)
if !ok {
t.Fatalf("expected types.Bytes, got %T", result)
}
if string(got) != string(tc.want) {
t.Errorf("hexdecode(%q) = %v, want %v", tc.input, []byte(got), tc.want)
}
}
})
t.Run("hexdecode_invalid", func(t *testing.T) {
result := overloads[idx["hexdecode"]].Unary(types.String("zz"))
if !types.IsError(result) {
t.Error("expected error for invalid hex input")
}
})
t.Run("hexdecode_wrong_type", func(t *testing.T) {
result := overloads[idx["hexdecode"]].Unary(types.Int(99))
if !types.IsError(result) {
t.Error("expected error for non-String input")
}
})
}
+53
View File
@@ -0,0 +1,53 @@
package lib
import (
"testing"
"unicode"
"github.com/google/cel-go/common/types"
)
func TestRegisterMiscImplementations_TongdaDate(t *testing.T) {
overloads := registerMiscImplementations()
idx := make(map[string]int, len(overloads))
for i, o := range overloads {
idx[o.Operator] = i
}
i, ok := idx["tongda_date"]
if !ok {
t.Fatal("overload tongda_date not found")
}
fn := overloads[i].Function
if fn == nil {
t.Fatal("tongda_date Function field is nil")
}
result := fn()
if types.IsError(result) {
t.Fatalf("unexpected error: %v", result)
}
got, ok := result.(types.String)
if !ok {
t.Fatalf("expected types.String, got %T", result)
}
s := string(got)
t.Run("length_is_4", func(t *testing.T) {
if len(s) != 4 {
t.Errorf("tongda_date returned %q, want 4-char string", s)
}
})
t.Run("all_digits", func(t *testing.T) {
for _, r := range s {
if !unicode.IsDigit(r) {
t.Errorf("tongda_date returned %q, contains non-digit char %q", s, r)
}
}
})
}
+264
View File
@@ -0,0 +1,264 @@
package lib
import (
"testing"
"unicode"
"github.com/google/cel-go/common/types"
)
func TestRegisterRandomImplementations(t *testing.T) {
overloads := registerRandomImplementations()
idx := make(map[string]int, len(overloads))
for i, o := range overloads {
idx[o.Operator] = i
}
t.Run("randomInt_int_int", func(t *testing.T) {
i, ok := idx["randomInt_int_int"]
if !ok {
t.Fatal("overload randomInt_int_int not found")
}
binary := overloads[i].Binary
t.Run("returns_Int_type", func(t *testing.T) {
result := binary(types.Int(0), types.Int(100))
if types.IsError(result) {
t.Fatalf("unexpected error: %v", result)
}
if _, ok := result.(types.Int); !ok {
t.Errorf("expected types.Int, got %T", result)
}
})
t.Run("value_in_range", func(t *testing.T) {
min, max := types.Int(10), types.Int(20)
for range 50 {
result := binary(min, max)
if types.IsError(result) {
t.Fatalf("unexpected error: %v", result)
}
v := int64(result.(types.Int))
if v < 10 || v >= 20 {
t.Errorf("randomInt(10,20) = %d, out of [10,20)", v)
}
}
})
t.Run("max_le_min_returns_error", func(t *testing.T) {
result := binary(types.Int(5), types.Int(5))
if !types.IsError(result) {
t.Errorf("expected error when max == min, got %v", result)
}
})
t.Run("wrong_lhs_type", func(t *testing.T) {
result := binary(types.String("x"), types.Int(10))
if !types.IsError(result) {
t.Error("expected error for non-Int lhs")
}
})
t.Run("wrong_rhs_type", func(t *testing.T) {
result := binary(types.Int(0), types.String("x"))
if !types.IsError(result) {
t.Error("expected error for non-Int rhs")
}
})
})
t.Run("randomLowercase_int", func(t *testing.T) {
i, ok := idx["randomLowercase_int"]
if !ok {
t.Fatal("overload randomLowercase_int not found")
}
unary := overloads[i].Unary
t.Run("returns_String_type", func(t *testing.T) {
result := unary(types.Int(8))
if types.IsError(result) {
t.Fatalf("unexpected error: %v", result)
}
if _, ok := result.(types.String); !ok {
t.Errorf("expected types.String, got %T", result)
}
})
t.Run("correct_length", func(t *testing.T) {
for _, n := range []int{0, 1, 8, 16} {
result := unary(types.Int(n))
if types.IsError(result) {
t.Fatalf("unexpected error for n=%d: %v", n, result)
}
got := string(result.(types.String))
if len(got) != n {
t.Errorf("randomLowercase(%d) returned length %d", n, len(got))
}
}
})
t.Run("all_lowercase", func(t *testing.T) {
result := unary(types.Int(32))
got := string(result.(types.String))
for _, r := range got {
if !unicode.IsLower(r) {
t.Errorf("randomLowercase returned non-lowercase char %q in %q", r, got)
}
}
})
t.Run("invalid_length_negative", func(t *testing.T) {
result := unary(types.Int(-1))
if !types.IsError(result) {
t.Error("expected error for negative length")
}
})
t.Run("invalid_length_too_large", func(t *testing.T) {
result := unary(types.Int(maxRandomStringLength + 1))
if !types.IsError(result) {
t.Error("expected error for length > maxRandomStringLength")
}
})
t.Run("wrong_type", func(t *testing.T) {
result := unary(types.String("x"))
if !types.IsError(result) {
t.Error("expected error for non-Int input")
}
})
})
t.Run("randomUppercase_int", func(t *testing.T) {
i, ok := idx["randomUppercase_int"]
if !ok {
t.Fatal("overload randomUppercase_int not found")
}
unary := overloads[i].Unary
t.Run("returns_String_type", func(t *testing.T) {
result := unary(types.Int(8))
if types.IsError(result) {
t.Fatalf("unexpected error: %v", result)
}
if _, ok := result.(types.String); !ok {
t.Errorf("expected types.String, got %T", result)
}
})
t.Run("correct_length", func(t *testing.T) {
for _, n := range []int{0, 1, 8, 16} {
result := unary(types.Int(n))
if types.IsError(result) {
t.Fatalf("unexpected error for n=%d: %v", n, result)
}
got := string(result.(types.String))
if len(got) != n {
t.Errorf("randomUppercase(%d) returned length %d", n, len(got))
}
}
})
t.Run("all_uppercase", func(t *testing.T) {
result := unary(types.Int(32))
got := string(result.(types.String))
for _, r := range got {
if !unicode.IsUpper(r) {
t.Errorf("randomUppercase returned non-uppercase char %q in %q", r, got)
}
}
})
t.Run("wrong_type", func(t *testing.T) {
result := unary(types.String("x"))
if !types.IsError(result) {
t.Error("expected error for non-Int input")
}
})
})
t.Run("randomString_int", func(t *testing.T) {
i, ok := idx["randomString_int"]
if !ok {
t.Fatal("overload randomString_int not found")
}
unary := overloads[i].Unary
t.Run("returns_String_type", func(t *testing.T) {
result := unary(types.Int(8))
if types.IsError(result) {
t.Fatalf("unexpected error: %v", result)
}
if _, ok := result.(types.String); !ok {
t.Errorf("expected types.String, got %T", result)
}
})
t.Run("correct_length", func(t *testing.T) {
for _, n := range []int{0, 1, 8, 16} {
result := unary(types.Int(n))
if types.IsError(result) {
t.Fatalf("unexpected error for n=%d: %v", n, result)
}
got := string(result.(types.String))
if len(got) != n {
t.Errorf("randomString(%d) returned length %d", n, len(got))
}
}
})
t.Run("wrong_type", func(t *testing.T) {
result := unary(types.String("x"))
if !types.IsError(result) {
t.Error("expected error for non-Int input")
}
})
})
}
func TestRandomIntSpan(t *testing.T) {
t.Run("normal_range", func(t *testing.T) {
span, err := randomIntSpan(10, 20)
if err != nil {
t.Fatalf("unexpected error: %v", err)
}
if span != 10 {
t.Errorf("randomIntSpan(10,20) = %d, want 10", span)
}
})
t.Run("min_zero", func(t *testing.T) {
span, err := randomIntSpan(0, 100)
if err != nil {
t.Fatalf("unexpected error: %v", err)
}
if span != 100 {
t.Errorf("randomIntSpan(0,100) = %d, want 100", span)
}
})
t.Run("negative_min", func(t *testing.T) {
span, err := randomIntSpan(-5, 5)
if err != nil {
t.Fatalf("unexpected error: %v", err)
}
if span != 10 {
t.Errorf("randomIntSpan(-5,5) = %d, want 10", span)
}
})
t.Run("max_eq_min_returns_error", func(t *testing.T) {
_, err := randomIntSpan(7, 7)
if err == nil {
t.Error("expected error when max == min")
}
})
t.Run("max_lt_min_returns_error", func(t *testing.T) {
_, err := randomIntSpan(10, 5)
if err == nil {
t.Error("expected error when max < min")
}
})
}
+337
View File
@@ -0,0 +1,337 @@
package lib
import (
"testing"
"github.com/google/cel-go/common/types"
"github.com/google/cel-go/common/types/ref"
)
func TestRegisterStringImplementations(t *testing.T) {
overloads := registerStringImplementations()
idx := make(map[string]int, len(overloads))
for i, o := range overloads {
idx[o.Operator] = i
}
t.Run("bytes_bcontains_bytes", func(t *testing.T) {
binary := overloads[idx["bytes_bcontains_bytes"]].Binary
tests := []struct {
name string
lhs types.Bytes
rhs types.Bytes
want types.Bool
}{
{"contains", types.Bytes([]byte("hello world")), types.Bytes([]byte("world")), true},
{"not_contains", types.Bytes([]byte("hello world")), types.Bytes([]byte("xyz")), false},
{"empty_needle", types.Bytes([]byte("hello")), types.Bytes([]byte{}), true},
{"both_empty", types.Bytes([]byte{}), types.Bytes([]byte{}), true},
}
for _, tc := range tests {
t.Run(tc.name, func(t *testing.T) {
result := binary(tc.lhs, tc.rhs)
if types.IsError(result) {
t.Fatalf("unexpected error: %v", result)
}
got, ok := result.(types.Bool)
if !ok {
t.Fatalf("expected types.Bool, got %T", result)
}
if got != tc.want {
t.Errorf("bcontains = %v, want %v", got, tc.want)
}
})
}
})
t.Run("bytes_bcontains_bytes_wrong_lhs", func(t *testing.T) {
result := overloads[idx["bytes_bcontains_bytes"]].Binary(types.String("hello"), types.Bytes([]byte("x")))
if !types.IsError(result) {
t.Error("expected error for non-Bytes lhs")
}
})
t.Run("bytes_bcontains_bytes_wrong_rhs", func(t *testing.T) {
result := overloads[idx["bytes_bcontains_bytes"]].Binary(types.Bytes([]byte("hello")), types.String("x"))
if !types.IsError(result) {
t.Error("expected error for non-Bytes rhs")
}
})
t.Run("string_bmatches_bytes", func(t *testing.T) {
binary := overloads[idx["string_bmatches_bytes"]].Binary
tests := []struct {
name string
pattern types.String
input types.Bytes
want types.Bool
}{
{"digits_match", `\d+`, types.Bytes([]byte("abc123")), true},
{"digits_no_match", `\d+`, types.Bytes([]byte("abc")), false},
{"any", `.*`, types.Bytes([]byte("hello")), true},
{"empty_pattern", ``, types.Bytes([]byte("hello")), true},
}
for _, tc := range tests {
t.Run(tc.name, func(t *testing.T) {
result := binary(tc.pattern, tc.input)
if types.IsError(result) {
t.Fatalf("unexpected error: %v", result)
}
got, ok := result.(types.Bool)
if !ok {
t.Fatalf("expected types.Bool, got %T", result)
}
if got != tc.want {
t.Errorf("bmatches(%q, %q) = %v, want %v", tc.pattern, tc.input, got, tc.want)
}
})
}
})
t.Run("string_bmatches_bytes_invalid_regex", func(t *testing.T) {
result := overloads[idx["string_bmatches_bytes"]].Binary(types.String(`[invalid`), types.Bytes([]byte("hello")))
if !types.IsError(result) {
t.Error("expected error for invalid regex pattern")
}
})
t.Run("string_bmatches_bytes_wrong_lhs", func(t *testing.T) {
result := overloads[idx["string_bmatches_bytes"]].Binary(types.Int(0), types.Bytes([]byte("hello")))
if !types.IsError(result) {
t.Error("expected error for non-String lhs")
}
})
t.Run("string_bmatches_bytes_wrong_rhs", func(t *testing.T) {
result := overloads[idx["string_bmatches_bytes"]].Binary(types.String(`\d+`), types.String("123"))
if !types.IsError(result) {
t.Error("expected error for non-Bytes rhs")
}
})
t.Run("icontains_string", func(t *testing.T) {
binary := overloads[idx["icontains_string"]].Binary
tests := []struct {
name string
lhs types.String
rhs types.String
want types.Bool
}{
{"case_insensitive_match", "Hello World", "hello", true},
{"exact_match", "Hello World", "Hello", true},
{"upper_needle", "hello world", "WORLD", true},
{"not_contains", "hello world", "xyz", false},
{"empty_needle", "hello", "", true},
}
for _, tc := range tests {
t.Run(tc.name, func(t *testing.T) {
result := binary(tc.lhs, tc.rhs)
if types.IsError(result) {
t.Fatalf("unexpected error: %v", result)
}
got, ok := result.(types.Bool)
if !ok {
t.Fatalf("expected types.Bool, got %T", result)
}
if got != tc.want {
t.Errorf("icontains(%q, %q) = %v, want %v", tc.lhs, tc.rhs, got, tc.want)
}
})
}
})
t.Run("icontains_string_wrong_lhs", func(t *testing.T) {
result := overloads[idx["icontains_string"]].Binary(types.Bool(true), types.String("x"))
if !types.IsError(result) {
t.Error("expected error for non-String lhs")
}
})
t.Run("icontains_string_wrong_rhs", func(t *testing.T) {
result := overloads[idx["icontains_string"]].Binary(types.String("hello"), types.Int(1))
if !types.IsError(result) {
t.Error("expected error for non-String rhs")
}
})
t.Run("substr_string_int_int", func(t *testing.T) {
fn := overloads[idx["substr_string_int_int"]].Function
tests := []struct {
name string
str types.String
start types.Int
length types.Int
want types.String
}{
{"basic", "hello world", 0, 5, "hello"},
{"middle", "hello world", 6, 5, "world"},
{"single_char", "hello", 1, 1, "e"},
{"full", "hello", 0, 5, "hello"},
{"zero_length", "hello", 2, 0, ""},
}
for _, tc := range tests {
t.Run(tc.name, func(t *testing.T) {
result := fn(tc.str, tc.start, tc.length)
if types.IsError(result) {
t.Fatalf("substr(%q, %d, %d): unexpected error %v", tc.str, tc.start, tc.length, result)
}
got, ok := result.(types.String)
if !ok {
t.Fatalf("expected types.String, got %T", result)
}
if got != tc.want {
t.Errorf("substr(%q, %d, %d) = %q, want %q", tc.str, tc.start, tc.length, got, tc.want)
}
})
}
})
t.Run("substr_out_of_bounds", func(t *testing.T) {
fn := overloads[idx["substr_string_int_int"]].Function
oob := []struct {
name string
str types.String
start types.Int
length types.Int
}{
{"negative_start", "hello", -1, 2},
{"negative_length", "hello", 0, -1},
{"start_too_large", "hello", 10, 1},
{"length_overflow", "hello", 3, 10},
}
for _, tc := range oob {
t.Run(tc.name, func(t *testing.T) {
result := fn(tc.str, tc.start, tc.length)
if !types.IsError(result) {
t.Errorf("expected error for substr(%q, %d, %d), got %v", tc.str, tc.start, tc.length, result)
}
})
}
})
t.Run("substr_wrong_arg_count", func(t *testing.T) {
fn := overloads[idx["substr_string_int_int"]].Function
result := fn(types.String("hello"), types.Int(0))
if !types.IsError(result) {
t.Error("expected error for wrong argument count")
}
})
t.Run("substr_wrong_types", func(t *testing.T) {
fn := overloads[idx["substr_string_int_int"]].Function
cases := []struct {
name string
args []ref.Val
}{
{"wrong_str", []ref.Val{types.Int(0), types.Int(0), types.Int(1)}},
{"wrong_start", []ref.Val{types.String("hello"), types.String("x"), types.Int(1)}},
{"wrong_length", []ref.Val{types.String("hello"), types.Int(0), types.String("x")}},
}
for _, tc := range cases {
t.Run(tc.name, func(t *testing.T) {
result := fn(tc.args...)
if !types.IsError(result) {
t.Errorf("expected error, got %v", result)
}
})
}
})
t.Run("startsWith_bytes", func(t *testing.T) {
binary := overloads[idx["startsWith_bytes"]].Binary
tests := []struct {
name string
lhs types.Bytes
rhs types.Bytes
want types.Bool
}{
{"match", types.Bytes([]byte("hello world")), types.Bytes([]byte("hello")), true},
{"no_match", types.Bytes([]byte("hello world")), types.Bytes([]byte("world")), false},
{"empty_prefix", types.Bytes([]byte("hello")), types.Bytes([]byte{}), true},
{"exact", types.Bytes([]byte("hello")), types.Bytes([]byte("hello")), true},
}
for _, tc := range tests {
t.Run(tc.name, func(t *testing.T) {
result := binary(tc.lhs, tc.rhs)
if types.IsError(result) {
t.Fatalf("unexpected error: %v", result)
}
got, ok := result.(types.Bool)
if !ok {
t.Fatalf("expected types.Bool, got %T", result)
}
if got != tc.want {
t.Errorf("startsWith_bytes = %v, want %v", got, tc.want)
}
})
}
})
t.Run("startsWith_bytes_wrong_lhs", func(t *testing.T) {
result := overloads[idx["startsWith_bytes"]].Binary(types.String("hello"), types.Bytes([]byte("h")))
if !types.IsError(result) {
t.Error("expected error for non-Bytes lhs")
}
})
t.Run("startsWith_bytes_wrong_rhs", func(t *testing.T) {
result := overloads[idx["startsWith_bytes"]].Binary(types.Bytes([]byte("hello")), types.String("h"))
if !types.IsError(result) {
t.Error("expected error for non-Bytes rhs")
}
})
t.Run("startsWith_string", func(t *testing.T) {
binary := overloads[idx["startsWith_string"]].Binary
tests := []struct {
name string
lhs types.String
rhs types.String
want types.Bool
}{
{"case_insensitive_match", "Hello World", "hello", true},
{"upper_prefix", "hello world", "HELLO", true},
{"no_match", "hello world", "world", false},
{"empty_prefix", "hello", "", true},
{"exact", "Hello", "Hello", true},
}
for _, tc := range tests {
t.Run(tc.name, func(t *testing.T) {
result := binary(tc.lhs, tc.rhs)
if types.IsError(result) {
t.Fatalf("unexpected error: %v", result)
}
got, ok := result.(types.Bool)
if !ok {
t.Fatalf("expected types.Bool, got %T", result)
}
if got != tc.want {
t.Errorf("startsWith_string(%q, %q) = %v, want %v", tc.lhs, tc.rhs, got, tc.want)
}
})
}
})
t.Run("startsWith_string_wrong_lhs", func(t *testing.T) {
result := overloads[idx["startsWith_string"]].Binary(types.Int(0), types.String("h"))
if !types.IsError(result) {
t.Error("expected error for non-String lhs")
}
})
t.Run("startsWith_string_wrong_rhs", func(t *testing.T) {
result := overloads[idx["startsWith_string"]].Binary(types.String("hello"), types.Bool(true))
if !types.IsError(result) {
t.Error("expected error for non-String rhs")
}
})
}
+51
View File
@@ -1342,3 +1342,54 @@ func TestRandomStrRejectsNegativeLength(t *testing.T) {
t.Fatalf("RandomStr negative length = %q, want empty", got)
}
}
// =============================================================================
// MakeVarDecl 测试
// =============================================================================
func TestMakeVarDecl(t *testing.T) {
tests := []struct {
name string
key string
value string
wantIdent string // 期望 Decl.Name
wantKind string // "int" / "string" / "object"
}{
{"randomInt 前缀 -> Int", "myrand", "randomInt(1,100)", "myrand", "int"},
{"newReverse 前缀 -> Object", "myrev", "newReverse()", "myrev", "object"},
{"普通字符串 -> String", "myvar", "somevalue", "myvar", "string"},
{"空值 -> String", "empty", "", "empty", "string"},
{"randomIntExtra -> Int", "n", "randomInt(0, 65535)", "n", "int"},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
decl := MakeVarDecl(tt.key, tt.value)
if decl == nil {
t.Fatal("MakeVarDecl() returned nil")
}
if decl.Name != tt.wantIdent {
t.Errorf("Decl.Name = %q, want %q", decl.Name, tt.wantIdent)
}
// 通过 Type 字段判断类型种类
tp := decl.GetIdent().GetType()
if tp == nil {
t.Fatal("Decl.GetIdent().GetType() == nil")
}
switch tt.wantKind {
case "int":
if tp.GetPrimitive().String() != "INT64" {
t.Errorf("type = %v, want INT64", tp)
}
case "string":
if tp.GetPrimitive().String() != "STRING" {
t.Errorf("type = %v, want STRING", tp)
}
case "object":
if tp.GetMessageType() == "" {
t.Errorf("type = %v, want MessageType", tp)
}
}
})
}
}
+25
View File
@@ -2,9 +2,11 @@ package lib
import (
"crypto/md5" //nolint:gosec // G501: MD5用于POC规则去重,非加密用途
"errors"
"fmt"
"io"
"math/rand" //nolint:gosec // G404: math/rand用于生成测试数据,非加密用途
"net"
"net/http"
"net/url"
"os"
@@ -236,6 +238,9 @@ func executeRules(oReq *http.Request, p *Poc, variableMap map[string]interface{}
resp, err := DoRequest(newRequest, rule.FollowRedirects, session)
newRequest = nil
if err != nil {
if isTransportError(err) {
return false, nil
}
return false, err
}
@@ -793,6 +798,9 @@ func clustersend(oReq *http.Request, variableMap map[string]interface{}, req *Re
// 发送请求
resp, err := DoRequest(newRequest, rule.FollowRedirects, session)
if err != nil {
if isTransportError(err) {
return false, nil
}
return false, fmt.Errorf("%s: %w", i18n.GetText("webscan_request_send_error"), err)
}
@@ -936,3 +944,20 @@ func GetHeader(header map[string]string) string {
builder.WriteString("\r\n")
return builder.String()
}
func isTransportError(err error) bool {
if err == nil {
return false
}
var netErr *net.OpError
if errors.As(err, &netErr) {
return true
}
s := err.Error()
return strings.Contains(s, "malformed HTTP") ||
strings.Contains(s, "transport connection broken") ||
strings.Contains(s, "connection reset") ||
strings.Contains(s, "connection refused") ||
strings.Contains(s, "i/o timeout") ||
strings.Contains(s, "EOF")
}
+332
View File
@@ -517,6 +517,74 @@ func TestPocExecutorPureHelpers(t *testing.T) {
})
}
// =============================================================================
// isPlainLiteral 测试
// =============================================================================
func TestIsPlainLiteral_EmptyString(t *testing.T) {
if isPlainLiteral("", nil) {
t.Error("空字符串不是字面量")
}
}
func TestIsPlainLiteral_PlainWord(t *testing.T) {
if !isPlainLiteral("database", nil) {
t.Error("纯单词 'database' 应视为字面量")
}
}
func TestIsPlainLiteral_WithParens(t *testing.T) {
if isPlainLiteral("func()", nil) {
t.Error("含括号的表达式不是字面量")
}
}
func TestIsPlainLiteral_WithOperator(t *testing.T) {
for _, expr := range []string{"a+b", "a*b", "a==b", "a!=b", "a<b", "a>b", "a&&b", "a||b"} {
if isPlainLiteral(expr, nil) {
t.Errorf("含运算符的表达式 %q 不是字面量", expr)
}
}
}
func TestIsPlainLiteral_WithQuotes(t *testing.T) {
if isPlainLiteral(`"hello"`, nil) {
t.Error("含引号的表达式不是字面量")
}
if isPlainLiteral("'hello'", nil) {
t.Error("含单引号的表达式不是字面量")
}
}
func TestIsPlainLiteral_VariableRef(t *testing.T) {
// 如果 expr 是已声明变量的名字,应走 CEL 求值
varMap := map[string]interface{}{"token": "abc123"}
if isPlainLiteral("token", varMap) {
t.Error("已声明变量不应被视为字面量")
}
}
func TestIsPlainLiteral_UndeclaredVariable(t *testing.T) {
varMap := map[string]interface{}{"token": "abc123"}
// 未声明的变量名且无特殊字符 -> 字面量
if !isPlainLiteral("sql", varMap) {
t.Error("未声明的纯单词 'sql' 应视为字面量")
}
}
func TestIsPlainLiteral_WithBracket(t *testing.T) {
if isPlainLiteral("arr[0]", nil) {
t.Error("含方括号的表达式不是字面量")
}
}
func TestIsPlainLiteral_PathLike(t *testing.T) {
// 路径中可能含 /,但 / 不在排除字符中,视为字面量
if !isPlainLiteral("admin", nil) {
t.Error("纯字母字符串应为字面量")
}
}
func stringMatrixEqual(a, b [][]string) bool {
if len(a) != len(b) {
return false
@@ -533,3 +601,267 @@ func stringMatrixEqual(a, b [][]string) bool {
}
return true
}
// =============================================================================
// buildVulnDetails 测试
// =============================================================================
func TestBuildVulnDetails(t *testing.T) {
tests := []struct {
name string
pocDef *Poc
vulName string
params StrMap
wantKeys []string
wantNoKeys []string
wantVulnType string
wantVulnName string
wantParamVal string
wantParamKey string
}{
{
name: "最小Poc只有Name",
pocDef: &Poc{Name: "poc-yaml-test"},
vulName: "poc-yaml-test",
params: nil,
wantKeys: []string{"vulnerability_type", "vulnerability_name"},
wantNoKeys: []string{"author", "references", "description", "parameters"},
wantVulnType: "poc-yaml-test",
wantVulnName: "poc-yaml-test",
},
{
name: "完整Poc含Author+Links+Description",
pocDef: &Poc{
Name: "poc-yaml-full",
Detail: Detail{
Author: "kei",
Links: []string{"https://example.com"},
Description: "test vuln",
},
},
vulName: "Full Vuln",
params: nil,
wantKeys: []string{"vulnerability_type", "vulnerability_name", "author", "references", "description"},
wantNoKeys: []string{"parameters"},
wantVulnType: "poc-yaml-full",
wantVulnName: "Full Vuln",
},
{
name: "有params则details含parameters字段",
pocDef: &Poc{Name: "poc-yaml-params"},
vulName: "Params Vuln",
params: StrMap{
{Key: "user", Value: "admin"},
{Key: "pass", Value: "123456"},
},
wantKeys: []string{"vulnerability_type", "vulnerability_name", "parameters"},
wantNoKeys: []string{"author"},
wantParamKey: "user",
wantParamVal: "admin",
},
{
name: "空params不含parameters字段",
pocDef: &Poc{Name: "poc-yaml-empty-params"},
vulName: "Empty Params",
params: StrMap{},
wantKeys: []string{"vulnerability_type", "vulnerability_name"},
wantNoKeys: []string{"parameters"},
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
details := buildVulnDetails(tt.pocDef, tt.vulName, tt.params)
for _, k := range tt.wantKeys {
if _, ok := details[k]; !ok {
t.Errorf("details 缺少字段 %q", k)
}
}
for _, k := range tt.wantNoKeys {
if _, ok := details[k]; ok {
t.Errorf("details 不应含字段 %q", k)
}
}
if tt.wantVulnType != "" {
if got, _ := details["vulnerability_type"].(string); got != tt.wantVulnType {
t.Errorf("vulnerability_type = %q, want %q", got, tt.wantVulnType)
}
}
if tt.wantVulnName != "" {
if got, _ := details["vulnerability_name"].(string); got != tt.wantVulnName {
t.Errorf("vulnerability_name = %q, want %q", got, tt.wantVulnName)
}
}
if tt.wantParamKey != "" {
pm, ok := details["parameters"].(map[string]string)
if !ok {
t.Fatalf("parameters 类型错误,实际 %T", details["parameters"])
}
if got := pm[tt.wantParamKey]; got != tt.wantParamVal {
t.Errorf("parameters[%q] = %q, want %q", tt.wantParamKey, got, tt.wantParamVal)
}
}
})
}
}
// =============================================================================
// buildVulnLogMsg 测试
// =============================================================================
func TestBuildVulnLogMsg(t *testing.T) {
tests := []struct {
name string
targetURL string
pocDef *Poc
vulName string
params StrMap
}{
{
name: "backup-file名称走特殊模板",
targetURL: "http://example.com",
pocDef: &Poc{Name: "poc-yaml-backup-file"},
vulName: "poc-yaml-backup-file",
params: nil,
},
{
name: "sql-file名称走特殊模板",
targetURL: "http://example.com",
pocDef: &Poc{Name: "poc-yaml-sql-file"},
vulName: "poc-yaml-sql-file",
params: nil,
},
{
name: "有params走params模板",
targetURL: "http://example.com",
pocDef: &Poc{Name: "poc-yaml-rce"},
vulName: "RCE",
params: StrMap{{Key: "cmd", Value: "id"}},
},
{
name: "无params走detail_header模板",
targetURL: "http://example.com",
pocDef: &Poc{
Name: "poc-yaml-sqli",
Detail: Detail{
Author: "kei",
Links: []string{"https://cve.mitre.org/cgi-bin/cvename.cgi?name=CVE-2024-0001"},
Description: "SQL injection",
},
},
vulName: "SQLi",
params: nil,
},
{
name: "无params无detail只走header",
targetURL: "http://example.com",
pocDef: &Poc{Name: "poc-yaml-generic"},
vulName: "Generic",
params: nil,
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
msg := buildVulnLogMsg(tt.targetURL, tt.pocDef, tt.vulName, tt.params)
if msg == "" {
t.Errorf("buildVulnLogMsg() 返回空字符串")
}
})
}
}
// =============================================================================
// collectVarDeclarations 测试
// =============================================================================
func TestCollectVarDeclarations(t *testing.T) {
t.Run("空 POC 返回空切片", func(t *testing.T) {
p := &Poc{}
decls := collectVarDeclarations(p)
if len(decls) != 0 {
t.Fatalf("len = %d, want 0", len(decls))
}
})
t.Run("仅 Set 字段", func(t *testing.T) {
p := &Poc{
Set: StrMap{
{Key: "token", Value: "randomLowercase(8)"},
{Key: "port", Value: "randomInt(1000, 9000)"},
},
}
decls := collectVarDeclarations(p)
if len(decls) != 2 {
t.Fatalf("len = %d, want 2", len(decls))
}
if decls[0].Name != "token" {
t.Errorf("decls[0].Name = %q, want token", decls[0].Name)
}
if decls[1].Name != "port" {
t.Errorf("decls[1].Name = %q, want port", decls[1].Name)
}
})
t.Run("仅 Sets 字段", func(t *testing.T) {
p := &Poc{
Sets: ListMap{
{Key: "user", Value: []string{"admin", "root"}},
},
}
decls := collectVarDeclarations(p)
if len(decls) != 1 {
t.Fatalf("len = %d, want 1", len(decls))
}
if decls[0].Name != "user" {
t.Errorf("decls[0].Name = %q, want user", decls[0].Name)
}
})
t.Run("Sets 空值列表不 panic", func(t *testing.T) {
p := &Poc{
Sets: ListMap{
{Key: "empty", Value: []string{}},
},
}
decls := collectVarDeclarations(p)
if len(decls) != 1 {
t.Fatalf("len = %d, want 1", len(decls))
}
if decls[0].Name != "empty" {
t.Errorf("decls[0].Name = %q, want empty", decls[0].Name)
}
})
t.Run("Set 和 Sets 合并", func(t *testing.T) {
p := &Poc{
Set: StrMap{
{Key: "a", Value: "x"},
},
Sets: ListMap{
{Key: "b", Value: []string{"y"}},
},
}
decls := collectVarDeclarations(p)
if len(decls) != 2 {
t.Fatalf("len = %d, want 2", len(decls))
}
})
t.Run("newReverse 前缀推断 Object 类型", func(t *testing.T) {
p := &Poc{
Set: StrMap{
{Key: "rev", Value: "newReverse()"},
},
}
decls := collectVarDeclarations(p)
if len(decls) != 1 {
t.Fatalf("len = %d, want 1", len(decls))
}
tp := decls[0].GetIdent().GetType()
if tp == nil || tp.GetMessageType() == "" {
t.Errorf("期望 Object 类型,实际 %v", tp)
}
})
}
+237
View File
@@ -0,0 +1,237 @@
package lib
import (
"encoding/base64"
"strings"
"testing"
)
// =============================================================================
// Padding 测试
// =============================================================================
func TestPadding_BasicBlockAlignment(t *testing.T) {
tests := []struct {
name string
input []byte
blockSize int
wantLen int // 期望长度
}{
{
name: "空输入填充整个块",
input: []byte{},
blockSize: 16,
wantLen: 16,
},
{
name: "15字节填充1字节",
input: make([]byte, 15),
blockSize: 16,
wantLen: 16,
},
{
name: "整块对齐追加完整块",
input: make([]byte, 16),
blockSize: 16,
wantLen: 32,
},
{
name: "1字节填充15字节",
input: []byte{0x01},
blockSize: 16,
wantLen: 16,
},
{
name: "blockSize=8时的对齐",
input: make([]byte, 5),
blockSize: 8,
wantLen: 8,
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
result := Padding(tt.input, tt.blockSize)
if len(result) != tt.wantLen {
t.Errorf("Padding() len=%d, want %d", len(result), tt.wantLen)
}
// 验证填充字节值符合 PKCS7 规范
if len(result) > 0 {
padLen := int(result[len(result)-1])
if padLen == 0 || padLen > tt.blockSize {
t.Errorf("填充字节值 %d 超出 blockSize=%d", padLen, tt.blockSize)
}
// 验证所有填充字节相同
for i := len(result) - padLen; i < len(result); i++ {
if result[i] != byte(padLen) {
t.Errorf("填充字节[%d]=%d 不等于 padLen=%d", i, result[i], padLen)
}
}
}
})
}
}
func TestPadding_ResultLength(t *testing.T) {
// 任意长度输入,结果都应该是 blockSize 的整数倍
blockSize := 16
for inputLen := 0; inputLen < 50; inputLen++ {
input := make([]byte, inputLen)
result := Padding(input, blockSize)
if len(result)%blockSize != 0 {
t.Errorf("输入长度 %d: 填充后长度 %d 不是 %d 的倍数", inputLen, len(result), blockSize)
}
}
}
// =============================================================================
// AESCBCEncrypt 测试
// =============================================================================
func TestAESCBCEncrypt_ValidKey128(t *testing.T) {
// 128-bit AES key (16 bytes)
key := base64.StdEncoding.EncodeToString(make([]byte, 16))
result := AESCBCEncrypt(key)
if result == "" {
t.Error("有效的128位密钥应返回非空结果")
}
// 结果应为有效的 base64
_, err := base64.StdEncoding.DecodeString(result)
if err != nil {
t.Errorf("AESCBCEncrypt 结果应为有效 base64: %v", err)
}
}
func TestAESCBCEncrypt_ValidKey256(t *testing.T) {
// 256-bit AES key (32 bytes)
key := base64.StdEncoding.EncodeToString(make([]byte, 32))
result := AESCBCEncrypt(key)
if result == "" {
t.Error("有效的256位密钥应返回非空结果")
}
}
func TestAESCBCEncrypt_InvalidBase64Key(t *testing.T) {
result := AESCBCEncrypt("!!!not-valid-base64!!!")
if result != "" {
t.Error("无效 base64 密钥应返回空字符串")
}
}
func TestAESCBCEncrypt_InvalidKeySize(t *testing.T) {
// AES 要求密钥为 16/24/32 字节,10 字节无效
key := base64.StdEncoding.EncodeToString(make([]byte, 10))
result := AESCBCEncrypt(key)
if result != "" {
t.Error("无效密钥长度应返回空字符串")
}
}
func TestAESCBCEncrypt_NonDeterministic(t *testing.T) {
// 因为 IV 是随机的,两次加密结果应不同
key := base64.StdEncoding.EncodeToString(make([]byte, 16))
r1 := AESCBCEncrypt(key)
r2 := AESCBCEncrypt(key)
if r1 == r2 {
// 极小概率相同,记录即可
t.Log("两次加密结果相同(极低概率事件)")
}
}
// =============================================================================
// AESGCMEncrypt 测试
// =============================================================================
func TestAESGCMEncrypt_ValidKey128(t *testing.T) {
key := base64.StdEncoding.EncodeToString(make([]byte, 16))
result := AESGCMEncrypt(key)
if result == "" {
t.Error("有效的128位密钥应返回非空结果")
}
_, err := base64.StdEncoding.DecodeString(result)
if err != nil {
t.Errorf("AESGCMEncrypt 结果应为有效 base64: %v", err)
}
}
func TestAESGCMEncrypt_InvalidKey(t *testing.T) {
result := AESGCMEncrypt("invalid-base64!!!")
if result != "" {
t.Error("无效 base64 密钥应返回空字符串")
}
}
func TestAESGCMEncrypt_NonDeterministic(t *testing.T) {
key := base64.StdEncoding.EncodeToString(make([]byte, 16))
r1 := AESGCMEncrypt(key)
r2 := AESGCMEncrypt(key)
// GCM nonce 随机,结果不应相同
if r1 == r2 {
t.Log("两次 GCM 加密结果相同(极低概率事件)")
}
}
// =============================================================================
// GetShrioCookie 测试
// =============================================================================
func TestGetShrioCookie_CBCMode(t *testing.T) {
key := base64.StdEncoding.EncodeToString(make([]byte, 16))
result := GetShrioCookie(key, "cbc")
if result == "" {
t.Error("CBC 模式应返回非空 cookie")
}
}
func TestGetShrioCookie_GCMMode(t *testing.T) {
key := base64.StdEncoding.EncodeToString(make([]byte, 16))
result := GetShrioCookie(key, "gcm")
if result == "" {
t.Error("GCM 模式应返回非空 cookie")
}
}
func TestGetShrioCookie_DefaultMode(t *testing.T) {
// 非 gcm 模式走 CBC
key := base64.StdEncoding.EncodeToString(make([]byte, 16))
result := GetShrioCookie(key, "other")
cbcResult := AESCBCEncrypt(key)
// 两个结果都应为非空 base64,但由于随机 IV 不一定相同
if result == "" {
t.Error("默认(非gcm)模式应使用 CBC 加密并返回非空结果")
}
_ = cbcResult
}
func TestGetShrioCookie_RealShiroKey(t *testing.T) {
// 使用真实的 Shiro 默认密钥
shiroDefaultKey := "kPH+bIxk5D2deZiIxcaaaA=="
result := GetShrioCookie(shiroDefaultKey, "cbc")
if result == "" {
t.Error("使用默认 Shiro 密钥应能生成有效 cookie")
}
// 验证结果是 base64 编码
decoded, err := base64.StdEncoding.DecodeString(result)
if err != nil {
t.Errorf("结果应为有效 base64: %v", err)
}
// CBC 模式:IV(16字节) + 密文,结果至少 32 字节
if len(decoded) < 32 {
t.Errorf("CBC 加密结果太短: %d 字节", len(decoded))
}
}
func TestGetShrioCookie_ResultIsBase64(t *testing.T) {
key := base64.StdEncoding.EncodeToString(make([]byte, 16))
for _, mode := range []string{"cbc", "gcm"} {
result := GetShrioCookie(key, mode)
if result == "" {
t.Errorf("mode=%s: 结果不应为空", mode)
continue
}
// base64 只含 [A-Za-z0-9+/=]
if strings.ContainsAny(result, " \t\n\r") {
t.Errorf("mode=%s: base64 结果不应含空白字符", mode)
}
}
}