121 lines
3.7 KiB
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
|
|
}
|