Files
jiang13-bbs/backend/service/operations_security.go

121 lines
3.7 KiB
Go

package service
import (
"context"
"crypto/sha256"
"encoding/hex"
"errors"
"fmt"
"github.com/freefire/jiang13-bbs/model"
"gorm.io/gorm"
"gorm.io/gorm/clause"
"net"
"net/netip"
"os"
"strings"
"time"
)
func counterKey(key string) string { x := sha256.Sum256([]byte(key)); return hex.EncodeToString(x[:]) }
// PostgreSQL row locks provide shared atomic quotas across application instances.
func (o *Operations) Quota(key string, limit, seconds int) (int, error) {
if seconds <= 0 {
return 0, nil
}
wait := 0
e := o.db.Transaction(func(tx *gorm.DB) error {
key = counterKey(key)
if e := tx.Clauses(clause.OnConflict{DoNothing: true}).Create(&model.ActionCounter{Key: key, ExpiresAt: time.Now().Add(time.Duration(seconds) * time.Second)}).Error; e != nil {
return e
}
var c model.ActionCounter
if e := tx.Clauses(clause.Locking{Strength: "UPDATE"}).First(&c, "key = ?", key).Error; e != nil {
return e
}
now := time.Now()
if !now.Before(c.ExpiresAt) {
c.Count = 0
c.ExpiresAt = now.Add(time.Duration(seconds) * time.Second)
}
if c.Count >= limit {
wait = int(time.Until(c.ExpiresAt).Seconds()) + 1
return nil
}
c.Count++
return tx.Save(&c).Error
})
return wait, e
}
func (o *Operations) FailureWait(account, ip string, c SecurityConfig) (int, error) {
var rows []model.ActionCounter
if e := o.db.Where("key IN ? AND expires_at > ?", []string{counterKey("failure:account:" + strings.ToLower(strings.TrimSpace(account))), counterKey("failure:ip:" + ip)}, time.Now()).Find(&rows).Error; e != nil {
return 0, e
}
for _, r := range rows {
limit := c.LoginFailures
if r.Key == counterKey("failure:ip:"+ip) {
limit *= 4
}
if r.Count >= limit {
return int(time.Until(r.ExpiresAt).Seconds()) + 1, nil
}
}
return 0, nil
}
func (o *Operations) RecordFailure(account, ip string, c SecurityConfig) {
_, _ = o.Quota("failure:account:"+strings.ToLower(strings.TrimSpace(account)), c.LoginFailures, c.LoginWindow)
_, _ = o.Quota("failure:ip:"+ip, c.LoginFailures*4, c.LoginWindow)
}
func (o *Operations) ClearFailure(account string) {
o.db.Delete(&model.ActionCounter{}, "key = ?", counterKey("failure:account:"+strings.ToLower(strings.TrimSpace(account))))
}
// Resolve once, validate every address, and dial the validated IP to prevent DNS rebinding.
// Private targets require an exact deployment allowlist entry, not a settings checkbox.
func safeDial(ctx context.Context, network, address string) (net.Conn, error) {
host, port, e := net.SplitHostPort(address)
if e != nil {
return nil, errors.New("连接地址无效")
}
allow := false
for _, v := range strings.Split(os.Getenv("SERVICE_PRIVATE_HOSTS"), ",") {
if strings.EqualFold(strings.TrimSpace(v), host) {
allow = true
}
}
ips, e := net.DefaultResolver.LookupIPAddr(ctx, host)
if e != nil || len(ips) == 0 {
return nil, errors.New("地址解析失败")
}
for _, a := range ips {
if !allow && (restrictedServiceIP(a.IP) || !a.IP.IsGlobalUnicast() || a.IP.IsPrivate() || a.IP.IsLoopback() || a.IP.IsLinkLocalUnicast() || a.IP.IsUnspecified()) {
return nil, errors.New("目标地址受限;私网服务需部署 SERVICE_PRIVATE_HOSTS")
}
}
var last error
for _, a := range ips {
c, e := (&net.Dialer{Timeout: 10 * time.Second}).DialContext(ctx, network, net.JoinHostPort(a.IP.String(), port))
if e == nil {
return c, nil
}
last = e
}
_ = last
return nil, fmt.Errorf("连接失败或超时")
}
func restrictedServiceIP(ip net.IP) bool {
a, ok := netip.AddrFromSlice(ip)
if !ok {
return true
}
a = a.Unmap()
for _, p := range []string{"100.64.0.0/10", "192.0.0.0/24", "198.18.0.0/15", "2001:db8::/32"} {
if netip.MustParsePrefix(p).Contains(a) {
return true
}
}
return false
}