Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
71 changes: 71 additions & 0 deletions api/dms/service/v1/access_restriction.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,71 @@
package v1

import base "github.com/actiontech/dms/pkg/dms-common/api/base/v1"

type AccessWhitelistRuleItem struct {
UID string `json:"uid"`
Source string `json:"source"`
PolicyType string `json:"policy_type"`
Remark string `json:"remark"`
UpdatedAt string `json:"updated_at"`
}

type AccessRestrictionConfig struct {
Enabled bool `json:"enabled"`
Rules []AccessWhitelistRuleItem `json:"rules"`
}

// swagger:model GetAccessRestrictionReply
type GetAccessRestrictionReply struct {
Data AccessRestrictionConfig `json:"data"`
base.GenericResp
}

// swagger:model
type UpdateAccessRestrictionReq struct {
Enabled *bool `json:"enabled" validate:"required"`
}

// swagger:model
type CreateAccessWhitelistRuleReq struct {
Source string `json:"source" validate:"required"`
Remark string `json:"remark"`
PolicyType string `json:"policy_type"`
}

// swagger:model CreateAccessWhitelistRuleReply
type CreateAccessWhitelistRuleReply struct {
Data AccessWhitelistRuleItem `json:"data"`
base.GenericResp
}

// swagger:parameters UpdateAccessWhitelistRuleReq
type UpdateAccessWhitelistRuleReq struct {
// in:path
RuleUID string `param:"rule_uid" json:"rule_uid" validate:"required"`
Source string `json:"source" validate:"required"`
Remark string `json:"remark"`
PolicyType string `json:"policy_type"`
}

// swagger:model UpdateAccessWhitelistRuleReply
type UpdateAccessWhitelistRuleReply struct {
Data AccessWhitelistRuleItem `json:"data"`
base.GenericResp
}

// swagger:parameters DeleteAccessWhitelistRuleReq
type DeleteAccessWhitelistRuleReq struct {
// in:path
RuleUID string `param:"rule_uid" json:"rule_uid" validate:"required"`
}

// swagger:model GetAccessRestrictionClientIPReply
type GetAccessRestrictionClientIPReply struct {
Data AccessRestrictionClientIP `json:"data"`
base.GenericResp
}

type AccessRestrictionClientIP struct {
ClientIP string `json:"client_ip"`
}
67 changes: 67 additions & 0 deletions internal/apiserver/middleware/access_restriction.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,67 @@
package middleware

import (
"fmt"
"net/http"
"strings"

"github.com/actiontech/dms/internal/dms/biz"
dmsV1 "github.com/actiontech/dms/pkg/dms-common/api/dms/v1"
"github.com/labstack/echo/v4"
)

const accessRestrictionDenyMsg = "禁止访问:来源 IP 不在访问白名单"

// AccessRestriction enforces IP/CIDR access restriction for all DMS entry points.
// Order (must not reorder short-circuits):
// 1. never-block register channel POST /v1/dms/proxys
// 2. switch off → allow
// 3. switch on → registered ProxyTarget host IP → allow
// 4. whitelist hit → allow
// 5. else HTTP 403 (not 401)
//
// Loopback is NOT auto-allowed. CloudBeaver paths are not exempted.
func AccessRestriction(u *biz.AccessRestrictionUsecase, proxy *biz.DmsProxyUsecase) echo.MiddlewareFunc {
return func(next echo.HandlerFunc) echo.HandlerFunc {
return func(c echo.Context) error {
if isNeverBlockRegisterPath(c) {
return next(c)
}
if u == nil {
return next(c)
}
enabled, err := u.IsEnabled(c.Request().Context())
if err != nil || !enabled {
// Fail-open on read error; AC-003: disabled → allow.
return next(c)
}

clientIP := biz.ExtractClientIP(c.Request())
if proxy != nil && proxy.IsRegisteredServiceIP(clientIP) {
return next(c)
}
matched, err := u.MatchClientIP(c.Request().Context(), clientIP)
if err != nil {
// Fail-open on match errors to avoid locking out operators on transient DB faults.
return next(c)
}
if matched {
return next(c)
}
msg := accessRestrictionDenyMsg
if clientIP != "" {
msg = fmt.Sprintf("%s(识别 IP:%s)", accessRestrictionDenyMsg, clientIP)
}
return echo.NewHTTPError(http.StatusForbidden, msg)
}
}
}

