diff --git a/README.md b/README.md index 9c0a02b..733ade2 100644 --- a/README.md +++ b/README.md @@ -85,7 +85,8 @@ Proxy Pool 用 Controller 协调这些变化,并让 Gateway 数据面只消费 CIDR、可信代理、严格请求解析和敏感信息最小化;Admin 的读写权限与 Distribution 提取权限可按命中凭据分别收敛。Distribution 凭据还可限制单次 提取数量、可访问 Upstream 与地区,Gateway 凭据可限制可访问 Routing;未配置时 - 保持既有全范围行为。 + 保持既有全范围行为。Gateway 凭据还可在每个 Worker 内限制每分钟请求数, + 不把该热路径计数写入 Redis 或 PostgreSQL。 ## 架构概览 diff --git a/docs/configuration/reference.md b/docs/configuration/reference.md index af68f45..bbf7689 100644 --- a/docs/configuration/reference.md +++ b/docs/configuration/reference.md @@ -181,14 +181,16 @@ distribution: ### 3.4 Gateway Client 路由约束 -Gateway 的 auth.client 只接受 allowedRoutings。认证成功后,Gateway 先完成目标 -地址策略和本地 Routing 匹配,再于本地 Dispatcher 前检查该集合;不匹配返回 403, -不会尝试选择或预留 Proxy。该检查只读取请求上下文和当前本地 Snapshot,不访问 -Redis、PostgreSQL 或 Provider。 +Gateway 的 auth.client 支持 allowedRoutings 和 requestsPerMinute。认证成功后, +Gateway 先完成目标地址策略和本地 Routing 匹配,再于本地 Dispatcher 前检查 +Routing 集合;不匹配返回 403,不会尝试选择或预留 Proxy。速率限制以认证 Client +为计数键,在每个 Gateway Worker 内执行,超过限制返回 429。两项检查都只读取 +请求上下文和当前本地状态,不访问 Redis、PostgreSQL 或 Provider。 allowedRoutings 中的每个名称必须对应一个启用的 Gateway Routing。Gateway 上的 maxExtractCount、allowedUpstreams 和 allowedRegions,以及 Distribution 上的 -allowedRoutings 均会被配置校验拒绝。Admin 不支持任何 client 约束。 +allowedRoutings 和 requestsPerMinute 均会被配置校验拒绝。Admin 不支持任何 +client 约束。 ~~~yaml gateway: @@ -199,6 +201,7 @@ gateway: valueFile: /run/secrets/checkout-gateway-token client: allowedRoutings: [checkout] + requestsPerMinute: 600 ~~~ Gateway、Distribution、Admin 与 Provider API 是独立认证边界。改变其中一套 diff --git a/docs/security/security-model.md b/docs/security/security-model.md index 02e20d7..d13918c 100644 --- a/docs/security/security-model.md +++ b/docs/security/security-model.md @@ -26,6 +26,8 @@ 原始会话标识不写入 PostgreSQL、Redis、日志或指标,也不会转发给目标站点。 绑定仅保留在当前 Worker 的有界内存中,并在代理不再符合当前 Snapshot、到达 代理可用期或转发失败时失效。 +- Gateway 凭据的 requestsPerMinute 限制按认证 Client 在本 Worker 内执行,不把 + 认证身份、会话或请求计数写入 PostgreSQL、Redis 或 Prometheus 标签。 ## 3. 目标地址策略 diff --git a/internal/config/config_test.go b/internal/config/config_test.go index 9778b3d..298f678 100644 --- a/internal/config/config_test.go +++ b/internal/config/config_test.go @@ -407,6 +407,17 @@ func TestValidateDistributionCredentialClientPolicy(t *testing.T) { if err := Validate(cfg); err == nil || !strings.Contains(err.Error(), "only supported on distribution") { t.Fatalf("Validate(gateway extraction policy) error = %v", err) } + + cfg.Gateway.Auth.ClientPolicy = clientpolicy.Policy{RequestsPerMinute: 60} + if err := Validate(cfg); err != nil { + t.Fatalf("Validate(gateway request rate policy) error = %v", err) + } + + cfg.Gateway.Auth.ClientPolicy = clientpolicy.Policy{} + cfg.Distribution.Auth.ClientPolicy = clientpolicy.Policy{RequestsPerMinute: 60} + if err := Validate(cfg); err == nil || !strings.Contains(err.Error(), "only supported on gateway") { + t.Fatalf("Validate(distribution request rate policy) error = %v", err) + } } func TestValidateGatewayStickySession(t *testing.T) { diff --git a/internal/config/validate.go b/internal/config/validate.go index ccb0680..2c67dae 100644 --- a/internal/config/validate.go +++ b/internal/config/validate.go @@ -552,8 +552,8 @@ func validateDistributionClientPolicies(auth Auth, upstreams map[string]Upstream } func validateDistributionClientPolicy(scope string, policy clientpolicy.Policy, upstreams map[string]Upstream) error { - if len(policy.AllowedRoutings) != 0 { - return fmt.Errorf("validate %s.allowedRoutings: routing access is only supported on gateway", scope) + if len(policy.AllowedRoutings) != 0 || policy.RequestsPerMinute != 0 { + return fmt.Errorf("validate %s: routing access and request rate limits are only supported on gateway", scope) } for _, upstream := range policy.AllowedUpstreams { if _, exists := upstreams[upstream]; !exists { @@ -581,6 +581,9 @@ func validateGatewayClientPolicy(scope string, policy clientpolicy.Policy, routi if policy.MaxExtractCount != 0 || len(policy.AllowedUpstreams) != 0 || len(policy.AllowedRegions) != 0 { return fmt.Errorf("validate %s: extraction constraints are only supported on distribution", scope) } + if int64(policy.RequestsPerMinute) > MaximumExactCounter { + return fmt.Errorf("validate %s.requestsPerMinute: exceeds exact counter range", scope) + } for _, routing := range policy.AllowedRoutings { if _, exists := routings[routing]; !exists { return fmt.Errorf("validate %s.allowedRoutings: routing %q does not exist or is not an enabled gateway routing", scope, routing) diff --git a/internal/domain/clientpolicy/policy.go b/internal/domain/clientpolicy/policy.go index 5d2e46b..2a448c8 100644 --- a/internal/domain/clientpolicy/policy.go +++ b/internal/domain/clientpolicy/policy.go @@ -12,19 +12,20 @@ var ErrInvalidPolicy = errors.New("invalid client policy") // Policy limits an authenticated client's permitted Distribution extraction or // Gateway routing. Zero values preserve the existing unrestricted behavior. type Policy struct { - MaxExtractCount int `yaml:"maxExtractCount"` - AllowedUpstreams []string `yaml:"allowedUpstreams"` - AllowedRegions []string `yaml:"allowedRegions"` - AllowedRoutings []string `yaml:"allowedRoutings"` + MaxExtractCount int `yaml:"maxExtractCount"` + RequestsPerMinute int `yaml:"requestsPerMinute"` + AllowedUpstreams []string `yaml:"allowedUpstreams"` + AllowedRegions []string `yaml:"allowedRegions"` + AllowedRoutings []string `yaml:"allowedRoutings"` } func (policy Policy) IsZero() bool { - return policy.MaxExtractCount == 0 && len(policy.AllowedUpstreams) == 0 && + return policy.MaxExtractCount == 0 && policy.RequestsPerMinute == 0 && len(policy.AllowedUpstreams) == 0 && len(policy.AllowedRegions) == 0 && len(policy.AllowedRoutings) == 0 } func (policy Policy) Validate() error { - if policy.MaxExtractCount < 0 || + if policy.MaxExtractCount < 0 || policy.RequestsPerMinute < 0 || !validUniqueValues(policy.AllowedUpstreams) || !validUniqueValues(policy.AllowedRegions) || !validUniqueValues(policy.AllowedRoutings) { diff --git a/internal/domain/clientpolicy/policy_test.go b/internal/domain/clientpolicy/policy_test.go index e7c3d99..d776abb 100644 --- a/internal/domain/clientpolicy/policy_test.go +++ b/internal/domain/clientpolicy/policy_test.go @@ -16,12 +16,14 @@ func TestValidate(t *testing.T) { }{ {name: "empty policy", policy: Policy{}}, {name: "bounded policy", policy: Policy{ - MaxExtractCount: 10, - AllowedUpstreams: []string{"provider-a", "provider-b"}, - AllowedRegions: []string{"shanghai", "beijing"}, - AllowedRoutings: []string{"checkout", "catalog"}, + MaxExtractCount: 10, + RequestsPerMinute: 600, + AllowedUpstreams: []string{"provider-a", "provider-b"}, + AllowedRegions: []string{"shanghai", "beijing"}, + AllowedRoutings: []string{"checkout", "catalog"}, }}, {name: "negative max count", policy: Policy{MaxExtractCount: -1}, wantErr: true}, + {name: "negative request rate", policy: Policy{RequestsPerMinute: -1}, wantErr: true}, {name: "empty upstream", policy: Policy{AllowedUpstreams: []string{""}}, wantErr: true}, {name: "whitespace upstream", policy: Policy{AllowedUpstreams: []string{" provider-a"}}, wantErr: true}, {name: "duplicate upstream", policy: Policy{AllowedUpstreams: []string{"provider-a", "provider-a"}}, wantErr: true}, diff --git a/internal/gateway/server/bootstrap.go b/internal/gateway/server/bootstrap.go index 79fb532..d7e6126 100644 --- a/internal/gateway/server/bootstrap.go +++ b/internal/gateway/server/bootstrap.go @@ -5,6 +5,7 @@ import ( "time" "proxy-pool/internal/config" + "proxy-pool/internal/domain/clientpolicy" "proxy-pool/internal/gateway/policy" platformAdmission "proxy-pool/internal/platform/admission" "proxy-pool/internal/platform/httpsecurity" @@ -31,7 +32,7 @@ func BuildProtection(listener config.Listener) (Protection, error) { return Protection{}, err } - var admission Guard + var admissionGuards []Guard if listener.Limits.RequestsPerMinute > 0 || listener.Limits.RequestsPerMinutePerClient > 0 { limiter, limiterErr := platformAdmission.NewFixedWindow(platformAdmission.FixedWindowConfig{ Window: time.Minute, @@ -41,9 +42,45 @@ func BuildProtection(listener config.Listener) (Protection, error) { if limiterErr != nil { return Protection{}, fmt.Errorf("build gateway admission: %w", limiterErr) } - admission = NewAdmissionGuard(clientIPs, limiter) + admissionGuards = append(admissionGuards, NewAdmissionGuard(clientIPs, limiter)) } - return Protection{Auth: auth, Access: access, Admission: admission, ClientIPs: clientIPs}, nil + credentialAdmission, err := buildCredentialAdmission(listener.Auth) + if err != nil { + return Protection{}, err + } + admissionGuards = append(admissionGuards, credentialAdmission) + return Protection{Auth: auth, Access: access, Admission: chainGuards(admissionGuards...), ClientIPs: clientIPs}, nil +} + +func buildCredentialAdmission(authentication config.Auth) (*CredentialAdmissionGuard, error) { + limits := make(map[int]Admitter) + for _, policy := range credentialPolicies(authentication) { + limit := policy.RequestsPerMinute + if limit == 0 { + continue + } + if _, exists := limits[limit]; exists { + continue + } + limiter, err := platformAdmission.NewFixedWindow(platformAdmission.FixedWindowConfig{ + Window: time.Minute, + PerKey: limit, + }) + if err != nil { + return nil, fmt.Errorf("build gateway credential admission: %w", err) + } + limits[limit] = limiter + } + return NewCredentialAdmissionGuard(limits), nil +} + +func credentialPolicies(authentication config.Auth) []clientpolicy.Policy { + policies := make([]clientpolicy.Policy, 0, len(authentication.Methods)+1) + policies = append(policies, authentication.ClientPolicy) + for _, method := range authentication.Methods { + policies = append(policies, method.ClientPolicy) + } + return policies } func ConfigFromListener(listener config.Listener) Config { diff --git a/internal/gateway/server/bootstrap_test.go b/internal/gateway/server/bootstrap_test.go index 59c0da4..ddc5c12 100644 --- a/internal/gateway/server/bootstrap_test.go +++ b/internal/gateway/server/bootstrap_test.go @@ -9,6 +9,7 @@ import ( "time" "proxy-pool/internal/config" + "proxy-pool/internal/domain/clientpolicy" "proxy-pool/internal/gateway/policy" "proxy-pool/internal/platform/httpsecurity" ) @@ -66,6 +67,38 @@ func TestBuildProtectionSupportsBearerProxyAuthentication(t *testing.T) { } } +func TestBuildProtectionAppliesCredentialRequestRateLimit(t *testing.T) { + t.Parallel() + + protection, err := BuildProtection(config.Listener{ + Auth: config.Auth{ + Mode: "bearer", Token: "gateway-token", + ClientPolicy: clientpolicy.Policy{RequestsPerMinute: 1}, + }, + }) + if err != nil { + t.Fatalf("BuildProtection() error = %v", err) + } + for attempt := range 2 { + request := httptest.NewRequest(http.MethodGet, "http://example.test", nil) + request.RemoteAddr = "198.51.100.8:1234" + request.Header.Set("Proxy-Authorization", "Bearer gateway-token") + if err := protection.Auth.Check(context.Background(), request); err != nil { + t.Fatalf("Auth.Check(%d) error = %v", attempt, err) + } + err := protection.Admission.Check(context.Background(), request) + if attempt == 0 && err != nil { + t.Fatalf("Admission.Check(first) error = %v", err) + } + if attempt == 1 { + var admissionError *HTTPError + if !errors.As(err, &admissionError) || admissionError.StatusCode != http.StatusTooManyRequests { + t.Fatalf("Admission.Check(second) error = %T %v, want 429", err, err) + } + } + } +} + func TestBuildProtectionAnyPreservesIPWhitelistRejectionIndependentOfOrder(t *testing.T) { t.Parallel() methods := [][]config.AuthMethod{ diff --git a/internal/gateway/server/protection.go b/internal/gateway/server/protection.go index 63a4d5f..421cf5d 100644 --- a/internal/gateway/server/protection.go +++ b/internal/gateway/server/protection.go @@ -83,3 +83,67 @@ func (guard *AdmissionGuard) Check(ctx context.Context, request *http.Request) e } return nil } + +// CredentialAdmissionGuard enforces the request rate embedded in the +// authenticated credential policy. It executes after authentication has +// attached an immutable Client identity to the request. +type CredentialAdmissionGuard struct { + limiters map[int]Admitter +} + +func NewCredentialAdmissionGuard(limiters map[int]Admitter) *CredentialAdmissionGuard { + if len(limiters) == 0 { + return nil + } + cloned := make(map[int]Admitter, len(limiters)) + for limit, limiter := range limiters { + if limit > 0 && limiter != nil { + cloned[limit] = limiter + } + } + if len(cloned) == 0 { + return nil + } + return &CredentialAdmissionGuard{limiters: cloned} +} + +func (guard *CredentialAdmissionGuard) Check(ctx context.Context, request *http.Request) error { + if guard == nil { + return nil + } + identity, authenticated := httpsecurity.IdentityFromRequest(request) + if !authenticated || identity.ClientPolicy.RequestsPerMinute == 0 { + return nil + } + limiter, exists := guard.limiters[identity.ClientPolicy.RequestsPerMinute] + if !exists || limiter == nil || identity.ClientID == "" { + return &HTTPError{StatusCode: http.StatusInternalServerError, Cause: errors.New("gateway credential admission is not configured")} + } + if err := limiter.Admit(ctx, identity.ClientID); err != nil { + return &HTTPError{StatusCode: http.StatusTooManyRequests, Cause: err} + } + return nil +} + +func chainGuards(guards ...Guard) Guard { + filtered := make([]Guard, 0, len(guards)) + for _, guard := range guards { + if guard != nil { + filtered = append(filtered, guard) + } + } + if len(filtered) == 0 { + return nil + } + if len(filtered) == 1 { + return filtered[0] + } + return GuardFunc(func(ctx context.Context, request *http.Request) error { + for _, guard := range filtered { + if err := guard.Check(ctx, request); err != nil { + return err + } + } + return nil + }) +}