Files

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())
}