This commit is contained in:
@@ -15,7 +15,10 @@ func CorsMiddleware(c *gin.Context) {
|
||||
}
|
||||
// c.Writer.Header().Set("Access-Control-Allow-Origin", c.Request.Host)
|
||||
c.Writer.Header().Set("Access-Control-Allow-Methods", "POST, GET, OPTIONS, PUT, DELETE, UPDATE")
|
||||
c.Writer.Header().Set("Access-Control-Allow-Headers", "Content-Type, Origin, X-CSRF-Token, Authorization, AccessToken, Token, Range, api-header")
|
||||
c.Writer.Header().Set(
|
||||
"Access-Control-Allow-Headers",
|
||||
"Content-Type, Origin, X-CSRF-Token, Authorization, AccessToken, Token, Range, api-header, X-Signature-Enabled, X-App-Id, X-Timestamp, X-Nonce, X-Signature",
|
||||
)
|
||||
c.Writer.Header().Set("Access-Control-Expose-Headers", "Content-Length, Access-Control-Allow-Origin, Access-Control-Allow-Headers")
|
||||
c.Writer.Header().Set("Access-Control-Allow-Credentials", "true")
|
||||
c.Writer.Header().Set("Access-Control-Max-Age", "172800")
|
||||
|
||||
@@ -0,0 +1,137 @@
|
||||
package middleware
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"io"
|
||||
"strings"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
"github.com/perfect-panel/server/internal/svc"
|
||||
"github.com/perfect-panel/server/pkg/result"
|
||||
"github.com/perfect-panel/server/pkg/signature"
|
||||
"github.com/perfect-panel/server/pkg/xerr"
|
||||
)
|
||||
|
||||
var (
|
||||
publicSignaturePrefixes = []string{
|
||||
"/v1/auth",
|
||||
"/v1/auth/oauth",
|
||||
"/v1/common",
|
||||
"/v1/public",
|
||||
}
|
||||
defaultSignatureSkipPrefixes = []string{
|
||||
"/v1/notify/",
|
||||
"/v1/iap/notifications",
|
||||
"/v1/telegram/webhook",
|
||||
"/v1/subscribe/config",
|
||||
}
|
||||
)
|
||||
|
||||
func SignatureMiddleware(srvCtx *svc.ServiceContext) func(c *gin.Context) {
|
||||
skipPrefixes := collectSignatureSkipPrefixes(srvCtx)
|
||||
return func(c *gin.Context) {
|
||||
path := c.Request.URL.Path
|
||||
if !isPublicSignaturePath(path) {
|
||||
c.Next()
|
||||
return
|
||||
}
|
||||
|
||||
for _, prefix := range skipPrefixes {
|
||||
if strings.HasPrefix(path, prefix) {
|
||||
c.Next()
|
||||
return
|
||||
}
|
||||
}
|
||||
if !srvCtx.Config.Signature.EnableSignature {
|
||||
c.Next()
|
||||
return
|
||||
}
|
||||
if c.GetHeader("X-Signature-Enabled") != "1" {
|
||||
c.Next()
|
||||
return
|
||||
}
|
||||
|
||||
appId := c.GetHeader("X-App-Id")
|
||||
if appId == "" {
|
||||
result.HttpResult(c, nil, xerr.NewErrCode(xerr.InvalidAccess))
|
||||
c.Abort()
|
||||
return
|
||||
}
|
||||
|
||||
timestamp := c.GetHeader("X-Timestamp")
|
||||
nonce := c.GetHeader("X-Nonce")
|
||||
sig := c.GetHeader("X-Signature")
|
||||
if timestamp == "" || nonce == "" || sig == "" {
|
||||
result.HttpResult(c, nil, xerr.NewErrCode(xerr.SignatureMissing))
|
||||
c.Abort()
|
||||
return
|
||||
}
|
||||
|
||||
var bodyBytes []byte
|
||||
if c.Request.Body != nil {
|
||||
bodyBytes, _ = io.ReadAll(c.Request.Body)
|
||||
c.Request.Body = io.NopCloser(bytes.NewBuffer(bodyBytes))
|
||||
}
|
||||
|
||||
sts := signature.BuildStringToSign(
|
||||
c.Request.Method,
|
||||
path,
|
||||
c.Request.URL.RawQuery,
|
||||
bodyBytes,
|
||||
appId,
|
||||
timestamp,
|
||||
nonce,
|
||||
)
|
||||
if err := srvCtx.SignatureValidator.Validate(c.Request.Context(), appId, timestamp, nonce, sig, sts); err != nil {
|
||||
result.HttpResult(c, nil, xerr.NewErrCode(mapSignatureErr(err)))
|
||||
c.Abort()
|
||||
return
|
||||
}
|
||||
|
||||
c.Next()
|
||||
}
|
||||
}
|
||||
|
||||
func isPublicSignaturePath(path string) bool {
|
||||
for _, prefix := range publicSignaturePrefixes {
|
||||
if strings.HasPrefix(path, prefix) {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
func collectSignatureSkipPrefixes(srvCtx *svc.ServiceContext) []string {
|
||||
prefixSet := make(map[string]struct{}, len(defaultSignatureSkipPrefixes)+len(srvCtx.Config.AppSignature.SkipPrefixes)+1)
|
||||
for _, prefix := range defaultSignatureSkipPrefixes {
|
||||
prefixSet[prefix] = struct{}{}
|
||||
}
|
||||
for _, prefix := range srvCtx.Config.AppSignature.SkipPrefixes {
|
||||
if strings.TrimSpace(prefix) == "" {
|
||||
continue
|
||||
}
|
||||
prefixSet[prefix] = struct{}{}
|
||||
}
|
||||
if path := strings.TrimSpace(srvCtx.Config.Subscribe.SubscribePath); path != "" {
|
||||
prefixSet[path] = struct{}{}
|
||||
}
|
||||
|
||||
prefixes := make([]string, 0, len(prefixSet))
|
||||
for prefix := range prefixSet {
|
||||
prefixes = append(prefixes, prefix)
|
||||
}
|
||||
return prefixes
|
||||
}
|
||||
|
||||
func mapSignatureErr(err error) uint32 {
|
||||
switch err {
|
||||
case signature.ErrSignatureMissing:
|
||||
return xerr.SignatureMissing
|
||||
case signature.ErrSignatureExpired:
|
||||
return xerr.SignatureExpired
|
||||
case signature.ErrSignatureReplay:
|
||||
return xerr.SignatureReplay
|
||||
default:
|
||||
return xerr.SignatureInvalid
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,239 @@
|
||||
package middleware
|
||||
|
||||
import (
|
||||
"context"
|
||||
"crypto/hmac"
|
||||
"crypto/sha256"
|
||||
"encoding/hex"
|
||||
"encoding/json"
|
||||
"io"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"strconv"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
"github.com/perfect-panel/server/internal/config"
|
||||
"github.com/perfect-panel/server/internal/svc"
|
||||
"github.com/perfect-panel/server/pkg/signature"
|
||||
"github.com/perfect-panel/server/pkg/xerr"
|
||||
)
|
||||
|
||||
type testNonceStore struct {
|
||||
seen map[string]bool
|
||||
}
|
||||
|
||||
func newTestNonceStore() *testNonceStore {
|
||||
return &testNonceStore{seen: map[string]bool{}}
|
||||
}
|
||||
|
||||
func (s *testNonceStore) SetIfNotExists(_ context.Context, appId, nonce string, _ int64) (bool, error) {
|
||||
key := appId + ":" + nonce
|
||||
if s.seen[key] {
|
||||
return true, nil
|
||||
}
|
||||
s.seen[key] = true
|
||||
return false, nil
|
||||
}
|
||||
|
||||
func makeTestSignature(secret, sts string) string {
|
||||
mac := hmac.New(sha256.New, []byte(secret))
|
||||
mac.Write([]byte(sts))
|
||||
return hex.EncodeToString(mac.Sum(nil))
|
||||
}
|
||||
|
||||
func newTestServiceContext() *svc.ServiceContext {
|
||||
conf := config.Config{}
|
||||
conf.Signature.EnableSignature = true
|
||||
conf.AppSignature = signature.SignatureConf{
|
||||
AppSecrets: map[string]string{
|
||||
"web-client": "test-secret",
|
||||
},
|
||||
ValidWindowSeconds: 300,
|
||||
SkipPrefixes: []string{
|
||||
"/v1/public/health",
|
||||
},
|
||||
}
|
||||
return &svc.ServiceContext{
|
||||
Config: conf,
|
||||
SignatureValidator: signature.NewValidator(conf.AppSignature, newTestNonceStore()),
|
||||
}
|
||||
}
|
||||
|
||||
func newTestServiceContextWithSwitch(enabled bool) *svc.ServiceContext {
|
||||
svcCtx := newTestServiceContext()
|
||||
svcCtx.Config.Signature.EnableSignature = enabled
|
||||
return svcCtx
|
||||
}
|
||||
|
||||
func decodeCode(t *testing.T, body []byte) uint32 {
|
||||
t.Helper()
|
||||
var resp struct {
|
||||
Code uint32 `json:"code"`
|
||||
}
|
||||
if err := json.Unmarshal(body, &resp); err != nil {
|
||||
t.Fatalf("unmarshal response failed: %v", err)
|
||||
}
|
||||
return resp.Code
|
||||
}
|
||||
|
||||
func TestSignatureMiddlewareMissingAppID(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
svcCtx := newTestServiceContext()
|
||||
r := gin.New()
|
||||
r.Use(SignatureMiddleware(svcCtx))
|
||||
r.GET("/v1/public/ping", func(c *gin.Context) {
|
||||
c.String(http.StatusOK, "ok")
|
||||
})
|
||||
|
||||
req := httptest.NewRequest(http.MethodGet, "/v1/public/ping", nil)
|
||||
req.Header.Set("X-Signature-Enabled", "1")
|
||||
resp := httptest.NewRecorder()
|
||||
r.ServeHTTP(resp, req)
|
||||
|
||||
if code := decodeCode(t, resp.Body.Bytes()); code != xerr.InvalidAccess {
|
||||
t.Fatalf("expected InvalidAccess(%d), got %d", xerr.InvalidAccess, code)
|
||||
}
|
||||
}
|
||||
|
||||
func TestSignatureMiddlewareMissingSignatureHeaders(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
svcCtx := newTestServiceContext()
|
||||
r := gin.New()
|
||||
r.Use(SignatureMiddleware(svcCtx))
|
||||
r.GET("/v1/public/ping", func(c *gin.Context) {
|
||||
c.String(http.StatusOK, "ok")
|
||||
})
|
||||
|
||||
req := httptest.NewRequest(http.MethodGet, "/v1/public/ping", nil)
|
||||
req.Header.Set("X-Signature-Enabled", "1")
|
||||
req.Header.Set("X-App-Id", "web-client")
|
||||
resp := httptest.NewRecorder()
|
||||
r.ServeHTTP(resp, req)
|
||||
|
||||
if code := decodeCode(t, resp.Body.Bytes()); code != xerr.SignatureMissing {
|
||||
t.Fatalf("expected SignatureMissing(%d), got %d", xerr.SignatureMissing, code)
|
||||
}
|
||||
}
|
||||
|
||||
func TestSignatureMiddlewarePassesWhenSignatureHeaderMissing(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
svcCtx := newTestServiceContext()
|
||||
r := gin.New()
|
||||
r.Use(SignatureMiddleware(svcCtx))
|
||||
r.GET("/v1/public/ping", func(c *gin.Context) {
|
||||
c.String(http.StatusOK, "ok")
|
||||
})
|
||||
|
||||
req := httptest.NewRequest(http.MethodGet, "/v1/public/ping", nil)
|
||||
resp := httptest.NewRecorder()
|
||||
r.ServeHTTP(resp, req)
|
||||
|
||||
if resp.Code != http.StatusOK || resp.Body.String() != "ok" {
|
||||
t.Fatalf("expected pass-through without X-Signature-Enabled, got code=%d body=%s", resp.Code, resp.Body.String())
|
||||
}
|
||||
}
|
||||
|
||||
func TestSignatureMiddlewarePassesWhenSignatureHeaderIsZero(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
svcCtx := newTestServiceContext()
|
||||
r := gin.New()
|
||||
r.Use(SignatureMiddleware(svcCtx))
|
||||
r.GET("/v1/public/ping", func(c *gin.Context) {
|
||||
c.String(http.StatusOK, "ok")
|
||||
})
|
||||
|
||||
req := httptest.NewRequest(http.MethodGet, "/v1/public/ping", nil)
|
||||
req.Header.Set("X-Signature-Enabled", "0")
|
||||
resp := httptest.NewRecorder()
|
||||
r.ServeHTTP(resp, req)
|
||||
|
||||
if resp.Code != http.StatusOK || resp.Body.String() != "ok" {
|
||||
t.Fatalf("expected pass-through when X-Signature-Enabled=0, got code=%d body=%s", resp.Code, resp.Body.String())
|
||||
}
|
||||
}
|
||||
|
||||
func TestSignatureMiddlewarePassesWhenSystemSwitchDisabled(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
svcCtx := newTestServiceContextWithSwitch(false)
|
||||
r := gin.New()
|
||||
r.Use(SignatureMiddleware(svcCtx))
|
||||
r.GET("/v1/public/ping", func(c *gin.Context) {
|
||||
c.String(http.StatusOK, "ok")
|
||||
})
|
||||
|
||||
req := httptest.NewRequest(http.MethodGet, "/v1/public/ping", nil)
|
||||
req.Header.Set("X-Signature-Enabled", "1")
|
||||
resp := httptest.NewRecorder()
|
||||
r.ServeHTTP(resp, req)
|
||||
|
||||
if resp.Code != http.StatusOK || resp.Body.String() != "ok" {
|
||||
t.Fatalf("expected pass-through when system switch is disabled, got code=%d body=%s", resp.Code, resp.Body.String())
|
||||
}
|
||||
}
|
||||
|
||||
func TestSignatureMiddlewareSkipsNonPublicPath(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
svcCtx := newTestServiceContext()
|
||||
r := gin.New()
|
||||
r.Use(SignatureMiddleware(svcCtx))
|
||||
r.GET("/v1/admin/ping", func(c *gin.Context) {
|
||||
c.String(http.StatusOK, "ok")
|
||||
})
|
||||
|
||||
req := httptest.NewRequest(http.MethodGet, "/v1/admin/ping", nil)
|
||||
resp := httptest.NewRecorder()
|
||||
r.ServeHTTP(resp, req)
|
||||
|
||||
if resp.Code != http.StatusOK || resp.Body.String() != "ok" {
|
||||
t.Fatalf("expected pass-through for non-public path, got code=%d body=%s", resp.Code, resp.Body.String())
|
||||
}
|
||||
}
|
||||
|
||||
func TestSignatureMiddlewareHonorsSkipPrefix(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
svcCtx := newTestServiceContext()
|
||||
r := gin.New()
|
||||
r.Use(SignatureMiddleware(svcCtx))
|
||||
r.GET("/v1/public/healthz", func(c *gin.Context) {
|
||||
c.String(http.StatusOK, "ok")
|
||||
})
|
||||
|
||||
req := httptest.NewRequest(http.MethodGet, "/v1/public/healthz", nil)
|
||||
resp := httptest.NewRecorder()
|
||||
r.ServeHTTP(resp, req)
|
||||
|
||||
if resp.Code != http.StatusOK || resp.Body.String() != "ok" {
|
||||
t.Fatalf("expected skip-prefix pass-through, got code=%d body=%s", resp.Code, resp.Body.String())
|
||||
}
|
||||
}
|
||||
|
||||
func TestSignatureMiddlewareRestoresBodyAfterVerify(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
svcCtx := newTestServiceContext()
|
||||
r := gin.New()
|
||||
r.Use(SignatureMiddleware(svcCtx))
|
||||
r.POST("/v1/public/body", func(c *gin.Context) {
|
||||
body, _ := io.ReadAll(c.Request.Body)
|
||||
c.String(http.StatusOK, string(body))
|
||||
})
|
||||
|
||||
body := `{"hello":"world"}`
|
||||
req := httptest.NewRequest(http.MethodPost, "/v1/public/body?a=1&b=2", strings.NewReader(body))
|
||||
ts := strconv.FormatInt(time.Now().Unix(), 10)
|
||||
nonce := "nonce-body-1"
|
||||
sts := signature.BuildStringToSign(http.MethodPost, "/v1/public/body", "a=1&b=2", []byte(body), "web-client", ts, nonce)
|
||||
req.Header.Set("X-Signature-Enabled", "1")
|
||||
req.Header.Set("X-App-Id", "web-client")
|
||||
req.Header.Set("X-Timestamp", ts)
|
||||
req.Header.Set("X-Nonce", nonce)
|
||||
req.Header.Set("X-Signature", makeTestSignature("test-secret", sts))
|
||||
resp := httptest.NewRecorder()
|
||||
r.ServeHTTP(resp, req)
|
||||
|
||||
if resp.Code != http.StatusOK || resp.Body.String() != body {
|
||||
t.Fatalf("expected restored body, got code=%d body=%s", resp.Code, resp.Body.String())
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user