func isNeverBlockRegisterPath(c echo.Context) bool {
if c.Request().Method != http.MethodPost {
return false
}
path := strings.TrimSuffix(c.Request().URL.Path, "/")
// Exact register channel: /v1/dms/proxys (ProxyRouterGroup under /v1).
return path == "/v1"+dmsV1.ProxyRouterGroup
}
177 changes: 177 additions & 0 deletions internal/apiserver/service/dms_controller.go
Original file line number Diff line number Diff line change
Expand Up @@ -5026,6 +5026,183 @@ func (ctl *DMSController) UpdateSystemVariables(c echo.Context) error {
return NewOkResp(c)
}

// swagger:route GET /v1/dms/configurations/access_restriction Configuration GetAccessRestriction
//
// Get access restriction configuration.
//
// responses:
// 200: body:GetAccessRestrictionReply
// default: body:GenericResp
func (ctl *DMSController) GetAccessRestriction(c echo.Context) error {
currentUserUid, err := jwt.GetUserUidStrFromContext(c)
if err != nil {
return NewErrResp(c, err, apiError.DMSServiceErr)
}
reply, err := ctl.DMS.GetAccessRestriction(c.Request().Context(), currentUserUid)
if err != nil {
return NewErrResp(c, err, apiError.DMSServiceErr)
}
return NewOkRespWithReply(c, reply)
}

// swagger:operation PATCH /v1/dms/configurations/access_restriction Configuration UpdateAccessRestriction
//
// Update access restriction switch.
//
// ---
// parameters:
// - name: access_restriction
// in: body
// required: true
// schema:
// "$ref": "#/definitions/UpdateAccessRestrictionReq"
// responses:
// '200':
// description: GenericResp
// schema:
// "$ref": "#/definitions/GenericResp"
// default:
// description: GenericResp
// schema:
// "$ref": "#/definitions/GenericResp"
func (ctl *DMSController) UpdateAccessRestriction(c echo.Context) error {
req := new(aV1.UpdateAccessRestrictionReq)
err := bindAndValidateReq(c, req)
if err != nil {
return NewErrResp(c, err, apiError.BadRequestErr)
}
currentUserUid, err := jwt.GetUserUidStrFromContext(c)
if err != nil {
return NewErrResp(c, err, apiError.DMSServiceErr)
}
err = ctl.DMS.UpdateAccessRestriction(c.Request().Context(), currentUserUid, req, biz.ExtractClientIP(c.Request()))
if err != nil {
return NewErrResp(c, err, apiError.DMSServiceErr)
}
return NewOkResp(c)
}

// swagger:operation POST /v1/dms/configurations/access_restriction/rules Configuration CreateAccessWhitelistRule
//
// Create access whitelist rule.
//
// ---
// parameters:
// - name: rule
// in: body
// required: true
// schema:
// "$ref": "#/definitions/CreateAccessWhitelistRuleReq"
// responses:
// '200':
// description: CreateAccessWhitelistRuleReply
// schema:
// "$ref": "#/definitions/CreateAccessWhitelistRuleReply"
// default:
// description: GenericResp
// schema:
// "$ref": "#/definitions/GenericResp"
func (ctl *DMSController) CreateAccessWhitelistRule(c echo.Context) error {
req := new(aV1.CreateAccessWhitelistRuleReq)
err := bindAndValidateReq(c, req)
if err != nil {
return NewErrResp(c, err, apiError.BadRequestErr)
}
currentUserUid, err := jwt.GetUserUidStrFromContext(c)
if err != nil {
return NewErrResp(c, err, apiError.DMSServiceErr)
}
reply, err := ctl.DMS.CreateAccessWhitelistRule(c.Request().Context(), currentUserUid, req)
if err != nil {
return NewErrResp(c, err, apiError.DMSServiceErr)
}
return NewOkRespWithReply(c, reply)
}

// swagger:operation PUT /v1/dms/configurations/access_restriction/rules/{rule_uid} Configuration UpdateAccessWhitelistRule
//
// Update access whitelist rule.
//
// ---
// parameters:
// - name: rule_uid
// in: path
// required: true
// type: string
// - name: rule
// in: body
// required: true
// schema:
// "$ref": "#/definitions/UpdateAccessWhitelistRuleReq"
// responses:
// '200':
// description: UpdateAccessWhitelistRuleReply
// schema:
// "$ref": "#/definitions/UpdateAccessWhitelistRuleReply"
// default:
// description: GenericResp
// schema:
// "$ref": "#/definitions/GenericResp"
func (ctl *DMSController) UpdateAccessWhitelistRule(c echo.Context) error {
req := new(aV1.UpdateAccessWhitelistRuleReq)
err := bindAndValidateReq(c, req)
if err != nil {
return NewErrResp(c, err, apiError.BadRequestErr)
}
currentUserUid, err := jwt.GetUserUidStrFromContext(c)
if err != nil {
return NewErrResp(c, err, apiError.DMSServiceErr)
}
reply, err := ctl.DMS.UpdateAccessWhitelistRule(c.Request().Context(), currentUserUid, req)
if err != nil {
return NewErrResp(c, err, apiError.DMSServiceErr)
}
return NewOkRespWithReply(c, reply)
}

