127 lines
4.1 KiB
Go
127 lines
4.1 KiB
Go
// Package middleware — HTTP-промежуточные слои API.
|
|
package middleware
|
|
|
|
import (
|
|
"crypto/hmac"
|
|
"crypto/sha256"
|
|
"encoding/base64"
|
|
"log/slog"
|
|
"net/http"
|
|
"strings"
|
|
)
|
|
|
|
// Auth — проверка Bearer-токена.
|
|
//
|
|
// authTestMode (env AUTH_TEST_MODE, дефолт false):
|
|
// - true — принимается любая строка без пробелов (для локальных тестов);
|
|
// - false — проверка JWT: sub + exp; если задан hmacSecret — обязательна
|
|
// HS256-подпись (HMAC-SHA256, env JWT_HMAC_SECRET). Без секрета —
|
|
// структурная проверка (периметр обеспечивает платформа).
|
|
func Auth(authTestMode bool, hmacSecret string, log *slog.Logger, next http.Handler) http.Handler {
|
|
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
|
header := r.Header.Get("Authorization")
|
|
if header == "" {
|
|
log.Warn("auth: missing authorization header", "remote", r.RemoteAddr, "path", r.URL.Path)
|
|
http.Error(w, `{"error":"authorization required"}`, http.StatusUnauthorized)
|
|
return
|
|
}
|
|
parts := strings.SplitN(header, " ", 2)
|
|
if len(parts) != 2 || !strings.EqualFold(parts[0], "bearer") {
|
|
log.Warn("auth: invalid authorization format", "remote", r.RemoteAddr)
|
|
http.Error(w, `{"error":"invalid authorization format, use Bearer <token>"}`, http.StatusUnauthorized)
|
|
return
|
|
}
|
|
token := parts[1]
|
|
|
|
if authTestMode && isPlainToken(token) {
|
|
log.Info("auth: test mode — plain token accepted", "remote", r.RemoteAddr, "path", r.URL.Path)
|
|
next.ServeHTTP(w, r)
|
|
return
|
|
}
|
|
|
|
if err := validateJWT(token, hmacSecret); err != nil {
|
|
log.Warn("auth: invalid token", "remote", r.RemoteAddr, "path", r.URL.Path, "reason", err.Error())
|
|
http.Error(w, `{"error":"invalid token"}`, http.StatusForbidden)
|
|
return
|
|
}
|
|
next.ServeHTTP(w, r)
|
|
})
|
|
}
|
|
|
|
// isPlainToken — непустая строка без пробелов, не похожая на JWT (три части через точку).
|
|
func isPlainToken(token string) bool {
|
|
if token == "" || strings.ContainsAny(token, " \t\n\r") {
|
|
return false
|
|
}
|
|
return len(strings.Split(token, ".")) != 3
|
|
}
|
|
|
|
// jwtError — структурная ошибка JWT.
|
|
type jwtError struct{ msg string }
|
|
|
|
func (e *jwtError) Error() string { return e.msg }
|
|
|
|
// validateJWT проверяет структуру JWT (sub + exp) и, если задан secret,
|
|
// HS256-подпись. Без secret подпись не проверяется — см. комментарий Auth.
|
|
func validateJWT(token, secret string) error {
|
|
jwtParts := strings.Split(token, ".")
|
|
if len(jwtParts) != 3 {
|
|
return &jwtError{"not a JWT: expected 3 parts"}
|
|
}
|
|
|
|
if secret != "" {
|
|
var header struct {
|
|
Alg string `json:"alg"`
|
|
}
|
|
headerBytes, err := base64.RawURLEncoding.DecodeString(jwtParts[0])
|
|
if err != nil {
|
|
headerBytes, err = base64.StdEncoding.DecodeString(jwtParts[0])
|
|
}
|
|
if err != nil || jsonUnmarshal(headerBytes, &header) != nil || header.Alg != "HS256" {
|
|
return &jwtError{"JWT must be HS256 when JWT_HMAC_SECRET is set"}
|
|
}
|
|
mac := hmac.New(sha256.New, []byte(secret))
|
|
mac.Write([]byte(jwtParts[0] + "." + jwtParts[1]))
|
|
expected := mac.Sum(nil)
|
|
sig, err := base64.RawURLEncoding.DecodeString(jwtParts[2])
|
|
if err != nil {
|
|
return &jwtError{"cannot decode JWT signature"}
|
|
}
|
|
if !hmac.Equal(expected, sig) {
|
|
return &jwtError{"JWT signature mismatch"}
|
|
}
|
|
}
|
|
|
|
payload := jwtParts[1]
|
|
switch len(payload) % 4 {
|
|
case 2:
|
|
payload += "=="
|
|
case 3:
|
|
payload += "="
|
|
}
|
|
decoded, err := base64.URLEncoding.DecodeString(payload)
|
|
if err != nil {
|
|
decoded, err = base64.StdEncoding.DecodeString(payload)
|
|
}
|
|
if err != nil {
|
|
return &jwtError{"cannot decode JWT payload"}
|
|
}
|
|
|
|
var claims struct {
|
|
Sub string `json:"sub"`
|
|
Exp float64 `json:"exp"`
|
|
}
|
|
if err := jsonUnmarshal(decoded, &claims); err != nil {
|
|
return &jwtError{"cannot parse JWT payload"}
|
|
}
|
|
if claims.Sub == "" {
|
|
return &jwtError{"JWT missing sub claim"}
|
|
}
|
|
if claims.Exp > 0 {
|
|
if now := float64(timeNow().Unix()); claims.Exp < now {
|
|
return &jwtError{"JWT expired"}
|
|
}
|
|
}
|
|
return nil
|
|
}
|