feat: add shared HTTP security protection

This commit is contained in:
youfak 2026-07-29 11:12:21 +08:00
parent 2ea5be5036
commit 68f685d89b
32 changed files with 1636 additions and 356 deletions

View File

@ -14,6 +14,7 @@ security:
- AdminApiKey: [] - AdminApiKey: []
- BasicAuth: [] - BasicAuth: []
- BearerAuth: [] - BearerAuth: []
- {}
paths: paths:
/api/v1/status: /api/v1/status:
get: get:
@ -217,6 +218,7 @@ components:
description: 管理入口认证失败 description: 管理入口认证失败
headers: headers:
X-Request-ID: {$ref: '#/components/headers/RequestID'} X-Request-ID: {$ref: '#/components/headers/RequestID'}
WWW-Authenticate: {schema: {type: string}}
content: content:
application/problem+json: application/problem+json:
schema: {$ref: '#/components/schemas/Problem'} schema: {$ref: '#/components/schemas/Problem'}

View File

@ -155,7 +155,8 @@ components:
required: false required: false
description: | description: |
同一客户端在幂等记录保留期内重用该键会得到首次提交结果,不会再次 同一客户端在幂等记录保留期内重用该键会得到首次提交结果,不会再次
提取。建议所有会自动重试的客户端提供。 提取。未认证客户端默认以可信代理链解析后的规范化来源 IP 标识。
建议所有会自动重试的客户端提供。
schema: schema:
type: string type: string
minLength: 8 minLength: 8
@ -336,6 +337,9 @@ components:
headers: headers:
X-Request-ID: X-Request-ID:
$ref: '#/components/headers/RequestID' $ref: '#/components/headers/RequestID'
WWW-Authenticate:
schema:
type: string
content: content:
application/problem+json: application/problem+json:
schema: schema:

View File

