diff --git a/api/dms/service/v1/access_restriction.go b/api/dms/service/v1/access_restriction.go new file mode 100644 index 00000000..83cabcdc --- /dev/null +++ b/api/dms/service/v1/access_restriction.go @@ -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"` +} diff --git a/internal/apiserver/middleware/access_restriction.go b/internal/apiserver/middleware/access_restriction.go new file mode 100644 index 00000000..8ef8234c --- /dev/null +++ b/internal/apiserver/middleware/access_restriction.go @@ -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 +} diff --git a/internal/apiserver/service/dms_controller.go b/internal/apiserver/service/dms_controller.go index baf13e75..cfae6670 100644 --- a/internal/apiserver/service/dms_controller.go +++ b/internal/apiserver/service/dms_controller.go @@ -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. diff --git a/internal/apiserver/service/router.go b/internal/apiserver/service/router.go index 61dc0e38..01aacdec 100644 --- a/internal/apiserver/service/router.go +++ b/internal/apiserver/service/router.go @@ -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()*/ @@ -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", diff --git a/internal/dms/biz/access_restriction.go b/internal/dms/biz/access_restriction.go new file mode 100644 index 00000000..b2966608 --- /dev/null +++ b/internal/dms/biz/access_restriction.go @@ -0,0 +1,272 @@ +package biz + +import ( + "context" + "fmt" + "net" + "net/http" + "strings" + "time" + + utilLog "github.com/actiontech/dms/pkg/dms-common/pkg/log" + pkgRand "github.com/actiontech/dms/pkg/rand" +) + +const ( + AccessRestrictionEnabledKey = "access_restriction_enabled" + AccessPolicyTypeWhitelist = "whitelist" +) + +type AccessWhitelistRule struct { + Base + UID string + Source string + PolicyType string + Remark string +} + +type AccessRestrictionRepo interface { + ListRules(ctx context.Context) ([]*AccessWhitelistRule, error) + GetRuleByUID(ctx context.Context, uid string) (*AccessWhitelistRule, error) + GetRuleBySource(ctx context.Context, source string) (*AccessWhitelistRule, error) + CreateRule(ctx context.Context, rule *AccessWhitelistRule) error + UpdateRule(ctx context.Context, rule *AccessWhitelistRule) error + DeleteRule(ctx context.Context, uid string) error + GetEnabled(ctx context.Context) (bool, error) + SetEnabled(ctx context.Context, enabled bool) error +} + +type AccessRestrictionUsecase struct { + repo AccessRestrictionRepo + log *utilLog.Helper +} + +func NewAccessRestrictionUsecase(log utilLog.Logger, repo AccessRestrictionRepo) *AccessRestrictionUsecase { + return &AccessRestrictionUsecase{ + repo: repo, + log: utilLog.NewHelper(log, utilLog.WithMessageKey("biz.access_restriction")), + } +} + +func (u *AccessRestrictionUsecase) GetConfig(ctx context.Context) (enabled bool, rules []*AccessWhitelistRule, err error) { + enabled, err = u.repo.GetEnabled(ctx) + if err != nil { + return false, nil, err + } + rules, err = u.repo.ListRules(ctx) + if err != nil { + return false, nil, err + } + return enabled, rules, nil +} + +// SetEnabled toggles access restriction. Enabling requires a non-empty whitelist +// and that clientIP matches a rule (same MatchClientIP as future middleware deny). +// Failure does not write enabled=true. Disabling only needs permission (caller). +// Loopback has no privilege: 127.0.0.1 must be explicitly listed. +func (u *AccessRestrictionUsecase) SetEnabled(ctx context.Context, enabled bool, clientIP string) error { + if !enabled { + return u.repo.SetEnabled(ctx, false) + } + + rules, err := u.repo.ListRules(ctx) + if err != nil { + return err + } + if len(rules) == 0 { + return fmt.Errorf("白名单为空,无法开启访问限制") + } + + matched, err := matchIPAgainstRules(clientIP, rules) + if err != nil { + return err + } + if !matched { + return fmt.Errorf("当前访问来源不在白名单,开启后将无法访问,请先添加当前 IP/网段(检测到的 IP:%s)", clientIP) + } + return u.repo.SetEnabled(ctx, true) +} + +func (u *AccessRestrictionUsecase) IsEnabled(ctx context.Context) (bool, error) { + return u.repo.GetEnabled(ctx) +} + +func (u *AccessRestrictionUsecase) CreateRule(ctx context.Context, source, remark, policyType string) (*AccessWhitelistRule, error) { + normalized, err := NormalizeIPv4OrCIDR(source) + if err != nil { + return nil, err + } + policy, err := normalizePolicyType(policyType) + if err != nil { + return nil, err + } + exist, err := u.repo.GetRuleBySource(ctx, normalized) + if err != nil { + return nil, err + } + if exist != nil { + return nil, fmt.Errorf("来源已存在") + } + uid, err := pkgRand.GenStrUid() + if err != nil { + return nil, err + } + rule := &AccessWhitelistRule{ + UID: uid, + Source: normalized, + PolicyType: policy, + Remark: remark, + } + if err := u.repo.CreateRule(ctx, rule); err != nil { + return nil, err + } + return u.repo.GetRuleByUID(ctx, uid) +} + +func (u *AccessRestrictionUsecase) UpdateRule(ctx context.Context, uid, source, remark, policyType string) (*AccessWhitelistRule, error) { + existing, err := u.repo.GetRuleByUID(ctx, uid) + if err != nil { + return nil, err + } + if existing == nil { + return nil, fmt.Errorf("规则不存在") + } + normalized, err := NormalizeIPv4OrCIDR(source) + if err != nil { + return nil, err + } + policy, err := normalizePolicyType(policyType) + if err != nil { + return nil, err + } + conflict, err := u.repo.GetRuleBySource(ctx, normalized) + if err != nil { + return nil, err + } + if conflict != nil && conflict.UID != uid { + return nil, fmt.Errorf("来源已存在") + } + existing.Source = normalized + existing.Remark = remark + existing.PolicyType = policy + if err := u.repo.UpdateRule(ctx, existing); err != nil { + return nil, err + } + return u.repo.GetRuleByUID(ctx, uid) +} + +func (u *AccessRestrictionUsecase) DeleteRule(ctx context.Context, uid string) error { + existing, err := u.repo.GetRuleByUID(ctx, uid) + if err != nil { + return err + } + if existing == nil { + return fmt.Errorf("规则不存在") + } + return u.repo.DeleteRule(ctx, uid) +} + +// MatchClientIP reports whether ip hits any whitelist rule (single IP = /32). +// Shared by enable-guard (S4) and access middleware deny path (S3/AC-004). +func (u *AccessRestrictionUsecase) MatchClientIP(ctx context.Context, ipStr string) (bool, error) { + rules, err := u.repo.ListRules(ctx) + if err != nil { + return false, err + } + return matchIPAgainstRules(ipStr, rules) +} + +func matchIPAgainstRules(ipStr string, rules []*AccessWhitelistRule) (bool, error) { + ip := net.ParseIP(ipStr) + if ip == nil || ip.To4() == nil { + return false, nil + } + for _, rule := range rules { + if rule == nil { + continue + } + if ruleMatchesIP(rule.Source, ip) { + return true, nil + } + } + return false, nil +} + +// ExtractClientIP returns the request source IPv4 from RemoteAddr only (MVP; no XFF trust). +func ExtractClientIP(r *http.Request) string { + if r == nil { + return "" + } + host := r.RemoteAddr + if h, _, err := net.SplitHostPort(r.RemoteAddr); err == nil { + host = h + } + ip := net.ParseIP(host) + if ip == nil || ip.To4() == nil { + return host + } + return ip.To4().String() +} + +func NormalizeIPv4OrCIDR(raw string) (string, error) { + s := strings.TrimSpace(raw) + if s == "" { + return "", fmt.Errorf("来源不能为空") + } + if strings.Contains(s, ":") { + return "", fmt.Errorf("不支持 IPv6,请填写 IPv4 或 IPv4 CIDR") + } + if strings.Contains(s, "/") { + _, network, err := net.ParseCIDR(s) + if err != nil { + return "", fmt.Errorf("来源格式非法,请填写合法 IPv4 或 IPv4 CIDR") + } + if network.IP.To4() == nil { + return "", fmt.Errorf("不支持 IPv6,请填写 IPv4 或 IPv4 CIDR") + } + ones, bits := network.Mask.Size() + if bits != 32 || ones < 0 || ones > 32 { + return "", fmt.Errorf("来源格式非法,请填写合法 IPv4 或 IPv4 CIDR") + } + return fmt.Sprintf("%s/%d", network.IP.Mask(network.Mask).String(), ones), nil + } + ip := net.ParseIP(s) + if ip == nil || ip.To4() == nil { + return "", fmt.Errorf("来源格式非法,请填写合法 IPv4 或 IPv4 CIDR") + } + return ip.To4().String(), nil +} + +func normalizePolicyType(policyType string) (string, error) { + p := strings.TrimSpace(policyType) + if p == "" { + return AccessPolicyTypeWhitelist, nil + } + if p != AccessPolicyTypeWhitelist { + return "", fmt.Errorf("访问策略仅支持白名单") + } + return AccessPolicyTypeWhitelist, nil +} + +func ruleMatchesIP(source string, ip net.IP) bool { + if strings.Contains(source, "/") { + _, network, err := net.ParseCIDR(source) + if err != nil { + return false + } + return network.Contains(ip) + } + ruleIP := net.ParseIP(source) + if ruleIP == nil { + return false + } + return ruleIP.Equal(ip) +} + +// Ensure UpdatedAt is refreshed on update path when storage omits auto-update in map. +func TouchUpdatedAt(rule *AccessWhitelistRule) { + if rule == nil { + return + } + rule.UpdatedAt = time.Now() +} diff --git a/internal/dms/biz/access_restriction_test.go b/internal/dms/biz/access_restriction_test.go new file mode 100644 index 00000000..d1c26365 --- /dev/null +++ b/internal/dms/biz/access_restriction_test.go @@ -0,0 +1,133 @@ +package biz + +import ( + "context" + "io" + "strings" + "testing" + + utilLog "github.com/actiontech/dms/pkg/dms-common/pkg/log" +) + +func TestNormalizeIPv4OrCIDR(t *testing.T) { + cases := []struct { + in string + want string + wantErr bool + }{ + {"192.168.1.1", "192.168.1.1", false}, + {"10.0.0.0/24", "10.0.0.0/24", false}, + {"10.0.0.8/24", "10.0.0.0/24", false}, + {"", "", true}, + {"not-an-ip", "", true}, + {"2001:db8::1", "", true}, + {"1.2.3.4/33", "", true}, + } + for _, c := range cases { + got, err := NormalizeIPv4OrCIDR(c.in) + if c.wantErr { + if err == nil { + t.Fatalf("NormalizeIPv4OrCIDR(%q) expected error", c.in) + } + continue + } + if err != nil { + t.Fatalf("NormalizeIPv4OrCIDR(%q) unexpected error: %v", c.in, err) + } + if got != c.want { + t.Fatalf("NormalizeIPv4OrCIDR(%q)=%q want %q", c.in, got, c.want) + } + } +} + +func TestNormalizePolicyType(t *testing.T) { + got, err := normalizePolicyType("") + if err != nil || got != AccessPolicyTypeWhitelist { + t.Fatalf("empty policy: got=%q err=%v", got, err) + } + if _, err := normalizePolicyType("blacklist"); err == nil { + t.Fatal("blacklist should be rejected") + } +} + +func TestMatchIPAgainstRules(t *testing.T) { + rules := []*AccessWhitelistRule{ + {Source: "192.168.1.100"}, + {Source: "10.8.0.0/24"}, + } + ok, err := matchIPAgainstRules("10.8.0.5", rules) + if err != nil || !ok { + t.Fatalf("CIDR hit: ok=%v err=%v", ok, err) + } + ok, err = matchIPAgainstRules("127.0.0.1", rules) + if err != nil || ok { + t.Fatalf("loopback not listed: ok=%v err=%v", ok, err) + } + ok, err = matchIPAgainstRules("not-ip", rules) + if err != nil || ok { + t.Fatalf("invalid ip: ok=%v err=%v", ok, err) + } +} + +type memAccessRestrictionRepo struct { + enabled bool + rules []*AccessWhitelistRule +} + +func (m *memAccessRestrictionRepo) ListRules(ctx context.Context) ([]*AccessWhitelistRule, error) { + return m.rules, nil +} +func (m *memAccessRestrictionRepo) GetRuleByUID(ctx context.Context, uid string) (*AccessWhitelistRule, error) { + return nil, nil +} +func (m *memAccessRestrictionRepo) GetRuleBySource(ctx context.Context, source string) (*AccessWhitelistRule, error) { + return nil, nil +} +func (m *memAccessRestrictionRepo) CreateRule(ctx context.Context, rule *AccessWhitelistRule) error { + return nil +} +func (m *memAccessRestrictionRepo) UpdateRule(ctx context.Context, rule *AccessWhitelistRule) error { + return nil +} +func (m *memAccessRestrictionRepo) DeleteRule(ctx context.Context, uid string) error { return nil } +func (m *memAccessRestrictionRepo) GetEnabled(ctx context.Context) (bool, error) { + return m.enabled, nil +} +func (m *memAccessRestrictionRepo) SetEnabled(ctx context.Context, enabled bool) error { + m.enabled = enabled + return nil +} + +func TestSetEnabledGuard(t *testing.T) { + repo := &memAccessRestrictionRepo{enabled: false} + u := NewAccessRestrictionUsecase(utilLog.NewMyLogger(io.Discard), repo) + + if err := u.SetEnabled(context.Background(), true, "127.0.0.1"); err == nil || !strings.Contains(err.Error(), "白名单为空") { + t.Fatalf("empty list enable: err=%v", err) + } + if repo.enabled { + t.Fatal("empty list must keep switch off") + } + + repo.rules = []*AccessWhitelistRule{{Source: "192.168.1.100"}} + if err := u.SetEnabled(context.Background(), true, "127.0.0.1"); err == nil || !strings.Contains(err.Error(), "检测到的 IP:127.0.0.1") { + t.Fatalf("miss enable: err=%v", err) + } + if repo.enabled { + t.Fatal("miss must keep switch off") + } + + repo.rules = []*AccessWhitelistRule{{Source: "127.0.0.1"}} + if err := u.SetEnabled(context.Background(), true, "127.0.0.1"); err != nil { + t.Fatalf("hit enable: %v", err) + } + if !repo.enabled { + t.Fatal("hit should enable") + } + if err := u.SetEnabled(context.Background(), false, "10.0.0.1"); err != nil { + t.Fatalf("disable: %v", err) + } + if repo.enabled { + t.Fatal("disable should turn off") + } +} diff --git a/internal/dms/biz/proxy.go b/internal/dms/biz/proxy.go index 9f081e9b..081ec53b 100644 --- a/internal/dms/biz/proxy.go +++ b/internal/dms/biz/proxy.go @@ -3,6 +3,7 @@ package biz import ( "context" "fmt" + "net" "net/url" "strings" "sync" @@ -19,6 +20,7 @@ import ( type ProxyTargetRepo interface { SaveProxyTarget(ctx context.Context, u *ProxyTarget) error UpdateProxyTarget(ctx context.Context, u *ProxyTarget) error + DeleteProxyTargetByName(ctx context.Context, name string) error ListProxyTargets(ctx context.Context) ([]*ProxyTarget, error) ListProxyTargetsByScenarios(ctx context.Context, scenarios []ProxyScenario) ([]*ProxyTarget, error) GetProxyTargetByName(ctx context.Context, name string) (*ProxyTarget, error) @@ -199,6 +201,59 @@ func (d *DmsProxyUsecase) ListProxyTargetsByScenarios(ctx context.Context, scena return d.repo.ListProxyTargetsByScenarios(ctx, scenarios) } +// IsRegisteredServiceIP reports whether clientIP equals the host IP of any registered ProxyTarget URL. +// Domains without a parseable IPv4 host do not exempt (no silent allow). defaultTargetSelf is excluded. +// Exemption follows registration lifecycle: present in memory after register; gone after DeleteProxyTargetByName. +func (d *DmsProxyUsecase) IsRegisteredServiceIP(clientIP string) bool { + if d == nil || clientIP == "" { + return false + } + ip := net.ParseIP(clientIP) + if ip == nil || ip.To4() == nil { + return false + } + client := ip.To4().String() + + d.mutex.RLock() + defer d.mutex.RUnlock() + for _, t := range d.targets { + if t == nil || t.URL == nil { + continue + } + host := t.URL.Hostname() + hostIP := net.ParseIP(host) + if hostIP == nil || hostIP.To4() == nil { + continue + } + if hostIP.To4().String() == client { + return true + } + } + return false +} + +// DeleteProxyTargetByName removes a registration from DB and memory so IP exemption is cancelled immediately. +func (d *DmsProxyUsecase) DeleteProxyTargetByName(ctx context.Context, name string) error { + if name == "" { + return fmt.Errorf("proxy target name is empty") + } + d.mutex.Lock() + defer d.mutex.Unlock() + + if err := d.repo.DeleteProxyTargetByName(ctx, name); err != nil { + return err + } + next := make([]*ProxyTarget, 0, len(d.targets)) + for _, t := range d.targets { + if t == nil || t.Name == name { + continue + } + next = append(next, t) + } + d.targets = next + return nil +} + // AddTarget实现echo的ProxyBalancer接口, 没有实际意义 func (d *DmsProxyUsecase) AddTarget(target *middleware.ProxyTarget) bool { return true diff --git a/internal/dms/biz/proxy_access_restriction_test.go b/internal/dms/biz/proxy_access_restriction_test.go new file mode 100644 index 00000000..9c2b5cd8 --- /dev/null +++ b/internal/dms/biz/proxy_access_restriction_test.go @@ -0,0 +1,28 @@ +package biz + +import ( + "net/url" + "testing" + + "github.com/labstack/echo/v4/middleware" +) + +func TestIsRegisteredServiceIP(t *testing.T) { + u1, _ := url.Parse("http://10.1.2.3:5432") + u2, _ := url.Parse("http://odc.example.com:8989") // domain → no exempt + d := &DmsProxyUsecase{ + targets: []*ProxyTarget{ + {ProxyTarget: middleware.ProxyTarget{Name: "sqle", URL: u1}}, + {ProxyTarget: middleware.ProxyTarget{Name: "odc-dns", URL: u2}}, + }, + } + if !d.IsRegisteredServiceIP("10.1.2.3") { + t.Fatal("registered IP should exempt") + } + if d.IsRegisteredServiceIP("127.0.0.1") { + t.Fatal("unregistered loopback must not auto-exempt") + } + if d.IsRegisteredServiceIP("odc.example.com") { + t.Fatal("hostname string is not client IP") + } +} diff --git a/internal/dms/service/access_restriction.go b/internal/dms/service/access_restriction.go new file mode 100644 index 00000000..3e3e1450 --- /dev/null +++ b/internal/dms/service/access_restriction.go @@ -0,0 +1,131 @@ +package service + +import ( + "context" + "fmt" + "net/http" + "time" + + dmsV1 "github.com/actiontech/dms/api/dms/service/v1" + "github.com/actiontech/dms/internal/dms/biz" +) + +func (d *DMSService) GetAccessRestriction(ctx context.Context, currentUserUid string) (*dmsV1.GetAccessRestrictionReply, error) { + canView, err := d.OpPermissionVerifyUsecase.CanViewGlobal(ctx, currentUserUid) + if err != nil { + return nil, fmt.Errorf("检查权限失败: %v", err) + } + if !canView { + return nil, fmt.Errorf("无权限查看访问限制配置") + } + + enabled, rules, err := d.AccessRestrictionUsecase.GetConfig(ctx) + if err != nil { + return nil, err + } + return &dmsV1.GetAccessRestrictionReply{ + Data: dmsV1.AccessRestrictionConfig{ + Enabled: enabled, + Rules: toAccessWhitelistRuleItems(rules), + }, + }, nil +} + +func (d *DMSService) UpdateAccessRestriction(ctx context.Context, currentUserUid string, req *dmsV1.UpdateAccessRestrictionReq, clientIP string) error { + canOp, err := d.OpPermissionVerifyUsecase.CanOpGlobal(ctx, currentUserUid, false) + if err != nil { + return fmt.Errorf("检查权限失败: %v", err) + } + if !canOp { + return fmt.Errorf("无权限修改访问限制配置") + } + if req.Enabled == nil { + return fmt.Errorf("enabled 不能为空") + } + return d.AccessRestrictionUsecase.SetEnabled(ctx, *req.Enabled, clientIP) +} + +func (d *DMSService) CreateAccessWhitelistRule(ctx context.Context, currentUserUid string, req *dmsV1.CreateAccessWhitelistRuleReq) (*dmsV1.CreateAccessWhitelistRuleReply, error) { + canOp, err := d.OpPermissionVerifyUsecase.CanOpGlobal(ctx, currentUserUid, false) + if err != nil { + return nil, fmt.Errorf("检查权限失败: %v", err) + } + if !canOp { + return nil, fmt.Errorf("无权限修改访问限制配置") + } + rule, err := d.AccessRestrictionUsecase.CreateRule(ctx, req.Source, req.Remark, req.PolicyType) + if err != nil { + return nil, err + } + return &dmsV1.CreateAccessWhitelistRuleReply{ + Data: toAccessWhitelistRuleItem(rule), + }, nil +} + +func (d *DMSService) UpdateAccessWhitelistRule(ctx context.Context, currentUserUid string, req *dmsV1.UpdateAccessWhitelistRuleReq) (*dmsV1.UpdateAccessWhitelistRuleReply, error) { + canOp, err := d.OpPermissionVerifyUsecase.CanOpGlobal(ctx, currentUserUid, false) + if err != nil { + return nil, fmt.Errorf("检查权限失败: %v", err) + } + if !canOp { + return nil, fmt.Errorf("无权限修改访问限制配置") + } + rule, err := d.AccessRestrictionUsecase.UpdateRule(ctx, req.RuleUID, req.Source, req.Remark, req.PolicyType) + if err != nil { + return nil, err + } + return &dmsV1.UpdateAccessWhitelistRuleReply{ + Data: toAccessWhitelistRuleItem(rule), + }, nil +} + +func (d *DMSService) DeleteAccessWhitelistRule(ctx context.Context, currentUserUid string, req *dmsV1.DeleteAccessWhitelistRuleReq) error { + canOp, err := d.OpPermissionVerifyUsecase.CanOpGlobal(ctx, currentUserUid, false) + if err != nil { + return fmt.Errorf("检查权限失败: %v", err) + } + if !canOp { + return fmt.Errorf("无权限修改访问限制配置") + } + return d.AccessRestrictionUsecase.DeleteRule(ctx, req.RuleUID) +} + +func (d *DMSService) GetAccessRestrictionClientIP(ctx context.Context, currentUserUid string, r *http.Request) (*dmsV1.GetAccessRestrictionClientIPReply, error) { + canView, err := d.OpPermissionVerifyUsecase.CanViewGlobal(ctx, currentUserUid) + if err != nil { + return nil, fmt.Errorf("检查权限失败: %v", err) + } + if !canView { + return nil, fmt.Errorf("无权限查看访问限制配置") + } + return &dmsV1.GetAccessRestrictionClientIPReply{ + Data: dmsV1.AccessRestrictionClientIP{ + ClientIP: biz.ExtractClientIP(r), + }, + }, nil +} + +func toAccessWhitelistRuleItems(rules []*biz.AccessWhitelistRule) []dmsV1.AccessWhitelistRuleItem { + out := make([]dmsV1.AccessWhitelistRuleItem, 0, len(rules)) + for _, rule := range rules { + out = append(out, toAccessWhitelistRuleItem(rule)) + } + return out +} + +func toAccessWhitelistRuleItem(rule *biz.AccessWhitelistRule) dmsV1.AccessWhitelistRuleItem { + if rule == nil { + return dmsV1.AccessWhitelistRuleItem{} + } + updatedAt := "" + if !rule.UpdatedAt.IsZero() { + updatedAt = rule.UpdatedAt.Format(time.RFC3339) + } + return dmsV1.AccessWhitelistRuleItem{ + UID: rule.UID, + Source: rule.Source, + PolicyType: rule.PolicyType, + Remark: rule.Remark, + UpdatedAt: updatedAt, + } +} diff --git a/internal/dms/service/service.go b/internal/dms/service/service.go index 85031eaa..be35fa38 100644 --- a/internal/dms/service/service.go +++ b/internal/dms/service/service.go @@ -53,6 +53,7 @@ type DMSService struct { OperationRecordUsecase *biz.OperationRecordUsecase MaintenanceTimeUsecase *biz.MaintenanceTimeUsecase UserActivityUsecase *biz.UserActivityUsecase + AccessRestrictionUsecase *biz.AccessRestrictionUsecase log *utilLog.Helper shutdownCallback func() error } @@ -148,6 +149,7 @@ func NewAndInitDMSService(logger utilLog.Logger, opts *conf.DMSOptions) (*DMSSer swaggerUseCase := biz.NewSwaggerUseCase(logger, dmsProxyUsecase) systemVariableUsecase := biz.NewSystemVariableUsecase(logger, storage.NewSystemVariableRepo(logger, st)) + accessRestrictionUsecase := biz.NewAccessRestrictionUsecase(logger, storage.NewAccessRestrictionRepo(logger, st)) maintenanceTimeUsecase := biz.NewMaintenanceTimeUsecase(logger, opPermissionVerifyUsecase) operationRecordRepo := storage.NewOperationRecordRepo(logger, st) operationRecordUsecase := biz.NewOperationRecordUsecase(logger, operationRecordRepo, systemVariableUsecase) @@ -224,6 +226,7 @@ func NewAndInitDMSService(logger utilLog.Logger, opts *conf.DMSOptions) (*DMSSer OperationRecordUsecase: operationRecordUsecase, MaintenanceTimeUsecase: maintenanceTimeUsecase, UserActivityUsecase: userActivityUsecase, + AccessRestrictionUsecase: accessRestrictionUsecase, log: utilLog.NewHelper(logger, utilLog.WithMessageKey("dms.service")), shutdownCallback: func() error { stopDataMaskingScheduler() diff --git a/internal/dms/storage/access_restriction.go b/internal/dms/storage/access_restriction.go new file mode 100644 index 00000000..2af1a3a6 --- /dev/null +++ b/internal/dms/storage/access_restriction.go @@ -0,0 +1,156 @@ +package storage + +import ( + "context" + "errors" + "fmt" + + "github.com/actiontech/dms/internal/dms/biz" + pkgErr "github.com/actiontech/dms/internal/dms/pkg/errors" + "github.com/actiontech/dms/internal/dms/storage/model" + utilLog "github.com/actiontech/dms/pkg/dms-common/pkg/log" + "gorm.io/gorm" +) + +var _ biz.AccessRestrictionRepo = (*AccessRestrictionRepo)(nil) + +type AccessRestrictionRepo struct { + *Storage + log *utilLog.Helper +} + +func NewAccessRestrictionRepo(log utilLog.Logger, s *Storage) *AccessRestrictionRepo { + return &AccessRestrictionRepo{ + Storage: s, + log: utilLog.NewHelper(log, utilLog.WithMessageKey("storage.access_restriction")), + } +} + +func (r *AccessRestrictionRepo) ListRules(ctx context.Context) ([]*biz.AccessWhitelistRule, error) { + var rows []*model.AccessWhitelistRule + if err := transaction(r.log, ctx, r.db, func(tx *gorm.DB) error { + if err := tx.WithContext(ctx).Order("updated_at DESC").Find(&rows).Error; err != nil { + return fmt.Errorf("failed to list access whitelist rules: %v", err) + } + return nil + }); err != nil { + return nil, err + } + out := make([]*biz.AccessWhitelistRule, 0, len(rows)) + for _, row := range rows { + out = append(out, convertModelAccessWhitelistRule(row)) + } + return out, nil +} + +func (r *AccessRestrictionRepo) GetRuleByUID(ctx context.Context, uid string) (*biz.AccessWhitelistRule, error) { + var row model.AccessWhitelistRule + if err := r.db.WithContext(ctx).Where("uid = ?", uid).First(&row).Error; err != nil { + if errors.Is(err, gorm.ErrRecordNotFound) { + return nil, nil + } + return nil, pkgErr.WrapStorageErr(r.log, fmt.Errorf("failed to get access whitelist rule: %v", err)) + } + return convertModelAccessWhitelistRule(&row), nil +} + +func (r *AccessRestrictionRepo) GetRuleBySource(ctx context.Context, source string) (*biz.AccessWhitelistRule, error) { + var row model.AccessWhitelistRule + if err := r.db.WithContext(ctx).Where("source = ?", source).First(&row).Error; err != nil { + if errors.Is(err, gorm.ErrRecordNotFound) { + return nil, nil + } + return nil, pkgErr.WrapStorageErr(r.log, fmt.Errorf("failed to get access whitelist rule by source: %v", err)) + } + return convertModelAccessWhitelistRule(&row), nil +} + +func (r *AccessRestrictionRepo) CreateRule(ctx context.Context, rule *biz.AccessWhitelistRule) error { + return transaction(r.log, ctx, r.db, func(tx *gorm.DB) error { + if err := tx.WithContext(ctx).Create(convertBizAccessWhitelistRule(rule)).Error; err != nil { + return pkgErr.WrapStorageErr(r.log, fmt.Errorf("failed to create access whitelist rule: %v", err)) + } + return nil + }) +} + +func (r *AccessRestrictionRepo) UpdateRule(ctx context.Context, rule *biz.AccessWhitelistRule) error { + return transaction(r.log, ctx, r.db, func(tx *gorm.DB) error { + result := tx.WithContext(ctx).Model(&model.AccessWhitelistRule{}).Where("uid = ?", rule.UID).Updates(map[string]interface{}{ + "source": rule.Source, + "policy_type": rule.PolicyType, + "remark": rule.Remark, + }) + if result.Error != nil { + return pkgErr.WrapStorageErr(r.log, fmt.Errorf("failed to update access whitelist rule: %v", result.Error)) + } + if result.RowsAffected == 0 { + return fmt.Errorf("规则不存在") + } + return nil + }) +} + +func (r *AccessRestrictionRepo) DeleteRule(ctx context.Context, uid string) error { + return transaction(r.log, ctx, r.db, func(tx *gorm.DB) error { + result := tx.WithContext(ctx).Where("uid = ?", uid).Delete(&model.AccessWhitelistRule{}) + if result.Error != nil { + return pkgErr.WrapStorageErr(r.log, fmt.Errorf("failed to delete access whitelist rule: %v", result.Error)) + } + if result.RowsAffected == 0 { + return fmt.Errorf("规则不存在") + } + return nil + }) +} + +func (r *AccessRestrictionRepo) GetEnabled(ctx context.Context) (bool, error) { + var row model.SystemVariable + if err := r.db.WithContext(ctx).Where("`key` = ?", biz.AccessRestrictionEnabledKey).First(&row).Error; err != nil { + if errors.Is(err, gorm.ErrRecordNotFound) { + return false, nil + } + return false, pkgErr.WrapStorageErr(r.log, fmt.Errorf("failed to get access restriction switch: %v", err)) + } + return row.Value == "true", nil +} + +func (r *AccessRestrictionRepo) SetEnabled(ctx context.Context, enabled bool) error { + value := "false" + if enabled { + value = "true" + } + return transaction(r.log, ctx, r.db, func(tx *gorm.DB) error { + row := &model.SystemVariable{ + Key: biz.AccessRestrictionEnabledKey, + Value: value, + } + if err := tx.WithContext(ctx).Save(row).Error; err != nil { + return pkgErr.WrapStorageErr(r.log, fmt.Errorf("failed to set access restriction switch: %v", err)) + } + return nil + }) +} + +func convertBizAccessWhitelistRule(b *biz.AccessWhitelistRule) *model.AccessWhitelistRule { + return &model.AccessWhitelistRule{ + Model: model.Model{ + UID: b.UID, + CreatedAt: b.CreatedAt, + UpdatedAt: b.UpdatedAt, + }, + Source: b.Source, + PolicyType: b.PolicyType, + Remark: b.Remark, + } +} + +func convertModelAccessWhitelistRule(m *model.AccessWhitelistRule) *biz.AccessWhitelistRule { + return &biz.AccessWhitelistRule{ + Base: convertBase(m.Model), + UID: m.UID, + Source: m.Source, + PolicyType: m.PolicyType, + Remark: m.Remark, + } +} diff --git a/internal/dms/storage/model/model.go b/internal/dms/storage/model/model.go index 3d899aea..3a83c2ef 100644 --- a/internal/dms/storage/model/model.go +++ b/internal/dms/storage/model/model.go @@ -61,6 +61,7 @@ var AutoMigrateList = []interface{}{ Gateway{}, SystemVariable{}, OperationRecord{}, + AccessWhitelistRule{}, } type Model struct { @@ -746,6 +747,18 @@ type SystemVariable struct { Value string `gorm:"not null;type:text"` } +// AccessWhitelistRule stores IP/CIDR whitelist entries for access restriction. +type AccessWhitelistRule struct { + Model + Source string `json:"source" gorm:"column:source;size:64;not null;uniqueIndex"` + PolicyType string `json:"policy_type" gorm:"column:policy_type;size:32;not null;default:whitelist"` + Remark string `json:"remark" gorm:"column:remark;size:255"` +} + +func (AccessWhitelistRule) TableName() string { + return "access_whitelist_rules" +} + type OperationRecord struct { ID uint `json:"id" gorm:"primary_key" example:"1"` CreatedAt time.Time `json:"created_at" gorm:"default:current_timestamp(3)" example:"2018-10-21T16:40:23+08:00"` diff --git a/internal/dms/storage/proxy.go b/internal/dms/storage/proxy.go index 6b35f505..959a6b47 100644 --- a/internal/dms/storage/proxy.go +++ b/internal/dms/storage/proxy.go @@ -153,3 +153,12 @@ func (d *ProxyTargetRepo) GetProxyTargetByName(ctx context.Context, name string) return t, nil } + +func (d *ProxyTargetRepo) DeleteProxyTargetByName(ctx context.Context, name string) error { + return transaction(d.log, ctx, d.db, func(tx *gorm.DB) error { + if err := tx.WithContext(ctx).Where("name = ?", name).Delete(&model.ProxyTarget{}).Error; err != nil { + return fmt.Errorf("failed to delete proxy target: %v", err) + } + return nil + }) +}