diff --git a/xy_sh/internal/middleware/middleware.go b/xy_sh/internal/middleware/middleware.go new file mode 100644 index 0000000..595df51 --- /dev/null +++ b/xy_sh/internal/middleware/middleware.go @@ -0,0 +1,120 @@ +package middleware + +import ( + "encoding/json" + + "xy_sh/internal/entities" + "xy_sh/pkg/crypto" + + "github.com/gofiber/fiber/v2" +) + +func SignVerifyMiddleware() fiber.Handler { + return func(c *fiber.Ctx) error { + timestamp := c.Get("timestamp") + sign := c.Get("sign") + + if timestamp == "" || sign == "" { + return c.Status(fiber.StatusUnauthorized).JSON(entities.CommonResponse{ + Code: entities.CodeFailed, + Msg: "缺少timestamp或sign请求头", + Data: nil, + }) + } + + body := c.Body() + + var req entities.EncryptedRequest + if err := json.Unmarshal(body, &req); err != nil { + return c.Status(fiber.StatusBadRequest).JSON(entities.CommonResponse{ + Code: entities.CodeFailed, + Msg: "请求体格式错误", + Data: nil, + }) + } + + if req.EncryptedData == "" { + return c.Status(fiber.StatusBadRequest).JSON(entities.CommonResponse{ + Code: entities.CodeFailed, + Msg: "缺少encryptedData参数", + Data: nil, + }) + } + + if !crypto.VerifySign(timestamp, req.EncryptedData, sign) { + return c.Status(fiber.StatusUnauthorized).JSON(entities.CommonResponse{ + Code: entities.CodeFailed, + Msg: "签名验证失败", + Data: nil, + }) + } + + decryptedData, err := crypto.SM4CBCDecrypt(req.EncryptedData) + if err != nil { + return c.Status(fiber.StatusBadRequest).JSON(entities.CommonResponse{ + Code: entities.CodeFailed, + Msg: "数据解密失败: " + err.Error(), + Data: nil, + }) + } + + c.Locals("decryptedData", decryptedData) + c.Locals("timestamp", timestamp) + + return c.Next() + } +} + +func EncryptResponseMiddleware() fiber.Handler { + return func(c *fiber.Ctx) error { + err := c.Next() + if err != nil { + return err + } + + body := c.Response().Body() + if len(body) == 0 { + return nil + } + + var resp entities.CommonResponse + if err := json.Unmarshal(body, &resp); err != nil { + return nil + } + + if resp.Data != nil { + dataBytes, err := json.Marshal(resp.Data) + if err != nil { + return c.Status(fiber.StatusInternalServerError).JSON(entities.CommonResponse{ + Code: entities.CodeFailed, + Msg: "响应数据序列化失败", + Data: nil, + }) + } + + encryptedData, err := crypto.SM4CBCEncrypt(dataBytes) + if err != nil { + return c.Status(fiber.StatusInternalServerError).JSON(entities.CommonResponse{ + Code: entities.CodeFailed, + Msg: "响应数据加密失败", + Data: nil, + }) + } + + resp.Data = encryptedData + + newBody, err := json.Marshal(resp) + if err != nil { + return c.Status(fiber.StatusInternalServerError).JSON(entities.CommonResponse{ + Code: entities.CodeFailed, + Msg: "响应序列化失败", + Data: nil, + }) + } + + c.Response().SetBody(newBody) + } + + return nil + } +} \ No newline at end of file