@ -12,7 +12,7 @@ gateway:
auth: auth:
mode: usernamePassword mode: usernamePassword
username: local-gateway username: local-gateway
password: env:PROXY_POOL_GATEWAY_PASSWORD password: "${PROXY_POOL_GATEWAY_PASSWORD}"
limits: limits:
maxConcurrentConnections: 20000 maxConcurrentConnections: 20000
requestsPerMinutePerClient: 60000 requestsPerMinutePerClient: 60000
@ -34,12 +34,12 @@ distribution:
auth: auth:
mode: apiKey mode: apiKey
header: X-API-Key header: X-API-Key
token: env:PROXY_POOL_EXTRACT_TOKEN token: "${PROXY_POOL_EXTRACT_TOKEN}"
limits: limits:
requestsPerMinute: 6000 requestsPerMinute: 6000
requestsPerMinutePerClient: 600 requestsPerMinutePerClient: 600
clientIdentification: clientIdentification:
mode: trustedProxyOrRemoteIP mode: sourceIP
extraction: extraction:
fulfillment: partial fulfillment: partial
maxCountPerRequest: 100 maxCountPerRequest: 100
@ -54,8 +54,8 @@ admin:
allowCIDRs: [172.16.0.0/12] allowCIDRs: [172.16.0.0/12]
auth: auth:
mode: apiKey mode: apiKey
header: X-Admin-Token header: X-Admin-Key
token: env:PROXY_POOL_ADMIN_TOKEN token: "${PROXY_POOL_ADMIN_TOKEN}"
metrics: metrics:
enabled: true enabled: true
@ -98,12 +98,13 @@ upstreams:
url: https://provider-a.invalid/api/proxies url: https://provider-a.invalid/api/proxies
method: GET method: GET
auth: auth:
mode: apiKey type: apiKey
header: Authorization location: header
token: env:PROVIDER_A_TOKEN name: Authorization
value: "${PROVIDER_A_TOKEN}"
template: '{{ . }}' template: '{{ . }}'
proxyAuth: proxyAuth:
mode: response type: response
pool: pool:
maxSize: 5000 maxSize: 5000
shrinkDelay: 30s shrinkDelay: 30s
@ -142,12 +143,13 @@ upstreams:
url: https://provider-b.invalid/api/proxies url: https://provider-b.invalid/api/proxies
method: GET method: GET
auth: auth:
mode: apiKey type: apiKey
header: Authorization location: header
token: env:PROVIDER_B_TOKEN name: Authorization
value: "${PROVIDER_B_TOKEN}"
template: '{{ . }}' template: '{{ . }}'
proxyAuth: proxyAuth:
mode: response type: response
pool: pool:
maxSize: 5000 maxSize: 5000
shrinkDelay: 30s shrinkDelay: 30s
@ -176,4 +178,3 @@ upstreams:
maxAttempts: 2 maxAttempts: 2
maxConsecutiveFailures: 3 maxConsecutiveFailures: 3
urls: [https://example.com/] urls: [https://example.com/]

View File

@ -16,8 +16,8 @@ data:
trustedProxies: [] trustedProxies: []
auth: auth:
mode: usernamePassword mode: usernamePassword
username: env:PROXY_POOL_GATEWAY_USERNAME username: "${PROXY_POOL_GATEWAY_USERNAME}"
password: env:PROXY_POOL_GATEWAY_PASSWORD password: "${PROXY_POOL_GATEWAY_PASSWORD}"
limits: limits:
maxConcurrentConnections: 100000 maxConcurrentConnections: 100000
requestsPerMinutePerClient: 60000 requestsPerMinutePerClient: 60000
@ -38,12 +38,12 @@ data:
auth: auth:
mode: apiKey mode: apiKey
header: X-API-Key header: X-API-Key
token: env:PROXY_POOL_EXTRACT_TOKEN token: "${PROXY_POOL_EXTRACT_TOKEN}"
limits: limits:
requestsPerMinute: 30000 requestsPerMinute: 30000
requestsPerMinutePerClient: 3000 requestsPerMinutePerClient: 3000
clientIdentification: clientIdentification:
mode: trustedProxyOrRemoteIP mode: sourceIP
extraction: extraction:
fulfillment: partial fulfillment: partial
maxCountPerRequest: 100 maxCountPerRequest: 100
@ -57,14 +57,14 @@ data:
allowCIDRs: [10.0.0.0/8] allowCIDRs: [10.0.0.0/8]
auth: auth:
mode: apiKey mode: apiKey
header: X-Admin-Token header: X-Admin-Key
token: env:PROXY_POOL_ADMIN_TOKEN token: "${PROXY_POOL_ADMIN_TOKEN}"
metrics: metrics:
enabled: true enabled: true
listen: 0.0.0.0:9090 listen: 0.0.0.0:9090
storage: storage:
postgresURL: env:PROXY_POOL_POSTGRES_URL postgresURL: "${PROXY_POOL_POSTGRES_URL}"
redisURL: env:PROXY_POOL_REDIS_URL redisURL: "${PROXY_POOL_REDIS_URL}"
routing: routing:
- name: gateway-default - name: gateway-default
enabled: true enabled: true
@ -97,12 +97,13 @@ data:
url: https://PROVIDER_A_HOST/api/proxies url: https://PROVIDER_A_HOST/api/proxies
method: GET method: GET
auth: auth:
mode: apiKey type: apiKey
header: Authorization location: header
token: env:PROVIDER_A_TOKEN name: Authorization
value: "${PROVIDER_A_TOKEN}"
template: '{{ . }}' template: '{{ . }}'
proxyAuth: proxyAuth:
mode: response type: response
pool: pool:
maxSize: 25000 maxSize: 25000
shrinkDelay: 30s shrinkDelay: 30s
@ -138,12 +139,13 @@ data:
url: https://PROVIDER_B_HOST/api/proxies url: https://PROVIDER_B_HOST/api/proxies
method: GET method: GET
auth: auth:
mode: apiKey type: apiKey
header: Authorization location: header
token: env:PROVIDER_B_TOKEN name: Authorization
value: "${PROVIDER_B_TOKEN}"
template: '{{ . }}' template: '{{ . }}'
proxyAuth: proxyAuth:
mode: response type: response
pool: pool:
maxSize: 25000 maxSize: 25000
shrinkDelay: 30s shrinkDelay: 30s
@ -169,4 +171,3 @@ data:
maxAttempts: 2 maxAttempts: 2
maxConsecutiveFailures: 3 maxConsecutiveFailures: 3
urls: [https://example.com/] urls: [https://example.com/]

View File

@ -21,7 +21,7 @@ func main() {
} }
defer file.Close() defer file.Close()
if _, err := config.Load(file); err != nil { if _, err := config.LoadResolved(file, config.OSResolver{}); err != nil {
fmt.Fprintf(os.Stderr, "invalid config: %v\n", err) fmt.Fprintf(os.Stderr, "invalid config: %v\n", err)
os.Exit(1) os.Exit(1)
} }

View File

@ -25,9 +25,9 @@ Admin API 使用独立监听器与权限,契约位于 `api/openapi/admin.yaml`
不能混用配置格式版本或单 Worker Snapshot 版本。 不能混用配置格式版本或单 Worker Snapshot 版本。
严格 JSON、请求体上限、Request ID、JSON/Problem 响应由 严格 JSON、请求体上限、Request ID、JSON/Problem 响应由
`platform/httpapi` 公用实现提供。Admin Handler 必须部署在独立监听器,并由 `platform/httpapi` 公用实现提供。Admin Handler 必须注入 `Authorizer`,标准
外层认证与授权中间件保护;网关使用的 `Proxy-Authorization`/407 语义不得复用 装配使用 `httpsecurity.Protection`,并在路由匹配前完成保护。网关使用的
到 Admin 的 `Authorization`/401 语义。 `Proxy-Authorization`/407 语义不得复用到 Admin 的 `Authorization`/401 语义。
除契约中的 401/403/404/409/422 外,运行时还明确返回: 除契约中的 401/403/404/409/422 外,运行时还明确返回:

View File

@ -199,9 +199,9 @@ Extraction Record 必须与状态更新处在相同事务边界或由同一权
proxyId, clientId, sourceIP, requestId, upstream, extractedAt, expiresAt proxyId, clientId, sourceIP, requestId, upstream, extractedAt, expiresAt
``` ```
无认证时 `clientId` 使用 `anonymous` 或稳定匿名标识并保留 `sourceIP`。记录只 无认证时 `clientId` 使用可信代理链解析后的规范化来源 IP 稳定标识,并保留
用于审计、排错和计费事实不承担资源归还语义。Proxy 到期后可以清理运行 `sourceIP`。记录只用于审计、排错和计费事实不承担资源归还语义。Proxy
记录,但 Extraction Record 按审计保留策略归档。 到期后可以清理运行记录,但 Extraction Record 按审计保留策略归档。
## 11. 运行时实现边界 ## 11. 运行时实现边界
@ -209,6 +209,7 @@ proxyId, clientId, sourceIP, requestId, upstream, extractedAt, expiresAt
身份结果注入、错误映射和健康探针。独占提取、TTL、Gateway 预留及幂等事务 身份结果注入、错误映射和健康探针。独占提取、TTL、Gateway 预留及幂等事务
继续由 `extraction.Service` 和持久化 Store 承担。 继续由 `extraction.Service` 和持久化 Store 承担。
请求体解码、Request ID 与 Problem JSON 统一复用 `platform/httpapi`。身份解析 请求体解码、Request ID 与 Problem JSON 统一复用 `platform/httpapi`。认证、
通过 `IdentityResolver` 注入;进程装配必须在 Handler 外层完成认证、可信代理 可信代理、来源控制、Client ID 和准入限流由必需的 `IdentityResolver` 注入,
来源解析与权限控制,且解析结果至少包含稳定 Client ID 或 Source IP。 标准装配使用 `httpsecurity.Protection`;解析结果至少包含稳定 Client ID 或
Source IP且安全检查先于请求体解析。

View File

@ -90,11 +90,14 @@ limits:
- `none`:无身份认证,访问控制与限流仍生效。 - `none`:无身份认证,访问控制与限流仍生效。
- `usernamePassword`:使用 `username``password` - `usernamePassword`:使用 `username``password`
- `apiKey`:使用 `header``token` - `apiKey`:使用 `header``token`
- `bearer`:使用 `Authorization: Bearer TOKEN`Gateway 对应
`Proxy-Authorization: Bearer TOKEN`
- `ipWhitelist`:使用 `cidrs` - `ipWhitelist`:使用 `cidrs`
- `any``methods` 中任一方法成功即可;方法字段名仍是 `mode` - `any``methods` 中任一方法成功即可Bearer 方法的 Secret 使用
`value`/`valueFile`。
Gateway、Distribution 与 Provider API 的认证三套独立边界。改变其中一套 Gateway、Distribution、Admin 与 Provider API 是独立认证边界。改变其中一套
不得连带改变另外两套 不得连带改变其他入口
## 4. Gateway ## 4. Gateway
@ -157,6 +160,13 @@ Extraction 是固定的一次性独占行为,**没有** `mode`、`leaseDuratio
`partial` 会提交实际可得数量;`allOrNothing` 数量不足时事务回滚,一个也不 `partial` 会提交实际可得数量;`allOrNothing` 数量不足时事务回滚,一个也不
提取。认证关闭时仍应使用 `sourceIP` 识别匿名 Client 并执行全局/来源限流。 提取。认证关闭时仍应使用 `sourceIP` 识别匿名 Client 并执行全局/来源限流。
`clientIdentification.mode` 支持:
- `sourceIP`:使用可信代理链解析后的来源地址;省略配置时采用此模式。
- `authenticatedClient`:使用 Basic 用户名或 Token 的稳定不可逆摘要;要求
启用认证。
- `authenticatedClientOrSourceIP`:优先认证主体,无主体时回退来源地址。
## 6. Routing ## 6. Routing
```yaml ```yaml

View File

@ -388,7 +388,7 @@ flowchart LR
## 15. 安全 ## 15. 安全
- Gateway、Distribution、Admin 认证相互独立。 - Gateway、Distribution、Admin 认证相互独立。
- `auth.mode` 支持 none、usernamePassword、apiKey、ipWhitelist 与组合 any。 - `auth.mode` 支持 none、usernamePassword、apiKey、bearer、ipWhitelist 与组合 any。
- `access.allowCIDRs` 独立于认证;代理头只在来源属于 trustedProxies 时接受。 - `access.allowCIDRs` 独立于认证;代理头只在来源属于 trustedProxies 时接受。
- 严格模式下,非回环监听且 auth=none、allowCIDRs 为空时启动失败。 - 严格模式下,非回环监听且 auth=none、allowCIDRs 为空时启动失败。
- 目的地址解析前后都拒绝 loopback、private、link-local、metadata 和配置禁区。 - 目的地址解析前后都拒绝 loopback、private、link-local、metadata 和配置禁区。
@ -425,7 +425,8 @@ CPU、内存、网络、Go 版本、配置和上游响应模型下测得。
### 17.2 热路径预算 ### 17.2 热路径预算
- Dispatch 无 I/O、无全局锁100k Proxy Snapshot 下 p99 小于 100 微秒。 - Dispatch 无 I/O、无全局锁100k Proxy Snapshot 下的设计预算为 p99 小于
100 微秒,仍需分位数基准验证。
- 所有队列、buffer、重试和日志均有界。 - 所有队列、buffer、重试和日志均有界。
- Listener、Client、Routing、Worker 和 Proxy 均有独立准入限制。 - Listener、Client、Routing、Worker 和 Proxy 均有独立准入限制。
- 过载在路由/建连前快速拒绝,不允许请求堆积耗尽内存。 - 过载在路由/建连前快速拒绝,不允许请求堆积耗尽内存。
@ -455,4 +456,3 @@ CPU、内存、网络、Go 版本、配置和上游响应模型下测得。
项目架构必须包含四个命令、领域模块、Gateway/Controller/Checker 模块、 项目架构必须包含四个命令、领域模块、Gateway/Controller/Checker 模块、
存储/协议 Adapter、OpenAPI/Proto、配置样例、Compose/Kubernetes、监控、 存储/协议 Adapter、OpenAPI/Proto、配置样例、Compose/Kubernetes、监控、
迁移、测试 fixture、负载场景和开发文档。目录存在但没有契约或测试不算完成。 迁移、测试 fixture、负载场景和开发文档。目录存在但没有契约或测试不算完成。

View File

@ -166,7 +166,9 @@ test/{fixtures,integration,e2e,load}/
当前进度2026-07-29已实现共享 `platform/httpapi`、Distribution 当前进度2026-07-29已实现共享 `platform/httpapi`、Distribution
extract/live/ready Handler 与 Admin status/enable/disable/switch/reload Handler extract/live/ready Handler 与 Admin status/enable/disable/switch/reload Handler
定向契约测试已覆盖严格 JSON、Body 上限、Request ID、幂等 Header、DTO 映射、 定向契约测试已覆盖严格 JSON、Body 上限、Request ID、幂等 Header、DTO 映射、
404/405 及业务错误映射。端点正式勾选仍等待独立监听器装配、认证/授权中间件、 404/405 及业务错误映射。共享 `platform/httpsecurity` 已补齐 Basic/API Key/
Bearer/CIDR、可信代理、Client ID、本地准入和 API 401/Gateway 407 差异,并作为
Admin/Distribution 必需依赖。端点正式勾选仍等待独立监听器装配、
PostgreSQL/Redis Adapter 与 Compose 集成测试。 PostgreSQL/Redis Adapter 与 Compose 集成测试。
## Task 11: Checker and Health Reducer ## Task 11: Checker and Health Reducer

View File

@ -38,7 +38,7 @@
- `PROVIDER-*`Provider HTTP Client、严格响应上限、模板解析安全边界、凭据 - `PROVIDER-*`Provider HTTP Client、严格响应上限、模板解析安全边界、凭据
引用 Store 与 Reconciler Adapter 已实现。 引用 Store 与 Reconciler Adapter 已实现。
- `DIST/Admin HTTP`:严格 JSON、Request ID、Problem 响应及 Distribution/Admin - `DIST/Admin HTTP`:严格 JSON、Request ID、Problem 响应及 Distribution/Admin
Handler 已实现;独立进程装配与认证授权仍在后续范围 Handler 已实现;共享认证、CIDR、可信代理、Client ID 与本地准入保护链已接入
## 2. 已执行验证 ## 2. 已执行验证
@ -68,7 +68,7 @@ CI 已配置 Linux race job。Docker/Kubernetes 仅完成静态验证,没有
5. Redis Leader、速率限制、心跳与可重建协调适配器。 5. Redis Leader、速率限制、心跳与可重建协调适配器。
6. Worker ownership drain/ACK/过期回收和网络快照流。 6. Worker ownership drain/ACK/过期回收和网络快照流。
7. Checker 调度、探测器和健康 reducer。 7. Checker 调度、探测器和健康 reducer。
8. Admin/Distribution 独立监听器装配、鉴权授权、分布式限流和审计查询。 8. Admin/Distribution 独立监听器装配、细粒度授权、分布式限流和审计查询。
9. 真实 Compose/Kubernetes 集成、故障演练和代表性集群负载测试。 9. 真实 Compose/Kubernetes 集成、故障演练和代表性集群负载测试。
## 4. 容量结论 ## 4. 容量结论

View File

@ -13,6 +13,11 @@
- 认证关闭不代表匿名状态消失:仍按可信代理链解析来源并形成 Client ID。 - 认证关闭不代表匿名状态消失:仍按可信代理链解析来源并形成 Client ID。
- Access、Auth、Rate Limit 和 Client Identification 相互独立。 - Access、Auth、Rate Limit 和 Client Identification 相互独立。
- Admin 使用独立凭据,不能复用普通 Gateway 或 Distribution 凭据。 - Admin 使用独立凭据,不能复用普通 Gateway 或 Distribution 凭据。
- `httpsecurity.Protection` 统一实现 Basic、API Key、Bearer、CIDR、可信代理链、
Client ID 与入口准入API 使用 401/`WWW-Authenticate`Gateway 使用
407/`Proxy-Authenticate`。
- Token Client ID 使用 SHA-256 的 128 位摘要前缀,不把 Token 本身写入领域、
日志或审计键。
## 3. 目标地址策略 ## 3. 目标地址策略

View File

@ -6,6 +6,8 @@ import (
"path/filepath" "path/filepath"
"strings" "strings"
"testing" "testing"
"go.yaml.in/yaml/v4"
) )
const validConfig = ` const validConfig = `
@ -170,6 +172,76 @@ func TestShippedConfigurationsAreValid(t *testing.T) {
} }
} }
func TestShippedDeploymentConfigurationsResolveEnvironment(t *testing.T) {
resolver := fixtureResolver{environment: map[string]string{
"PROXY_POOL_GATEWAY_USERNAME": "resolved-gateway-user",
"PROXY_POOL_GATEWAY_PASSWORD": "resolved-gateway-password",
"PROXY_POOL_EXTRACT_TOKEN": "resolved-extract-token",
"PROXY_POOL_ADMIN_TOKEN": "resolved-admin-token",
"PROXY_POOL_POSTGRES_URL": "postgres://resolved",
"PROXY_POOL_REDIS_URL": "redis://resolved",
"PROVIDER_A_TOKEN": "resolved-provider-a-token",
"PROVIDER_B_TOKEN": "resolved-provider-b-token",
}}
tests := []struct {
name string
path string
configMap bool
wantGatewayUser string
wantPostgresURL string
wantRedisURL string
}{
{
name: "local",
path: filepath.Join("..", "..", "deploy", "config", "local.yaml"),
wantGatewayUser: "local-gateway",
wantPostgresURL: "postgres://proxy_pool:local-only-change-me@postgres:5432/proxy_pool?sslmode=disable",
wantRedisURL: "redis://redis:6379/0",
},
{
name: "kubernetes",
path: filepath.Join("..", "..", "deploy", "kubernetes", "base", "configmap.yaml"),
configMap: true,
wantGatewayUser: "resolved-gateway-user",
wantPostgresURL: "postgres://resolved",
wantRedisURL: "redis://resolved",
},
}
for _, test := range tests {
test := test
t.Run(test.name, func(t *testing.T) {
content, err := os.ReadFile(test.path)
if err != nil {
t.Fatalf("ReadFile(): %v", err)
}
configuration := string(content)
if test.configMap {
var manifest struct {
Data map[string]string `yaml:"data"`
}
if err := yaml.Unmarshal(content, &manifest); err != nil {
t.Fatalf("Unmarshal(): %v", err)
}
configuration = manifest.Data["config.yaml"]
}
cfg, err := LoadResolved(strings.NewReader(configuration), resolver)
if err != nil {
t.Fatalf("LoadResolved(): %v", err)
}
if cfg.Gateway.Auth.Username != test.wantGatewayUser ||
cfg.Gateway.Auth.Password != "resolved-gateway-password" ||
cfg.Distribution.Auth.Token != "resolved-extract-token" ||
cfg.Admin.Auth.Token != "resolved-admin-token" ||
cfg.Storage.PostgresURL != test.wantPostgresURL ||
cfg.Storage.RedisURL != test.wantRedisURL ||
cfg.Upstreams["provider-a"].API.Auth.Value != "resolved-provider-a-token" ||
cfg.Upstreams["provider-b"].API.Auth.Value != "resolved-provider-b-token" {
t.Fatalf("deployment values were not resolved: %+v", cfg.Redacted())
}
})
}
}
func TestLoadResolvedExpandsEnvironmentWithoutChangingTemplateVariables(t *testing.T) { func TestLoadResolvedExpandsEnvironmentWithoutChangingTemplateVariables(t *testing.T) {
configured := strings.Replace(validConfig, ` auth: configured := strings.Replace(validConfig, ` auth:
type: none`, ` auth: type: none`, ` auth:
@ -230,6 +302,20 @@ func TestValidateRejectsUnsupportedListenerAuthMode(t *testing.T) {
} }
} }
func TestValidateAcceptsBearerListenerAuthentication(t *testing.T) {
cfg := mustLoadValidConfig(t)
cfg.Distribution.Auth = Auth{Mode: "bearer", Token: "resolved-token"}
cfg.Distribution.ClientIdentification.Mode = "authenticatedClient"
if err := Validate(cfg); err != nil {
t.Fatalf("Validate(bearer) error = %v", err)
}
cfg.Distribution.Auth = Auth{Mode: "any", Methods: []AuthMethod{{Mode: "bearer", Value: "method-token"}}}
if err := Validate(cfg); err != nil {
t.Fatalf("Validate(any bearer) error = %v", err)
}
}
func TestValidateRejectsInvalidConfigurationMatrix(t *testing.T) { func TestValidateRejectsInvalidConfigurationMatrix(t *testing.T) {
tests := []struct { tests := []struct {
name string name string
@ -283,6 +369,28 @@ func TestValidateRejectsInvalidConfigurationMatrix(t *testing.T) {
}, },
want: "trustedProxies", want: "trustedProxies",
}, },
{
name: "negative listener request limit",
mutate: func(cfg *Config) {
cfg.Distribution.Limits.RequestsPerMinute = -1
},
want: "requestsPerMinute",
},
{
name: "invalid client identification mode",
mutate: func(cfg *Config) {
cfg.Distribution.ClientIdentification.Mode = "header"
},
want: "clientIdentification.mode",
},
{
name: "authenticated client without authentication",
mutate: func(cfg *Config) {
cfg.Distribution.Auth = Auth{Mode: "none"}
cfg.Distribution.ClientIdentification.Mode = "authenticatedClient"
},
want: "authenticatedClient requires authentication",
},
{ {
name: "invalid destination deny CIDR", name: "invalid destination deny CIDR",
mutate: func(cfg *Config) { mutate: func(cfg *Config) {

View File

@ -55,6 +55,17 @@ func Validate(cfg *Config) error {
} }
} }
if cfg.Distribution.Enabled { if cfg.Distribution.Enabled {
clientIdentificationMode := cfg.Distribution.ClientIdentification.Mode
if clientIdentificationMode == "" {
clientIdentificationMode = "sourceIP"
}
if err := validateEnum("distribution.clientIdentification.mode", clientIdentificationMode,
"sourceIP", "authenticatedClient", "authenticatedClientOrSourceIP"); err != nil {
return err
}
if clientIdentificationMode == "authenticatedClient" && cfg.Distribution.Auth.Mode == "none" {
return fmt.Errorf("validate distribution clientIdentification.mode: authenticatedClient requires authentication")
}
if err := requirePositive("distribution.maxCountPerRequest", cfg.Distribution.Extraction.MaxCountPerRequest); err != nil { if err := requirePositive("distribution.maxCountPerRequest", cfg.Distribution.Extraction.MaxCountPerRequest); err != nil {
return err return err
} }
@ -84,6 +95,18 @@ func validateListener(name string, listener Listener, security Security) error {
if err := validateListenerAuth(name, listener.Auth); err != nil { if err := validateListenerAuth(name, listener.Auth); err != nil {
return err return err
} }
for _, limit := range []struct {
name string
value int
}{
{name: "maxConcurrentConnections", value: listener.Limits.MaxConcurrentConnections},
{name: "requestsPerMinute", value: listener.Limits.RequestsPerMinute},
{name: "requestsPerMinutePerClient", value: listener.Limits.RequestsPerMinutePerClient},
} {
if limit.value < 0 {
return fmt.Errorf("validate %s limits.%s: must be non-negative", name, limit.name)
}
}
host, _, err := net.SplitHostPort(listener.Listen) host, _, err := net.SplitHostPort(listener.Listen)
if err != nil { if err != nil {
return fmt.Errorf("validate %s listen: %w", name, err) return fmt.Errorf("validate %s listen: %w", name, err)
@ -210,6 +233,10 @@ func validateListenerAuth(listener string, auth Auth) error {
if auth.Header == "" || (auth.Token == "" && auth.TokenFile == "") { if auth.Header == "" || (auth.Token == "" && auth.TokenFile == "") {
return fmt.Errorf("validate %s auth.mode apiKey: header and token are required", listener) return fmt.Errorf("validate %s auth.mode apiKey: header and token are required", listener)
} }
case "bearer":
if auth.Token == "" && auth.TokenFile == "" {
return fmt.Errorf("validate %s auth.mode bearer: token is required", listener)
}
case "ipWhitelist": case "ipWhitelist":
if len(auth.CIDRs) == 0 { if len(auth.CIDRs) == 0 {
return fmt.Errorf("validate %s auth.mode ipWhitelist: cidrs are required", listener) return fmt.Errorf("validate %s auth.mode ipWhitelist: cidrs are required", listener)
@ -242,6 +269,10 @@ func validateAuthMethod(listener string, index int, method AuthMethod) error {
if method.Header == "" || (method.Value == "" && method.ValueFile == "") { if method.Header == "" || (method.Value == "" && method.ValueFile == "") {
return fmt.Errorf("validate %s auth.methods[%d]: header and value are required", listener, index) return fmt.Errorf("validate %s auth.methods[%d]: header and value are required", listener, index)
} }
case "bearer":
if method.Value == "" && method.ValueFile == "" {
return fmt.Errorf("validate %s auth.methods[%d]: bearer value is required", listener, index)
}
case "ipWhitelist": case "ipWhitelist":
if len(method.CIDRs) == 0 { if len(method.CIDRs) == 0 {
return fmt.Errorf("validate %s auth.methods[%d]: cidrs are required", listener, index) return fmt.Errorf("validate %s auth.methods[%d]: cidrs are required", listener, index)

View File

@ -8,6 +8,7 @@ import (
"strings" "strings"
"github.com/proxy-pool/proxy-pool/internal/platform/httpapi" "github.com/proxy-pool/proxy-pool/internal/platform/httpapi"
"github.com/proxy-pool/proxy-pool/internal/platform/httpsecurity"
) )
const ( const (
@ -33,6 +34,12 @@ type Service interface {
ReloadConfiguration(context.Context, ReloadCommand) (MutationResult, error) ReloadConfiguration(context.Context, ReloadCommand) (MutationResult, error)
} }
type Authorizer interface {
Check(context.Context, *http.Request) error
}
var _ Authorizer = (*httpsecurity.Protection)(nil)
type Options struct { type Options struct {
MaxBodyBytes int64 MaxBodyBytes int64
} }
@ -91,14 +98,15 @@ type ReloadCommand struct {
type Handler struct { type Handler struct {
service Service service Service
authorizer Authorizer
maxBodyBytes int64 maxBodyBytes int64
} }
func NewHandler(service Service, options Options) (*Handler, error) { func NewHandler(service Service, authorizer Authorizer, options Options) (*Handler, error) {
if service == nil || options.MaxBodyBytes <= 0 { if service == nil || authorizer == nil || options.MaxBodyBytes <= 0 {
return nil, ErrInvalidHandler return nil, ErrInvalidHandler
} }
return &Handler{service: service, maxBodyBytes: options.MaxBodyBytes}, nil return &Handler{service: service, authorizer: authorizer, maxBodyBytes: options.MaxBodyBytes}, nil
} }
func (handler *Handler) ServeHTTP(writer http.ResponseWriter, request *http.Request) { func (handler *Handler) ServeHTTP(writer http.ResponseWriter, request *http.Request) {
@ -107,6 +115,12 @@ func (handler *Handler) ServeHTTP(writer http.ResponseWriter, request *http.Requ
writeTransportProblem(writer, http.StatusBadRequest, "INVALID_REQUEST_ID", "Invalid request ID", "X-Request-ID is invalid", requestID) writeTransportProblem(writer, http.StatusBadRequest, "INVALID_REQUEST_ID", "Invalid request ID", "X-Request-ID is invalid", requestID)
return return
} }
if err := handler.authorizer.Check(request.Context(), request); err != nil {
if !httpsecurity.WriteProblem(writer, requestID, err) {
writeTransportProblem(writer, http.StatusInternalServerError, "INTERNAL_ERROR", "Internal server error", "the request could not be completed", requestID)
}
return
}
switch request.URL.Path { switch request.URL.Path {
case statusPath: case statusPath:

View File

@ -10,6 +10,7 @@ import (
"testing" "testing"
"github.com/proxy-pool/proxy-pool/internal/platform/httpapi" "github.com/proxy-pool/proxy-pool/internal/platform/httpapi"
"github.com/proxy-pool/proxy-pool/internal/platform/httpsecurity"
) )
func TestHandlerReturnsStatusWithoutSensitiveDetails(t *testing.T) { func TestHandlerReturnsStatusWithoutSensitiveDetails(t *testing.T) {
@ -39,12 +40,68 @@ func TestHandlerReturnsStatusWithoutSensitiveDetails(t *testing.T) {
func TestNewHandlerRejectsMissingDependenciesAndInvalidLimit(t *testing.T) { func TestNewHandlerRejectsMissingDependenciesAndInvalidLimit(t *testing.T) {
t.Parallel() t.Parallel()
if _, err := NewHandler(nil, Options{MaxBodyBytes: 1024}); !errors.Is(err, ErrInvalidHandler) { if _, err := NewHandler(nil, allowAuthorizer{}, Options{MaxBodyBytes: 1024}); !errors.Is(err, ErrInvalidHandler) {
t.Fatalf("NewHandler(nil) error = %v, want %v", err, ErrInvalidHandler) t.Fatalf("NewHandler(nil) error = %v, want %v", err, ErrInvalidHandler)
} }
if _, err := NewHandler(&stubService{}, Options{}); !errors.Is(err, ErrInvalidHandler) { if _, err := NewHandler(&stubService{}, allowAuthorizer{}, Options{}); !errors.Is(err, ErrInvalidHandler) {
t.Fatalf("NewHandler(zero limit) error = %v, want %v", err, ErrInvalidHandler) t.Fatalf("NewHandler(zero limit) error = %v, want %v", err, ErrInvalidHandler)
} }
if _, err := NewHandler(&stubService{}, nil, Options{MaxBodyBytes: 1024}); !errors.Is(err, ErrInvalidHandler) {
t.Fatalf("NewHandler(nil authorizer) error = %v, want %v", err, ErrInvalidHandler)
}
}
func TestHandlerAuthorizesBeforeRouting(t *testing.T) {
t.Parallel()
service := &stubService{}
handler, err := NewHandler(service, rejectAuthorizer{}, Options{MaxBodyBytes: 1024})
if err != nil {
t.Fatalf("NewHandler() error = %v", err)
}
recorder := httptest.NewRecorder()
handler.ServeHTTP(recorder, httptest.NewRequest(http.MethodGet, "/missing", nil))
if recorder.Code != http.StatusUnauthorized {
t.Fatalf("status = %d, want %d; body=%s", recorder.Code, http.StatusUnauthorized, recorder.Body.String())
}
if recorder.Header().Get("WWW-Authenticate") != `Basic realm="proxy-pool"` {
t.Fatalf("challenge = %q", recorder.Header().Get("WWW-Authenticate"))
}
if service.statusCalls != 0 {
t.Fatalf("status calls = %d, want 0", service.statusCalls)
}
}
func TestHandlerUsesHTTPProtectionAuthenticationContract(t *testing.T) {
t.Parallel()
service := &stubService{status: Status{ConfigVersion: "cfg-1"}}
protection, err := httpsecurity.New(httpsecurity.Config{
Authentication: httpsecurity.Authentication{Mode: httpsecurity.ModeBearer, Token: "admin-token"},
}, nil)
if err != nil {
t.Fatalf("httpsecurity.New() error = %v", err)
}
handler, err := NewHandler(service, protection, Options{MaxBodyBytes: 1024})
if err != nil {
t.Fatalf("NewHandler() error = %v", err)
}
unauthorized := httptest.NewRecorder()
handler.ServeHTTP(unauthorized, httptest.NewRequest(http.MethodGet, "/api/v1/status", nil))
if unauthorized.Code != http.StatusUnauthorized || unauthorized.Header().Get("WWW-Authenticate") != `Bearer realm="proxy-pool"` {
t.Fatalf("unauthorized response = status %d challenge %q", unauthorized.Code, unauthorized.Header().Get("WWW-Authenticate"))
}
if service.statusCalls != 0 {
t.Fatalf("status calls after rejection = %d, want 0", service.statusCalls)
}
authorized := httptest.NewRecorder()
request := httptest.NewRequest(http.MethodGet, "/api/v1/status", nil)
request.Header.Set("Authorization", "Bearer admin-token")
handler.ServeHTTP(authorized, request)
if authorized.Code != http.StatusOK || service.statusCalls != 1 {
t.Fatalf("authorized response = status %d calls %d", authorized.Code, service.statusCalls)
}
} }
func TestHandlerEnablesAndDisablesUpstream(t *testing.T) { func TestHandlerEnablesAndDisablesUpstream(t *testing.T) {
@ -202,7 +259,7 @@ func TestHandlerRejectsInvalidTransportRequests(t *testing.T) {
func mustHandler(t *testing.T, service Service) *Handler { func mustHandler(t *testing.T, service Service) *Handler {
t.Helper() t.Helper()
handler, err := NewHandler(service, Options{MaxBodyBytes: 1024}) handler, err := NewHandler(service, allowAuthorizer{}, Options{MaxBodyBytes: 1024})
if err != nil { if err != nil {
t.Fatalf("NewHandler() error = %v", err) t.Fatalf("NewHandler() error = %v", err)
} }
@ -216,9 +273,11 @@ type stubService struct {
lastUpstream SetUpstreamCommand lastUpstream SetUpstreamCommand
lastSwitch SwitchCommand lastSwitch SwitchCommand
lastReload ReloadCommand lastReload ReloadCommand
statusCalls int
} }
func (service *stubService) Status(context.Context) (Status, error) { func (service *stubService) Status(context.Context) (Status, error) {
service.statusCalls++
return service.status, service.err return service.status, service.err
} }
@ -236,3 +295,18 @@ func (service *stubService) ReloadConfiguration(_ context.Context, command Reloa
service.lastReload = command service.lastReload = command
return service.mutation, service.err return service.mutation, service.err
} }
type allowAuthorizer struct{}
func (allowAuthorizer) Check(context.Context, *http.Request) error { return nil }
type rejectAuthorizer struct{}
func (rejectAuthorizer) Check(context.Context, *http.Request) error {
return &httpsecurity.HTTPError{
StatusCode: http.StatusUnauthorized,
Code: "UNAUTHORIZED",
Header: http.Header{"WWW-Authenticate": []string{`Basic realm="proxy-pool"`}},
Cause: errors.New("credential secret"),
}
}

View File

@ -10,6 +10,7 @@ import (
controllerExtraction "github.com/proxy-pool/proxy-pool/internal/controller/extraction" controllerExtraction "github.com/proxy-pool/proxy-pool/internal/controller/extraction"
domainExtraction "github.com/proxy-pool/proxy-pool/internal/domain/extraction" domainExtraction "github.com/proxy-pool/proxy-pool/internal/domain/extraction"
"github.com/proxy-pool/proxy-pool/internal/platform/httpapi" "github.com/proxy-pool/proxy-pool/internal/platform/httpapi"
"github.com/proxy-pool/proxy-pool/internal/platform/httpsecurity"
) )
const ( const (
@ -38,15 +39,14 @@ type Extractor interface {
Extract(context.Context, controllerExtraction.Request) (controllerExtraction.Response, error) Extract(context.Context, controllerExtraction.Request) (controllerExtraction.Response, error)
} }
type Identity struct { type Identity = httpsecurity.Identity
ClientID string
SourceIP string
}
type IdentityResolver interface { type IdentityResolver interface {
Resolve(*http.Request) (Identity, error) Resolve(*http.Request) (Identity, error)
} }
var _ IdentityResolver = (*httpsecurity.Protection)(nil)
type ReadinessChecker interface { type ReadinessChecker interface {
Ready(context.Context) error Ready(context.Context) error
} }
@ -156,6 +156,24 @@ func (h *Handler) ServeHTTP(writer http.ResponseWriter, request *http.Request) {
} }
func (h *Handler) handleExtract(writer http.ResponseWriter, request *http.Request, requestID string) { func (h *Handler) handleExtract(writer http.ResponseWriter, request *http.Request, requestID string) {
identity, err := h.identity.Resolve(request)
if err != nil {
if !httpsecurity.WriteProblem(writer, requestID, err) {
h.writeProblem(writer, httpapi.NewProblem(
http.StatusInternalServerError,
"INTERNAL_ERROR",
"Internal server error",
"",
requestID,
))
}
return
}
if strings.TrimSpace(identity.ClientID) == "" && strings.TrimSpace(identity.SourceIP) == "" {
h.writeProblem(writer, problemBadRequest(requestID, "INVALID_REQUEST", "Invalid request", "", nil))
return
}
idempotencyKey, err := validateIdempotencyKey(request.Header.Values(headerIdempotencyKey)) idempotencyKey, err := validateIdempotencyKey(request.Header.Values(headerIdempotencyKey))
if err != nil { if err != nil {
h.writeProblem(writer, problemBadRequest(requestID, "INVALID_HEADER", "Invalid request header", "", []httpapi.InvalidParam{{ h.writeProblem(writer, problemBadRequest(requestID, "INVALID_HEADER", "Invalid request header", "", []httpapi.InvalidParam{{
@ -183,12 +201,6 @@ func (h *Handler) handleExtract(writer http.ResponseWriter, request *http.Reques
return return
} }
identity, err := h.identity.Resolve(request)
if err != nil || (strings.TrimSpace(identity.ClientID) == "" && strings.TrimSpace(identity.SourceIP) == "") {
h.writeProblem(writer, problemBadRequest(requestID, "INVALID_REQUEST", "Invalid request", "", nil))
return
}
filters := payload.filtersOrZero() filters := payload.filtersOrZero()
serviceResponse, err := h.extractor.Extract(request.Context(), controllerExtraction.Request{ serviceResponse, err := h.extractor.Extract(request.Context(), controllerExtraction.Request{
RequestID: requestID, RequestID: requestID,

View File

@ -14,6 +14,7 @@ import (
controllerExtraction "github.com/proxy-pool/proxy-pool/internal/controller/extraction" controllerExtraction "github.com/proxy-pool/proxy-pool/internal/controller/extraction"
domainExtraction "github.com/proxy-pool/proxy-pool/internal/domain/extraction" domainExtraction "github.com/proxy-pool/proxy-pool/internal/domain/extraction"
"github.com/proxy-pool/proxy-pool/internal/platform/httpapi" "github.com/proxy-pool/proxy-pool/internal/platform/httpapi"
"github.com/proxy-pool/proxy-pool/internal/platform/httpsecurity"
) )
func TestNewHandlerValidatesDependenciesAndBodyLimit(t *testing.T) { func TestNewHandlerValidatesDependenciesAndBodyLimit(t *testing.T) {
@ -202,6 +203,65 @@ func TestHandlerRejectsDuplicateIdempotencyHeader(t *testing.T) {
} }
} }
func TestHandlerMapsSecurityFailureBeforeParsingBody(t *testing.T) {
t.Parallel()
extractor := &fakeExtractor{}
handler := mustNewHandler(t, Config{BodyLimitBytes: 1024}, Dependencies{
Extractor: extractor,
Identity: fakeIdentityResolver{err: &httpsecurity.HTTPError{
StatusCode: http.StatusUnauthorized,
Code: "UNAUTHORIZED",
Header: http.Header{"WWW-Authenticate": []string{`Bearer realm="proxy-pool"`}},
Cause: errors.New("token=secret"),
}},
Readiness: fakeReadinessChecker{},
})
request := httptest.NewRequest(http.MethodPost, "/api/v1/proxies/extract", strings.NewReader(`not-json`))
response := httptest.NewRecorder()
handler.ServeHTTP(response, request)
if response.Code != http.StatusUnauthorized {
t.Fatalf("status = %d, want %d; body=%s", response.Code, http.StatusUnauthorized, response.Body.String())
}
if challenge := response.Header().Get("WWW-Authenticate"); challenge != `Bearer realm="proxy-pool"` {
t.Fatalf("challenge = %q", challenge)
}
if strings.Contains(response.Body.String(), "secret") {
t.Fatalf("security response leaked cause: %s", response.Body.String())
}
if extractor.calls != 0 {
t.Fatalf("extractor calls = %d, want 0", extractor.calls)
}
}
func TestHandlerMapsUnexpectedIdentityFailureToInternalError(t *testing.T) {
t.Parallel()
extractor := &fakeExtractor{}
handler := mustNewHandler(t, Config{BodyLimitBytes: 1024}, Dependencies{
Extractor: extractor,
Identity: fakeIdentityResolver{err: errors.New("credential store secret")},
Readiness: fakeReadinessChecker{},
})
request := httptest.NewRequest(http.MethodPost, "/api/v1/proxies/extract", strings.NewReader(`{"count":1}`))
response := httptest.NewRecorder()
handler.ServeHTTP(response, request)
if response.Code != http.StatusInternalServerError {
t.Fatalf("status = %d, want %d; body=%s", response.Code, http.StatusInternalServerError, response.Body.String())
}
if !strings.Contains(response.Body.String(), `"code":"INTERNAL_ERROR"`) {
t.Fatalf("body = %s, want INTERNAL_ERROR", response.Body.String())
}
if strings.Contains(response.Body.String(), "secret") {
t.Fatalf("response leaked dependency error: %s", response.Body.String())
}
if extractor.calls != 0 {
t.Fatalf("extractor calls = %d, want 0", extractor.calls)
}
}
func TestHandlerExtractMapsErrorsToProblemResponsesWithoutSensitiveLeakage(t *testing.T) { func TestHandlerExtractMapsErrorsToProblemResponsesWithoutSensitiveLeakage(t *testing.T) {
t.Parallel() t.Parallel()
tooManyRegions := `{"count":1,"filters":{"regions":["` + strings.Join(makeUniqueValues(65), `","`) + `"]}}` tooManyRegions := `{"count":1,"filters":{"regions":["` + strings.Join(makeUniqueValues(65), `","`) + `"]}}`

View File

@ -1,13 +1,13 @@
package server package server
import ( import (
"errors"
"fmt" "fmt"
"time" "time"
"github.com/proxy-pool/proxy-pool/internal/config" "github.com/proxy-pool/proxy-pool/internal/config"
"github.com/proxy-pool/proxy-pool/internal/gateway/policy" "github.com/proxy-pool/proxy-pool/internal/gateway/policy"
platformAdmission "github.com/proxy-pool/proxy-pool/internal/platform/admission" platformAdmission "github.com/proxy-pool/proxy-pool/internal/platform/admission"
"github.com/proxy-pool/proxy-pool/internal/platform/httpsecurity"
) )
type Protection struct { type Protection struct {
@ -26,7 +26,7 @@ func BuildProtection(listener config.Listener) (Protection, error) {
if err != nil { if err != nil {
return Protection{}, err return Protection{}, err
} }
auth, err := buildConfiguredAuth(listener.Auth, clientIPs) auth, err := buildConfiguredAuth(listener)
if err != nil { if err != nil {
return Protection{}, err return Protection{}, err
} }
@ -69,40 +69,25 @@ func explicitlyAllowed(deny *bool) bool {
return deny != nil && !*deny return deny != nil && !*deny
} }
func buildConfiguredAuth(auth config.Auth, clientIPs *ClientIPResolver) (Guard, error) { func buildConfiguredAuth(listener config.Listener) (Guard, error) {
switch auth.Mode { switch listener.Auth.Mode {
case "", "none": case "", "none":
return nil, nil return nil, nil
case "usernamePassword": }
return NewBasicAuthGuard(auth.Username, auth.Password), nil authListener := config.Listener{
case "apiKey": Access: config.Access{
return NewAPIKeyGuard(auth.Header, auth.Token), nil TrustedProxies: append([]string(nil), listener.Access.TrustedProxies...),
case "ipWhitelist": },
return NewAccessGuard(clientIPs, auth.CIDRs) Auth: listener.Auth,
case "any": }
methods := make([]Guard, 0, len(auth.Methods)) protection, err := httpsecurity.NewFromListener(
for index, method := range auth.Methods { authListener,
guard, err := buildConfiguredMethod(method, clientIPs) httpsecurity.ClientSourceIP,
httpsecurity.ProxySemantics,
nil,
)
if err != nil { if err != nil {
return nil, fmt.Errorf("build gateway auth method %d: %w", index, err) return nil, fmt.Errorf("build gateway authentication: %w", err)
}
methods = append(methods, guard)
}
return NewAnyGuard(methods...), nil
default:
return nil, fmt.Errorf("build gateway auth: unsupported mode %q", auth.Mode)
}
}
func buildConfiguredMethod(method config.AuthMethod, clientIPs *ClientIPResolver) (Guard, error) {
switch method.Mode {
case "usernamePassword":
return NewBasicAuthGuard(method.Username, method.Password), nil
case "apiKey":
return NewAPIKeyGuard(method.Header, method.Value), nil
case "ipWhitelist":
return NewAccessGuard(clientIPs, method.CIDRs)
default:
return nil, errors.New("unsupported authentication method: " + method.Mode)
} }
return protection, nil
} }

View File

@ -9,6 +9,7 @@ import (
"github.com/proxy-pool/proxy-pool/internal/config" "github.com/proxy-pool/proxy-pool/internal/config"
"github.com/proxy-pool/proxy-pool/internal/gateway/policy" "github.com/proxy-pool/proxy-pool/internal/gateway/policy"
"github.com/proxy-pool/proxy-pool/internal/platform/httpsecurity"
) )
func TestBuildProtectionFromListenerConfig(t *testing.T) { func TestBuildProtectionFromListenerConfig(t *testing.T) {
@ -38,6 +39,62 @@ func TestBuildProtectionFromListenerConfig(t *testing.T) {
} }
} }
func TestBuildProtectionSupportsBearerProxyAuthentication(t *testing.T) {
t.Parallel()
protection, err := BuildProtection(config.Listener{Auth: config.Auth{Mode: "bearer", Token: "proxy-token"}})
if err != nil {
t.Fatalf("BuildProtection() error = %v", err)
}
request := httptest.NewRequest(http.MethodGet, "http://example.test", nil)
request.RemoteAddr = "198.51.100.8:1234"
request.Header.Set("Proxy-Authorization", "Bearer proxy-token")
if err := protection.Auth.Check(context.Background(), request); err != nil {
t.Fatalf("auth.Check() error = %v", err)
}
request.Header.Set("Proxy-Authorization", "Bearer wrong")
err = protection.Auth.Check(context.Background(), request)
var securityError *httpsecurity.HTTPError
if !errors.As(err, &securityError) || securityError.StatusCode != http.StatusProxyAuthRequired {
t.Fatalf("auth.Check(wrong) error = %T %v", err, err)
}
recorder := httptest.NewRecorder()
writeGatewayError(recorder, err)
if recorder.Code != http.StatusProxyAuthRequired || recorder.Header().Get("Proxy-Authenticate") == "" {
t.Fatalf("gateway response = status %d headers %v", recorder.Code, recorder.Header())
}
}
func TestBuildProtectionAnyPreservesIPWhitelistRejectionIndependentOfOrder(t *testing.T) {
t.Parallel()
methods := [][]config.AuthMethod{
{
{Mode: "ipWhitelist", CIDRs: []string{"10.0.0.0/8"}},
{Mode: "apiKey", Header: "X-Proxy-Key", Value: "secret"},
},
{
{Mode: "apiKey", Header: "X-Proxy-Key", Value: "secret"},
{Mode: "ipWhitelist", CIDRs: []string{"10.0.0.0/8"}},
},
}
for index, configuredMethods := range methods {
protection, err := BuildProtection(config.Listener{
Auth: config.Auth{Mode: "any", Methods: configuredMethods},
})
if err != nil {
t.Fatalf("BuildProtection(%d) error = %v", index, err)
}
request := httptest.NewRequest(http.MethodGet, "http://example.test", nil)
request.RemoteAddr = "198.51.100.8:1234"
err = protection.Auth.Check(context.Background(), request)
recorder := httptest.NewRecorder()
writeGatewayError(recorder, err)
if recorder.Code != http.StatusForbidden {
t.Fatalf("method order %d status = %d, want 403", index, recorder.Code)
}
}
}
func TestTargetPolicyFromOmittedConfigDefaultsToDeny(t *testing.T) { func TestTargetPolicyFromOmittedConfigDefaultsToDeny(t *testing.T) {
t.Parallel() t.Parallel()

View File

@ -18,6 +18,7 @@ import (
"github.com/proxy-pool/proxy-pool/internal/gateway/dispatch" "github.com/proxy-pool/proxy-pool/internal/gateway/dispatch"
"github.com/proxy-pool/proxy-pool/internal/gateway/policy" "github.com/proxy-pool/proxy-pool/internal/gateway/policy"
transportDomain "github.com/proxy-pool/proxy-pool/internal/gateway/transport" transportDomain "github.com/proxy-pool/proxy-pool/internal/gateway/transport"
"github.com/proxy-pool/proxy-pool/internal/platform/httpsecurity"
) )
type Config struct { type Config struct {
@ -523,6 +524,11 @@ func writeGatewayError(writer http.ResponseWriter, err error) {
if errors.As(err, &httpError) { if errors.As(err, &httpError) {
status = httpError.StatusCode status = httpError.StatusCode
copyHeaders(writer.Header(), httpError.Header) copyHeaders(writer.Header(), httpError.Header)
} else {
var securityError *httpsecurity.HTTPError
if errors.As(err, &securityError) {
status = securityError.StatusCode
copyHeaders(writer.Header(), securityError.Header)
} else { } else {
switch { switch {
case errors.Is(err, policy.ErrInvalidAuthority): case errors.Is(err, policy.ErrInvalidAuthority):
@ -537,6 +543,7 @@ func writeGatewayError(writer http.ResponseWriter, err error) {
status = http.StatusGatewayTimeout status = http.StatusGatewayTimeout
} }
} }
}
http.Error(writer, http.StatusText(status), status) http.Error(writer, http.StatusText(status), status)
} }

View File

@ -2,17 +2,12 @@ package server
import ( import (
"context" "context"
"crypto/subtle"
"encoding/base64"
"errors" "errors"
"fmt" "fmt"
"net"
"net/http" "net/http"
"net/netip"
"strconv"
"strings"
"github.com/proxy-pool/proxy-pool/internal/gateway/policy" "github.com/proxy-pool/proxy-pool/internal/gateway/policy"
"github.com/proxy-pool/proxy-pool/internal/platform/httpsecurity"
) )
type HTTPError struct { type HTTPError struct {
@ -30,190 +25,9 @@ func (err *HTTPError) Error() string {
func (err *HTTPError) Unwrap() error { return err.Cause } func (err *HTTPError) Unwrap() error { return err.Cause }
type BasicAuthGuard struct { type ClientIPResolver = httpsecurity.ClientIPResolver
username string
password string
}
func NewBasicAuthGuard(username, password string) *BasicAuthGuard { var NewClientIPResolver = httpsecurity.NewClientIPResolver
return &BasicAuthGuard{username: username, password: password}
}
func (guard *BasicAuthGuard) Check(_ context.Context, request *http.Request) error {
username, password, ok := parseBasicCredentials(request.Header.Get("Proxy-Authorization"))
if ok && constantTimeEqual(username, guard.username) && constantTimeEqual(password, guard.password) {
return nil
}
return &HTTPError{
StatusCode: http.StatusProxyAuthRequired,
Header: http.Header{"Proxy-Authenticate": []string{`Basic realm="proxy"`}},
Cause: errors.New("proxy authentication failed"),
}
}
type APIKeyGuard struct {
header string
value string
}
func NewAPIKeyGuard(header, value string) *APIKeyGuard {
if strings.TrimSpace(header) == "" {
header = "X-API-Key"
}
return &APIKeyGuard{header: header, value: value}
}
func (guard *APIKeyGuard) Check(_ context.Context, request *http.Request) error {
if constantTimeEqual(request.Header.Get(guard.header), guard.value) {
return nil
}
return &HTTPError{StatusCode: http.StatusProxyAuthRequired, Cause: errors.New("proxy API key authentication failed")}
}
type AnyGuard struct {
guards []Guard
}
func NewAnyGuard(guards ...Guard) *AnyGuard {
return &AnyGuard{guards: append([]Guard(nil), guards...)}
}
func (guard *AnyGuard) Check(ctx context.Context, request *http.Request) error {
var lastErr error
for _, candidate := range guard.guards {
if candidate == nil {
continue
}
if err := candidate.Check(ctx, request); err == nil {
return nil
} else {
lastErr = err
}
}
if lastErr == nil {
lastErr = errors.New("no authentication method is configured")
}
return lastErr
}
type ClientIPResolver struct {
trusted policy.CIDRMatcher
}
func NewClientIPResolver(trustedCIDRs []string) (*ClientIPResolver, error) {
trusted, err := policy.NewCIDRMatcher(trustedCIDRs)
if err != nil {
return nil, fmt.Errorf("create client IP resolver: %w", err)
}
return &ClientIPResolver{trusted: trusted}, nil
}
func (resolver *ClientIPResolver) Resolve(request *http.Request) (netip.Addr, error) {
if resolver == nil || request == nil {
return netip.Addr{}, errors.New("resolve client IP: resolver and request are required")
}
peer, err := parseRemoteAddress(request.RemoteAddr)
if err != nil {
return netip.Addr{}, err
}
if !resolver.trusted.Match(peer) {
return peer, nil
}
chain, present, err := parseForwardedChain(request.Header.Values("Forwarded"))
if err != nil {
return netip.Addr{}, err
}
if !present {
chain, err = parseXForwardedFor(request.Header.Values("X-Forwarded-For"))
if err != nil {
return netip.Addr{}, err
}
}
if len(chain) == 0 {
return peer, nil
}
for index := len(chain) - 1; index >= 0; index-- {
if !resolver.trusted.Match(chain[index]) {
return chain[index], nil
}
}
if len(chain) > 0 {
return chain[0], nil
}
return peer, nil
}
func parseXForwardedFor(fields []string) ([]netip.Addr, error) {
chain := make([]netip.Addr, 0, len(fields)+1)
for _, field := range fields {
for value := range strings.SplitSeq(field, ",") {
address, err := netip.ParseAddr(strings.TrimSpace(value))
if err != nil {
return nil, fmt.Errorf("resolve client IP: invalid X-Forwarded-For address %q", value)
}
chain = append(chain, address.Unmap())
}
}
return chain, nil
}
func parseForwardedChain(fields []string) ([]netip.Addr, bool, error) {
if len(fields) == 0 {
return nil, false, nil
}
chain := make([]netip.Addr, 0, len(fields)+1)
for _, field := range fields {
for element := range strings.SplitSeq(field, ",") {
found := false
for parameter := range strings.SplitSeq(element, ";") {
name, value, ok := strings.Cut(strings.TrimSpace(parameter), "=")
if !ok || !strings.EqualFold(name, "for") {
continue
}
address, err := parseForwardedIdentifier(value)
if err != nil {
return nil, true, err
}
chain = append(chain, address)
found = true
break
}
if !found {
return nil, true, errors.New("resolve client IP: Forwarded element is missing for parameter")
}
}
}
return chain, true, nil
}
func parseForwardedIdentifier(raw string) (netip.Addr, error) {
value := strings.TrimSpace(raw)
if strings.HasPrefix(value, `"`) {
unquoted, err := strconv.Unquote(value)
if err != nil {
return netip.Addr{}, fmt.Errorf("resolve client IP: invalid quoted Forwarded identifier")
}
value = unquoted
}
if strings.EqualFold(value, "unknown") || strings.HasPrefix(value, "_") {
return netip.Addr{}, fmt.Errorf("resolve client IP: non-IP Forwarded identifier")
}
if strings.HasPrefix(value, "[") {
closing := strings.IndexByte(value, ']')
if closing < 0 {
return netip.Addr{}, fmt.Errorf("resolve client IP: invalid Forwarded IPv6 identifier")
}
value = value[1:closing]
} else if host, _, err := net.SplitHostPort(value); err == nil {
value = host
}
address, err := netip.ParseAddr(value)
if err != nil {
return netip.Addr{}, fmt.Errorf("resolve client IP: invalid Forwarded address")
}
return address.Unmap(), nil
}
type AccessGuard struct { type AccessGuard struct {
resolver *ClientIPResolver resolver *ClientIPResolver
@ -269,32 +83,3 @@ func (guard *AdmissionGuard) Check(ctx context.Context, request *http.Request) e
} }
return nil return nil
} }
func parseBasicCredentials(value string) (string, string, bool) {
scheme, encoded, ok := strings.Cut(strings.TrimSpace(value), " ")
if !ok || !strings.EqualFold(scheme, "Basic") {
return "", "", false
}
decoded, err := base64.StdEncoding.DecodeString(strings.TrimSpace(encoded))
if err != nil {
return "", "", false
}
username, password, ok := strings.Cut(string(decoded), ":")
return username, password, ok
}
func constantTimeEqual(actual, expected string) bool {
return subtle.ConstantTimeCompare([]byte(actual), []byte(expected)) == 1
}
func parseRemoteAddress(remote string) (netip.Addr, error) {
host, _, err := net.SplitHostPort(strings.TrimSpace(remote))
if err != nil {
host = strings.TrimSpace(remote)
}
address, err := netip.ParseAddr(host)
if err != nil {
return netip.Addr{}, fmt.Errorf("resolve client IP: invalid remote address %q", remote)
}
return address.Unmap(), nil
}

