package config import ( "fmt" "os" "time" "gopkg.in/yaml.v3" ) // FileDoc wraps the parsed YAML node tree of a broker config file so that // database entries can be added, removed, or toggled while leaving the // rest of the file (formatting, comments, unrelated keys) intact. type FileDoc struct { path string root *yaml.Node } // databaseFileConfig mirrors DatabaseConfig with duration fields represented // as strings, which is the form used by the human-edited YAML config. YAML // unmarshalling does not parse strings into time.Duration automatically. type databaseFileConfig struct { Name string `yaml:"name"` Host string `yaml:"host"` Port int `yaml:"port"` Database string `yaml:"database"` User string `yaml:"user"` Password string `yaml:"password"` SSLMode string `yaml:"sslmode"` MaxOpenConns int `yaml:"max_open_conns"` MaxIdleConns int `yaml:"max_idle_conns"` ConnMaxLifetime string `yaml:"conn_max_lifetime"` ConnMaxIdleTime string `yaml:"conn_max_idle_time"` QueueCount int `yaml:"queue_count"` TenantID string `yaml:"tenant_id"` AutoMigrate bool `yaml:"auto_migrate"` Disabled bool `yaml:"disabled"` } func decodeDatabase(node *yaml.Node) (DatabaseConfig, error) { var raw databaseFileConfig if err := node.Decode(&raw); err != nil { return DatabaseConfig{}, err } db := DatabaseConfig{ Name: raw.Name, Host: raw.Host, Port: raw.Port, Database: raw.Database, User: raw.User, Password: raw.Password, SSLMode: raw.SSLMode, MaxOpenConns: raw.MaxOpenConns, MaxIdleConns: raw.MaxIdleConns, QueueCount: raw.QueueCount, TenantID: raw.TenantID, AutoMigrate: raw.AutoMigrate, Disabled: raw.Disabled, } var err error if raw.ConnMaxLifetime != "" { db.ConnMaxLifetime, err = time.ParseDuration(raw.ConnMaxLifetime) if err != nil { return DatabaseConfig{}, fmt.Errorf("conn_max_lifetime: %w", err) } } if raw.ConnMaxIdleTime != "" { db.ConnMaxIdleTime, err = time.ParseDuration(raw.ConnMaxIdleTime) if err != nil { return DatabaseConfig{}, fmt.Errorf("conn_max_idle_time: %w", err) } } return db, nil } // LoadFileDoc reads and parses the YAML config file at path for editing. // A missing file is treated as an empty document so `db add` can be used // to create a new config file from scratch. func LoadFileDoc(path string) (*FileDoc, error) { data, err := os.ReadFile(path) if err != nil { if !os.IsNotExist(err) { return nil, fmt.Errorf("failed to read config file %s: %w", path, err) } data = nil } var root yaml.Node if len(data) > 0 { if err := yaml.Unmarshal(data, &root); err != nil { return nil, fmt.Errorf("failed to parse config file %s: %w", path, err) } } if root.Kind == 0 { root.Kind = yaml.DocumentNode root.Content = []*yaml.Node{{Kind: yaml.MappingNode, Tag: "!!map"}} } return &FileDoc{path: path, root: &root}, nil } // Save writes the document back to its original path. func (f *FileDoc) Save() error { data, err := yaml.Marshal(f.root) if err != nil { return fmt.Errorf("failed to marshal config: %w", err) } if err := os.WriteFile(f.path, data, 0o644); err != nil { return fmt.Errorf("failed to write config file %s: %w", f.path, err) } return nil } func (f *FileDoc) mappingRoot() (*yaml.Node, error) { if len(f.root.Content) == 0 { return nil, fmt.Errorf("config file is empty") } m := f.root.Content[0] if m.Kind != yaml.MappingNode { return nil, fmt.Errorf("config file root is not a mapping") } return m, nil } // databasesSeq returns the "databases" sequence node, creating it (and the // key) if it isn't already present. func (f *FileDoc) databasesSeq() (*yaml.Node, error) { m, err := f.mappingRoot() if err != nil { return nil, err } for i := 0; i+1 < len(m.Content); i += 2 { if m.Content[i].Value == "databases" { seq := m.Content[i+1] if seq.Kind != yaml.SequenceNode { return nil, fmt.Errorf("'databases' key in config is not a list") } return seq, nil } } keyNode := &yaml.Node{Kind: yaml.ScalarNode, Value: "databases"} seqNode := &yaml.Node{Kind: yaml.SequenceNode, Tag: "!!seq"} m.Content = append(m.Content, keyNode, seqNode) return seqNode, nil } // ListDatabases decodes all database entries currently in the document, in // file order. func (f *FileDoc) ListDatabases() ([]DatabaseConfig, error) { seq, err := f.databasesSeq() if err != nil { return nil, err } dbs := make([]DatabaseConfig, 0, len(seq.Content)) for _, item := range seq.Content { db, err := decodeDatabase(item) if err != nil { return nil, fmt.Errorf("failed to decode database entry: %w", err) } dbs = append(dbs, db) } return dbs, nil } // FindDatabase returns the decoded entry with the given name. func (f *FileDoc) FindDatabase(name string) (DatabaseConfig, bool, error) { dbs, err := f.ListDatabases() if err != nil { return DatabaseConfig{}, false, err } for i := range dbs { db := &dbs[i] if db.Name == name { return *db, true, nil } } return DatabaseConfig{}, false, nil } // LastDatabase returns the last database entry in the file, used to prime // defaults for a new `db add` when no --from instance is given. func (f *FileDoc) LastDatabase() (DatabaseConfig, bool, error) { dbs, err := f.ListDatabases() if err != nil { return DatabaseConfig{}, false, err } if len(dbs) == 0 { return DatabaseConfig{}, false, nil } return dbs[len(dbs)-1], true, nil } // AddDatabase appends a new database entry. It fails if an entry with the // same name already exists. func (f *FileDoc) AddDatabase(db DatabaseConfig) error { seq, err := f.databasesSeq() if err != nil { return err } for _, item := range seq.Content { existing, err := decodeDatabase(item) if err != nil { return fmt.Errorf("failed to decode database entry: %w", err) } if existing.Name == db.Name { return fmt.Errorf("a database named %q already exists", db.Name) } } seq.Content = append(seq.Content, dbToNode(db)) return nil } // RemoveDatabase deletes the entry with the given name, reporting whether // it was found. func (f *FileDoc) RemoveDatabase(name string) (bool, error) { seq, err := f.databasesSeq() if err != nil { return false, err } for i, item := range seq.Content { existing, err := decodeDatabase(item) if err != nil { return false, fmt.Errorf("failed to decode database entry: %w", err) } if existing.Name == name { seq.Content = append(seq.Content[:i], seq.Content[i+1:]...) return true, nil } } return false, nil } // SetDisabled sets (or clears) the disabled flag on the named entry, // reporting whether the entry was found. func (f *FileDoc) SetDisabled(name string, disabled bool) (bool, error) { seq, err := f.databasesSeq() if err != nil { return false, err } for _, item := range seq.Content { existing, err := decodeDatabase(item) if err != nil { return false, fmt.Errorf("failed to decode database entry: %w", err) } if existing.Name != name { continue } if disabled { setMappingKey(item, "disabled", &yaml.Node{Kind: yaml.ScalarNode, Tag: "!!bool", Value: "true"}) } else { removeMappingKey(item, "disabled") } return true, nil } return false, nil } // setMappingKey sets key to value within a mapping node, appending the // pair if the key isn't already present. func setMappingKey(m *yaml.Node, key string, value *yaml.Node) { for i := 0; i+1 < len(m.Content); i += 2 { if m.Content[i].Value == key { m.Content[i+1] = value return } } m.Content = append(m.Content, &yaml.Node{Kind: yaml.ScalarNode, Value: key}, value) } // removeMappingKey removes key from a mapping node if present. func removeMappingKey(m *yaml.Node, key string) { for i := 0; i+1 < len(m.Content); i += 2 { if m.Content[i].Value == key { m.Content = append(m.Content[:i], m.Content[i+2:]...) return } } } // dbToNode builds a YAML mapping node for a database entry, omitting // fields left at their zero value so the on-disk defaults from // applyDatabaseDefaults keep applying (matching broker.example.yaml // style). Field order mirrors DatabaseConfig / broker.example.yaml. func dbToNode(db DatabaseConfig) *yaml.Node { m := &yaml.Node{Kind: yaml.MappingNode, Tag: "!!map"} add := func(key, value string) { m.Content = append(m.Content, &yaml.Node{Kind: yaml.ScalarNode, Value: key}, &yaml.Node{Kind: yaml.ScalarNode, Value: value}, ) } addBool := func(key string, value bool) { m.Content = append(m.Content, &yaml.Node{Kind: yaml.ScalarNode, Value: key}, &yaml.Node{Kind: yaml.ScalarNode, Tag: "!!bool", Value: fmt.Sprintf("%t", value)}, ) } addInt := func(key string, value int) { m.Content = append(m.Content, &yaml.Node{Kind: yaml.ScalarNode, Value: key}, &yaml.Node{Kind: yaml.ScalarNode, Tag: "!!int", Value: fmt.Sprintf("%d", value)}, ) } add("name", db.Name) add("host", db.Host) if db.Port != 0 { addInt("port", db.Port) } add("database", db.Database) add("user", db.User) if db.Password != "" { add("password", db.Password) } if db.SSLMode != "" { add("sslmode", db.SSLMode) } if db.MaxOpenConns != 0 { addInt("max_open_conns", db.MaxOpenConns) } if db.MaxIdleConns != 0 { addInt("max_idle_conns", db.MaxIdleConns) } if db.ConnMaxLifetime != 0 { add("conn_max_lifetime", db.ConnMaxLifetime.String()) } if db.ConnMaxIdleTime != 0 { add("conn_max_idle_time", db.ConnMaxIdleTime.String()) } if db.QueueCount != 0 { addInt("queue_count", db.QueueCount) } if db.TenantID != "" { add("tenant_id", db.TenantID) } if db.AutoMigrate { addBool("auto_migrate", db.AutoMigrate) } if db.Disabled { addBool("disabled", db.Disabled) } return m }