135 lines
3.2 KiB
Go
135 lines
3.2 KiB
Go
package config
|
|
|
|
import (
|
|
"fmt"
|
|
"os"
|
|
"regexp"
|
|
"strings"
|
|
|
|
"gopkg.in/yaml.v3"
|
|
)
|
|
|
|
// FilterMode 决定如何应用 patterns。
|
|
type FilterMode string
|
|
|
|
const (
|
|
FilterWhitelist FilterMode = "whitelist"
|
|
FilterBlacklist FilterMode = "blacklist"
|
|
FilterDisabled FilterMode = "disabled"
|
|
)
|
|
|
|
// Filter 表示路径匹配规则。
|
|
type Filter struct {
|
|
Mode FilterMode `yaml:"mode"`
|
|
Patterns []string `yaml:"patterns"`
|
|
|
|
compiled []*regexp.Regexp
|
|
}
|
|
|
|
// LoadFilter 从 yaml 文件加载过滤配置;文件不存在则默认全部记录。
|
|
//
|
|
// Pattern 语法(glob 风格):
|
|
// - `*` 匹配单个路径段内除 `/` 之外的任意字符(包括零个)。
|
|
// - `**` 匹配任意字符,含 `/`,可跨段。
|
|
// - `?` 匹配单个非 `/` 字符。
|
|
// - 其它字符按字面匹配。
|
|
//
|
|
// 示例:
|
|
// - `/v1/audio/*` 匹配 /v1/audio/speech、/v1/audio/transcriptions
|
|
// - `/v1/videos/**` 匹配 /v1/videos/任意子路径
|
|
// - `/v1beta/models/*:generateContent` 匹配 Gemini 风格端点
|
|
func LoadFilter(filePath string) (*Filter, error) {
|
|
data, err := os.ReadFile(filePath)
|
|
if err != nil {
|
|
if os.IsNotExist(err) {
|
|
return &Filter{Mode: FilterDisabled}, nil
|
|
}
|
|
return nil, err
|
|
}
|
|
var f Filter
|
|
if err := yaml.Unmarshal(data, &f); err != nil {
|
|
return nil, fmt.Errorf("parse filter yaml: %w", err)
|
|
}
|
|
switch f.Mode {
|
|
case FilterWhitelist, FilterBlacklist, FilterDisabled:
|
|
case "":
|
|
f.Mode = FilterDisabled
|
|
default:
|
|
return nil, fmt.Errorf("unknown filter mode: %q", f.Mode)
|
|
}
|
|
if err := f.compile(); err != nil {
|
|
return nil, err
|
|
}
|
|
return &f, nil
|
|
}
|
|
|
|
// NewFilter 程序化构造一个 Filter(主要供测试使用)。
|
|
func NewFilter(mode FilterMode, patterns []string) (*Filter, error) {
|
|
f := &Filter{Mode: mode, Patterns: append([]string(nil), patterns...)}
|
|
if err := f.compile(); err != nil {
|
|
return nil, err
|
|
}
|
|
return f, nil
|
|
}
|
|
|
|
func (f *Filter) compile() error {
|
|
f.compiled = make([]*regexp.Regexp, 0, len(f.Patterns))
|
|
for _, p := range f.Patterns {
|
|
re, err := CompileGlob(p)
|
|
if err != nil {
|
|
return fmt.Errorf("invalid pattern %q: %w", p, err)
|
|
}
|
|
f.compiled = append(f.compiled, re)
|
|
}
|
|
return nil
|
|
}
|
|
|
|
// ShouldLog 决定一个请求 path 是否需要被记录。
|
|
func (f *Filter) ShouldLog(reqPath string) bool {
|
|
if f == nil || f.Mode == FilterDisabled {
|
|
return true
|
|
}
|
|
matched := false
|
|
for _, re := range f.compiled {
|
|
if re.MatchString(reqPath) {
|
|
matched = true
|
|
break
|
|
}
|
|
}
|
|
switch f.Mode {
|
|
case FilterWhitelist:
|
|
return matched
|
|
case FilterBlacklist:
|
|
return !matched
|
|
default:
|
|
return true
|
|
}
|
|
}
|
|
|
|
// CompileGlob 将 glob 风格 pattern 转换为 anchored 正则表达式。
|
|
func CompileGlob(pattern string) (*regexp.Regexp, error) {
|
|
var sb strings.Builder
|
|
sb.WriteString("^")
|
|
for i := 0; i < len(pattern); i++ {
|
|
c := pattern[i]
|
|
switch c {
|
|
case '*':
|
|
if i+1 < len(pattern) && pattern[i+1] == '*' {
|
|
sb.WriteString(".*")
|
|
i++
|
|
} else {
|
|
sb.WriteString("[^/]*")
|
|
}
|
|
case '?':
|
|
sb.WriteString("[^/]")
|
|
case '.', '+', '(', ')', '|', '^', '$', '{', '}', '[', ']', '\\':
|
|
sb.WriteByte('\\')
|
|
sb.WriteByte(c)
|
|
default:
|
|
sb.WriteByte(c)
|
|
}
|
|
}
|
|
sb.WriteString("$")
|
|
return regexp.Compile(sb.String())
|
|
}
|