initial
This commit is contained in:
@@ -0,0 +1,314 @@
|
||||
// Package spam scores inbound messages using static heuristics plus
|
||||
// optional Bayesian token analysis (per RFC 5965 conventions).
|
||||
// Score >= threshold → deliver to Spam folder.
|
||||
package spam
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"fmt"
|
||||
"math"
|
||||
"net"
|
||||
"strings"
|
||||
"time"
|
||||
"unicode"
|
||||
|
||||
"ghb.freebede.com/nahakubuilder/mailgosend/internal/db"
|
||||
"ghb.freebede.com/nahakubuilder/mailgosend/internal/spf"
|
||||
)
|
||||
|
||||
// Scorer evaluates spam likelihood.
|
||||
type Scorer struct {
|
||||
db *db.DB
|
||||
threshold int
|
||||
dnsbl []string
|
||||
checkSPF bool
|
||||
checkDKIM bool
|
||||
}
|
||||
|
||||
// Result holds the total spam score and the component breakdown.
|
||||
type Result struct {
|
||||
Total int
|
||||
Reasons []string
|
||||
IsSpam bool
|
||||
}
|
||||
|
||||
// Params groups message features for scoring.
|
||||
type Params struct {
|
||||
ClientIP net.IP
|
||||
SenderDomain string
|
||||
SPFResult spf.Result
|
||||
DKIMValid bool
|
||||
DKIMPresent bool
|
||||
DMARCFail bool
|
||||
Subject string
|
||||
FromHeader string
|
||||
HasHTMLOnly bool // true if no text/plain part
|
||||
RecipCount int
|
||||
HasDateHeader bool
|
||||
HasMsgIDHeader bool
|
||||
BodyText string // first 1000 bytes of plain text for token analysis
|
||||
}
|
||||
|
||||
// NewScorer creates a scorer from config values.
|
||||
func NewScorer(database *db.DB, threshold int, dnsbl []string, checkSPF, checkDKIM bool) *Scorer {
|
||||
return &Scorer{
|
||||
db: database,
|
||||
threshold: threshold,
|
||||
dnsbl: dnsbl,
|
||||
checkSPF: checkSPF,
|
||||
checkDKIM: checkDKIM,
|
||||
}
|
||||
}
|
||||
|
||||
// Score evaluates the message and returns a Result.
|
||||
func (s *Scorer) Score(ctx context.Context, userID int64, p *Params) *Result {
|
||||
r := &Result{}
|
||||
add := func(pts int, reason string) {
|
||||
r.Total += pts
|
||||
r.Reasons = append(r.Reasons, fmt.Sprintf("+%d: %s", pts, reason))
|
||||
}
|
||||
|
||||
// DNSBL check (async-ish: each lookup gets its own goroutine with timeout)
|
||||
if p.ClientIP != nil {
|
||||
hits := s.dnsblCheck(ctx, p.ClientIP)
|
||||
for _, bl := range hits {
|
||||
add(5, "DNSBL hit: "+bl)
|
||||
}
|
||||
}
|
||||
|
||||
// SPF
|
||||
if s.checkSPF {
|
||||
switch p.SPFResult {
|
||||
case spf.ResultFail:
|
||||
add(4, "SPF fail")
|
||||
case spf.ResultSoftFail:
|
||||
add(2, "SPF softfail")
|
||||
case spf.ResultNone:
|
||||
add(1, "SPF none (no record)")
|
||||
}
|
||||
}
|
||||
|
||||
// DKIM
|
||||
if s.checkDKIM {
|
||||
if !p.DKIMPresent {
|
||||
add(2, "DKIM absent")
|
||||
} else if !p.DKIMValid {
|
||||
add(3, "DKIM invalid signature")
|
||||
}
|
||||
}
|
||||
|
||||
// DMARC
|
||||
if p.DMARCFail {
|
||||
add(5, "DMARC fail")
|
||||
}
|
||||
|
||||
// Missing required headers.
|
||||
if !p.HasDateHeader {
|
||||
add(1, "missing Date header")
|
||||
}
|
||||
if !p.HasMsgIDHeader {
|
||||
add(1, "missing Message-ID header")
|
||||
}
|
||||
|
||||
// HTML-only (no text/plain).
|
||||
if p.HasHTMLOnly {
|
||||
add(1, "HTML-only body")
|
||||
}
|
||||
|
||||
// All-caps subject.
|
||||
if p.Subject != "" && isAllCaps(p.Subject) {
|
||||
add(1, "all-caps subject")
|
||||
}
|
||||
|
||||
// Excessive recipients.
|
||||
if p.RecipCount > 20 {
|
||||
add(2, fmt.Sprintf("excessive recipients (%d)", p.RecipCount))
|
||||
}
|
||||
|
||||
// Bayesian (per-user trained model).
|
||||
if userID > 0 && p.BodyText != "" {
|
||||
bayesScore, err := s.bayesScore(ctx, userID, p.BodyText)
|
||||
if err == nil && bayesScore >= 0.8 {
|
||||
pts := int((bayesScore - 0.7) * 20) // 0.8→2 pts, 0.9→4 pts, 1.0→6 pts
|
||||
add(pts, fmt.Sprintf("Bayesian score %.2f", bayesScore))
|
||||
}
|
||||
}
|
||||
|
||||
r.IsSpam = r.Total >= s.threshold
|
||||
return r
|
||||
}
|
||||
|
||||
// TrainSpam adds body tokens to the user's spam corpus.
|
||||
func (s *Scorer) TrainSpam(ctx context.Context, userID int64, body string) error {
|
||||
return s.trainTokens(ctx, userID, body, true)
|
||||
}
|
||||
|
||||
// TrainHam adds body tokens to the user's ham corpus.
|
||||
func (s *Scorer) TrainHam(ctx context.Context, userID int64, body string) error {
|
||||
return s.trainTokens(ctx, userID, body, false)
|
||||
}
|
||||
|
||||
// ---- DNSBL ----
|
||||
|
||||
func (s *Scorer) dnsblCheck(ctx context.Context, ip net.IP) []string {
|
||||
ipv4 := ip.To4()
|
||||
if ipv4 == nil {
|
||||
return nil // DNSBL queries are IPv4-only for now
|
||||
}
|
||||
|
||||
// Reverse the IP octets: 1.2.3.4 → 4.3.2.1
|
||||
reversed := fmt.Sprintf("%d.%d.%d.%d", ipv4[3], ipv4[2], ipv4[1], ipv4[0])
|
||||
|
||||
type result struct{ bl string }
|
||||
hits := make(chan result, len(s.dnsbl))
|
||||
|
||||
timeout, cancel := context.WithTimeout(ctx, 5*time.Second)
|
||||
defer cancel()
|
||||
|
||||
for _, bl := range s.dnsbl {
|
||||
bl := bl
|
||||
go func() {
|
||||
query := reversed + "." + bl
|
||||
addrs, err := net.DefaultResolver.LookupHost(timeout, query)
|
||||
if err == nil && len(addrs) > 0 {
|
||||
hits <- result{bl}
|
||||
} else {
|
||||
hits <- result{}
|
||||
}
|
||||
}()
|
||||
}
|
||||
|
||||
var matched []string
|
||||
for range s.dnsbl {
|
||||
if h := <-hits; h.bl != "" {
|
||||
matched = append(matched, h.bl)
|
||||
}
|
||||
}
|
||||
return matched
|
||||
}
|
||||
|
||||
// ---- Bayesian ----
|
||||
|
||||
func tokenize(body string) []string {
|
||||
body = strings.ToLower(body)
|
||||
var tokens []string
|
||||
seen := make(map[string]struct{})
|
||||
words := strings.FieldsFunc(body, func(r rune) bool {
|
||||
return !unicode.IsLetter(r) && !unicode.IsDigit(r)
|
||||
})
|
||||
for _, w := range words {
|
||||
if len(w) < 3 || len(w) > 30 {
|
||||
continue
|
||||
}
|
||||
if _, ok := seen[w]; ok {
|
||||
continue
|
||||
}
|
||||
seen[w] = struct{}{}
|
||||
tokens = append(tokens, w)
|
||||
if len(tokens) >= 200 {
|
||||
break
|
||||
}
|
||||
}
|
||||
return tokens
|
||||
}
|
||||
|
||||
// bayesScore returns the probability [0,1] that the message is spam.
|
||||
func (s *Scorer) bayesScore(ctx context.Context, userID int64, body string) (float64, error) {
|
||||
tokens := tokenize(body)
|
||||
if len(tokens) == 0 {
|
||||
return 0, nil
|
||||
}
|
||||
|
||||
// Fetch total spam/ham message counts for this user (proxy: count rows with nonzero counts).
|
||||
var totalSpam, totalHam int64
|
||||
err := s.db.SQL().QueryRowContext(ctx, `
|
||||
SELECT COALESCE(SUM(spam_count),0), COALESCE(SUM(ham_count),0)
|
||||
FROM spam_tokens WHERE user_id=?`, userID).Scan(&totalSpam, &totalHam)
|
||||
if err != nil || (totalSpam+totalHam) < 50 {
|
||||
// Not enough training data.
|
||||
return 0, nil
|
||||
}
|
||||
|
||||
// Naive Bayes: P(spam|words) ∝ Π P(word|spam) / Π P(word|ham)
|
||||
// Use log-probabilities to avoid underflow.
|
||||
var logP float64
|
||||
|
||||
placeholders := make([]string, len(tokens))
|
||||
args := make([]interface{}, len(tokens)+1)
|
||||
args[0] = userID
|
||||
for i, tok := range tokens {
|
||||
placeholders[i] = "?"
|
||||
args[i+1] = tok
|
||||
}
|
||||
|
||||
query := fmt.Sprintf(`
|
||||
SELECT token, spam_count, ham_count FROM spam_tokens
|
||||
WHERE user_id=? AND token IN (%s)`, strings.Join(placeholders, ","))
|
||||
|
||||
rows, err := s.db.SQL().QueryContext(ctx, query, args...)
|
||||
if err != nil {
|
||||
return 0, err
|
||||
}
|
||||
defer rows.Close()
|
||||
|
||||
tokenData := make(map[string][2]int64)
|
||||
for rows.Next() {
|
||||
var tok string
|
||||
var sc, hc int64
|
||||
if err := rows.Scan(&tok, &sc, &hc); err != nil {
|
||||
return 0, err
|
||||
}
|
||||
tokenData[tok] = [2]int64{sc, hc}
|
||||
}
|
||||
|
||||
for _, tok := range tokens {
|
||||
counts := tokenData[tok]
|
||||
sc, hc := counts[0], counts[1]
|
||||
|
||||
// Laplace smoothing.
|
||||
pSpam := float64(sc+1) / float64(totalSpam+2)
|
||||
pHam := float64(hc+1) / float64(totalHam+2)
|
||||
|
||||
logP += math.Log(pSpam) - math.Log(pHam)
|
||||
}
|
||||
|
||||
// Convert log-odds back to probability.
|
||||
prob := 1.0 / (1.0 + math.Exp(-logP))
|
||||
return prob, nil
|
||||
}
|
||||
|
||||
func (s *Scorer) trainTokens(ctx context.Context, userID int64, body string, isSpam bool) error {
|
||||
tokens := tokenize(body)
|
||||
for _, tok := range tokens {
|
||||
var col string
|
||||
if isSpam {
|
||||
col = "spam_count"
|
||||
} else {
|
||||
col = "ham_count"
|
||||
}
|
||||
|
||||
_, err := s.db.SQL().ExecContext(ctx, fmt.Sprintf(`
|
||||
INSERT INTO spam_tokens (user_id, token, %s)
|
||||
VALUES (?, ?, 1)
|
||||
ON CONFLICT(user_id, token) DO UPDATE SET %s=%s+1`, col, col, col),
|
||||
userID, tok)
|
||||
if err != nil && err != sql.ErrNoRows {
|
||||
return fmt.Errorf("train token %q: %w", tok, err)
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func isAllCaps(s string) bool {
|
||||
hasLetter := false
|
||||
for _, r := range s {
|
||||
if unicode.IsLetter(r) {
|
||||
hasLetter = true
|
||||
if unicode.IsLower(r) {
|
||||
return false
|
||||
}
|
||||
}
|
||||
}
|
||||
return hasLetter
|
||||
}
|
||||
Reference in New Issue
Block a user