security: Add rate limiting, input validation, and tenant isolation on all handlers
- Add rate limiting per endpoint (login: 5/min, API: 100/min) - Add input validation helpers (email, UUID, string, int) - Add tenant isolation to all handlers - Remove old validation.go, replace with input.go - Fix service/customer.go to use new validation functions - Build successful
This commit is contained in:
@@ -0,0 +1,132 @@
|
||||
package middleware
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"regexp"
|
||||
"strconv"
|
||||
"strings"
|
||||
)
|
||||
|
||||
var (
|
||||
// Email regex
|
||||
emailRegex = regexp.MustCompile(`^[a-zA-Z0-9._%+-]+@[a-zA-Z0-9.-]+\.[a-zA-Z]{2,}$`)
|
||||
|
||||
// UUID regex
|
||||
uuidRegex = regexp.MustCompile(`^[0-9a-fA-F]{8}-[0-9a-fA-F]{4}-[0-9a-fA-F]{4}-[0-9a-fA-F]{4}-[0-9a-fA-F]{12}$`)
|
||||
|
||||
// Safe string regex (tillåtna tecken)
|
||||
safeStringRegex = regexp.MustCompile(`^[a-zA-Z0-9\s\-_\.@,;:()\[\]{}]+$`)
|
||||
)
|
||||
|
||||
// ValidateEmail kontrollerar email-format
|
||||
func ValidateEmail(email string) error {
|
||||
if email == "" {
|
||||
return fmt.Errorf("email is required")
|
||||
}
|
||||
if len(email) > 254 {
|
||||
return fmt.Errorf("email too long")
|
||||
}
|
||||
if !emailRegex.MatchString(email) {
|
||||
return fmt.Errorf("invalid email format")
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// ValidateUUID kontrollerar UUID-format
|
||||
func ValidateUUID(id string) error {
|
||||
if id == "" {
|
||||
return fmt.Errorf("id is required")
|
||||
}
|
||||
if !uuidRegex.MatchString(id) {
|
||||
return fmt.Errorf("invalid UUID format")
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// ValidateString kontrollerar sträng-input
|
||||
func ValidateString(s string, minLen, maxLen int, required bool) error {
|
||||
if required && strings.TrimSpace(s) == "" {
|
||||
return fmt.Errorf("field is required")
|
||||
}
|
||||
if s != "" {
|
||||
if len(s) < minLen {
|
||||
return fmt.Errorf("must be at least %d characters", minLen)
|
||||
}
|
||||
if len(s) > maxLen {
|
||||
return fmt.Errorf("must be at most %d characters", maxLen)
|
||||
}
|
||||
// Kontrollera farliga tecken
|
||||
if !safeStringRegex.MatchString(s) {
|
||||
return fmt.Errorf("contains invalid characters")
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// ValidateInt kontrollerar heltal
|
||||
func ValidateInt(val int, min, max int, required bool) error {
|
||||
if required && val == 0 {
|
||||
return fmt.Errorf("field is required")
|
||||
}
|
||||
if val != 0 {
|
||||
if val < min {
|
||||
return fmt.Errorf("must be at least %d", min)
|
||||
}
|
||||
if val > max {
|
||||
return fmt.Errorf("must be at most %d", max)
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// ValidatePagination kontrollerar pagination-parametrar
|
||||
func ValidatePagination(page, limit int) (int, int, error) {
|
||||
if page < 1 {
|
||||
page = 1
|
||||
}
|
||||
if limit < 1 {
|
||||
limit = 20
|
||||
}
|
||||
if limit > 100 {
|
||||
limit = 100
|
||||
}
|
||||
return page, limit, nil
|
||||
}
|
||||
|
||||
// SanitizeString tar bort farliga tecken
|
||||
func SanitizeString(s string) string {
|
||||
// Ta bort null bytes
|
||||
s = strings.ReplaceAll(s, "\x00", "")
|
||||
// Ta bort kontrolltecken
|
||||
var result strings.Builder
|
||||
for _, r := range s {
|
||||
if r >= 32 || r == '\t' || r == '\n' || r == '\r' {
|
||||
result.WriteRune(r)
|
||||
}
|
||||
}
|
||||
return strings.TrimSpace(result.String())
|
||||
}
|
||||
|
||||
// ParseInt parse en sträng till heltal med validering
|
||||
func ParseInt(s string, defaultVal int) int {
|
||||
if s == "" {
|
||||
return defaultVal
|
||||
}
|
||||
val, err := strconv.Atoi(s)
|
||||
if err != nil {
|
||||
return defaultVal
|
||||
}
|
||||
return val
|
||||
}
|
||||
|
||||
// ParseFloat parse en sträng till float med validering
|
||||
func ParseFloat(s string, defaultVal float64) float64 {
|
||||
if s == "" {
|
||||
return defaultVal
|
||||
}
|
||||
val, err := strconv.ParseFloat(s, 64)
|
||||
if err != nil {
|
||||
return defaultVal
|
||||
}
|
||||
return val
|
||||
}
|
||||
@@ -0,0 +1,125 @@
|
||||
package middleware
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"net/http"
|
||||
"strings"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"golang.org/x/time/rate"
|
||||
)
|
||||
|
||||
// RateLimiter hanterar rate limiting per endpoint
|
||||
type RateLimiter struct {
|
||||
limiters map[string]*rate.Limiter
|
||||
mu sync.RWMutex
|
||||
// Default rates
|
||||
defaultRate rate.Limit
|
||||
defaultBurst int
|
||||
}
|
||||
|
||||
// NewRateLimiter skapar en ny rate limiter
|
||||
func NewRateLimiter() *RateLimiter {
|
||||
return &RateLimiter{
|
||||
limiters: make(map[string]*rate.Limiter),
|
||||
defaultRate: rate.Every(time.Second), // 1 request per second
|
||||
defaultBurst: 10,
|
||||
}
|
||||
}
|
||||
|
||||
// getLimiter hämtar eller skapar en limiter för en given nyckel
|
||||
func (rl *RateLimiter) getLimiter(key string, r rate.Limit, b int) *rate.Limiter {
|
||||
rl.mu.RLock()
|
||||
limiter, exists := rl.limiters[key]
|
||||
rl.mu.RUnlock()
|
||||
|
||||
if exists {
|
||||
return limiter
|
||||
}
|
||||
|
||||
rl.mu.Lock()
|
||||
defer rl.mu.Unlock()
|
||||
|
||||
// Dubbelkolla efter lås
|
||||
limiter, exists = rl.limiters[key]
|
||||
if exists {
|
||||
return limiter
|
||||
}
|
||||
|
||||
limiter = rate.NewLimiter(r, b)
|
||||
rl.limiters[key] = limiter
|
||||
return limiter
|
||||
}
|
||||
|
||||
// RateLimit middleware med anpassade gränser per endpoint
|
||||
func RateLimit(rl *RateLimiter) func(http.Handler) http.Handler {
|
||||
return func(next http.Handler) http.Handler {
|
||||
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
// Skapa nyckel baserat på IP + endpoint
|
||||
clientIP := getClientIP(r)
|
||||
endpoint := r.Method + " " + r.URL.Path
|
||||
key := clientIP + ":" + endpoint
|
||||
|
||||
// Anpassade gränser per endpoint
|
||||
var limit rate.Limit
|
||||
var burst int
|
||||
|
||||
switch {
|
||||
// Login: 5 försök per minut
|
||||
case strings.Contains(endpoint, "/auth/login"):
|
||||
limit = rate.Every(12 * time.Second)
|
||||
burst = 5
|
||||
|
||||
// API endpoints: 100 per minut
|
||||
case strings.HasPrefix(r.URL.Path, "/api/v1/"):
|
||||
limit = rate.Every(600 * time.Millisecond)
|
||||
burst = 100
|
||||
|
||||
// Default: 10 per sekund
|
||||
default:
|
||||
limit = rl.defaultRate
|
||||
burst = rl.defaultBurst
|
||||
}
|
||||
|
||||
limiter := rl.getLimiter(key, limit, burst)
|
||||
|
||||
if !limiter.Allow() {
|
||||
w.Header().Set("Retry-After", "60")
|
||||
w.Header().Set("X-RateLimit-Limit", fmt.Sprintf("%v", limit))
|
||||
w.Header().Set("X-RateLimit-Remaining", "0")
|
||||
writeError(w, http.StatusTooManyRequests, "rate limit exceeded")
|
||||
return
|
||||
}
|
||||
|
||||
next.ServeHTTP(w, r)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// getClientIP hämtar klientens IP-adress
|
||||
func getClientIP(r *http.Request) string {
|
||||
// Kolla X-Forwarded-For header (för reverse proxies)
|
||||
xff := r.Header.Get("X-Forwarded-For")
|
||||
if xff != "" {
|
||||
// Ta första IP:et
|
||||
parts := strings.Split(xff, ",")
|
||||
if len(parts) > 0 {
|
||||
return strings.TrimSpace(parts[0])
|
||||
}
|
||||
}
|
||||
|
||||
// Kolla X-Real-IP
|
||||
xri := r.Header.Get("X-Real-IP")
|
||||
if xri != "" {
|
||||
return xri
|
||||
}
|
||||
|
||||
// Fallback till RemoteAddr
|
||||
ip := r.RemoteAddr
|
||||
// Ta bort port
|
||||
if idx := strings.LastIndex(ip, ":"); idx != -1 {
|
||||
ip = ip[:idx]
|
||||
}
|
||||
return ip
|
||||
}
|
||||
@@ -1,281 +0,0 @@
|
||||
package middleware
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"net/http"
|
||||
"reflect"
|
||||
"regexp"
|
||||
"strconv"
|
||||
"strings"
|
||||
)
|
||||
|
||||
var (
|
||||
emailRegex = regexp.MustCompile(`^[a-zA-Z0-9._%+-]+@[a-zA-Z0-9.-]+\.[a-zA-Z]{2,}$`)
|
||||
uuidRegex = regexp.MustCompile(`^[0-9a-fA-F]{8}-[0-9a-fA-F]{4}-[0-9a-fA-F]{4}-[0-9a-fA-F]{4}-[0-9a-fA-F]{12}$`)
|
||||
)
|
||||
|
||||
// ValidateEmail kontrollerar email-format
|
||||
func IsValidEmail(email string) bool {
|
||||
return emailRegex.MatchString(email)
|
||||
}
|
||||
|
||||
// ValidateUUID kontrollerar UUID-format (legacy alias)
|
||||
func ValidateUUID(uuid string) bool {
|
||||
return uuidRegex.MatchString(uuid)
|
||||
}
|
||||
|
||||
// ValidateEmail kontrollerar email-format (legacy alias)
|
||||
func ValidateEmail(email string) bool {
|
||||
return emailRegex.MatchString(email)
|
||||
}
|
||||
|
||||
// ValidatePhone kontrollerar telefonnummer
|
||||
func ValidatePhone(phone string) bool {
|
||||
// Tillåt +, siffror, mellanslag och bindestreck
|
||||
cleaned := strings.ReplaceAll(phone, " ", "")
|
||||
cleaned = strings.ReplaceAll(cleaned, "-", "")
|
||||
return len(cleaned) >= 8 && len(cleaned) <= 15
|
||||
}
|
||||
|
||||
// ValidateOrgNumber kontrollerar svenskt organisationsnummer
|
||||
func ValidateOrgNumber(org string) bool {
|
||||
// Ta bort mellanslag och bindestreck
|
||||
cleaned := strings.ReplaceAll(org, " ", "")
|
||||
cleaned = strings.ReplaceAll(cleaned, "-", "")
|
||||
|
||||
if len(cleaned) != 10 {
|
||||
return false
|
||||
}
|
||||
|
||||
// Kontrollera att det bara är siffror
|
||||
for _, c := range cleaned {
|
||||
if c < '0' || c > '9' {
|
||||
return false
|
||||
}
|
||||
}
|
||||
|
||||
return true
|
||||
}
|
||||
|
||||
// SanitizeString tar bort farliga tecken från strängar
|
||||
func SanitizeString(s string) string {
|
||||
// Ta bort null bytes och kontrolltecken
|
||||
var result strings.Builder
|
||||
for _, r := range s {
|
||||
if r >= 32 || r == '\t' || r == '\n' || r == '\r' {
|
||||
result.WriteRune(r)
|
||||
}
|
||||
}
|
||||
return strings.TrimSpace(result.String())
|
||||
}
|
||||
|
||||
// ── Request Validation ────────────────────────────────────────────────────
|
||||
|
||||
type Validator struct {
|
||||
errors map[string][]string
|
||||
}
|
||||
|
||||
func NewValidator() *Validator {
|
||||
return &Validator{errors: make(map[string][]string)}
|
||||
}
|
||||
|
||||
func (v *Validator) AddError(field, message string) {
|
||||
v.errors[field] = append(v.errors[field], message)
|
||||
}
|
||||
|
||||
func (v *Validator) HasErrors() bool {
|
||||
return len(v.errors) > 0
|
||||
}
|
||||
|
||||
func (v *Validator) Errors() map[string][]string {
|
||||
return v.errors
|
||||
}
|
||||
|
||||
func (v *Validator) ErrorResponse() map[string]interface{} {
|
||||
return map[string]interface{}{
|
||||
"error": "validation failed",
|
||||
"details": v.errors,
|
||||
}
|
||||
}
|
||||
|
||||
// ValidateString kontrollerar strängfält
|
||||
func (v *Validator) ValidateString(field, value string, minLen, maxLen int, required bool) {
|
||||
if required && strings.TrimSpace(value) == "" {
|
||||
v.AddError(field, "is required")
|
||||
return
|
||||
}
|
||||
|
||||
if value != "" {
|
||||
if len(value) < minLen {
|
||||
v.AddError(field, fmt.Sprintf("must be at least %d characters", minLen))
|
||||
}
|
||||
if len(value) > maxLen {
|
||||
v.AddError(field, fmt.Sprintf("must be at most %d characters", maxLen))
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// ValidateEmail kontrollerar email
|
||||
func (v *Validator) ValidateEmail(field, value string, required bool) {
|
||||
if required && strings.TrimSpace(value) == "" {
|
||||
v.AddError(field, "is required")
|
||||
return
|
||||
}
|
||||
|
||||
if value != "" && !IsValidEmail(value) {
|
||||
v.AddError(field, "invalid email format")
|
||||
}
|
||||
}
|
||||
|
||||
// ValidateUUID kontrollerar UUID
|
||||
func (v *Validator) ValidateUUID(field, value string, required bool) {
|
||||
if required && strings.TrimSpace(value) == "" {
|
||||
v.AddError(field, "is required")
|
||||
return
|
||||
}
|
||||
|
||||
if value != "" && !ValidateUUID(value) {
|
||||
v.AddError(field, "invalid UUID format")
|
||||
}
|
||||
}
|
||||
|
||||
// ValidateInt kontrollerar heltal
|
||||
func (v *Validator) ValidateInt(field string, value int, min, max int, required bool) {
|
||||
if required && value == 0 {
|
||||
v.AddError(field, "is required")
|
||||
return
|
||||
}
|
||||
|
||||
if value != 0 {
|
||||
if value < min {
|
||||
v.AddError(field, fmt.Sprintf("must be at least %d", min))
|
||||
}
|
||||
if value > max {
|
||||
v.AddError(field, fmt.Sprintf("must be at most %d", max))
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// ValidateFloat kontrollerar decimaltal
|
||||
func (v *Validator) ValidateFloat(field string, value float64, min, max float64, required bool) {
|
||||
if required && value == 0 {
|
||||
v.AddError(field, "is required")
|
||||
return
|
||||
}
|
||||
|
||||
if value != 0 {
|
||||
if value < min {
|
||||
v.AddError(field, fmt.Sprintf("must be at least %.2f", min))
|
||||
}
|
||||
if value > max {
|
||||
v.AddError(field, fmt.Sprintf("must be at most %.2f", max))
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// ValidateEnum kontrollerar att värdet finns i tillåtna värden
|
||||
func (v *Validator) ValidateEnum(field, value string, allowed []string, required bool) {
|
||||
if required && strings.TrimSpace(value) == "" {
|
||||
v.AddError(field, "is required")
|
||||
return
|
||||
}
|
||||
|
||||
if value != "" {
|
||||
found := false
|
||||
for _, a := range allowed {
|
||||
if a == value {
|
||||
found = true
|
||||
break
|
||||
}
|
||||
}
|
||||
if !found {
|
||||
v.AddError(field, fmt.Sprintf("must be one of: %s", strings.Join(allowed, ", ")))
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// ── Validation Middleware ─────────────────────────────────────────────────
|
||||
|
||||
// ValidateBody validerar request body mot en struct
|
||||
func ValidateBody(dst interface{}) func(http.Handler) http.Handler {
|
||||
return func(next http.Handler) http.Handler {
|
||||
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
if r.Body == nil {
|
||||
http.Error(w, `{"error":"request body required"}`, http.StatusBadRequest)
|
||||
return
|
||||
}
|
||||
|
||||
decoder := json.NewDecoder(r.Body)
|
||||
decoder.DisallowUnknownFields()
|
||||
|
||||
if err := decoder.Decode(dst); err != nil {
|
||||
http.Error(w, fmt.Sprintf(`{"error":"invalid request body: %s"}`, err.Error()), http.StatusBadRequest)
|
||||
return
|
||||
}
|
||||
|
||||
// Validera fält
|
||||
validator := NewValidator()
|
||||
validateStruct(validator, dst)
|
||||
|
||||
if validator.HasErrors() {
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
w.WriteHeader(http.StatusBadRequest)
|
||||
json.NewEncoder(w).Encode(validator.ErrorResponse())
|
||||
return
|
||||
}
|
||||
|
||||
// Spara validerad struct i context
|
||||
ctx := context.WithValue(r.Context(), "validated_body", dst)
|
||||
next.ServeHTTP(w, r.WithContext(ctx))
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// validateStruct validerar en struct baserat på tags
|
||||
func validateStruct(v *Validator, s interface{}) {
|
||||
val := reflect.ValueOf(s)
|
||||
if val.Kind() == reflect.Ptr {
|
||||
val = val.Elem()
|
||||
}
|
||||
|
||||
typ := val.Type()
|
||||
for i := 0; i < val.NumField(); i++ {
|
||||
field := val.Field(i)
|
||||
fieldType := typ.Field(i)
|
||||
|
||||
// Hämta validation tags
|
||||
tag := fieldType.Tag.Get("validate")
|
||||
if tag == "" {
|
||||
continue
|
||||
}
|
||||
|
||||
// Parsa tag
|
||||
parts := strings.Split(tag, ",")
|
||||
required := false
|
||||
minLen := 0
|
||||
maxLen := 255
|
||||
|
||||
for _, part := range parts {
|
||||
part = strings.TrimSpace(part)
|
||||
if part == "required" {
|
||||
required = true
|
||||
} else if strings.HasPrefix(part, "min=") {
|
||||
minLen, _ = strconv.Atoi(strings.TrimPrefix(part, "min="))
|
||||
} else if strings.HasPrefix(part, "max=") {
|
||||
maxLen, _ = strconv.Atoi(strings.TrimPrefix(part, "max="))
|
||||
}
|
||||
}
|
||||
|
||||
// Validera baserat på typ
|
||||
switch field.Kind() {
|
||||
case reflect.String:
|
||||
v.ValidateString(fieldType.Name, field.String(), minLen, maxLen, required)
|
||||
case reflect.Int, reflect.Int64:
|
||||
v.ValidateInt(fieldType.Name, int(field.Int()), 0, 999999, required)
|
||||
case reflect.Float64:
|
||||
v.ValidateFloat(fieldType.Name, field.Float(), 0, 999999999, required)
|
||||
}
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user