diff --git a/README.md b/README.md index 49fcc58..4d530e5 100644 --- a/README.md +++ b/README.md @@ -369,6 +369,27 @@ The `databases` array can contain multiple database configurations. Each entry s | `conn_max_idle_time` | Connection max idle time | `10m` | | `queue_count` | Number of queues for this database | `4` | | `tenant_id` | Tenant id set for this connection (RLS) | `default` | +| `auto_migrate` | Apply pending migrations during `start` | `false` | +| `disabled` | Exclude this database from `start` | `false` | + +### Database Management Commands + +The configuration can be managed without editing YAML by hand: + +```bash +pgsql-broker db list --config broker.yaml +pgsql-broker db add replica --from db1 --host db.example --database jobs_replica --user broker --password secret --yes --non-interactive --config broker.yaml +pgsql-broker db disable replica --config broker.yaml +pgsql-broker db enable replica --config broker.yaml +pgsql-broker db remove replica --yes --config broker.yaml +``` + +`db add` accepts flags for every database setting. If a setting is omitted in +interactive mode, it is pre-filled from the last database in the file, or +from the instance named by `--from`; press Enter to keep that value. Use +`--non-interactive` to reject missing required values instead of prompting. +The add command asks for confirmation unless `--yes` is supplied. Remove also +requires `--yes` when stdin is not a terminal. ### Broker Settings diff --git a/cmd/broker/db.go b/cmd/broker/db.go new file mode 100644 index 0000000..f7e8a37 --- /dev/null +++ b/cmd/broker/db.go @@ -0,0 +1,512 @@ +package main + +import ( + "bufio" + "fmt" + "os" + "strconv" + "strings" + "time" + + "github.com/spf13/cobra" + "golang.org/x/term" + + "git.warky.dev/wdevs/pgsql-broker/pkg/broker/config" +) + +// Flags shared by the `db` subcommands. +var ( + dbConfigPath string + + dbAddFrom string + dbAddName string + dbAddHost string + dbAddPort int + dbAddDatabase string + dbAddUser string + dbAddPassword string + dbAddSSLMode string + dbAddMaxOpenConns int + dbAddMaxIdleConns int + dbAddConnMaxLifetime string + dbAddConnMaxIdleTime string + dbAddQueueCount int + dbAddTenantID string + dbAddAutoMigrate bool + dbAddDisabled bool + dbAddNonInteractive bool + dbYes bool +) + +var dbCmd = &cobra.Command{ + Use: "db", + Short: "Manage databases in the broker config file", + Long: `Add, remove, enable, disable, or list the databases configured in the broker config file.`, +} + +var dbAddCmd = &cobra.Command{ + Use: "add [name]", + Short: "Add a database to the config file", + Long: `Add a database entry to the config file. + +Every field can be supplied via flags for fully non-interactive use. Any +field left unset falls back to an interactive prompt (when stdin is a +terminal) that is pre-filled with the value from --from (a named existing +instance) or, if --from is not given, the last database currently in the +file. Press Enter at a prompt to accept the shown value.`, + Args: cobra.MaximumNArgs(1), + RunE: func(cmd *cobra.Command, args []string) error { + if len(args) == 1 && dbAddName == "" { + dbAddName = args[0] + } + return runDBAdd(cmd) + }, +} + +var dbRemoveCmd = &cobra.Command{ + Use: "remove ", + Short: "Remove a database from the config file", + Args: cobra.ExactArgs(1), + RunE: func(cmd *cobra.Command, args []string) error { + return runDBRemove(args[0]) + }, +} + +var dbEnableCmd = &cobra.Command{ + Use: "enable ", + Short: "Enable a database (removes the disabled flag)", + Args: cobra.ExactArgs(1), + RunE: func(cmd *cobra.Command, args []string) error { + return runDBSetDisabled(args[0], false) + }, +} + +var dbDisableCmd = &cobra.Command{ + Use: "disable ", + Short: "Disable a database (excluded from `start` until re-enabled)", + Args: cobra.ExactArgs(1), + RunE: func(cmd *cobra.Command, args []string) error { + return runDBSetDisabled(args[0], true) + }, +} + +var dbListCmd = &cobra.Command{ + Use: "list", + Short: "List databases configured in the config file", + Args: cobra.NoArgs, + RunE: func(cmd *cobra.Command, args []string) error { + return runDBList() + }, +} + +func init() { + rootCmd.AddCommand(dbCmd) + dbCmd.AddCommand(dbAddCmd, dbRemoveCmd, dbEnableCmd, dbDisableCmd, dbListCmd) + + dbCmd.PersistentFlags().StringVar(&dbConfigPath, "config", "broker.yaml", "config file to edit") + + dbAddCmd.Flags().StringVar(&dbAddFrom, "from", "", "prime defaults from this existing database name (default: last database in the file)") + dbAddCmd.Flags().StringVar(&dbAddName, "name", "", "unique name for the new database (required)") + dbAddCmd.Flags().StringVar(&dbAddHost, "host", "", "PostgreSQL host") + dbAddCmd.Flags().IntVar(&dbAddPort, "port", 0, "PostgreSQL port (default 5432)") + dbAddCmd.Flags().StringVar(&dbAddDatabase, "database", "", "database name") + dbAddCmd.Flags().StringVar(&dbAddUser, "user", "", "database user") + dbAddCmd.Flags().StringVar(&dbAddPassword, "password", "", "database password") + dbAddCmd.Flags().StringVar(&dbAddSSLMode, "sslmode", "", "SSL mode (disable, require, verify-ca, verify-full)") + dbAddCmd.Flags().IntVar(&dbAddMaxOpenConns, "max-open-conns", 0, "max open connections (default 25)") + dbAddCmd.Flags().IntVar(&dbAddMaxIdleConns, "max-idle-conns", 0, "max idle connections (default 5)") + dbAddCmd.Flags().StringVar(&dbAddConnMaxLifetime, "conn-max-lifetime", "", "connection max lifetime (e.g. 5m, default 5m)") + dbAddCmd.Flags().StringVar(&dbAddConnMaxIdleTime, "conn-max-idle-time", "", "connection max idle time (e.g. 10m, default 10m)") + dbAddCmd.Flags().IntVar(&dbAddQueueCount, "queue-count", 0, "number of concurrent queues (default 4)") + dbAddCmd.Flags().StringVar(&dbAddTenantID, "tenant-id", "", "RLS tenant id (default \"default\")") + dbAddCmd.Flags().BoolVar(&dbAddAutoMigrate, "auto-migrate", false, "apply pending migrations automatically on connect") + dbAddCmd.Flags().BoolVar(&dbAddDisabled, "disabled", false, "add the database in a disabled state") + dbAddCmd.Flags().BoolVar(&dbAddNonInteractive, "non-interactive", false, "fail instead of prompting for missing fields") + dbAddCmd.Flags().BoolVarP(&dbYes, "yes", "y", false, "skip the confirmation prompt") + + dbRemoveCmd.Flags().BoolVarP(&dbYes, "yes", "y", false, "skip the confirmation prompt") +} + +// runDBAdd builds a new DatabaseConfig by layering (in priority order) +// explicit flags, then interactive prompts primed from --from / the last +// database in the file, then falls back to leaving optional fields unset +// so config.applyDatabaseDefaults keeps handling defaults at load time. +func runDBAdd(cmd *cobra.Command) error { + doc, err := config.LoadFileDoc(dbConfigPath) + if err != nil { + return err + } + + var prime config.DatabaseConfig + if dbAddFrom != "" { + found, ok, err := doc.FindDatabase(dbAddFrom) + if err != nil { + return err + } + if !ok { + return fmt.Errorf("--from instance %q not found in %s", dbAddFrom, dbConfigPath) + } + prime = found + } else if last, ok, err := doc.LastDatabase(); err != nil { + return err + } else if ok { + prime = last + } + // Never prime the new entry's name from an existing one. + prime.Name = "" + + interactive := !dbAddNonInteractive && term.IsTerminal(int(os.Stdin.Fd())) + reader := bufio.NewReader(os.Stdin) + + if interactive { + fmt.Fprintln(os.Stderr, "Adding a new database. Press Enter to accept the shown value.") + } + + db := config.DatabaseConfig{} + + if dbAddName != "" { + db.Name = dbAddName + } else { + db.Name, err = resolveStringField(cmd, reader, interactive, "name", "", true, dbAddName) + if err != nil { + return err + } + } + if db.Name == "" { + return fmt.Errorf("name is required (--name, a positional argument, or run interactively)") + } + db.Host, err = resolveStringField(cmd, reader, interactive, "host", prime.Host, true, dbAddHost) + if err != nil { + return err + } + db.Database, err = resolveStringField(cmd, reader, interactive, "database", prime.Database, true, dbAddDatabase) + if err != nil { + return err + } + db.User, err = resolveStringField(cmd, reader, interactive, "user", prime.User, true, dbAddUser) + if err != nil { + return err + } + db.Password, err = resolveStringField(cmd, reader, interactive, "password", prime.Password, false, dbAddPassword) + if err != nil { + return err + } + db.SSLMode, err = resolveStringField(cmd, reader, interactive, "sslmode", prime.SSLMode, false, dbAddSSLMode) + if err != nil { + return err + } + db.TenantID, err = resolveStringField(cmd, reader, interactive, "tenant_id", prime.TenantID, false, dbAddTenantID) + if err != nil { + return err + } + + port, err := resolveIntField(cmd, reader, interactive, "port", prime.Port, "port", dbAddPort) + if err != nil { + return err + } + db.Port = port + + maxOpen, err := resolveIntField(cmd, reader, interactive, "max_open_conns", prime.MaxOpenConns, "max-open-conns", dbAddMaxOpenConns) + if err != nil { + return err + } + db.MaxOpenConns = maxOpen + + maxIdle, err := resolveIntField(cmd, reader, interactive, "max_idle_conns", prime.MaxIdleConns, "max-idle-conns", dbAddMaxIdleConns) + if err != nil { + return err + } + db.MaxIdleConns = maxIdle + + queueCount, err := resolveIntField(cmd, reader, interactive, "queue_count", prime.QueueCount, "queue-count", dbAddQueueCount) + if err != nil { + return err + } + db.QueueCount = queueCount + + connMaxLifetime, err := resolveDurationField(cmd, reader, interactive, "conn_max_lifetime", prime.ConnMaxLifetime, "conn-max-lifetime", dbAddConnMaxLifetime) + if err != nil { + return err + } + db.ConnMaxLifetime = connMaxLifetime + + connMaxIdleTime, err := resolveDurationField(cmd, reader, interactive, "conn_max_idle_time", prime.ConnMaxIdleTime, "conn-max-idle-time", dbAddConnMaxIdleTime) + if err != nil { + return err + } + db.ConnMaxIdleTime = connMaxIdleTime + + if cmd.Flags().Changed("auto-migrate") { + db.AutoMigrate = dbAddAutoMigrate + } else { + db.AutoMigrate = prime.AutoMigrate + } + if cmd.Flags().Changed("disabled") { + db.Disabled = dbAddDisabled + } else { + db.Disabled = prime.Disabled + } + + for field, value := range map[string]string{ + "host": db.Host, + "database": db.Database, + "user": db.User, + } { + if strings.TrimSpace(value) == "" { + return fmt.Errorf("%s is required (--%s or an interactive value)", field, strings.ReplaceAll(field, "_", "-")) + } + } + + if err := doc.AddDatabase(db); err != nil { + return err + } + + if !dbYes && interactive { + fmt.Fprintf(os.Stderr, "\nAbout to add database %q to %s:\n", db.Name, dbConfigPath) + printDBSummary(os.Stderr, db) + fmt.Fprint(os.Stderr, "Proceed? [Y/n]: ") + line, _ := reader.ReadString('\n') + line = strings.TrimSpace(strings.ToLower(line)) + if line == "n" || line == "no" { + fmt.Println("Aborted; config file not modified.") + return nil + } + } + + if err := doc.Save(); err != nil { + return err + } + fmt.Printf("Added database %q to %s\n", db.Name, dbConfigPath) + return nil +} + +func runDBRemove(name string) error { + doc, err := config.LoadFileDoc(dbConfigPath) + if err != nil { + return err + } + db, ok, err := doc.FindDatabase(name) + if err != nil { + return err + } + if !ok { + return fmt.Errorf("database %q not found in %s", name, dbConfigPath) + } + + if !dbYes { + if !term.IsTerminal(int(os.Stdin.Fd())) { + return fmt.Errorf("refusing to remove %q non-interactively without --yes", name) + } + fmt.Fprintf(os.Stderr, "Remove database %q (%s@%s:%d/%s) from %s? [y/N]: ", + name, db.User, db.Host, db.Port, db.Database, dbConfigPath) + reader := bufio.NewReader(os.Stdin) + line, _ := reader.ReadString('\n') + line = strings.TrimSpace(strings.ToLower(line)) + if line != "y" && line != "yes" { + fmt.Println("Aborted; config file not modified.") + return nil + } + } + + if _, err := doc.RemoveDatabase(name); err != nil { + return err + } + if err := doc.Save(); err != nil { + return err + } + fmt.Printf("Removed database %q from %s\n", name, dbConfigPath) + return nil +} + +func runDBSetDisabled(name string, disabled bool) error { + doc, err := config.LoadFileDoc(dbConfigPath) + if err != nil { + return err + } + found, err := doc.SetDisabled(name, disabled) + if err != nil { + return err + } + if !found { + return fmt.Errorf("database %q not found in %s", name, dbConfigPath) + } + if err := doc.Save(); err != nil { + return err + } + verb := "Enabled" + if disabled { + verb = "Disabled" + } + fmt.Printf("%s database %q in %s\n", verb, name, dbConfigPath) + return nil +} + +func runDBList() error { + doc, err := config.LoadFileDoc(dbConfigPath) + if err != nil { + return err + } + dbs, err := doc.ListDatabases() + if err != nil { + return err + } + if len(dbs) == 0 { + fmt.Printf("No databases configured in %s\n", dbConfigPath) + return nil + } + fmt.Printf("%-20s %-20s %-6s %-16s %-16s %-10s\n", "NAME", "HOST", "PORT", "DATABASE", "USER", "STATUS") + for _, db := range dbs { + status := "enabled" + if db.Disabled { + status = "disabled" + } + port := db.Port + if port == 0 { + port = 5432 + } + fmt.Printf("%-20s %-20s %-6d %-16s %-16s %-10s\n", db.Name, db.Host, port, db.Database, db.User, status) + } + return nil +} + +func printDBSummary(w *os.File, db config.DatabaseConfig) { + password := "" + if db.Password != "" { + password = "(set)" + } + fmt.Fprintf(w, " host: %s\n", db.Host) + fmt.Fprintf(w, " port: %s\n", intOrDefault(db.Port, 5432)) + fmt.Fprintf(w, " database: %s\n", db.Database) + fmt.Fprintf(w, " user: %s\n", db.User) + fmt.Fprintf(w, " password: %s\n", password) + fmt.Fprintf(w, " sslmode: %s\n", stringOrDefault(db.SSLMode, "disable")) + fmt.Fprintf(w, " max_open_conns: %s\n", intOrDefault(db.MaxOpenConns, 25)) + fmt.Fprintf(w, " max_idle_conns: %s\n", intOrDefault(db.MaxIdleConns, 5)) + fmt.Fprintf(w, " conn_max_lifetime: %s\n", durationOrDefault(db.ConnMaxLifetime, 5*time.Minute)) + fmt.Fprintf(w, " conn_max_idle_time:%s\n", durationOrDefault(db.ConnMaxIdleTime, 10*time.Minute)) + fmt.Fprintf(w, " queue_count: %s\n", intOrDefault(db.QueueCount, 4)) + fmt.Fprintf(w, " tenant_id: %s\n", stringOrDefault(db.TenantID, "default")) + fmt.Fprintf(w, " auto_migrate: %t\n", db.AutoMigrate) + fmt.Fprintf(w, " disabled: %t\n", db.Disabled) +} + +func intOrDefault(v, def int) string { + if v == 0 { + return fmt.Sprintf("%d (default)", def) + } + return strconv.Itoa(v) +} + +func stringOrDefault(v, def string) string { + if v == "" { + return fmt.Sprintf("%s (default)", def) + } + return v +} + +func durationOrDefault(v, def time.Duration) string { + if v == 0 { + return fmt.Sprintf("%s (default)", def) + } + return v.String() +} + +// resolveStringField returns, in priority order: the flag value if the +// flag was explicitly set, otherwise an interactive prompt (pre-filled +// with defaultVal) when interactive, otherwise defaultVal. When required +// and still empty after all of that, an error is returned by the caller. +func resolveStringField(cmd *cobra.Command, reader *bufio.Reader, interactive bool, label, defaultVal string, required bool, flagVal string) (string, error) { + flagName := flagNameFor(label) + if cmd.Flags().Changed(flagName) { + return flagVal, nil + } + if !interactive { + return defaultVal, nil + } + for { + fmt.Fprintf(os.Stderr, "%s [%s]: ", label, defaultVal) + line, _ := reader.ReadString('\n') + line = strings.TrimSpace(line) + if line == "" { + line = defaultVal + } + if line == "" && required { + fmt.Fprintln(os.Stderr, "value is required") + continue + } + return line, nil + } +} + +func resolveIntField(cmd *cobra.Command, reader *bufio.Reader, interactive bool, label string, defaultVal int, flagName string, flagVal int) (int, error) { + if cmd.Flags().Changed(flagName) { + return flagVal, nil + } + if !interactive { + return defaultVal, nil + } + for { + fmt.Fprintf(os.Stderr, "%s [%d]: ", label, defaultVal) + line, _ := reader.ReadString('\n') + line = strings.TrimSpace(line) + if line == "" { + return defaultVal, nil + } + n, err := strconv.Atoi(line) + if err != nil { + fmt.Fprintln(os.Stderr, "must be a whole number") + continue + } + return n, nil + } +} + +func resolveDurationField(cmd *cobra.Command, reader *bufio.Reader, interactive bool, label string, defaultVal time.Duration, flagName, flagVal string) (time.Duration, error) { + if cmd.Flags().Changed(flagName) { + d, err := time.ParseDuration(flagVal) + if err != nil { + return 0, fmt.Errorf("--%s: %w", flagName, err) + } + return d, nil + } + if !interactive { + return defaultVal, nil + } + for { + fmt.Fprintf(os.Stderr, "%s [%s]: ", label, defaultVal) + line, _ := reader.ReadString('\n') + line = strings.TrimSpace(line) + if line == "" { + return defaultVal, nil + } + d, err := time.ParseDuration(line) + if err != nil { + fmt.Fprintln(os.Stderr, "must be a duration like 5m, 30s") + continue + } + return d, nil + } +} + +// flagNameFor maps a prompt label to its cobra flag name for the string +// fields prompted in runDBAdd. +func flagNameFor(label string) string { + switch label { + case "name": + return "name" + case "host": + return "host" + case "database": + return "database" + case "user": + return "user" + case "password": + return "password" + case "sslmode": + return "sslmode" + case "tenant_id": + return "tenant-id" + default: + return label + } +} diff --git a/go.mod b/go.mod index 63a8274..d62e5d8 100644 --- a/go.mod +++ b/go.mod @@ -8,6 +8,7 @@ require ( github.com/spf13/viper v1.21.0 github.com/stretchr/testify v1.11.1 golang.org/x/term v0.46.0 + gopkg.in/yaml.v3 v3.0.1 ) require ( @@ -26,5 +27,4 @@ require ( go.yaml.in/yaml/v3 v3.0.4 // indirect golang.org/x/sys v0.48.0 // indirect golang.org/x/text v0.28.0 // indirect - gopkg.in/yaml.v3 v3.0.1 // indirect ) diff --git a/pkg/broker/broker.go b/pkg/broker/broker.go index 1a51605..3501ea8 100644 --- a/pkg/broker/broker.go +++ b/pkg/broker/broker.go @@ -41,9 +41,13 @@ func New(cfg *config.Config, logger adapter.Logger, version string) (*Broker, er func (b *Broker) Start() error { b.logger.Info("starting broker", "database_count", len(b.config.Databases)) - // Create and start an instance for each database + // Create and start an instance for each enabled database for i := range b.config.Databases { dbCfg := &b.config.Databases[i] + if dbCfg.Disabled { + b.logger.Info("skipping disabled database", "name", dbCfg.Name) + continue + } b.logger.Info("starting database instance", "name", dbCfg.Name, "host", dbCfg.Host, "database", dbCfg.Database) // Create database adapter @@ -68,6 +72,10 @@ func (b *Broker) Start() error { b.logger.Info("database instance started", "name", dbCfg.Name, "instance_id", instance.ID) } + if len(b.instances) == 0 { + return fmt.Errorf("no enabled databases configured (all %d database(s) are disabled)", len(b.config.Databases)) + } + b.logger.Info("broker started successfully", "database_instances", len(b.instances)) return nil } diff --git a/pkg/broker/config/config.go b/pkg/broker/config/config.go index 5a901c3..dfea8aa 100644 --- a/pkg/broker/config/config.go +++ b/pkg/broker/config/config.go @@ -39,6 +39,10 @@ type DatabaseConfig struct { // fast if the schema is behind, naming the missing migrations and // pointing at `pgsql-broker install`. AutoMigrate bool `mapstructure:"auto_migrate"` + // Disabled, when true, excludes this database from `start` (no + // instance is created for it). Defaults to false so existing config + // files without this field are unaffected. + Disabled bool `mapstructure:"disabled"` } // BrokerConfig holds broker-specific settings diff --git a/pkg/broker/config/dbedit.go b/pkg/broker/config/dbedit.go new file mode 100644 index 0000000..82ca476 --- /dev/null +++ b/pkg/broker/config/dbedit.go @@ -0,0 +1,336 @@ +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 _, db := range dbs { + 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 +} diff --git a/pkg/broker/config/dbedit_test.go b/pkg/broker/config/dbedit_test.go new file mode 100644 index 0000000..29a0319 --- /dev/null +++ b/pkg/broker/config/dbedit_test.go @@ -0,0 +1,102 @@ +package config + +import ( + "os" + "path/filepath" + "testing" + "time" + + "github.com/stretchr/testify/require" +) + +func TestFileDocDatabaseManagement(t *testing.T) { + path := filepath.Join(t.TempDir(), "broker.yaml") + require.NoError(t, os.WriteFile(path, []byte(`broker: + name: test +# keep this comment +databases: + - name: primary + host: db.example + port: 5433 + database: jobs + user: broker + password: secret + sslmode: verify-full + max_open_conns: 12 + max_idle_conns: 4 + conn_max_lifetime: 2m + conn_max_idle_time: 30s + queue_count: 3 + tenant_id: tenant-a + auto_migrate: true +`), 0o600)) + + doc, err := LoadFileDoc(path) + require.NoError(t, err) + primary, found, err := doc.FindDatabase("primary") + require.NoError(t, err) + require.True(t, found) + require.Equal(t, 2*time.Minute, primary.ConnMaxLifetime) + require.Equal(t, 30*time.Second, primary.ConnMaxIdleTime) + require.True(t, primary.AutoMigrate) + + copy := primary + copy.Name = "replica" + require.NoError(t, doc.AddDatabase(copy)) + found, err = doc.SetDisabled("replica", true) + require.NoError(t, err) + require.True(t, found) + require.NoError(t, doc.Save()) + + doc, err = LoadFileDoc(path) + require.NoError(t, err) + replica, found, err := doc.FindDatabase("replica") + require.NoError(t, err) + require.True(t, found) + require.True(t, replica.Disabled) + require.NoError(t, func() error { + found, err := doc.SetDisabled("replica", false) + require.True(t, found) + return err + }()) + require.NoError(t, doc.Save()) + + doc, err = LoadFileDoc(path) + require.NoError(t, err) + replica, found, err = doc.FindDatabase("replica") + require.NoError(t, err) + require.True(t, found) + require.False(t, replica.Disabled) + + removed, err := doc.RemoveDatabase("primary") + require.NoError(t, err) + require.True(t, removed) + require.NoError(t, doc.Save()) + contents, err := os.ReadFile(path) + require.NoError(t, err) + require.Contains(t, string(contents), "# keep this comment") + + _, found, err = doc.FindDatabase("primary") + require.NoError(t, err) + require.False(t, found) +} + +func TestFileDocMissingFileCanAddDatabase(t *testing.T) { + path := filepath.Join(t.TempDir(), "new.yaml") + doc, err := LoadFileDoc(path) + require.NoError(t, err) + require.NoError(t, doc.AddDatabase(DatabaseConfig{ + Name: "new", + Host: "localhost", + Database: "jobs", + User: "broker", + })) + require.NoError(t, doc.Save()) + + reloaded, err := LoadFileDoc(path) + require.NoError(t, err) + dbs, err := reloaded.ListDatabases() + require.NoError(t, err) + require.Len(t, dbs, 1) + require.Equal(t, "new", dbs[0].Name) +}