// swagger:route DELETE /v1/dms/configurations/access_restriction/rules/{rule_uid} Configuration DeleteAccessWhitelistRule
//
// Delete access whitelist rule.
//
// responses:
// 200: body:GenericResp
// default: body:GenericResp
func (ctl *DMSController) DeleteAccessWhitelistRule(c echo.Context) error {
req := new(aV1.DeleteAccessWhitelistRuleReq)
err := bindAndValidateReq(c, req)
if err != nil {
return NewErrResp(c, err, apiError.BadRequestErr)
}
currentUserUid, err := jwt.GetUserUidStrFromContext(c)
if err != nil {
return NewErrResp(c, err, apiError.DMSServiceErr)
}
err = ctl.DMS.DeleteAccessWhitelistRule(c.Request().Context(), currentUserUid, req)
if err != nil {
return NewErrResp(c, err, apiError.DMSServiceErr)
}
return NewOkResp(c)
}

// swagger:route GET /v1/dms/configurations/access_restriction/client_ip Configuration GetAccessRestrictionClientIP
//
// Get current request client IP for access restriction.
//
// responses:
// 200: body:GetAccessRestrictionClientIPReply
// default: body:GenericResp
func (ctl *DMSController) GetAccessRestrictionClientIP(c echo.Context) error {
currentUserUid, err := jwt.GetUserUidStrFromContext(c)
if err != nil {
return NewErrResp(c, err, apiError.DMSServiceErr)
}
reply, err := ctl.DMS.GetAccessRestrictionClientIP(c.Request().Context(), currentUserUid, c.Request())
if err != nil {
return NewErrResp(c, err, apiError.DMSServiceErr)
}
return NewOkRespWithReply(c, reply)
}

// swagger:operation POST /v1/dms/operation_records OperationRecord AddOperationRecord
//
// Add operation record.
Expand Down
9 changes: 9 additions & 0 deletions internal/apiserver/service/router.go
Original file line number Diff line number Diff line change
Expand Up @@ -218,6 +218,12 @@ func (s *APIServer) initRouter() error {
configurationV1.GET("/license/usage", s.DMSController.GetLicenseUsage) /* TODO AdminUserAllowed()*/
configurationV1.GET("/system_variables", s.DMSController.GetSystemVariables) /* TODO AdminUserAllowed()*/
configurationV1.PATCH("/system_variables", s.DMSController.UpdateSystemVariables) /* TODO AdminUserAllowed()*/
configurationV1.GET("/access_restriction", s.DMSController.GetAccessRestriction)
configurationV1.PATCH("/access_restriction", s.DMSController.UpdateAccessRestriction)
configurationV1.POST("/access_restriction/rules", s.DMSController.CreateAccessWhitelistRule)
configurationV1.PUT("/access_restriction/rules/:rule_uid", s.DMSController.UpdateAccessWhitelistRule)
configurationV1.DELETE("/access_restriction/rules/:rule_uid", s.DMSController.DeleteAccessWhitelistRule)
configurationV1.GET("/access_restriction/client_ip", s.DMSController.GetAccessRestrictionClientIP)
// notify
notificationV1 := v1.Group(dmsV1.NotificationRouterGroup)
notificationV1.POST("", s.DMSController.Notify) /* TODO AdminUserAllowed()*/
Expand Down Expand Up @@ -384,6 +390,9 @@ func (s *APIServer) installMiddleware() error {
}
}(allowedMethods))

// Access restriction: early global gate (before JWT). Order: register never-block → off allow → registered IP → whitelist → 403.
s.echo.Use(dmsMiddleware.AccessRestriction(s.DMSController.DMS.AccessRestrictionUsecase, s.DMSController.DMS.DmsProxyUsecase))

var skipJWTPaths = []string{
dmsV1.SessionRouterGroup,
"/v1/dms/sessions/refresh",
Expand Down
Loading