195 lines
4.9 KiB
Go
195 lines
4.9 KiB
Go
|
|
package config
|
||
|
|
|
||
|
|
import (
|
||
|
|
"bytes"
|
||
|
|
"errors"
|
||
|
|
"fmt"
|
||
|
|
"io"
|
||
|
|
"net"
|
||
|
|
"os"
|
||
|
|
"path/filepath"
|
||
|
|
"strconv"
|
||
|
|
"strings"
|
||
|
|
|
||
|
|
"gopkg.in/yaml.v3"
|
||
|
|
)
|
||
|
|
|
||
|
|
type Config struct {
|
||
|
|
ListenAddr string
|
||
|
|
DatabaseDSN string
|
||
|
|
SigningPrivateKeyFile string
|
||
|
|
SigningCertificateFile string
|
||
|
|
}
|
||
|
|
|
||
|
|
type fileConfig struct {
|
||
|
|
Server struct {
|
||
|
|
ListenAddr strictString `yaml:"listen_addr"`
|
||
|
|
} `yaml:"server"`
|
||
|
|
Database struct {
|
||
|
|
DSN strictString `yaml:"dsn"`
|
||
|
|
} `yaml:"database"`
|
||
|
|
Signing struct {
|
||
|
|
PrivateKeyFile strictString `yaml:"private_key_file"`
|
||
|
|
CertificateFile strictString `yaml:"certificate_file"`
|
||
|
|
} `yaml:"signing"`
|
||
|
|
}
|
||
|
|
|
||
|
|
type strictString string
|
||
|
|
|
||
|
|
func (value *strictString) UnmarshalYAML(node *yaml.Node) error {
|
||
|
|
if node.Kind != yaml.ScalarNode || node.Tag != "!!str" {
|
||
|
|
return errors.New("必须是字符串")
|
||
|
|
}
|
||
|
|
*value = strictString(node.Value)
|
||
|
|
return nil
|
||
|
|
}
|
||
|
|
|
||
|
|
func LoadFile(filePath string) (Config, error) {
|
||
|
|
if filePath == "" {
|
||
|
|
return Config{}, errors.New("配置文件路径不能为空")
|
||
|
|
}
|
||
|
|
|
||
|
|
content, err := os.ReadFile(filePath)
|
||
|
|
if err != nil {
|
||
|
|
return Config{}, fmt.Errorf("读取配置文件 %q: %w", filePath, err)
|
||
|
|
}
|
||
|
|
if err := rejectYAMLMergeKeys(content); err != nil {
|
||
|
|
return Config{}, err
|
||
|
|
}
|
||
|
|
|
||
|
|
var raw fileConfig
|
||
|
|
decoder := yaml.NewDecoder(bytes.NewReader(content))
|
||
|
|
decoder.KnownFields(true)
|
||
|
|
if err := decoder.Decode(&raw); err != nil {
|
||
|
|
if errors.Is(err, io.EOF) {
|
||
|
|
return Config{}, errors.New("配置文件不能为空")
|
||
|
|
}
|
||
|
|
return Config{}, fmt.Errorf("解析配置文件: %w", err)
|
||
|
|
}
|
||
|
|
|
||
|
|
var extra any
|
||
|
|
if err := decoder.Decode(&extra); err == nil {
|
||
|
|
return Config{}, errors.New("配置文件只允许包含一个 YAML 文档")
|
||
|
|
} else if !errors.Is(err, io.EOF) {
|
||
|
|
return Config{}, fmt.Errorf("解析配置文件的尾部内容: %w", err)
|
||
|
|
}
|
||
|
|
|
||
|
|
return raw.build()
|
||
|
|
}
|
||
|
|
|
||
|
|
func rejectYAMLMergeKeys(content []byte) error {
|
||
|
|
var document yaml.Node
|
||
|
|
decoder := yaml.NewDecoder(bytes.NewReader(content))
|
||
|
|
if err := decoder.Decode(&document); err != nil {
|
||
|
|
if errors.Is(err, io.EOF) {
|
||
|
|
return nil
|
||
|
|
}
|
||
|
|
return fmt.Errorf("解析配置文件: %w", err)
|
||
|
|
}
|
||
|
|
|
||
|
|
return rejectYAMLMergeNode(&document, make(map[*yaml.Node]struct{}))
|
||
|
|
}
|
||
|
|
|
||
|
|
func rejectYAMLMergeNode(node *yaml.Node, visited map[*yaml.Node]struct{}) error {
|
||
|
|
if node == nil {
|
||
|
|
return nil
|
||
|
|
}
|
||
|
|
if _, exists := visited[node]; exists {
|
||
|
|
return nil
|
||
|
|
}
|
||
|
|
visited[node] = struct{}{}
|
||
|
|
|
||
|
|
if node.Kind == yaml.MappingNode {
|
||
|
|
for index := 0; index+1 < len(node.Content); index += 2 {
|
||
|
|
key := node.Content[index]
|
||
|
|
if key.Value == "<<" || key.Tag == "!!merge" || key.Tag == "tag:yaml.org,2002:merge" {
|
||
|
|
return fmt.Errorf("配置文件不允许使用 YAML merge key <<(第 %d 行,第 %d 列)", key.Line, key.Column)
|
||
|
|
}
|
||
|
|
if err := rejectYAMLMergeNode(key, visited); err != nil {
|
||
|
|
return err
|
||
|
|
}
|
||
|
|
if err := rejectYAMLMergeNode(node.Content[index+1], visited); err != nil {
|
||
|
|
return err
|
||
|
|
}
|
||
|
|
}
|
||
|
|
} else {
|
||
|
|
for _, child := range node.Content {
|
||
|
|
if err := rejectYAMLMergeNode(child, visited); err != nil {
|
||
|
|
return err
|
||
|
|
}
|
||
|
|
}
|
||
|
|
}
|
||
|
|
|
||
|
|
if node.Kind == yaml.AliasNode {
|
||
|
|
return rejectYAMLMergeNode(node.Alias, visited)
|
||
|
|
}
|
||
|
|
return nil
|
||
|
|
}
|
||
|
|
|
||
|
|
func (raw fileConfig) build() (Config, error) {
|
||
|
|
listenAddr, err := requireValue("server.listen_addr", raw.Server.ListenAddr)
|
||
|
|
if err != nil {
|
||
|
|
return Config{}, err
|
||
|
|
}
|
||
|
|
if err := validateLoopbackListenAddr(listenAddr); err != nil {
|
||
|
|
return Config{}, err
|
||
|
|
}
|
||
|
|
|
||
|
|
databaseDSN, err := requireValue("database.dsn", raw.Database.DSN)
|
||
|
|
if err != nil {
|
||
|
|
return Config{}, err
|
||
|
|
}
|
||
|
|
signingPrivateKeyFile, err := requireAbsolutePath("signing.private_key_file", raw.Signing.PrivateKeyFile)
|
||
|
|
if err != nil {
|
||
|
|
return Config{}, err
|
||
|
|
}
|
||
|
|
signingCertificateFile, err := requireAbsolutePath("signing.certificate_file", raw.Signing.CertificateFile)
|
||
|
|
if err != nil {
|
||
|
|
return Config{}, err
|
||
|
|
}
|
||
|
|
|
||
|
|
return Config{
|
||
|
|
ListenAddr: listenAddr,
|
||
|
|
DatabaseDSN: databaseDSN,
|
||
|
|
SigningPrivateKeyFile: signingPrivateKeyFile,
|
||
|
|
SigningCertificateFile: signingCertificateFile,
|
||
|
|
}, nil
|
||
|
|
}
|
||
|
|
|
||
|
|
func requireValue(name string, value strictString) (string, error) {
|
||
|
|
text := string(value)
|
||
|
|
if strings.TrimSpace(text) == "" {
|
||
|
|
return "", fmt.Errorf("配置项 %s 不能为空", name)
|
||
|
|
}
|
||
|
|
return text, nil
|
||
|
|
}
|
||
|
|
|
||
|
|
func requireAbsolutePath(name string, value strictString) (string, error) {
|
||
|
|
filePath, err := requireValue(name, value)
|
||
|
|
if err != nil {
|
||
|
|
return "", err
|
||
|
|
}
|
||
|
|
if !filepath.IsAbs(filePath) {
|
||
|
|
return "", fmt.Errorf("配置项 %s 必须是绝对路径", name)
|
||
|
|
}
|
||
|
|
return filePath, nil
|
||
|
|
}
|
||
|
|
|
||
|
|
func validateLoopbackListenAddr(listenAddr string) error {
|
||
|
|
host, port, err := net.SplitHostPort(listenAddr)
|
||
|
|
if err != nil {
|
||
|
|
return fmt.Errorf("配置项 server.listen_addr 必须是主机和端口: %w", err)
|
||
|
|
}
|
||
|
|
|
||
|
|
ip := net.ParseIP(host)
|
||
|
|
if ip == nil || !ip.IsLoopback() {
|
||
|
|
return errors.New("配置项 server.listen_addr 必须使用回环 IP 地址")
|
||
|
|
}
|
||
|
|
|
||
|
|
portNumber, err := strconv.ParseUint(port, 10, 16)
|
||
|
|
if err != nil || portNumber == 0 {
|
||
|
|
return errors.New("配置项 server.listen_addr 的端口无效")
|
||
|
|
}
|
||
|
|
return nil
|
||
|
|
}
|