Files
gochat/backend/internal/security/security_test.go
T
Rogeeandrogee 2b182f9956 H-300: wire Captain Skills into Web runtime (#48)
* H-300: wire Captain Skills into Web runtime

* H-300: enforce effective model and conservative skill budget

* H-300: fix CI gosec step

* ci: extend golangci-lint timeout

* fix lint findings across backend

* fix(push): resolve delivery protocol blockers

* test(repository): close SQLite test databases

* test(repository): reuse SQLite schema per package

* H-307: restore backend Go cache in CI

* H-307: prefetch modules before cold lint

* H-307: resolve govulncheck security gate

* H-307: build lint with patched Go toolchain

* H-307: clear remaining security scan findings

---------

Co-authored-by: Rogee <rogee@ipao.vip>
2026-08-19 07:08:14 +08:00

1059 lines
32 KiB
Go

package security
import (
"context"
"crypto/hmac"
"crypto/sha256"
"encoding/base64"
"encoding/hex"
"fmt"
"net"
"net/http"
"net/http/httptest"
"net/url"
"strings"
"testing"
"time"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
)
// --- Encryption Tests ---
func TestDefaultEncryptionConfig(t *testing.T) {
cfg := DefaultEncryptionConfig()
assert.False(t, cfg.Enabled)
assert.Equal(t, "", cfg.AESKey)
assert.Equal(t, 1, cfg.KeyVersion)
}
func TestNewEncryptor_Disabled(t *testing.T) {
e, err := NewEncryptor(DefaultEncryptionConfig())
require.NoError(t, err)
require.NotNil(t, e)
assert.False(t, e.IsEnabled())
assert.Equal(t, 1, e.KeyVersion())
}
func TestNewEncryptor_Enabled_ValidKey(t *testing.T) {
key := base64.StdEncoding.EncodeToString(make([]byte, 32))
cfg := EncryptionConfig{AESKey: key, KeyVersion: 2, Enabled: true}
e, err := NewEncryptor(cfg)
require.NoError(t, err)
require.NotNil(t, e)
assert.True(t, e.IsEnabled())
assert.Equal(t, 2, e.KeyVersion())
}
func TestNewEncryptor_InvalidBase64(t *testing.T) {
cfg := EncryptionConfig{AESKey: "!!!invalid-base64!!!", Enabled: true}
_, err := NewEncryptor(cfg)
require.Error(t, err)
assert.Contains(t, err.Error(), "decode base64")
}
func TestNewEncryptor_WrongKeySize(t *testing.T) {
cfg := EncryptionConfig{AESKey: base64.StdEncoding.EncodeToString([]byte("short")), Enabled: true}
_, err := NewEncryptor(cfg)
require.Error(t, err)
assert.Contains(t, err.Error(), "32 bytes")
}
func TestEncryptor_EncryptDecrypt_RoundTrip(t *testing.T) {
key := base64.StdEncoding.EncodeToString(make([]byte, 32))
e, err := NewEncryptor(EncryptionConfig{AESKey: key, KeyVersion: 1, Enabled: true})
require.NoError(t, err)
plaintext := "my-secret-token"
encrypted, err := e.Encrypt(plaintext)
require.NoError(t, err)
assert.NotEqual(t, plaintext, encrypted)
assert.True(t, strings.HasPrefix(encrypted, "enc:v1:"))
decrypted, err := e.Decrypt(encrypted)
require.NoError(t, err)
assert.Equal(t, plaintext, decrypted)
}
func TestEncryptor_EncryptDisabled_ReturnsPlaintext(t *testing.T) {
e, err := NewEncryptor(DefaultEncryptionConfig())
require.NoError(t, err)
encrypted, err := e.Encrypt("secret")
require.NoError(t, err)
assert.Equal(t, "secret", encrypted)
}
func TestEncryptor_DecryptDisabled_ReturnsCiphertext(t *testing.T) {
e, err := NewEncryptor(DefaultEncryptionConfig())
require.NoError(t, err)
decrypted, err := e.Decrypt("some-ciphertext")
require.NoError(t, err)
assert.Equal(t, "some-ciphertext", decrypted)
}
func TestEncryptor_EncryptEmptyString(t *testing.T) {
key := base64.StdEncoding.EncodeToString(make([]byte, 32))
e, err := NewEncryptor(EncryptionConfig{AESKey: key, Enabled: true})
require.NoError(t, err)
encrypted, err := e.Encrypt("")
require.NoError(t, err)
assert.Equal(t, "", encrypted)
}
func TestEncryptor_DecryptEmptyString(t *testing.T) {
key := base64.StdEncoding.EncodeToString(make([]byte, 32))
e, err := NewEncryptor(EncryptionConfig{AESKey: key, Enabled: true})
require.NoError(t, err)
decrypted, err := e.Decrypt("")
require.NoError(t, err)
assert.Equal(t, "", decrypted)
}
func TestEncryptor_Decrypt_InvalidBase64(t *testing.T) {
key := base64.StdEncoding.EncodeToString(make([]byte, 32))
e, err := NewEncryptor(EncryptionConfig{AESKey: key, Enabled: true})
require.NoError(t, err)
_, err = e.Decrypt("enc:v1:!!!invalid!!!")
require.Error(t, err)
assert.Contains(t, err.Error(), "base64")
}
func TestEncryptor_Decrypt_CorruptedData(t *testing.T) {
key := base64.StdEncoding.EncodeToString(make([]byte, 32))
e, err := NewEncryptor(EncryptionConfig{AESKey: key, Enabled: true})
require.NoError(t, err)
// Create a too-short ciphertext (less than nonce size)
shortData := base64.StdEncoding.EncodeToString([]byte("short"))
_, err = e.Decrypt(shortData)
require.Error(t, err)
assert.Contains(t, err.Error(), "too short")
}
func TestEncryptor_Decrypt_TamperedData(t *testing.T) {
key := base64.StdEncoding.EncodeToString(make([]byte, 32))
e, err := NewEncryptor(EncryptionConfig{AESKey: key, Enabled: true})
require.NoError(t, err)
encrypted, err := e.Encrypt("secret-data")
require.NoError(t, err)
// Tamper with the encrypted data
tampered := encrypted[:len(encrypted)-2] + "XX"
_, err = e.Decrypt(tampered)
require.Error(t, err)
assert.Contains(t, err.Error(), "decryption failed")
}
func TestEncryptField_DecryptField_RoundTrip(t *testing.T) {
key := base64.StdEncoding.EncodeToString(make([]byte, 32))
e, err := NewEncryptor(EncryptionConfig{AESKey: key, Enabled: true})
require.NoError(t, err)
encrypted, err := e.EncryptField("my-api-key", FieldTypeAPIKey)
require.NoError(t, err)
decrypted, err := e.DecryptField(encrypted, FieldTypeAPIKey)
require.NoError(t, err)
assert.Equal(t, "my-api-key", decrypted)
}
func TestIsEncrypted(t *testing.T) {
assert.True(t, IsEncrypted("enc:v1:somedata"))
assert.True(t, IsEncrypted("enc:v2:abcdef"))
assert.False(t, IsEncrypted("plaintext"))
assert.False(t, IsEncrypted("enc:"))
assert.False(t, IsEncrypted(""))
}
func TestStripEncryptedPrefix(t *testing.T) {
assert.Equal(t, "payload", stripEncryptedPrefix("enc:v1:payload"))
assert.Equal(t, "data", stripEncryptedPrefix("enc:v10:data"))
// Malformed prefix (no colon after version)
assert.Equal(t, "enc:v1nopayload", stripEncryptedPrefix("enc:v1nopayload"))
}
func TestGenerateAESKey(t *testing.T) {
key, err := GenerateAESKey()
require.NoError(t, err)
assert.NotEmpty(t, key)
decoded, err := base64.StdEncoding.DecodeString(key)
require.NoError(t, err)
assert.Len(t, decoded, 32)
// Generate another key — should be different
key2, err := GenerateAESKey()
require.NoError(t, err)
assert.NotEqual(t, key, key2)
}
func TestEncryptor_DifferentKeyVersions(t *testing.T) {
key := base64.StdEncoding.EncodeToString(make([]byte, 32))
e1, _ := NewEncryptor(EncryptionConfig{AESKey: key, KeyVersion: 1, Enabled: true})
e2, _ := NewEncryptor(EncryptionConfig{AESKey: key, KeyVersion: 2, Enabled: true})
enc1, _ := e1.Encrypt("test")
enc2, _ := e2.Encrypt("test")
assert.True(t, strings.HasPrefix(enc1, "enc:v1:"))
assert.True(t, strings.HasPrefix(enc2, "enc:v2:"))
// Both should decrypt with the same key
dec1, _ := e1.Decrypt(enc1)
dec2, _ := e2.Decrypt(enc2)
assert.Equal(t, "test", dec1)
assert.Equal(t, "test", dec2)
}
// --- JWT Security Tests ---
func TestHashToken(t *testing.T) {
hash1 := hashToken("token1")
hash2 := hashToken("token2")
hash1Again := hashToken("token1")
assert.NotEqual(t, hash1, hash2)
assert.Equal(t, hash1, hash1Again)
assert.Len(t, hash1, 64) // SHA-256 hex = 64 chars
}
func TestParseUnixTimestamp(t *testing.T) {
ts, err := parseUnixTimestamp("1700000000")
require.NoError(t, err)
assert.Equal(t, int64(1700000000), ts.Unix())
}
func TestParseUnixTimestamp_Invalid(t *testing.T) {
_, err := parseUnixTimestamp("not-a-number")
assert.Error(t, err)
}
func TestClaims_StructFields(t *testing.T) {
c := Claims{
UserID: 1,
AccountID: 2,
Role: "admin",
}
assert.Equal(t, uint(1), c.UserID)
assert.Equal(t, uint(2), c.AccountID)
assert.Equal(t, "admin", c.Role)
}
// --- SQL Safety Tests ---
func TestDefaultColumnWhitelists(t *testing.T) {
wl := DefaultColumnWhitelists()
assert.NotEmpty(t, wl)
// Check some known tables exist
tableMap := make(map[string]bool)
for _, w := range wl {
tableMap[w.Table] = true
}
assert.True(t, tableMap["conversations"])
assert.True(t, tableMap["messages"])
assert.True(t, tableMap["contacts"])
assert.True(t, tableMap["accounts"])
assert.True(t, tableMap["users"])
}
func TestNewSQLInjectionValidator_NilWhitelists(t *testing.T) {
v := NewSQLInjectionValidator(nil)
require.NotNil(t, v)
// Should use default whitelists
assert.NotEmpty(t, v.whitelists)
}
func TestNewSQLInjectionValidator_CustomWhitelists(t *testing.T) {
custom := []ColumnWhitelist{
{Table: "custom_table", Columns: []string{"id", "name"}},
}
v := NewSQLInjectionValidator(custom)
require.NotNil(t, v)
// Should use custom whitelist
_, err := v.ValidateSortColumn("custom_table", "id")
assert.NoError(t, err)
// Default tables should not be whitelisted
_, err = v.ValidateSortColumn("conversations", "id")
assert.Error(t, err)
}
func TestValidateSortColumn_ValidColumn(t *testing.T) {
v := NewSQLInjectionValidator(nil)
col, err := v.ValidateSortColumn("conversations", "id")
require.NoError(t, err)
assert.Equal(t, "id", col)
col, err = v.ValidateSortColumn("conversations", "created_at DESC")
require.NoError(t, err)
assert.Equal(t, "created_at DESC", col)
col, err = v.ValidateSortColumn("conversations", "status ASC")
require.NoError(t, err)
assert.Equal(t, "status ASC", col)
}
func TestValidateSortColumn_EmptyInput(t *testing.T) {
v := NewSQLInjectionValidator(nil)
col, err := v.ValidateSortColumn("conversations", "")
require.NoError(t, err)
assert.Equal(t, "", col)
}
func TestValidateSortColumn_NotWhitelisted(t *testing.T) {
v := NewSQLInjectionValidator(nil)
_, err := v.ValidateSortColumn("conversations", "password_hash")
assert.Error(t, err)
assert.Contains(t, err.Error(), "not allowed")
}
func TestValidateSortColumn_UnknownTable(t *testing.T) {
v := NewSQLInjectionValidator(nil)
_, err := v.ValidateSortColumn("unknown_table", "id")
assert.Error(t, err)
}
func TestValidateSortColumn_InvalidDirection(t *testing.T) {
v := NewSQLInjectionValidator(nil)
_, err := v.ValidateSortColumn("conversations", "id RANDOM")
assert.Error(t, err)
assert.Contains(t, err.Error(), "invalid sort direction")
}
func TestValidateColumnName_Valid(t *testing.T) {
v := NewSQLInjectionValidator(nil)
err := v.ValidateColumnName("conversations", "id")
assert.NoError(t, err)
err = v.ValidateColumnName("messages", "conversation_id")
assert.NoError(t, err)
}
func TestValidateColumnName_Empty(t *testing.T) {
v := NewSQLInjectionValidator(nil)
err := v.ValidateColumnName("conversations", "")
assert.Error(t, err)
}
func TestValidateColumnName_InvalidChars(t *testing.T) {
v := NewSQLInjectionValidator(nil)
err := v.ValidateColumnName("conversations", "id; DROP TABLE")
assert.Error(t, err)
assert.Contains(t, err.Error(), "invalid characters")
}
func TestValidateColumnName_NotWhitelisted(t *testing.T) {
v := NewSQLInjectionValidator(nil)
err := v.ValidateColumnName("conversations", "some_random_col")
assert.Error(t, err)
}
func TestValidateTableName_Valid(t *testing.T) {
err := ValidateTableName("conversations")
assert.NoError(t, err)
err = ValidateTableName("messages_2024")
assert.NoError(t, err)
}
func TestValidateTableName_Empty(t *testing.T) {
err := ValidateTableName("")
assert.Error(t, err)
}
func TestValidateTableName_InvalidChars(t *testing.T) {
err := ValidateTableName("conversations; DROP TABLE users")
assert.Error(t, err)
assert.Contains(t, err.Error(), "invalid characters")
}
func TestEscapeLikeWildcards(t *testing.T) {
input := "test%value_here"
result := EscapeLikeWildcards(input)
assert.Equal(t, `test\%value\_here`, result)
}
func TestEscapeLikeWildcards_Backslash(t *testing.T) {
input := `test\value`
result := EscapeLikeWildcards(input)
assert.Equal(t, `test\\value`, result)
}
func TestEscapeLikeWildcards_Empty(t *testing.T) {
result := EscapeLikeWildcards("")
assert.Equal(t, "", result)
}
func TestEscapeLikeWildcards_NoWildcards(t *testing.T) {
result := EscapeLikeWildcards("plain-text")
assert.Equal(t, "plain-text", result)
}
func TestValidateIDParameter_Valid(t *testing.T) {
id, err := ValidateIDParameter("123")
require.NoError(t, err)
assert.Equal(t, int64(123), id)
}
func TestValidateIDParameter_Empty(t *testing.T) {
_, err := ValidateIDParameter("")
assert.Error(t, err)
}
func TestValidateIDParameter_NonDigit(t *testing.T) {
_, err := ValidateIDParameter("12a3")
assert.Error(t, err)
assert.Contains(t, err.Error(), "non-digit")
}
func TestValidateIDParameter_Zero(t *testing.T) {
_, err := ValidateIDParameter("0")
assert.Error(t, err)
assert.Contains(t, err.Error(), "positive integer")
}
func TestValidateIDParameter_Negative(t *testing.T) {
_, err := ValidateIDParameter("-1")
assert.Error(t, err)
}
func TestValidateInClauseValues_Empty(t *testing.T) {
err := ValidateInClauseValues("users", []interface{}{}, 100)
assert.Error(t, err)
assert.Contains(t, err.Error(), "cannot be empty")
}
func TestValidateInClauseValues_ValidSize(t *testing.T) {
err := ValidateInClauseValues("users", []interface{}{1, 2, 3}, 100)
assert.NoError(t, err)
}
func TestValidateInClauseValues_ExceedsMax(t *testing.T) {
values := make([]interface{}, 101)
for i := range values {
values[i] = i
}
err := ValidateInClauseValues("users", values, 100)
assert.Error(t, err)
assert.Contains(t, err.Error(), "exceeds maximum")
}
func TestValidateInClauseValues_DefaultMax(t *testing.T) {
values := make([]interface{}, 10)
for i := range values {
values[i] = i
}
err := ValidateInClauseValues("users", values, 0) // default max = 1000
assert.NoError(t, err)
}
func TestAuditRawQuery_NoIssues(t *testing.T) {
warnings := AuditRawQuery("SELECT * FROM users WHERE id = ?")
assert.Empty(t, warnings)
}
func TestAuditRawQuery_SprintfPattern(t *testing.T) {
warnings := AuditRawQuery("SELECT * FROM users WHERE name = '%s'")
assert.NotEmpty(t, warnings)
assert.Contains(t, warnings[0], "CRITICAL")
}
func TestAuditRawQuery_SQLKeywordsWithoutPlaceholder(t *testing.T) {
warnings := AuditRawQuery("DELETE FROM users WHERE name = 'test'")
assert.NotEmpty(t, warnings)
}
func TestAuditRawQuery_Semicolon(t *testing.T) {
warnings := AuditRawQuery("SELECT * FROM users; DROP TABLE users")
assert.NotEmpty(t, warnings)
found := false
for _, w := range warnings {
if strings.Contains(w, "Semicolon") {
found = true
}
}
assert.True(t, found)
}
func TestAuditRawQuery_CommentPattern(t *testing.T) {
warnings := AuditRawQuery("SELECT * FROM users -- comment")
assert.NotEmpty(t, warnings)
}
func TestGORMSafePattern_NotEmpty(t *testing.T) {
assert.NotEmpty(t, GORMSafePattern)
var safeCount, unsafeCount int
for _, p := range GORMSafePattern {
if p.Safe {
safeCount++
} else {
unsafeCount++
}
}
assert.True(t, safeCount > 0)
assert.True(t, unsafeCount > 0)
}
func TestSafeQueryBuilder_SafeWhere_Valid(t *testing.T) {
v := NewSQLInjectionValidator(nil)
b := NewSafeQueryBuilder(nil, v)
b.SafeWhere("column = ?", "value")
assert.False(t, b.HasErrors())
}
func TestSafeQueryBuilder_SafeWhere_Unsafe(t *testing.T) {
v := NewSQLInjectionValidator(nil)
b := NewSafeQueryBuilder(nil, v)
b.SafeWhere("column = 'value'")
assert.True(t, b.HasErrors())
errs := b.GetErrors()
assert.NotEmpty(t, errs)
}
func TestSafeQueryBuilder_SafeOrder_Valid(t *testing.T) {
v := NewSQLInjectionValidator(nil)
b := NewSafeQueryBuilder(nil, v)
b.SafeOrder("conversations", "id")
assert.False(t, b.HasErrors())
}
func TestSafeQueryBuilder_SafeOrder_Invalid(t *testing.T) {
v := NewSQLInjectionValidator(nil)
b := NewSafeQueryBuilder(nil, v)
b.SafeOrder("conversations", "malicious_column")
assert.True(t, b.HasErrors())
}
func TestSafeQueryBuilder_GetErrors(t *testing.T) {
v := NewSQLInjectionValidator(nil)
b := NewSafeQueryBuilder(nil, v)
b.SafeWhere("bad = 'value'")
b.SafeOrder("conversations", "bad_col")
errs := b.GetErrors()
assert.Len(t, errs, 2)
}
// --- SSRF Protection Tests ---
func TestDefaultSSRFConfig(t *testing.T) {
cfg := DefaultSSRFConfig()
assert.NotEmpty(t, cfg.AllowedDomains)
assert.NotEmpty(t, cfg.BlockedCIDRs)
assert.Equal(t, 3, cfg.MaxRedirects)
assert.False(t, cfg.RequireTLS)
assert.Contains(t, cfg.BlockedCIDRs, "10.0.0.0/8")
assert.Contains(t, cfg.BlockedCIDRs, "127.0.0.0/8")
assert.Contains(t, cfg.BlockedCIDRs, "192.168.0.0/16")
}
func TestIsBlockedIP_PrivateIPs(t *testing.T) {
cfg := DefaultSSRFConfig()
blocked := []string{"10.0.0.1", "172.16.0.1", "192.168.1.1", "127.0.0.1", "0.0.0.0", "169.254.1.1"}
for _, ip := range blocked {
assert.True(t, isBlockedIP(net.ParseIP(ip), cfg.BlockedCIDRs), "expected %s to be blocked", ip)
}
}
func TestIsBlockedIP_PublicIPs(t *testing.T) {
cfg := DefaultSSRFConfig()
assert.False(t, isBlockedIP(net.ParseIP("8.8.8.8"), cfg.BlockedCIDRs), "expected 8.8.8.8 to NOT be blocked")
assert.False(t, isBlockedIP(net.ParseIP("1.1.1.1"), cfg.BlockedCIDRs), "expected 1.1.1.1 to NOT be blocked")
}
func TestIsBlockedIP_InvalidCIDR(t *testing.T) {
assert.False(t, isBlockedIP(net.ParseIP("8.8.8.8"), []string{"invalid-cidr"}))
}
func TestSameIPSets_Equal(t *testing.T) {
a := []net.IPAddr{{IP: net.ParseIP("1.2.3.4")}, {IP: net.ParseIP("5.6.7.8")}}
b := []net.IPAddr{{IP: net.ParseIP("5.6.7.8")}, {IP: net.ParseIP("1.2.3.4")}}
assert.True(t, sameIPSets(a, b))
}
func TestSameIPSets_DifferentLength(t *testing.T) {
a := []net.IPAddr{{IP: net.ParseIP("1.2.3.4")}}
b := []net.IPAddr{{IP: net.ParseIP("1.2.3.4")}, {IP: net.ParseIP("5.6.7.8")}}
assert.False(t, sameIPSets(a, b))
}
func TestSameIPSets_DifferentIPs(t *testing.T) {
a := []net.IPAddr{{IP: net.ParseIP("1.2.3.4")}}
b := []net.IPAddr{{IP: net.ParseIP("5.6.7.8")}}
assert.False(t, sameIPSets(a, b))
}
func TestSafeRedirectCheck(t *testing.T) {
check := safeRedirectCheck(2)
// Under redirect limit — should pass (non-IP host)
req, _ := http.NewRequest("GET", "https://example.com", nil)
err := check(req, nil)
assert.NoError(t, err)
// At redirect limit — should fail
via := []*http.Request{req, req}
err = check(req, via)
assert.Error(t, err)
assert.Contains(t, err.Error(), "stopped after")
}
func TestSafeRedirectCheck_IPLiteral(t *testing.T) {
check := safeRedirectCheck(3)
req, _ := http.NewRequest("GET", "http://10.0.0.1", nil)
err := check(req, nil)
assert.Error(t, err)
assert.Contains(t, err.Error(), "IP literal")
}
func TestValidateURL_Empty(t *testing.T) {
err := ValidateURL("", DefaultSSRFConfig())
assert.Error(t, err)
assert.Contains(t, err.Error(), "empty")
}
func TestValidateURL_AllowedDomain(t *testing.T) {
err := ValidateURL("https://api.telegram.org/bot123/sendMessage", DefaultSSRFConfig())
assert.NoError(t, err)
}
func TestValidateURL_AllowedDomainSubdomain(t *testing.T) {
err := ValidateURL("https://sub.api.telegram.org/test", DefaultSSRFConfig())
assert.NoError(t, err)
}
func TestValidateURL_PublicDomain(t *testing.T) {
err := ValidateURL("https://example.com/test", DefaultSSRFConfig())
assert.NoError(t, err)
}
func TestValidateURL_PrivateIP(t *testing.T) {
err := ValidateURL("http://10.0.0.1/admin", DefaultSSRFConfig())
assert.Error(t, err)
assert.Contains(t, err.Error(), "private/reserved")
}
func TestValidateURL_LoopbackIP(t *testing.T) {
err := ValidateURL("http://127.0.0.1/admin", DefaultSSRFConfig())
assert.Error(t, err)
}
func TestValidateURL_InvalidURL(t *testing.T) {
err := ValidateURL("://invalid-url", DefaultSSRFConfig())
assert.Error(t, err)
}
func TestNewSafeHTTPClient(t *testing.T) {
cfg := DefaultSSRFConfig()
client := NewSafeHTTPClient(cfg)
require.NotNil(t, client)
require.NotNil(t, client.client)
assert.Equal(t, cfg, client.cfg)
}
func TestSafeHTTPClient_Do_AllowedDomain(t *testing.T) {
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
w.WriteHeader(http.StatusOK)
_, err := w.Write([]byte("hello"))
require.NoError(t, err)
}))
defer srv.Close()
parsed, _ := url.Parse(srv.URL)
host := parsed.Hostname()
// httptest.NewServer runs on 127.0.0.1 which SSRF blocks.
// Use empty BlockedCIDRs to allow localhost for testing.
cfg := SSRFConfig{
AllowedDomains: []string{host},
BlockedCIDRs: []string{},
MaxRedirects: 3,
}
client := NewSafeHTTPClient(cfg)
req, _ := http.NewRequest("GET", srv.URL, nil)
resp, err := client.Do(req)
require.NoError(t, err)
defer resp.Body.Close()
assert.Equal(t, 200, resp.StatusCode)
}
func TestSafeHTTPClient_Do_RequireTLS(t *testing.T) {
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
w.WriteHeader(http.StatusOK)
}))
defer srv.Close()
cfg := SSRFConfig{
AllowedDomains: nil,
BlockedCIDRs: DefaultSSRFConfig().BlockedCIDRs,
MaxRedirects: 3,
RequireTLS: true,
}
client := NewSafeHTTPClient(cfg)
req, _ := http.NewRequest("GET", srv.URL, nil)
_, err := client.Do(req)
assert.Error(t, err)
assert.Contains(t, err.Error(), "non-HTTPS")
}
func TestSafeHTTPClient_SafeFetchURL(t *testing.T) {
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
w.WriteHeader(http.StatusOK)
_, err := w.Write([]byte("fetched"))
require.NoError(t, err)
}))
defer srv.Close()
parsed, _ := url.Parse(srv.URL)
host := parsed.Hostname()
// httptest.NewServer runs on 127.0.0.1 which SSRF blocks.
// Use empty BlockedCIDRs to allow localhost for testing.
cfg := SSRFConfig{
AllowedDomains: []string{host},
BlockedCIDRs: []string{},
MaxRedirects: 3,
}
client := NewSafeHTTPClient(cfg)
resp, err := client.SafeFetchURL(context.Background(), srv.URL)
require.NoError(t, err)
defer resp.Body.Close()
assert.Equal(t, 200, resp.StatusCode)
}
func TestSafeHTTPClient_SafeFetchURL_InvalidURL(t *testing.T) {
cfg := DefaultSSRFConfig()
client := NewSafeHTTPClient(cfg)
_, err := client.SafeFetchURL(context.Background(), "://invalid")
assert.Error(t, err)
}
func TestTestWrite(t *testing.T) {
TestWrite()
}
// --- Webhook Signing Tests ---
func TestDefaultWebhookSignatureConfig(t *testing.T) {
cfg := DefaultWebhookSignatureConfig()
assert.Equal(t, int64(300), cfg.MaxAgeSeconds)
assert.Equal(t, int64(30), cfg.ClockSkewSeconds)
assert.NotEmpty(t, cfg.ChannelConfigs)
_, ok := cfg.ChannelConfigs["meta"]
assert.True(t, ok)
_, ok = cfg.ChannelConfigs["telegram"]
assert.True(t, ok)
_, ok = cfg.ChannelConfigs["api"]
assert.True(t, ok)
}
func TestNewWebhookSignatureService_DefaultConfig(t *testing.T) {
svc := NewWebhookSignatureService(WebhookSignatureConfig{})
assert.NotNil(t, svc)
assert.Equal(t, int64(300), svc.config.MaxAgeSeconds)
}
func TestNewWebhookSignatureService_CustomMaxAge(t *testing.T) {
svc := NewWebhookSignatureService(WebhookSignatureConfig{MaxAgeSeconds: 600, ClockSkewSeconds: 10})
assert.NotNil(t, svc)
assert.Equal(t, int64(600), svc.config.MaxAgeSeconds)
assert.Equal(t, int64(10), svc.config.ClockSkewSeconds)
}
func TestGenerateSignature_Success(t *testing.T) {
svc := NewWebhookSignatureService(DefaultWebhookSignatureConfig())
payload := []byte(`{"event":"test"}`)
ts := int64(1700000000)
result, err := svc.GenerateSignature("meta", "my-secret", payload, ts)
require.NoError(t, err)
require.NotNil(t, result)
assert.Equal(t, ts, result.Timestamp)
assert.Equal(t, "X-Hub-Signature-256", result.HeaderName)
assert.True(t, strings.HasPrefix(result.HeaderValue, "sha256="))
assert.NotEmpty(t, result.Signature)
// Verify the signature is correct
expectedMsg := fmt.Sprintf("%d.%s", ts, string(payload))
mac := hmac.New(sha256.New, []byte("my-secret"))
mac.Write([]byte(expectedMsg))
expectedSig := hex.EncodeToString(mac.Sum(nil))
assert.Equal(t, expectedSig, result.Signature)
assert.Equal(t, "sha256="+expectedSig, result.HeaderValue)
}
func TestGenerateSignature_EmptySecret(t *testing.T) {
svc := NewWebhookSignatureService(DefaultWebhookSignatureConfig())
_, err := svc.GenerateSignature("meta", "", []byte("payload"), 0)
assert.Error(t, err)
assert.Contains(t, err.Error(), "secret cannot be empty")
}
func TestGenerateSignature_EmptyPayload(t *testing.T) {
svc := NewWebhookSignatureService(DefaultWebhookSignatureConfig())
_, err := svc.GenerateSignature("meta", "secret", []byte{}, 0)
assert.Error(t, err)
assert.Contains(t, err.Error(), "payload cannot be empty")
}
func TestGenerateSignature_UnknownChannel(t *testing.T) {
svc := NewWebhookSignatureService(DefaultWebhookSignatureConfig())
_, err := svc.GenerateSignature("unknown_channel", "secret", []byte("payload"), 0)
assert.Error(t, err)
assert.Contains(t, err.Error(), "unknown channel type")
}
func TestGenerateSignature_DefaultTimestamp(t *testing.T) {
svc := NewWebhookSignatureService(DefaultWebhookSignatureConfig())
before := time.Now().Unix()
result, err := svc.GenerateSignature("api", "secret", []byte("payload"), 0)
require.NoError(t, err)
after := time.Now().Unix()
assert.True(t, result.Timestamp >= before && result.Timestamp <= after)
}
func TestGenerateSignature_Telegram(t *testing.T) {
svc := NewWebhookSignatureService(DefaultWebhookSignatureConfig())
result, err := svc.GenerateSignature("telegram", "secret", []byte("payload"), 1700000000)
require.NoError(t, err)
assert.Equal(t, "X-Telegram-Bot-Api-Secret-Token", result.HeaderName)
// Telegram has empty prefix, so HeaderValue is the raw hex HMAC signature
assert.NotEmpty(t, result.HeaderValue)
assert.True(t, strings.HasPrefix(result.HeaderValue, "") == true)
}
func TestVerifySignature_Success(t *testing.T) {
svc := NewWebhookSignatureService(DefaultWebhookSignatureConfig())
payload := []byte(`{"event":"test"}`)
ts := time.Now().Unix()
result, err := svc.GenerateSignature("api", "my-secret", payload, ts)
require.NoError(t, err)
err = svc.VerifySignature("api", "my-secret", payload, result.HeaderValue, fmt.Sprintf("%d", ts))
assert.NoError(t, err)
}
func TestVerifySignature_WrongSecret(t *testing.T) {
svc := NewWebhookSignatureService(DefaultWebhookSignatureConfig())
payload := []byte(`{"event":"test"}`)
ts := time.Now().Unix()
result, err := svc.GenerateSignature("api", "correct-secret", payload, ts)
require.NoError(t, err)
err = svc.VerifySignature("api", "wrong-secret", payload, result.HeaderValue, fmt.Sprintf("%d", ts))
assert.Error(t, err)
assert.Contains(t, err.Error(), "verification failed")
}
func TestVerifySignature_EmptySecret(t *testing.T) {
svc := NewWebhookSignatureService(DefaultWebhookSignatureConfig())
err := svc.VerifySignature("api", "", []byte("payload"), "sig", "123")
assert.Error(t, err)
}
func TestVerifySignature_EmptyPayload(t *testing.T) {
svc := NewWebhookSignatureService(DefaultWebhookSignatureConfig())
err := svc.VerifySignature("api", "secret", []byte{}, "sig", "123")
assert.Error(t, err)
}
func TestVerifySignature_EmptySignature(t *testing.T) {
svc := NewWebhookSignatureService(DefaultWebhookSignatureConfig())
err := svc.VerifySignature("api", "secret", []byte("payload"), "", "123")
assert.Error(t, err)
}
func TestVerifySignature_UnknownChannel(t *testing.T) {
svc := NewWebhookSignatureService(DefaultWebhookSignatureConfig())
err := svc.VerifySignature("unknown", "secret", []byte("payload"), "sig", "123")
assert.Error(t, err)
}
func TestVerifySignature_Telegram_SecretToken(t *testing.T) {
svc := NewWebhookSignatureService(DefaultWebhookSignatureConfig())
err := svc.VerifySignature("telegram", "my-secret-token", []byte{}, "my-secret-token", "")
assert.NoError(t, err)
}
func TestVerifySignature_Telegram_WrongToken(t *testing.T) {
svc := NewWebhookSignatureService(DefaultWebhookSignatureConfig())
err := svc.VerifySignature("telegram", "correct-token", []byte{}, "wrong-token", "")
assert.Error(t, err)
assert.Contains(t, err.Error(), "verification failed")
}
func TestVerifySignature_Telegram_WithTimestamp(t *testing.T) {
svc := NewWebhookSignatureService(DefaultWebhookSignatureConfig())
ts := time.Now().Unix()
err := svc.VerifySignature("telegram", "my-token", []byte{}, "my-token", fmt.Sprintf("%d", ts))
assert.NoError(t, err)
}
func TestVerifySignature_Telegram_ExpiredTimestamp(t *testing.T) {
svc := NewWebhookSignatureService(DefaultWebhookSignatureConfig())
oldTs := time.Now().Unix() - 600
err := svc.VerifySignature("telegram", "my-token", []byte{}, "my-token", fmt.Sprintf("%d", oldTs))
assert.Error(t, err)
}
func TestVerifySignature_ExpiredTimestamp(t *testing.T) {
svc := NewWebhookSignatureService(DefaultWebhookSignatureConfig())
payload := []byte(`{"event":"test"}`)
oldTs := time.Now().Unix() - 600
result, err := svc.GenerateSignature("api", "secret", payload, oldTs)
require.NoError(t, err)
err = svc.VerifySignature("api", "secret", payload, result.HeaderValue, fmt.Sprintf("%d", oldTs))
assert.Error(t, err)
assert.Contains(t, err.Error(), "expired")
}
func TestVerifySignature_FutureTimestamp(t *testing.T) {
svc := NewWebhookSignatureService(DefaultWebhookSignatureConfig())
futureTs := time.Now().Unix() + 120
err := svc.VerifySignature("api", "secret", []byte("payload"), "some-sig", fmt.Sprintf("%d", futureTs))
assert.Error(t, err)
assert.Contains(t, err.Error(), "future")
}
func TestVerifySignature_InvalidTimestamp(t *testing.T) {
svc := NewWebhookSignatureService(DefaultWebhookSignatureConfig())
err := svc.VerifySignature("api", "secret", []byte("payload"), "some-sig", "not-a-number")
assert.Error(t, err)
assert.Contains(t, err.Error(), "invalid timestamp")
}
func TestVerifySignature_Meta_PrefixMissing(t *testing.T) {
svc := NewWebhookSignatureService(DefaultWebhookSignatureConfig())
err := svc.VerifySignature("meta", "secret", []byte("payload"), "rawhexwithoutprefix", fmt.Sprintf("%d", time.Now().Unix()))
assert.Error(t, err)
assert.Contains(t, err.Error(), "prefix")
}
func TestVerifyTelegramWebhook(t *testing.T) {
svc := NewWebhookSignatureService(DefaultWebhookSignatureConfig())
err := svc.VerifyTelegramWebhook("my-token", "my-token", "")
assert.NoError(t, err)
}
func TestVerifyTelegramWebhook_WrongToken(t *testing.T) {
svc := NewWebhookSignatureService(DefaultWebhookSignatureConfig())
err := svc.VerifyTelegramWebhook("correct", "wrong", "")
assert.Error(t, err)
}
func TestVerifyAPIWebhook(t *testing.T) {
svc := NewWebhookSignatureService(DefaultWebhookSignatureConfig())
payload := []byte(`{"event":"test"}`)
ts := time.Now().Unix()
result, err := svc.GenerateSignature("api", "my-secret", payload, ts)
require.NoError(t, err)
err = svc.VerifyAPIWebhook("my-secret", payload, result.HeaderValue, fmt.Sprintf("%d", ts))
assert.NoError(t, err)
}
func TestVerifyWebWidgetWebhook(t *testing.T) {
svc := NewWebhookSignatureService(DefaultWebhookSignatureConfig())
payload := []byte(`{"event":"test"}`)
ts := time.Now().Unix()
result, err := svc.GenerateSignature("web_widget", "my-secret", payload, ts)
require.NoError(t, err)
err = svc.VerifyWebWidgetWebhook("my-secret", payload, result.HeaderValue, fmt.Sprintf("%d", ts))
assert.NoError(t, err)
}
func TestExtractSignature_WithPrefix(t *testing.T) {
sig, err := extractSignature("sha256=abcdef", "sha256=", false)
require.NoError(t, err)
assert.Equal(t, "abcdef", sig)
}
func TestExtractSignature_PrefixMissing(t *testing.T) {
_, err := extractSignature("rawhex", "sha256=", false)
assert.Error(t, err)
}
func TestExtractSignature_NoPrefix(t *testing.T) {
sig, err := extractSignature("rawhex", "", false)
require.NoError(t, err)
assert.Equal(t, "rawhex", sig)
}
func TestExtractSignature_SecretToken(t *testing.T) {
sig, err := extractSignature("my-secret-token", "", true)
require.NoError(t, err)
assert.Equal(t, "my-secret-token", sig)
}
func TestExtractSignature_WithWhitespace(t *testing.T) {
// extractSignature does not trim whitespace before checking prefix,
// so input with leading whitespace before the prefix will fail.
sig, err := extractSignature("sha256=abcdef ", "sha256=", false)
require.NoError(t, err)
assert.Equal(t, "abcdef", sig)
}
func TestComputeHMAC_WithTimestamp(t *testing.T) {
svc := NewWebhookSignatureService(DefaultWebhookSignatureConfig())
sig, err := svc.computeHMAC("secret", []byte("payload"), 1700000000)
require.NoError(t, err)
assert.NotEmpty(t, sig)
expectedMsg := "1700000000.payload"
mac := hmac.New(sha256.New, []byte("secret"))
mac.Write([]byte(expectedMsg))
expectedSig := hex.EncodeToString(mac.Sum(nil))
assert.Equal(t, expectedSig, sig)
}
func TestComputeHMAC_WithoutTimestamp(t *testing.T) {
svc := NewWebhookSignatureService(DefaultWebhookSignatureConfig())
sig, err := svc.computeHMAC("secret", []byte("payload"), 0)
require.NoError(t, err)
assert.NotEmpty(t, sig)
mac := hmac.New(sha256.New, []byte("secret"))
mac.Write([]byte("payload"))
expectedSig := hex.EncodeToString(mac.Sum(nil))
assert.Equal(t, expectedSig, sig)
}
func TestValidateTimestamp_Valid(t *testing.T) {
svc := NewWebhookSignatureService(DefaultWebhookSignatureConfig())
ts := time.Now().Unix()
err := svc.validateTimestamp(fmt.Sprintf("%d", ts))
assert.NoError(t, err)
}
func TestValidateTimestamp_Invalid(t *testing.T) {
svc := NewWebhookSignatureService(DefaultWebhookSignatureConfig())
err := svc.validateTimestamp("not-a-number")
assert.Error(t, err)
}