增加工具模块与可信代理配置

- internal/utils:ClientIP 按序枚举 CDN/代理头(含 RFC 7239 Forwarded),仅在可信代理来源时采信,否则回退直连 IP;RemoteIP 取直连地址;RandomString 生成安全随机串
- 新增 server.trusted_proxies 配置(IP/CIDR,ConfigVersion 2→3 自动补全),启动时同步应用到 gin 与 utils
- 初始管理员密码生成改用 utils.RandomString,原密码测试迁至 utils
This commit is contained in:
2026-09-21 20:28:38 +08:00
parent 6949b12af0
commit f440bb22b6
10 files changed
+453 -31

No files matched your search

+2 -1
View File
@@ -1,11 +1,12 @@
# rill 服务端配置
version: 2 # 配置版本,用于启动时自动补全缺失项,请勿手动修改
version: 3 # 配置版本,用于启动时自动补全缺失项,请勿手动修改
server:
host: "0.0.0.0" # 监听地址,0.0.0.0 表示所有网卡
port: 8080 #web 服务端口,为 0 则不使用 tcp html
sock: "web.sock" # unix socket 文件路径,留空表示不启用
mode: release # gin 运行模式: debug / release / test
trusted_proxies: [] # 可信代理/CDN 回源网段(IP 或 CIDR),仅这些来源的转发头会被采信;留空表示不信任任何代理
log:
+20 -8
View File
@@ -6,6 +6,7 @@ import (
"fmt"
"io/fs"
"log/slog"
"net"
"os"
"path/filepath"
"strings"
@@ -28,10 +29,11 @@ type Config struct {
}
type ServerConfig struct {
Host string `yaml:"host"`
Port int `yaml:"port"`
Sock string `yaml:"sock"`
Mode string `yaml:"mode"`
Host string `yaml:"host"`
Port int `yaml:"port"`
Sock string `yaml:"sock"`
Mode string `yaml:"mode"`
TrustedProxies []string `yaml:"trusted_proxies"`
}
// TCPEnabled 是否启用 TCP 监听(port 为 0 表示不启用)。
@@ -98,10 +100,11 @@ type MySQLConfig struct {
func defaultConfig() *Config {
return &Config{
Server: ServerConfig{
Host: "0.0.0.0",
Port: 8080,
Sock: "web.sock",
Mode: "release",
Host: "0.0.0.0",
Port: 8080,
Sock: "web.sock",
Mode: "release",
TrustedProxies: []string{},
},
Log: LogConfig{
Level: "info",
@@ -197,6 +200,15 @@ func (c *Config) validate() error {
if !c.Server.TCPEnabled() && !c.Server.SockEnabled() {
return fmt.Errorf("server.port 与 server.sock 至少需要启用一个")
}
for _, proxy := range c.Server.TrustedProxies {
value := strings.TrimSpace(proxy)
if value == "" || net.ParseIP(value) != nil {
continue
}
if _, _, err := net.ParseCIDR(value); err != nil {
return fmt.Errorf("server.trusted_proxies 无效: %q(需为 IP 或 CIDR", proxy)
}
}
if _, err := parseLogLevel(c.Log.Level); err != nil {
return err
}
+1
View File
@@ -44,6 +44,7 @@ custom:
`prefix: "/api"`,
"auth:",
`token_ttl: "24h"`,
"trusted_proxies",
} {
if !strings.Contains(out, want) {
t.Errorf("补全结果缺少 %q\n---\n%s", want, out)
+1 -1
View File
@@ -15,7 +15,7 @@ import (
)
// ConfigVersion 当前配置结构版本,新增配置项时递增。
const ConfigVersion = 2
const ConfigVersion = 3
// upgradeResult 描述一次配置自动补全的结果。
type upgradeResult struct {
+2 -21
View File
@@ -1,18 +1,16 @@
package database
import (
"crypto/rand"
"fmt"
"log/slog"
"math/big"
"os"
"path/filepath"
"strings"
"golang.org/x/crypto/bcrypt"
"gorm.io/gorm"
"rill/internal/model"
"rill/internal/utils"
)
const (
@@ -22,9 +20,6 @@ const (
adminPasswordFilename = "admin_password.txt"
)
// adminPasswordCharset 去掉了易混淆字符(0/O、1/l/I)。
const adminPasswordCharset = "abcdefghijkmnpqrstuvwxyzABCDEFGHJKLMNPQRSTUVWXYZ23456789"
// seedAdminUser 创建初始管理员用户并加入 admin 组,密码随机生成且仅在迁移时打印一次。
func seedAdminUser(tx *gorm.DB) error {
var count int64
@@ -36,7 +31,7 @@ func seedAdminUser(tx *gorm.DB) error {
return nil
}
password, err := generatePassword(adminPasswordLen)
password, err := utils.RandomString(adminPasswordLen)
if err != nil {
return err
}
@@ -103,17 +98,3 @@ func writeAdminPasswordFile(path, password string) error {
)
return os.WriteFile(path, []byte(content), 0o600)
}
func generatePassword(length int) (string, error) {
limit := big.NewInt(int64(len(adminPasswordCharset)))
var builder strings.Builder
builder.Grow(length)
for i := 0; i < length; i++ {
n, err := rand.Int(rand.Reader, limit)
if err != nil {
return "", fmt.Errorf("生成随机密码失败: %w", err)
}
builder.WriteByte(adminPasswordCharset[n.Int64()])
}
return builder.String(), nil
}
+184
View File
@@ -0,0 +1,184 @@
// Package utils 提供通用工具函数:客户端 IP 解析、随机字符串等。
package utils
import (
"fmt"
"net"
"net/http"
"strings"
"sync"
"github.com/gin-gonic/gin"
)
// clientIPHeaders 按优先级排列的单值客户端 IP 请求头;Forwarded 与 X-Forwarded-For 因需要链式解析单独处理。
var clientIPHeaders = []string{
"CF-Connecting-IP",
"True-Client-IP",
"Ali-CDN-Real-IP",
"X-Real-IP",
"X-Client-IP",
"Fastly-Client-IP",
}
var (
trustedMu sync.RWMutex
trustedCIDRs []*net.IPNet
)
// SetTrustedProxies 设置可信代理/CDN 回源网段,支持 IP 与 CIDR,启动时调用一次;并发安全。
func SetTrustedProxies(cidrs []string) error {
parsed := make([]*net.IPNet, 0, len(cidrs))
for _, raw := range cidrs {
value := strings.TrimSpace(raw)
if value == "" {
continue
}
if ip := net.ParseIP(value); ip != nil {
bits := 128
if ip.To4() != nil {
bits = 32
}
_, cidr, err := net.ParseCIDR(fmt.Sprintf("%s/%d", ip.String(), bits))
if err != nil {
return fmt.Errorf("invalid trusted proxy %q: %w", raw, err)
}
parsed = append(parsed, cidr)
continue
}
if _, cidr, err := net.ParseCIDR(value); err == nil {
parsed = append(parsed, cidr)
continue
}
return fmt.Errorf("invalid trusted proxy %q", raw)
}
trustedMu.Lock()
trustedCIDRs = parsed
trustedMu.Unlock()
return nil
}
// ClientIP 依次枚举 CDN/代理请求头,仅在直连地址可信时采信;否则返回直连 IP。
func ClientIP(c *gin.Context) string {
direct := RemoteIP(c)
if !isTrustedPeer(c, direct) {
return direct
}
for _, header := range clientIPHeaders {
if ip := normalizeIP(c.GetHeader(header)); ip != "" {
return ip
}
}
if ip := pickClientIP(parseForwarded(c.GetHeader("Forwarded"))); ip != "" {
return ip
}
if ip := pickClientIP(parseIPList(c.GetHeader("X-Forwarded-For"))); ip != "" {
return ip
}
return direct
}
// RemoteIP 返回直连地址(已去端口),不读取任何请求头。
func RemoteIP(c *gin.Context) string {
return normalizeIP(c.Request.RemoteAddr)
}
// pickClientIP 从右往左取第一个非可信代理 IP;全部可信时取最左值。
func pickClientIP(chain []string) string {
for i := len(chain) - 1; i >= 0; i-- {
if !isTrustedStringIP(chain[i]) {
return chain[i]
}
}
if len(chain) > 0 {
return chain[0]
}
return ""
}
func isTrustedPeer(c *gin.Context, direct string) bool {
if isUnixSocket(c) {
return true
}
ip := net.ParseIP(direct)
return ip != nil && isTrustedIP(ip)
}
func isTrustedIP(ip net.IP) bool {
trustedMu.RLock()
defer trustedMu.RUnlock()
for _, cidr := range trustedCIDRs {
if cidr.Contains(ip) {
return true
}
}
return false
}
func isTrustedStringIP(value string) bool {
ip := net.ParseIP(value)
return ip != nil && isTrustedIP(ip)
}
func isUnixSocket(c *gin.Context) bool {
addr, ok := c.Request.Context().Value(http.LocalAddrContextKey).(net.Addr)
return ok && strings.HasPrefix(addr.Network(), "unix")
}
// normalizeIP 清理引号与端口后校验为合法 IP,非法返回空串。
func normalizeIP(raw string) string {
value := strings.TrimSpace(raw)
if value == "" {
return ""
}
value = strings.Trim(value, `"`)
if host, _, err := net.SplitHostPort(value); err == nil {
value = host
}
value = strings.Trim(value, "[]")
ip := net.ParseIP(value)
if ip == nil {
return ""
}
return ip.String()
}
func parseIPList(value string) []string {
if strings.TrimSpace(value) == "" {
return nil
}
parts := strings.Split(value, ",")
ips := make([]string, 0, len(parts))
for _, part := range parts {
if ip := normalizeIP(part); ip != "" {
ips = append(ips, ip)
}
}
return ips
}
// parseForwarded 解析 RFC 7239 Forwarded 头中的 for= 参数,按出现顺序返回合法 IP。
func parseForwarded(value string) []string {
if strings.TrimSpace(value) == "" {
return nil
}
var ips []string
for _, element := range strings.Split(value, ",") {
for _, param := range strings.Split(element, ";") {
key, val, ok := strings.Cut(strings.TrimSpace(param), "=")
if !ok || !strings.EqualFold(strings.TrimSpace(key), "for") {
continue
}
candidate := strings.TrimSpace(val)
if candidate == "" || candidate == "_hidden" || strings.EqualFold(candidate, "unknown") {
continue
}
if ip := normalizeIP(candidate); ip != "" {
ips = append(ips, ip)
}
}
}
return ips
}
+178
View File
@@ -0,0 +1,178 @@
package utils
import (
"context"
"net"
"net/http"
"net/http/httptest"
"os"
"testing"
"github.com/gin-gonic/gin"
)
func TestMain(m *testing.M) {
gin.SetMode(gin.TestMode)
os.Exit(m.Run())
}
func newContext(t *testing.T, remoteAddr string, headers map[string]string) *gin.Context {
t.Helper()
req := httptest.NewRequest(http.MethodGet, "/", nil)
req.RemoteAddr = remoteAddr
for key, value := range headers {
req.Header.Set(key, value)
}
w := httptest.NewRecorder()
c, _ := gin.CreateTestContext(w)
c.Request = req
return c
}
func setTrusted(t *testing.T, cidrs ...string) {
t.Helper()
if err := SetTrustedProxies(cidrs); err != nil {
t.Fatalf("SetTrustedProxies(%v) 失败: %v", cidrs, err)
}
t.Cleanup(func() {
_ = SetTrustedProxies(nil)
})
}
func TestClientIPUntrustedPeer(t *testing.T) {
setTrusted(t)
c := newContext(t, "192.0.2.10:5555", map[string]string{
"CF-Connecting-IP": "203.0.113.9",
"X-Real-IP": "203.0.113.8",
})
if got := ClientIP(c); got != "192.0.2.10" {
t.Errorf("ClientIP = %q, 期望直连地址 192.0.2.10", got)
}
}
func TestClientIPHeaderPriority(t *testing.T) {
setTrusted(t, "192.0.2.0/24")
cases := []struct {
name string
headers map[string]string
want string
}{
{"cloudflare", map[string]string{"CF-Connecting-IP": "203.0.113.1", "X-Real-IP": "203.0.113.2"}, "203.0.113.1"},
{"true-client", map[string]string{"True-Client-IP": "203.0.113.3", "X-Real-IP": "203.0.113.2"}, "203.0.113.3"},
{"ali-cdn", map[string]string{"Ali-CDN-Real-IP": "203.0.113.4"}, "203.0.113.4"},
{"real-ip", map[string]string{"X-Real-IP": "203.0.113.5"}, "203.0.113.5"},
{"client-ip", map[string]string{"X-Client-IP": "203.0.113.6"}, "203.0.113.6"},
{"fastly", map[string]string{"Fastly-Client-IP": "203.0.113.7"}, "203.0.113.7"},
}
for _, tc := range cases {
c := newContext(t, "192.0.2.10:5555", tc.headers)
if got := ClientIP(c); got != tc.want {
t.Errorf("%s: ClientIP = %q, 期望 %q", tc.name, got, tc.want)
}
}
}
func TestClientIPInvalidValueFallsThrough(t *testing.T) {
setTrusted(t, "192.0.2.10")
c := newContext(t, "192.0.2.10:5555", map[string]string{
"CF-Connecting-IP": "not-an-ip",
"X-Real-IP": "203.0.113.5",
})
if got := ClientIP(c); got != "203.0.113.5" {
t.Errorf("ClientIP = %q, 期望 203.0.113.5", got)
}
}
func TestClientIPXForwardedFor(t *testing.T) {
setTrusted(t, "192.0.2.0/24", "10.0.0.0/8")
c := newContext(t, "192.0.2.10:5555", map[string]string{
"X-Forwarded-For": "198.51.100.1, 10.0.0.5",
})
if got := ClientIP(c); got != "198.51.100.1" {
t.Errorf("应取最右侧非可信 IP, ClientIP = %q, 期望 198.51.100.1", got)
}
c = newContext(t, "192.0.2.10:5555", map[string]string{
"X-Forwarded-For": "10.0.0.1, 10.0.0.2",
})
if got := ClientIP(c); got != "10.0.0.1" {
t.Errorf("全部可信时应取最左, ClientIP = %q, 期望 10.0.0.1", got)
}
}
func TestClientIPForwarded(t *testing.T) {
setTrusted(t, "192.0.2.10")
c := newContext(t, "192.0.2.10:5555", map[string]string{
"Forwarded": `for=203.0.113.9;proto=https, for="[2001:db8::1]:4711"`,
})
if got := ClientIP(c); got != "2001:db8::1" {
t.Errorf("ClientIP = %q, 期望 2001:db8::1", got)
}
}
func TestClientIPNormalization(t *testing.T) {
setTrusted(t, "192.0.2.10")
cases := []struct {
header string
value string
want string
}{
{"CF-Connecting-IP", "192.0.2.99:443", "192.0.2.99"},
{"X-Real-IP", `"198.51.100.7"`, "198.51.100.7"},
{"X-Real-IP", " 198.51.100.8 ", "198.51.100.8"},
{"X-Real-IP", "[2001:db8::2]:8080", "2001:db8::2"},
{"X-Real-IP", "2001:db8::3", "2001:db8::3"},
}
for _, tc := range cases {
c := newContext(t, "192.0.2.10:5555", map[string]string{tc.header: tc.value})
if got := ClientIP(c); got != tc.want {
t.Errorf("%s=%q: ClientIP = %q, 期望 %q", tc.header, tc.value, got, tc.want)
}
}
}
func TestClientIPUnixSocketTrusted(t *testing.T) {
setTrusted(t)
req := httptest.NewRequest(http.MethodGet, "/", nil)
req.RemoteAddr = "@"
req = req.WithContext(context.WithValue(req.Context(), http.LocalAddrContextKey,
&net.UnixAddr{Name: "web.sock", Net: "unix"}))
req.Header.Set("CF-Connecting-IP", "203.0.113.9")
w := httptest.NewRecorder()
c, _ := gin.CreateTestContext(w)
c.Request = req
if got := ClientIP(c); got != "203.0.113.9" {
t.Errorf("unix socket 应视为可信, ClientIP = %q, 期望 203.0.113.9", got)
}
}
func TestSetTrustedProxies(t *testing.T) {
if err := SetTrustedProxies([]string{"192.0.2.0/24", "198.51.100.7", "2001:db8::/32", " "}); err != nil {
t.Fatalf("合法输入不应报错: %v", err)
}
t.Cleanup(func() {
_ = SetTrustedProxies(nil)
})
if err := SetTrustedProxies([]string{"not-a-cidr"}); err == nil {
t.Error("非法 CIDR 应报错")
}
}
func TestRemoteIP(t *testing.T) {
if got := RemoteIP(newContext(t, "192.0.2.1:1234", nil)); got != "192.0.2.1" {
t.Errorf("RemoteIP = %q, 期望 192.0.2.1", got)
}
if got := RemoteIP(newContext(t, "@", nil)); got != "" {
t.Errorf("unix socket RemoteIP = %q, 期望空串", got)
}
}
+26
View File
@@ -0,0 +1,26 @@
package utils
import (
"crypto/rand"
"fmt"
"math/big"
"strings"
)
// randomCharset 去掉了易混淆字符(0/O、1/l/I)。
const randomCharset = "abcdefghijkmnpqrstuvwxyzABCDEFGHJKLMNPQRSTUVWXYZ23456789"
// RandomString 生成指定长度的随机字符串。
func RandomString(length int) (string, error) {
limit := big.NewInt(int64(len(randomCharset)))
var builder strings.Builder
builder.Grow(length)
for i := 0; i < length; i++ {
n, err := rand.Int(rand.Reader, limit)
if err != nil {
return "", fmt.Errorf("生成随机字符串失败: %w", err)
}
builder.WriteByte(randomCharset[n.Int64()])
}
return builder.String(), nil
}
+29
View File
@@ -0,0 +1,29 @@
package utils
import (
"strings"
"testing"
)
func TestRandomString(t *testing.T) {
value, err := RandomString(16)
if err != nil {
t.Fatalf("生成随机字符串失败: %v", err)
}
if len(value) != 16 {
t.Errorf("长度 = %d, 期望 16", len(value))
}
for _, r := range value {
if !strings.ContainsRune(randomCharset, r) {
t.Errorf("包含非法字符 %q", r)
}
}
other, err := RandomString(16)
if err != nil {
t.Fatalf("生成随机字符串失败: %v", err)
}
if value == other {
t.Error("两次生成的字符串不应相同")
}
}