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) }