1152 lines
35 KiB
Go
1152 lines
35 KiB
Go
package security
|
|
|
|
import (
|
|
"bufio"
|
|
"context"
|
|
"crypto/hmac"
|
|
"crypto/sha256"
|
|
"encoding/base64"
|
|
"encoding/hex"
|
|
"fmt"
|
|
"net"
|
|
"net/http"
|
|
"net/http/httptest"
|
|
"net/url"
|
|
"strings"
|
|
"sync/atomic"
|
|
"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 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 TestSafeRedirectCheckConfig_BlocksPrivateAndNonHTTPRedirects(t *testing.T) {
|
|
check := safeRedirectCheckConfig(DefaultSSRFConfig())
|
|
for _, target := range []string{
|
|
"http://169.254.169.254/latest/meta-data/",
|
|
"http://10.0.0.1/admin",
|
|
"ftp://example.com/file",
|
|
"https://example.com:8443/file",
|
|
} {
|
|
req, err := http.NewRequest(http.MethodGet, target, nil)
|
|
require.NoError(t, err)
|
|
assert.Error(t, check(req, nil), target)
|
|
}
|
|
}
|
|
|
|
func TestSafeHTTPClient_RevalidatesEveryRealRedirect(t *testing.T) {
|
|
var firstHits, secondHits, privateHits atomic.Int32
|
|
private := httptest.NewServer(http.HandlerFunc(func(http.ResponseWriter, *http.Request) {
|
|
privateHits.Add(1)
|
|
}))
|
|
defer private.Close()
|
|
second := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
|
secondHits.Add(1)
|
|
http.Redirect(w, r, "http://169.254.169.254/latest/meta-data/", http.StatusFound)
|
|
}))
|
|
defer second.Close()
|
|
first := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
|
firstHits.Add(1)
|
|
http.Redirect(w, r, "http://second.example/next", http.StatusFound)
|
|
}))
|
|
defer first.Close()
|
|
|
|
addresses := map[string]string{
|
|
"first.example": first.Listener.Addr().String(),
|
|
"second.example": second.Listener.Addr().String(),
|
|
"169.254.169.254": private.Listener.Addr().String(),
|
|
}
|
|
client := NewSafeHTTPClient(DefaultSSRFConfig())
|
|
client.client.Transport = &http.Transport{DialContext: func(ctx context.Context, network, addr string) (net.Conn, error) {
|
|
host, _, err := net.SplitHostPort(addr)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
return (&net.Dialer{}).DialContext(ctx, network, addresses[host])
|
|
}}
|
|
|
|
req, err := http.NewRequest(http.MethodGet, "http://first.example/start", nil)
|
|
require.NoError(t, err)
|
|
resp, err := client.Do(req)
|
|
if resp != nil {
|
|
resp.Body.Close()
|
|
}
|
|
require.ErrorContains(t, err, "private/reserved")
|
|
assert.Equal(t, int32(1), firstHits.Load())
|
|
assert.Equal(t, int32(1), secondHits.Load())
|
|
assert.Zero(t, privateHits.Load())
|
|
}
|
|
|
|
func TestSafeHTTPClient_DialsFirstValidatedDNSResult(t *testing.T) {
|
|
serverConn, clientConn := net.Pipe()
|
|
defer serverConn.Close()
|
|
serverDone := make(chan error, 1)
|
|
go func() {
|
|
request, err := http.ReadRequest(bufio.NewReader(serverConn))
|
|
if err == nil {
|
|
request.Body.Close()
|
|
_, err = serverConn.Write([]byte("HTTP/1.1 200 OK\r\nContent-Length: 2\r\nConnection: close\r\n\r\nok"))
|
|
}
|
|
serverDone <- err
|
|
}()
|
|
|
|
lookups := 0
|
|
dialedAddress := ""
|
|
client := newSafeHTTPClient(DefaultSSRFConfig(),
|
|
func(context.Context, string) ([]net.IPAddr, error) {
|
|
lookups++
|
|
if lookups == 1 {
|
|
return []net.IPAddr{{IP: net.ParseIP("93.184.216.34")}}, nil
|
|
}
|
|
return []net.IPAddr{{IP: net.ParseIP("127.0.0.1")}}, nil
|
|
},
|
|
func(_ context.Context, _, addr string) (net.Conn, error) {
|
|
dialedAddress = addr
|
|
return clientConn, nil
|
|
},
|
|
)
|
|
|
|
resp, err := client.SafeFetchURL(context.Background(), "http://rebind.example/file")
|
|
require.NoError(t, err)
|
|
defer resp.Body.Close()
|
|
assert.Equal(t, 1, lookups)
|
|
assert.Equal(t, "93.184.216.34:80", dialedAddress)
|
|
require.NoError(t, <-serverDone)
|
|
}
|
|
|
|
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_Do_BlocksResolvedLoopback(t *testing.T) {
|
|
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
|
t.Fatal("blocked target must not receive a request")
|
|
}))
|
|
defer srv.Close()
|
|
|
|
cfg := DefaultSSRFConfig()
|
|
cfg.AllowedPorts = nil
|
|
client := NewSafeHTTPClient(cfg)
|
|
req, _ := http.NewRequest(http.MethodGet, srv.URL, nil)
|
|
_, err := client.Do(req)
|
|
require.Error(t, err)
|
|
assert.Contains(t, err.Error(), "SSRF")
|
|
}
|
|
|
|
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)
|
|
}
|