View File

@ -9,32 +9,6 @@ import (
"testing" "testing"
) )
func TestBasicAuthGuardUsesProxyAuthorizationAndReturnsChallenge(t *testing.T) {
t.Parallel()
guard := NewBasicAuthGuard("client", "secret")
allowed := httptest.NewRequest(http.MethodGet, "http://example.test", nil)
allowed.SetBasicAuth("ignored", "ignored")
allowed.Header.Set("Proxy-Authorization", "Basic Y2xpZW50OnNlY3JldA==")
if err := guard.Check(context.Background(), allowed); err != nil {
t.Fatalf("Check(valid) error = %v", err)
}
denied := httptest.NewRequest(http.MethodGet, "http://example.test", nil)
denied.Header.Set("Proxy-Authorization", "Basic Y2xpZW50Ondyb25n")
err := guard.Check(context.Background(), denied)
var httpError *HTTPError
if !errors.As(err, &httpError) {
t.Fatalf("Check(invalid) error = %T %v, want *HTTPError", err, err)
}
if httpError.StatusCode != http.StatusProxyAuthRequired {
t.Fatalf("status = %d, want 407", httpError.StatusCode)
}
if got := httpError.Header.Get("Proxy-Authenticate"); got != `Basic realm="proxy"` {
t.Fatalf("challenge = %q", got)
}
}
func TestClientIPResolverOnlyTrustsForwardedChainFromTrustedPeer(t *testing.T) { func TestClientIPResolverOnlyTrustsForwardedChainFromTrustedPeer(t *testing.T) {
t.Parallel() t.Parallel()

View File

@ -0,0 +1,242 @@
package httpsecurity
import (
"crypto/sha256"
"crypto/subtle"
"encoding/base64"
"encoding/hex"
"errors"
"net/http"
"net/netip"
"strings"
)
var (
errCredentialRejected = errors.New("credential rejected")
errSourceRejected = errors.New("source rejected")
)
type authenticator interface {
authenticate(*http.Request, string) (string, error)
challenges() []string
}
type noAuthenticator struct{}
func (noAuthenticator) authenticate(*http.Request, string) (string, error) { return "", nil }
func (noAuthenticator) challenges() []string { return nil }
type basicAuthenticator struct {
header string
username string
password string
}
func (auth basicAuthenticator) authenticate(request *http.Request, _ string) (string, error) {
value, ok := singleHeader(request, auth.header)
username, password, parsed := parseBasicCredentials(value)
valid := subtle.ConstantTimeSelect(boolInt(ok && parsed), 1, 0)
valid &= secureEqual(username, auth.username)
valid &= secureEqual(password, auth.password)
if valid != 1 {
return "", errCredentialRejected
}
return "basic:" + auth.username, nil
}
func (basicAuthenticator) challenges() []string { return []string{`Basic realm="proxy-pool"`} }
type tokenAuthenticator struct {
mode string
header string
token string
}
func (auth tokenAuthenticator) authenticate(request *http.Request, _ string) (string, error) {
value, ok := singleHeader(request, auth.header)
if auth.mode == ModeBearer {
value, ok = parseScheme(value, "Bearer", ok)
}
if !ok || secureEqual(value, auth.token) != 1 {
return "", errCredentialRejected
}
return credentialSubject(auth.mode, auth.token), nil
}
func (auth tokenAuthenticator) challenges() []string {
if auth.mode == ModeBearer {
return []string{`Bearer realm="proxy-pool"`}
}
return []string{`ApiKey realm="proxy-pool", header="` + auth.header + `"`}
}
type ipAuthenticator struct{ allowed cidrMatcher }
func (auth ipAuthenticator) authenticate(_ *http.Request, source string) (string, error) {
address, err := netip.ParseAddr(source)
if err != nil || !auth.allowed.match(address) {
return "", errSourceRejected
}
return "source:" + source, nil
}
func (ipAuthenticator) challenges() []string { return nil }
type anyAuthenticator struct{ methods []authenticator }
func (auth anyAuthenticator) authenticate(request *http.Request, source string) (string, error) {
sourceRejected := false
for _, method := range auth.methods {
principal, err := method.authenticate(request, source)
if err == nil {
return principal, nil
}
if errors.Is(err, errSourceRejected) {
sourceRejected = true
}
}
if sourceRejected {
return "", errSourceRejected
}
return "", errCredentialRejected
}
func (auth anyAuthenticator) challenges() []string {
var result []string
for _, method := range auth.methods {
result = append(result, method.challenges()...)
}
return result
}
func buildAuthenticator(authentication Authentication, semantics Semantics) (authenticator, error) {
header := "Authorization"
if semantics == ProxySemantics {
header = "Proxy-Authorization"
}
switch authentication.Mode {
case "", ModeNone:
return noAuthenticator{}, nil
case ModeUsernamePassword:
if authentication.Username == "" || authentication.Password == "" {
return nil, ErrInvalidConfig
}
return basicAuthenticator{header: header, username: authentication.Username, password: authentication.Password}, nil
case ModeAPIKey:
if !validHeaderName(authentication.Header) || authentication.Token == "" {
return nil, ErrInvalidConfig
}
return tokenAuthenticator{mode: ModeAPIKey, header: authentication.Header, token: authentication.Token}, nil
case ModeBearer:
if authentication.Token == "" {
return nil, ErrInvalidConfig
}
return tokenAuthenticator{mode: ModeBearer, header: header, token: authentication.Token}, nil
case ModeIPWhitelist:
allowed, err := newCIDRMatcher(authentication.CIDRs)
if err != nil || len(authentication.CIDRs) == 0 {
return nil, ErrInvalidConfig
}
return ipAuthenticator{allowed: allowed}, nil
case ModeAny:
if len(authentication.Methods) == 0 {
return nil, ErrInvalidConfig
}
methods := make([]authenticator, 0, len(authentication.Methods))
for _, method := range authentication.Methods {
candidate, err := buildMethod(method, semantics)
if err != nil {
return nil, err
}
methods = append(methods, candidate)
}
return anyAuthenticator{methods: methods}, nil
default:
return nil, ErrInvalidConfig
}
}
func buildMethod(method Method, semantics Semantics) (authenticator, error) {
authentication := Authentication{
Mode: method.Mode, Username: method.Username, Password: method.Password,
Header: method.Header, Token: method.Value, CIDRs: method.CIDRs,
}
if method.Mode == ModeAny || method.Mode == ModeNone || method.Mode == "" {
return nil, ErrInvalidConfig
}
return buildAuthenticator(authentication, semantics)
}
func parseBasicCredentials(value string) (string, string, bool) {
encoded, ok := parseScheme(value, "Basic", true)
if !ok {
return "", "", false
}
decoded, err := base64.StdEncoding.DecodeString(encoded)
if err != nil {
return "", "", false
}
username, password, ok := strings.Cut(string(decoded), ":")
return username, password, ok
}
func parseScheme(value, expected string, present bool) (string, bool) {
if !present {
return "", false
}
fields := strings.Fields(value)
if len(fields) != 2 || !strings.EqualFold(fields[0], expected) {
return "", false
}
return fields[1], true
}
func singleHeader(request *http.Request, name string) (string, bool) {
if request == nil {
return "", false
}
values := request.Header.Values(name)
return first(values), len(values) == 1 && values[0] != ""
}
func first(values []string) string {
if len(values) == 0 {
return ""
}
return values[0]
}
func secureEqual(actual, expected string) int {
actualHash := sha256.Sum256([]byte(actual))
expectedHash := sha256.Sum256([]byte(expected))
return subtle.ConstantTimeCompare(actualHash[:], expectedHash[:])
}
func credentialSubject(kind, credential string) string {
digest := sha256.Sum256([]byte(credential))
return kind + ":" + hex.EncodeToString(digest[:16])
}
func boolInt(value bool) int {
if value {
return 1
}
return 0
}
func validHeaderName(value string) bool {
if value == "" {
return false
}
for _, character := range []byte(value) {
if !isTokenCharacter(character) {
return false
}
}
return true
}
func isTokenCharacter(character byte) bool {
return character >= 'a' && character <= 'z' || character >= 'A' && character <= 'Z' ||
character >= '0' && character <= '9' || strings.ContainsRune("!#$%&'*+-.^_`|~", rune(character))
}

View File

@ -0,0 +1,54 @@
package httpsecurity
import (
"fmt"
"time"
"github.com/proxy-pool/proxy-pool/internal/config"
platformAdmission "github.com/proxy-pool/proxy-pool/internal/platform/admission"
)
func BuildFromListener(listener config.Listener, clientIdentification string, semantics Semantics) (*Protection, error) {
var admitter Admitter
if listener.Limits.RequestsPerMinute > 0 || listener.Limits.RequestsPerMinutePerClient > 0 {
limiter, err := platformAdmission.NewFixedWindow(platformAdmission.FixedWindowConfig{
Window: time.Minute,
Global: listener.Limits.RequestsPerMinute,
PerKey: listener.Limits.RequestsPerMinutePerClient,
})
if err != nil {
return nil, fmt.Errorf("build HTTP security admission: %w", err)
}
admitter = limiter
}
return NewFromListener(listener, clientIdentification, semantics, admitter)
}
func NewFromListener(listener config.Listener, clientIdentification string, semantics Semantics, admitter Admitter) (*Protection, error) {
methods := make([]Method, 0, len(listener.Auth.Methods))
for _, method := range listener.Auth.Methods {
methods = append(methods, Method{
Mode: method.Mode,
Username: method.Username,
Password: method.Password,
Header: method.Header,
Value: method.Value,
CIDRs: append([]string(nil), method.CIDRs...),
})
}
return New(Config{
TrustedProxies: append([]string(nil), listener.Access.TrustedProxies...),
AllowCIDRs: append([]string(nil), listener.Access.AllowCIDRs...),
Authentication: Authentication{
Mode: listener.Auth.Mode,
Username: listener.Auth.Username,
Password: listener.Auth.Password,
Header: listener.Auth.Header,
Token: listener.Auth.Token,
CIDRs: append([]string(nil), listener.Auth.CIDRs...),
Methods: methods,
},
ClientIdentification: clientIdentification,
Semantics: semantics,
}, admitter)
}

View File

@ -0,0 +1,61 @@
package httpsecurity
import (
"errors"
"net/http"
"net/http/httptest"
"testing"
"github.com/proxy-pool/proxy-pool/internal/config"
)
func TestNewFromListenerMapsResolvedConfiguration(t *testing.T) {
t.Parallel()
listener := config.Listener{
Access: config.Access{
AllowCIDRs: []string{"198.51.100.0/24"},
TrustedProxies: []string{"10.0.0.0/8"},
},
Auth: config.Auth{Mode: ModeAny, Methods: []config.AuthMethod{
{Mode: ModeBearer, Value: "bearer-secret"},
{Mode: ModeAPIKey, Header: "X-API-Key", Value: "api-secret"},
}},
}
protection, err := NewFromListener(listener, ClientAuthenticated, APIAuthSemantics, nil)
if err != nil {
t.Fatalf("NewFromListener() error = %v", err)
}
request := newRequest()
request.Header.Set("Authorization", "Bearer bearer-secret")
identity, err := protection.Resolve(request)
if err != nil {
t.Fatalf("Resolve() error = %v", err)
}
if identity.ClientID != credentialSubject("bearer", "bearer-secret") {
t.Fatalf("identity = %+v", identity)
}
if _, err := protection.Resolve(httptest.NewRequest(http.MethodGet, "/", nil)); err == nil {
t.Fatal("Resolve(request without remote address) error = nil")
}
}
func TestBuildFromListenerAppliesConfiguredAdmissionLimits(t *testing.T) {
t.Parallel()
listener := config.Listener{
Auth: config.Auth{Mode: ModeNone},
Limits: config.Limits{RequestsPerMinute: 1, RequestsPerMinutePerClient: 1},
}
protection, err := BuildFromListener(listener, ClientSourceIP, APIAuthSemantics)
if err != nil {
t.Fatalf("BuildFromListener() error = %v", err)
}
if _, err := protection.Resolve(newRequest()); err != nil {
t.Fatalf("Resolve(first) error = %v", err)
}
_, err = protection.Resolve(newRequest())
var httpError *HTTPError
if !errors.As(err, &httpError) || httpError.StatusCode != http.StatusTooManyRequests {
t.Fatalf("Resolve(second) error = %T %v", err, err)
}
}

View File

@ -0,0 +1,122 @@
package httpsecurity
import (
"context"
"errors"
"net/http"
)
type Protection struct {
resolver *ClientIPResolver
allow cidrMatcher
allowAll bool
authentication authenticator
clientMode string
semantics Semantics
admitter Admitter
}
func New(config Config, admitter Admitter) (*Protection, error) {
if config.Semantics != APIAuthSemantics && config.Semantics != ProxySemantics {
return nil, ErrInvalidConfig
}
resolver, err := NewClientIPResolver(config.TrustedProxies)
if err != nil {
return nil, ErrInvalidConfig
}
allow, err := newCIDRMatcher(config.AllowCIDRs)
if err != nil {
return nil, ErrInvalidConfig
}
authentication, err := buildAuthenticator(config.Authentication, config.Semantics)
if err != nil {
return nil, ErrInvalidConfig
}
clientMode := config.ClientIdentification
if clientMode == "" {
clientMode = ClientSourceIP
}
if clientMode != ClientSourceIP && clientMode != ClientAuthenticated && clientMode != ClientAuthenticatedOrSourceIP {
return nil, ErrInvalidConfig
}
if clientMode == ClientAuthenticated && (config.Authentication.Mode == "" || config.Authentication.Mode == ModeNone) {
return nil, ErrInvalidConfig
}
return &Protection{
resolver: resolver, allow: allow, allowAll: len(config.AllowCIDRs) == 0,
authentication: authentication, clientMode: clientMode, semantics: config.Semantics,
admitter: admitter,
}, nil
}
func (protection *Protection) Resolve(request *http.Request) (Identity, error) {
if request == nil {
return Identity{}, newHTTPError(http.StatusBadRequest, "INVALID_SOURCE", nil, errors.New("request is required"))
}
return protection.evaluate(request.Context(), request)
}
func (protection *Protection) Check(ctx context.Context, request *http.Request) error {
_, err := protection.evaluate(ctx, request)
return err
}
func (protection *Protection) evaluate(ctx context.Context, request *http.Request) (Identity, error) {
if protection == nil || protection.resolver == nil || protection.authentication == nil || request == nil {
return Identity{}, newHTTPError(http.StatusInternalServerError, "SECURITY_NOT_CONFIGURED", nil, ErrInvalidConfig)
}
address, err := protection.resolver.Resolve(request)
if err != nil {
return Identity{}, newHTTPError(http.StatusBadRequest, "INVALID_SOURCE", nil, err)
}
if !protection.allowAll && !protection.allow.match(address) {
return Identity{}, newHTTPError(http.StatusForbidden, "FORBIDDEN", nil, errSourceRejected)
}
source := address.String()
principal, err := protection.authentication.authenticate(request, source)
if err != nil {
if errors.Is(err, errSourceRejected) {
return Identity{}, newHTTPError(http.StatusForbidden, "FORBIDDEN", nil, err)
}
return Identity{}, protection.unauthorized(err)
}
identity := Identity{SourceIP: source}
switch protection.clientMode {
case ClientSourceIP:
identity.ClientID = "source:" + source
case ClientAuthenticated:
if principal == "" {
return Identity{}, protection.unauthorized(errCredentialRejected)
}
identity.ClientID = principal
case ClientAuthenticatedOrSourceIP:
identity.ClientID = principal
if identity.ClientID == "" {
identity.ClientID = "source:" + source
}
}
if protection.admitter != nil {
if err := protection.admitter.Admit(ctx, identity.ClientID); err != nil {
return Identity{}, newHTTPError(http.StatusTooManyRequests, "RATE_LIMITED", nil, err)
}
}
return identity, nil
}
func (protection *Protection) unauthorized(cause error) *HTTPError {
status := http.StatusUnauthorized
headerName := "WWW-Authenticate"
if protection.semantics == ProxySemantics {
status = http.StatusProxyAuthRequired
headerName = "Proxy-Authenticate"
}
header := make(http.Header)
for _, challenge := range protection.authentication.challenges() {
header.Add(headerName, challenge)
}
return newHTTPError(status, "UNAUTHORIZED", header, cause)
}
func newHTTPError(status int, code string, header http.Header, cause error) *HTTPError {
return &HTTPError{StatusCode: status, Code: code, Header: header, Cause: cause}
}

View File

@ -0,0 +1,36 @@
package httpsecurity
import (
"errors"
"net/http"
"github.com/proxy-pool/proxy-pool/internal/platform/httpapi"
)
func WriteProblem(writer http.ResponseWriter, requestID string, err error) bool {
var securityError *HTTPError
if !errors.As(err, &securityError) || securityError.StatusCode < 400 || securityError.StatusCode > 599 {
return false
}
for name, values := range securityError.Header {
for _, value := range values {
writer.Header().Add(name, value)
}
}
code := securityError.Code
if code == "" {
code = "SECURITY_REJECTED"
}
title := http.StatusText(securityError.StatusCode)
if title == "" {
title = "Request rejected"
}
httpapi.WriteProblem(writer, httpapi.NewProblem(
securityError.StatusCode,
code,
title,
"",
requestID,
))
return true
}

View File

@ -0,0 +1,50 @@
package httpsecurity
import (
"errors"
"net/http"
"net/http/httptest"
"strings"
"testing"
)
func TestWriteProblemUsesSafeFallbackForCustomStatus(t *testing.T) {
t.Parallel()
recorder := httptest.NewRecorder()
err := &HTTPError{
StatusCode: 599,
Header: http.Header{"X-Security-Policy": []string{"local"}},
Cause: errors.New("secret backend details"),
}
if !WriteProblem(recorder, "request-123", err) {
t.Fatal("WriteProblem() = false, want true")
}
if recorder.Code != 599 {
t.Fatalf("status = %d, want 599", recorder.Code)
}
if got := recorder.Header().Get("X-Security-Policy"); got != "local" {
t.Fatalf("X-Security-Policy = %q, want local", got)
}
body := recorder.Body.String()
for _, want := range []string{`"code":"SECURITY_REJECTED"`, `"title":"Request rejected"`, `"requestId":"request-123"`} {
if !strings.Contains(body, want) {
t.Fatalf("body = %s, want %s", body, want)
}
}
if strings.Contains(body, "secret backend details") {
t.Fatalf("body leaked cause: %s", body)
}
}
func TestWriteProblemIgnoresNonSecurityErrors(t *testing.T) {
t.Parallel()
recorder := httptest.NewRecorder()
if WriteProblem(recorder, "request-123", errors.New("plain error")) {
t.Fatal("WriteProblem() = true, want false")
}
if recorder.Code != http.StatusOK || recorder.Body.Len() != 0 {
t.Fatalf("recorder = status %d body %q, want untouched", recorder.Code, recorder.Body.String())
}
}

View File

@ -0,0 +1,329 @@
package httpsecurity
import (
"context"
"errors"
"net/http"
"net/http/httptest"
"sync"
"sync/atomic"
"testing"
)
func TestProtectionAuthenticatesAPIRequestsAndBuildsStableIdentity(t *testing.T) {
t.Parallel()
tests := []struct {
name string
auth Authentication
configure func(*http.Request)
wantClient string
}{
{
name: "basic",
auth: Authentication{Mode: ModeUsernamePassword, Username: "alice", Password: "secret"},
configure: func(request *http.Request) {
request.SetBasicAuth("alice", "secret")
},
wantClient: "basic:alice",
},
{
name: "api key",
auth: Authentication{Mode: ModeAPIKey, Header: "X-API-Key", Token: "api-secret"},
configure: func(request *http.Request) {
request.Header.Set("X-API-Key", "api-secret")
},
wantClient: credentialSubject("apiKey", "api-secret"),
},
{
name: "bearer",
auth: Authentication{Mode: ModeBearer, Token: "bearer-secret"},
configure: func(request *http.Request) {
request.Header.Set("Authorization", "Bearer bearer-secret")
},
wantClient: credentialSubject("bearer", "bearer-secret"),
},
}
for _, test := range tests {
test := test
t.Run(test.name, func(t *testing.T) {
t.Parallel()
protection := mustProtection(t, Config{
Authentication: test.auth,
ClientIdentification: ClientAuthenticated,
}, nil)
request := newRequest()
test.configure(request)
identity, err := protection.Resolve(request)
if err != nil {
t.Fatalf("Resolve() error = %v", err)
}
if identity.ClientID != test.wantClient || identity.SourceIP != "198.51.100.8" {
t.Fatalf("identity = %+v", identity)
}
})
}
}
func TestProtectionSupportsAnyAuthentication(t *testing.T) {
t.Parallel()
protection := mustProtection(t, Config{
Authentication: Authentication{Mode: ModeAny, Methods: []Method{
{Mode: ModeUsernamePassword, Username: "alice", Password: "wrong-for-request"},
{Mode: ModeAPIKey, Header: "X-API-Key", Value: "api-secret"},
}},
ClientIdentification: ClientAuthenticatedOrSourceIP,
}, nil)
request := newRequest()
request.Header.Set("X-API-Key", "api-secret")
identity, err := protection.Resolve(request)
if err != nil {
t.Fatalf("Resolve() error = %v", err)
}
if identity.ClientID != credentialSubject("apiKey", "api-secret") {
t.Fatalf("client ID = %q", identity.ClientID)
}
}
func TestProtectionAnyPreservesSourceRejection(t *testing.T) {
t.Parallel()
tests := []struct {
name string
methods []Method
}{
{
name: "IP whitelist only",
methods: []Method{{Mode: ModeIPWhitelist, CIDRs: []string{"10.0.0.0/8"}}},
},
{
name: "IP whitelist and bearer",
methods: []Method{
{Mode: ModeIPWhitelist, CIDRs: []string{"10.0.0.0/8"}},
{Mode: ModeBearer, Value: "bearer-secret"},
},
},
}
for _, test := range tests {
test := test
t.Run(test.name, func(t *testing.T) {
t.Parallel()
protection := mustProtection(t, Config{
Authentication: Authentication{Mode: ModeAny, Methods: test.methods},
}, nil)
_, err := protection.Resolve(newRequest())
var httpError *HTTPError
if !errors.As(err, &httpError) || httpError.StatusCode != http.StatusForbidden || httpError.Code != "FORBIDDEN" {
t.Fatalf("Resolve() error = %+v, want 403 FORBIDDEN", httpError)
}
})
}
}
func TestProtectionReturnsSafeHTTPFailures(t *testing.T) {
t.Parallel()
tests := []struct {
name string
config Config
configure func(*http.Request)
admitter Admitter
wantStatus int
wantCode string
wantHeader string
}{
{
name: "unauthorized basic",
config: Config{Authentication: Authentication{Mode: ModeUsernamePassword, Username: "alice", Password: "secret"}},
wantStatus: http.StatusUnauthorized,
wantCode: "UNAUTHORIZED",
wantHeader: `Basic realm="proxy-pool"`,
},
{
name: "unauthorized API key",
config: Config{Authentication: Authentication{Mode: ModeAPIKey, Header: "X-API-Key", Token: "secret"}},
wantStatus: http.StatusUnauthorized,
wantCode: "UNAUTHORIZED",
wantHeader: `ApiKey realm="proxy-pool", header="X-API-Key"`,
},
{
name: "forbidden source",
config: Config{Authentication: Authentication{Mode: ModeNone}, AllowCIDRs: []string{"10.0.0.0/8"}},
wantStatus: http.StatusForbidden,
wantCode: "FORBIDDEN",
},
{
name: "invalid forwarded chain",
config: Config{Authentication: Authentication{Mode: ModeNone}, TrustedProxies: []string{"10.0.0.0/8"}},
configure: func(request *http.Request) {
request.RemoteAddr = "10.0.0.1:1"
request.Header.Set("Forwarded", "for=unknown")
},
wantStatus: http.StatusBadRequest,
wantCode: "INVALID_SOURCE",
},
{
name: "invalid forwarded IPv6 suffix",
config: Config{Authentication: Authentication{Mode: ModeNone}, TrustedProxies: []string{"10.0.0.0/8"}},
configure: func(request *http.Request) {
request.RemoteAddr = "10.0.0.1:1"
request.Header.Set("Forwarded", `for="[2001:db8::7]junk"`)
},
wantStatus: http.StatusBadRequest,
wantCode: "INVALID_SOURCE",
},
{
name: "rate limited",
config: Config{Authentication: Authentication{Mode: ModeNone}},
admitter: rejectAdmitter{},
wantStatus: http.StatusTooManyRequests,
wantCode: "RATE_LIMITED",
},
}
for _, test := range tests {
test := test
t.Run(test.name, func(t *testing.T) {
t.Parallel()
protection := mustProtection(t, test.config, test.admitter)
request := newRequest()
if test.configure != nil {
test.configure(request)
}
_, err := protection.Resolve(request)
var httpError *HTTPError
if !errors.As(err, &httpError) {
t.Fatalf("Resolve() error = %T %v, want *HTTPError", err, err)
}
if httpError.StatusCode != test.wantStatus || httpError.Code != test.wantCode {
t.Fatalf("HTTP error = %+v", httpError)
}
if test.wantHeader != "" && httpError.Header.Get("WWW-Authenticate") != test.wantHeader {
t.Fatalf("challenge = %q", httpError.Header.Get("WWW-Authenticate"))
}
if strings := err.Error(); strings == "" || strings == "secret" {
t.Fatalf("unsafe error = %q", strings)
}
})
}
}
func TestProtectionOnlyTrustsForwardedHeadersFromConfiguredPeers(t *testing.T) {
t.Parallel()
protection := mustProtection(t, Config{
Authentication: Authentication{Mode: ModeNone},
TrustedProxies: []string{"10.0.0.0/8", "192.0.2.0/24"},
ClientIdentification: ClientSourceIP,
}, nil)
request := newRequest()
request.RemoteAddr = "10.0.0.9:1234"
request.Header.Set("X-Forwarded-For", "198.51.100.7, 192.0.2.5")
identity, err := protection.Resolve(request)
if err != nil {
t.Fatalf("Resolve(trusted) error = %v", err)
}
if identity.SourceIP != "198.51.100.7" || identity.ClientID != "source:198.51.100.7" {
t.Fatalf("trusted identity = %+v", identity)
}
request.RemoteAddr = "203.0.113.9:1234"
identity, err = protection.Resolve(request)
if err != nil {
t.Fatalf("Resolve(untrusted) error = %v", err)
}
if identity.SourceIP != "203.0.113.9" {
t.Fatalf("untrusted identity = %+v", identity)
}
}
func TestProtectionUsesProxyAuthenticationSemantics(t *testing.T) {
t.Parallel()
protection := mustProtection(t, Config{
Authentication: Authentication{Mode: ModeBearer, Token: "proxy-token"},
Semantics: ProxySemantics,
}, nil)
request := newRequest()
request.Header.Set("Authorization", "Bearer proxy-token")
if _, err := protection.Resolve(request); err == nil {
t.Fatal("Resolve(Authorization) error = nil, want proxy authentication failure")
}
request.Header.Del("Authorization")
request.Header.Set("Proxy-Authorization", "Bearer proxy-token")
if _, err := protection.Resolve(request); err != nil {
t.Fatalf("Resolve(Proxy-Authorization) error = %v", err)
}
request.Header.Set("Proxy-Authorization", "Bearer wrong")
_, err := protection.Resolve(request)
var httpError *HTTPError
if !errors.As(err, &httpError) || httpError.StatusCode != http.StatusProxyAuthRequired ||
httpError.Header.Get("Proxy-Authenticate") != `Bearer realm="proxy-pool"` {
t.Fatalf("proxy error = %+v", httpError)
}
}
func TestNewProtectionRejectsInvalidConfiguration(t *testing.T) {
t.Parallel()
tests := []Config{
{Authentication: Authentication{Mode: "unknown"}},
{Authentication: Authentication{Mode: ModeAPIKey, Header: "Bad Header", Token: "secret"}},
{Authentication: Authentication{Mode: ModeNone}, TrustedProxies: []string{"invalid"}},
{Authentication: Authentication{Mode: ModeNone}, ClientIdentification: "unknown"},
{Authentication: Authentication{Mode: ModeNone}, ClientIdentification: ClientAuthenticated},
}
for _, config := range tests {
if _, err := New(config, nil); !errors.Is(err, ErrInvalidConfig) {
t.Fatalf("New(%+v) error = %v, want %v", config, err, ErrInvalidConfig)
}
}
}
func TestProtectionResolvesConcurrentRequestsWithoutSharedMutation(t *testing.T) {
t.Parallel()
protection := mustProtection(t, Config{
Authentication: Authentication{Mode: ModeAPIKey, Header: "X-API-Key", Token: "api-secret"},
ClientIdentification: ClientAuthenticated,
}, nil)
var failures atomic.Int64
var wait sync.WaitGroup
for range 1000 {
wait.Add(1)
go func() {
defer wait.Done()
request := newRequest()
request.Header.Set("X-API-Key", "api-secret")
identity, err := protection.Resolve(request)
if err != nil || identity.ClientID != credentialSubject("apiKey", "api-secret") {
failures.Add(1)
}
}()
}
wait.Wait()
if failures.Load() != 0 {
t.Fatalf("concurrent resolve failures = %d", failures.Load())
}
}
func mustProtection(t *testing.T, config Config, admitter Admitter) *Protection {
t.Helper()
protection, err := New(config, admitter)
if err != nil {
t.Fatalf("New() error = %v", err)
}
return protection
}
func newRequest() *http.Request {
request := httptest.NewRequest(http.MethodGet, "http://example.test", nil)
request.RemoteAddr = "198.51.100.8:1234"
return request
}
type rejectAdmitter struct{}
func (rejectAdmitter) Admit(context.Context, string) error { return errors.New("backend details") }

View File

@ -0,0 +1,167 @@
package httpsecurity
import (
"errors"
"fmt"
"net"
"net/http"
"net/netip"
"strconv"
"strings"
)
type cidrMatcher struct {
prefixes []netip.Prefix
}
func newCIDRMatcher(values []string) (cidrMatcher, error) {
matcher := cidrMatcher{prefixes: make([]netip.Prefix, 0, len(values))}
for _, value := range values {
prefix, err := netip.ParsePrefix(strings.TrimSpace(value))
if err != nil {
return cidrMatcher{}, fmt.Errorf("%w: invalid CIDR", ErrInvalidConfig)
}
matcher.prefixes = append(matcher.prefixes, prefix.Masked())
}
return matcher, nil
}
func (matcher cidrMatcher) match(address netip.Addr) bool {
address = address.Unmap()
for _, prefix := range matcher.prefixes {
if prefix.Contains(address) {
return true
}
}
return false
}
type ClientIPResolver struct {
trusted cidrMatcher
}
func NewClientIPResolver(trustedCIDRs []string) (*ClientIPResolver, error) {
trusted, err := newCIDRMatcher(trustedCIDRs)
if err != nil {
return nil, err
}
return &ClientIPResolver{trusted: trusted}, nil
}
func (resolver *ClientIPResolver) Resolve(request *http.Request) (netip.Addr, error) {
if resolver == nil || request == nil {
return netip.Addr{}, errors.New("client IP resolver and request are required")
}
peer, err := parseRemoteAddress(request.RemoteAddr)
if err != nil {
return netip.Addr{}, err
}
if !resolver.trusted.match(peer) {
return peer, nil
}
chain, present, err := parseForwardedChain(request.Header.Values("Forwarded"))
if err != nil {
return netip.Addr{}, err
}
if !present {
chain, err = parseXForwardedFor(request.Header.Values("X-Forwarded-For"))
if err != nil {
return netip.Addr{}, err
}
}
if len(chain) == 0 {
return peer, nil
}
for index := len(chain) - 1; index >= 0; index-- {
if !resolver.trusted.match(chain[index]) {
return chain[index], nil
}
}
return chain[0], nil
}
func parseRemoteAddress(remote string) (netip.Addr, error) {
host, _, err := net.SplitHostPort(strings.TrimSpace(remote))
if err != nil {
host = strings.TrimSpace(remote)
}
address, err := netip.ParseAddr(host)
if err != nil {
return netip.Addr{}, errors.New("invalid remote address")
}
return address.Unmap(), nil
}
func parseXForwardedFor(fields []string) ([]netip.Addr, error) {
chain := make([]netip.Addr, 0, len(fields)+1)
for _, field := range fields {
for value := range strings.SplitSeq(field, ",") {
address, err := netip.ParseAddr(strings.TrimSpace(value))
if err != nil {
return nil, errors.New("invalid X-Forwarded-For address")
}
chain = append(chain, address.Unmap())
}
}
return chain, nil
}
func parseForwardedChain(fields []string) ([]netip.Addr, bool, error) {
if len(fields) == 0 {
return nil, false, nil
}
chain := make([]netip.Addr, 0, len(fields)+1)
for _, field := range fields {
for element := range strings.SplitSeq(field, ",") {
found := false
for parameter := range strings.SplitSeq(element, ";") {
name, value, ok := strings.Cut(strings.TrimSpace(parameter), "=")
if !ok || !strings.EqualFold(name, "for") {
continue
}
address, err := parseForwardedIdentifier(value)
if err != nil {
return nil, true, err
}
chain = append(chain, address)
found = true
break
}
if !found {
return nil, true, errors.New("Forwarded element is missing for parameter")
}
}
}
return chain, true, nil
}
func parseForwardedIdentifier(raw string) (netip.Addr, error) {
value := strings.TrimSpace(raw)
if strings.HasPrefix(value, `"`) {
unquoted, err := strconv.Unquote(value)
if err != nil {
return netip.Addr{}, errors.New("invalid quoted Forwarded identifier")
}
value = unquoted
}
if strings.EqualFold(value, "unknown") || strings.HasPrefix(value, "_") {
return netip.Addr{}, errors.New("non-IP Forwarded identifier")
}
if strings.HasPrefix(value, "[") {
if addressPort, err := netip.ParseAddrPort(value); err == nil {
return addressPort.Addr().Unmap(), nil
}
if !strings.HasSuffix(value, "]") {
return netip.Addr{}, errors.New("invalid Forwarded IPv6 identifier")
}
value = strings.TrimSuffix(strings.TrimPrefix(value, "["), "]")
} else if host, _, err := net.SplitHostPort(value); err == nil {
value = host
}
address, err := netip.ParseAddr(value)
if err != nil {
return netip.Addr{}, errors.New("invalid Forwarded address")
}
return address.Unmap(), nil
}

View File

@ -0,0 +1,86 @@
package httpsecurity
import (
"context"
"errors"
"net/http"
)
const (
ModeNone = "none"
ModeUsernamePassword = "usernamePassword"
ModeAPIKey = "apiKey"
ModeBearer = "bearer"
ModeIPWhitelist = "ipWhitelist"
ModeAny = "any"
ClientSourceIP = "sourceIP"
ClientAuthenticated = "authenticatedClient"
ClientAuthenticatedOrSourceIP = "authenticatedClientOrSourceIP"
)
const (
APIAuthSemantics Semantics = iota
ProxySemantics
)
var ErrInvalidConfig = errors.New("invalid HTTP security configuration")
type Semantics uint8
type Authentication struct {
Mode string
Username string
Password string
Header string
Token string
CIDRs []string
Methods []Method
}
type Method struct {
Mode string
Username string
Password string
Header string
Value string
CIDRs []string
}
type Config struct {
TrustedProxies []string
AllowCIDRs []string
Authentication Authentication
ClientIdentification string
Semantics Semantics
}
type Identity struct {
ClientID string
SourceIP string
}
type Admitter interface {
Admit(context.Context, string) error
}
type HTTPError struct {
StatusCode int
Code string
Header http.Header
Cause error
}
func (err *HTTPError) Error() string {
if err == nil || err.Code == "" {
return "HTTP security request rejected"
}
return err.Code
}
func (err *HTTPError) Unwrap() error {
if err == nil {
return nil
}
return err.Cause
}