Files
selfpost/internal/web/handlers/handlers_ratelimit.go
T
mix c1ec4fbd79
test / test (push) Waiting to run
release: 1.6.0
Add 30-day send statistics and auto level-2 rate limits on the domain page. Close Unreleased; pin compose and docs to 1.6.0.

Co-authored-by: Cursor <cursoragent@cursor.com>
2026-08-18 22:26:51 +03:00

312 lines
9.4 KiB
Go
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
package handlers
import (
"fmt"
"net"
"net/http"
"strconv"
"strings"
"github.com/mixeme/selfpost/internal/store"
)
const defaultRateLimitWindowSeconds = 3600
type rateLimitInput struct {
clear bool
mode string
ips []string
maxMessages int
windowSeconds int
autoMultiplier float64
}
func (h *Handlers) l1Messages() int {
if h.cfg.RateLimitMessagesPerIP > 0 {
return h.cfg.RateLimitMessagesPerIP
}
return 100
}
func (h *Handlers) l1Window() int {
if h.cfg.RateLimitWindowSeconds > 0 {
return h.cfg.RateLimitWindowSeconds
}
return defaultRateLimitWindowSeconds
}
func parseAutoMultiplier(raw string) (float64, error) {
raw = strings.TrimSpace(raw)
if raw == "" {
return store.DefaultAutoMultiplier, nil
}
v, err := strconv.ParseFloat(raw, 64)
if err != nil {
return 0, fmt.Errorf("enter a valid multiplier (%.1f%.1f)", store.MinAutoMultiplier, store.MaxAutoMultiplier)
}
if v < store.MinAutoMultiplier || v > store.MaxAutoMultiplier {
return 0, fmt.Errorf("multiplier must be between %.1f and %.1f", store.MinAutoMultiplier, store.MaxAutoMultiplier)
}
return v, nil
}
func parseRateLimitMode(raw string) (string, error) {
raw = strings.TrimSpace(raw)
if raw == "" {
return store.RateLimitModeManual, nil
}
if raw != store.RateLimitModeManual && raw != store.RateLimitModeAuto {
return "", fmt.Errorf("choose manual or auto mode")
}
return raw, nil
}
func parseDomainRateLimitForm(r *http.Request, l1Max int) (rateLimitInput, error) {
if err := r.ParseForm(); err != nil {
return rateLimitInput{}, fmt.Errorf("invalid form submission")
}
if r.PostFormValue("clear") != "" {
return rateLimitInput{clear: true}, nil
}
mode, err := parseRateLimitMode(r.PostFormValue("mode"))
if err != nil {
return rateLimitInput{}, err
}
if mode == store.RateLimitModeAuto {
mult, err := parseAutoMultiplier(r.PostFormValue("auto_multiplier"))
if err != nil {
return rateLimitInput{}, err
}
return rateLimitInput{mode: mode, autoMultiplier: mult}, nil
}
rawMax := strings.TrimSpace(r.PostFormValue("max_messages"))
if rawMax == "" {
return rateLimitInput{clear: true}, nil
}
maxMessages, err := parsePositiveInt(rawMax, 0)
if err != nil || maxMessages <= 0 {
return rateLimitInput{}, fmt.Errorf("enter a message limit greater than zero")
}
if maxMessages > l1Max {
return rateLimitInput{}, fmt.Errorf("message limit cannot exceed the level-1 backstop (%d)", l1Max)
}
windowSeconds, err := parsePositiveInt(r.PostFormValue("window_seconds"), defaultRateLimitWindowSeconds)
if err != nil || windowSeconds <= 0 {
return rateLimitInput{}, fmt.Errorf("enter a time window greater than zero seconds")
}
return rateLimitInput{mode: store.RateLimitModeManual, maxMessages: maxMessages, windowSeconds: windowSeconds}, nil
}
func parseAppRateLimitForm(r *http.Request, l1Max, domainMax int, domainActive bool) (rateLimitInput, error) {
if err := r.ParseForm(); err != nil {
return rateLimitInput{}, fmt.Errorf("invalid form submission")
}
if r.PostFormValue("clear") != "" {
return rateLimitInput{clear: true}, nil
}
mode, err := parseRateLimitMode(r.PostFormValue("mode"))
if err != nil {
return rateLimitInput{}, err
}
ips, err := parseIPList(r.PostFormValue("allowed_ips"))
if err != nil {
return rateLimitInput{}, err
}
if len(ips) == 0 {
return rateLimitInput{}, fmt.Errorf("enter at least one trusted client IP for an application override")
}
if mode == store.RateLimitModeAuto {
mult, err := parseAutoMultiplier(r.PostFormValue("auto_multiplier"))
if err != nil {
return rateLimitInput{}, err
}
return rateLimitInput{mode: mode, ips: ips, autoMultiplier: mult}, nil
}
rawMax := strings.TrimSpace(r.PostFormValue("max_messages"))
if rawMax == "" {
return rateLimitInput{clear: true}, nil
}
maxMessages, err := parsePositiveInt(rawMax, 0)
if err != nil || maxMessages <= 0 {
return rateLimitInput{}, fmt.Errorf("enter a message limit greater than zero")
}
if maxMessages > l1Max {
return rateLimitInput{}, fmt.Errorf("message limit cannot exceed the level-1 backstop (%d)", l1Max)
}
if domainActive && maxMessages <= domainMax {
return rateLimitInput{}, fmt.Errorf("application override must be greater than the domain limit (%d)", domainMax)
}
windowSeconds, err := parsePositiveInt(r.PostFormValue("window_seconds"), defaultRateLimitWindowSeconds)
if err != nil || windowSeconds <= 0 {
return rateLimitInput{}, fmt.Errorf("enter a time window greater than zero seconds")
}
return rateLimitInput{mode: store.RateLimitModeManual, ips: ips, maxMessages: maxMessages, windowSeconds: windowSeconds}, nil
}
func parseIPList(raw string) ([]string, error) {
fields := strings.FieldsFunc(raw, func(r rune) bool {
return r == '\n' || r == '\r' || r == ',' || r == ' ' || r == '\t' || r == ';'
})
var out []string
seen := make(map[string]bool)
for _, f := range fields {
ip := net.ParseIP(f)
if ip == nil {
return nil, fmt.Errorf("%q is not a valid IP address", f)
}
c := ip.String()
if !seen[c] {
seen[c] = true
out = append(out, c)
}
}
return out, nil
}
func parsePositiveInt(raw string, def int) (int, error) {
raw = strings.TrimSpace(raw)
if raw == "" {
return def, nil
}
return strconv.Atoi(raw)
}
func (h *Handlers) HandleDomainRateLimit(w http.ResponseWriter, r *http.Request) {
d, ok := h.lookupDomain(w, r)
if !ok {
return
}
in, err := parseDomainRateLimitForm(r, h.l1Messages())
if err != nil {
h.renderDomainDetail(w, r, http.StatusBadRequest, d, detailView{
FormMode: store.AddressModeWildcard,
RateLimitErr: err.Error(),
})
return
}
if err := h.applyDomainRateLimit(in, d.ID); err != nil {
logf("panel: domain %d: save rate limit: %v", d.ID, err)
http.Error(w, "internal error", http.StatusInternalServerError)
return
}
http.Redirect(w, r, fmt.Sprintf("/domains/%d?ratelimit=1", d.ID), http.StatusSeeOther)
}
func (h *Handlers) HandleAppRateLimit(w http.ResponseWriter, r *http.Request) {
a, ok := h.lookupApplication(w, r)
if !ok {
return
}
d, err := h.domains.Get(a.DomainID)
if err != nil {
http.Error(w, "internal error", http.StatusInternalServerError)
return
}
domainRL, domainOK, err := h.domains.RateLimit(d.ID)
if err != nil {
logf("panel: domain %d: rate limit: %v", d.ID, err)
http.Error(w, "internal error", http.StatusInternalServerError)
return
}
domainActive := domainOK && domainRL.Active()
in, err := parseAppRateLimitForm(r, h.l1Messages(), domainRL.MaxMessages, domainActive)
if err != nil {
h.renderDomainDetail(w, r, http.StatusBadRequest, d, detailView{
FormMode: store.AddressModeWildcard,
RateLimitErr: fmt.Sprintf("%s: %s", a.Login, err.Error()),
})
return
}
if err := h.applyAppRateLimit(in, a.ID); err != nil {
logf("panel: application %d: save rate limit: %v", a.ID, err)
http.Error(w, "internal error", http.StatusInternalServerError)
return
}
http.Redirect(w, r, fmt.Sprintf("/domains/%d?ratelimit=1", a.DomainID), http.StatusSeeOther)
}
func (h *Handlers) HandleDomainRateLimitRecalc(w http.ResponseWriter, r *http.Request) {
d, ok := h.lookupDomain(w, r)
if !ok {
return
}
if err := h.recalcRateLimit(store.RateLimitScopeDomain, d.ID); err != nil {
logf("panel: domain %d: recalc rate limit: %v", d.ID, err)
h.renderDomainDetail(w, r, http.StatusBadRequest, d, detailView{
FormMode: store.AddressModeWildcard,
RateLimitErr: err.Error(),
})
return
}
http.Redirect(w, r, fmt.Sprintf("/domains/%d?recalculated=1", d.ID), http.StatusSeeOther)
}
func (h *Handlers) HandleAppRateLimitRecalc(w http.ResponseWriter, r *http.Request) {
a, ok := h.lookupApplication(w, r)
if !ok {
return
}
if err := h.recalcRateLimit(store.RateLimitScopeApp, a.ID); err != nil {
d, _ := h.domains.Get(a.DomainID)
h.renderDomainDetail(w, r, http.StatusBadRequest, d, detailView{
FormMode: store.AddressModeWildcard,
RateLimitErr: fmt.Sprintf("%s: %s", a.Login, err.Error()),
})
return
}
http.Redirect(w, r, fmt.Sprintf("/domains/%d?recalculated=1", a.DomainID), http.StatusSeeOther)
}
func (h *Handlers) recalcRateLimit(scope string, refID int64) error {
return h.store.RecalcAutoRateLimit(scope, refID, h.sendLogRetentionDays(), h.l1Messages(), h.l1Window())
}
func (h *Handlers) applyDomainRateLimit(in rateLimitInput, domainID int64) error {
if in.clear {
return h.domains.ClearRateLimit(domainID)
}
rl := store.RateLimit{
Scope: store.RateLimitScopeDomain,
RefID: domainID,
Mode: in.mode,
MaxMessages: in.maxMessages,
WindowSeconds: in.windowSeconds,
AutoMultiplier: in.autoMultiplier,
}
if in.mode == store.RateLimitModeAuto {
rl.WindowSeconds = h.l1Window()
if err := h.domains.SaveRateLimit(domainID, rl); err != nil {
return err
}
return h.recalcRateLimit(store.RateLimitScopeDomain, domainID)
}
return h.domains.SaveRateLimit(domainID, rl)
}
func (h *Handlers) applyAppRateLimit(in rateLimitInput, appID int64) error {
if in.clear {
return h.apps.ClearRateLimit(appID)
}
rl := store.RateLimit{
Scope: store.RateLimitScopeApp,
RefID: appID,
AllowedIPs: in.ips,
Mode: in.mode,
MaxMessages: in.maxMessages,
WindowSeconds: in.windowSeconds,
AutoMultiplier: in.autoMultiplier,
}
if in.mode == store.RateLimitModeAuto {
rl.WindowSeconds = h.l1Window()
if err := h.apps.SaveRateLimit(appID, rl); err != nil {
return err
}
return h.recalcRateLimit(store.RateLimitScopeApp, appID)
}
return h.apps.SaveRateLimit(appID, rl)
}