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 }