From 6f36d73bb52349365fd599060bcb84bf6f193a34 Mon Sep 17 00:00:00 2001 From: tjblackheart Date: Mon, 3 Aug 2026 07:33:05 +0200 Subject: [PATCH 1/2] refactor: move vault loading into registry --- cmd/andcli/main.go | 52 ++++------------ internal/config/config.go | 12 ++-- internal/config/config_test.go | 4 +- internal/config/flags.go | 2 +- internal/vaults/aegis/aegis.go | 4 +- internal/vaults/aegis/aegis_test.go | 16 ++--- internal/vaults/andotp/andotp.go | 6 +- internal/vaults/andotp/andotp_test.go | 16 ++--- internal/vaults/keepass/keepass.go | 4 +- internal/vaults/protonpass/protonpass.go | 14 +++-- internal/vaults/registry.go | 79 ++++++++++++++++++++++++ internal/vaults/registry_test.go | 57 +++++++++++++++++ internal/vaults/stratum/stratum.go | 6 +- internal/vaults/stratum/stratum_test.go | 16 ++--- internal/vaults/twofas/twofas.go | 4 +- internal/vaults/twofas/twofas_test.go | 16 ++--- internal/vaults/vault.go | 38 +----------- internal/vaults/vault_test.go | 35 ----------- 18 files changed, 214 insertions(+), 167 deletions(-) create mode 100644 internal/vaults/registry.go create mode 100644 internal/vaults/registry_test.go delete mode 100644 internal/vaults/vault_test.go diff --git a/cmd/andcli/main.go b/cmd/andcli/main.go index b70e0cc..ac3a8ca 100644 --- a/cmd/andcli/main.go +++ b/cmd/andcli/main.go @@ -1,6 +1,7 @@ package main import ( + "context" "fmt" "log" "os" @@ -13,12 +14,13 @@ import ( "github.com/tjblackheart/andcli/v2/internal/input" "github.com/tjblackheart/andcli/v2/internal/model" "github.com/tjblackheart/andcli/v2/internal/vaults" - "github.com/tjblackheart/andcli/v2/internal/vaults/aegis" - "github.com/tjblackheart/andcli/v2/internal/vaults/andotp" - "github.com/tjblackheart/andcli/v2/internal/vaults/keepass" - "github.com/tjblackheart/andcli/v2/internal/vaults/protonpass" - "github.com/tjblackheart/andcli/v2/internal/vaults/stratum" - "github.com/tjblackheart/andcli/v2/internal/vaults/twofas" + + _ "github.com/tjblackheart/andcli/v2/internal/vaults/aegis" + _ "github.com/tjblackheart/andcli/v2/internal/vaults/andotp" + _ "github.com/tjblackheart/andcli/v2/internal/vaults/keepass" + _ "github.com/tjblackheart/andcli/v2/internal/vaults/protonpass" + _ "github.com/tjblackheart/andcli/v2/internal/vaults/stratum" + _ "github.com/tjblackheart/andcli/v2/internal/vaults/twofas" ) func main() { @@ -70,42 +72,10 @@ func open(cfg *config.Config) (vaults.Vault, error) { return nil, err } - defer func() { - for i := range pw { - pw[i] = 0 - } - }() - - done := make(chan struct{}) - - var vault vaults.Vault - go func() { - switch cfg.Type { - case vaults.ANDOTP: - vault, err = andotp.Open(cfg.File, pw) - case vaults.AEGIS: - vault, err = aegis.Open(cfg.File, pw) - case vaults.TWOFAS: - vault, err = twofas.Open(cfg.File, pw) - case vaults.STRATUM: - vault, err = stratum.Open(cfg.File, pw) - case vaults.KEEPASS: - vault, err = keepass.Open(cfg.File, pw) - case vaults.PROTON: - vault, err = protonpass.Open(cfg.File, pw) - default: - vault, err = nil, fmt.Errorf("vault type %q: not implemented", cfg.Type) - } - done <- struct{}{} - }() - - select { - case <-done: - case <-time.After(cfg.DecryptionTimeoutD()): - return nil, fmt.Errorf("decrypt: operation timed out. wrong type?") - } + ctx, cancel := context.WithTimeout(context.Background(), cfg.DecryptionTimeoutD()) + defer cancel() - return vault, err + return vaults.Open(ctx, cfg.File, pw, cfg.Type) } func password(piped bool) ([]byte, error) { diff --git a/internal/config/config.go b/internal/config/config.go index d2d3ea7..bf85abe 100644 --- a/internal/config/config.go +++ b/internal/config/config.go @@ -16,12 +16,12 @@ import ( type ( Config struct { - File string `yaml:"file"` - Type vaults.Type `yaml:"type"` - ClipboardCmd string `yaml:"clipboard_cmd"` - Options *Opts `yaml:"options"` - Theme *Theme `yaml:"theme"` - SessionTimeout int `yaml:"session_timeout"` + File string `yaml:"file"` + Type vaults.VaultType `yaml:"type"` + ClipboardCmd string `yaml:"clipboard_cmd"` + Options *Opts `yaml:"options"` + Theme *Theme `yaml:"theme"` + SessionTimeout int `yaml:"session_timeout"` // path string passwordFromStdin bool diff --git a/internal/config/config_test.go b/internal/config/config_test.go index 1a6610f..d723207 100644 --- a/internal/config/config_test.go +++ b/internal/config/config_test.go @@ -235,7 +235,7 @@ theme: cfg := &Config{ File: "/new/vault.json", - Type: vaults.Type("2fas"), + Type: vaults.VaultType("2fas"), ClipboardCmd: "pbcopy", SessionTimeout: 300, Options: &Opts{ @@ -302,7 +302,7 @@ func Test_create(t *testing.T) { // default config want := &Config{ File: abs, - Type: vaults.Type(*vtype), + Type: vaults.VaultType(*vtype), SessionTimeout: 300, ClipboardCmd: "", Options: &Opts{ diff --git a/internal/config/flags.go b/internal/config/flags.go index d267286..b3a078c 100644 --- a/internal/config/flags.go +++ b/internal/config/flags.go @@ -54,7 +54,7 @@ func (cfg *Config) parseFlags() error { } if *vtype != "" { - cfg.Type = vaults.Type(*vtype) + cfg.Type = vaults.VaultType(*vtype) cfg.dirty = true } diff --git a/internal/vaults/aegis/aegis.go b/internal/vaults/aegis/aegis.go index 8f48266..0ebca76 100644 --- a/internal/vaults/aegis/aegis.go +++ b/internal/vaults/aegis/aegis.go @@ -58,6 +58,8 @@ type ( } ) +func init() { vaults.Register(vaultType, Open) } + func Open(filename string, pass []byte) (vaults.Vault, error) { var v aegis @@ -88,7 +90,7 @@ func Open(filename string, pass []byte) (vaults.Vault, error) { return nil, fmt.Errorf("%s: %w", vaultType, err) } - return v, nil + return &v, nil } func (v aegis) Entries() []vaults.Entry { diff --git a/internal/vaults/aegis/aegis_test.go b/internal/vaults/aegis/aegis_test.go index 2aec476..27d730c 100644 --- a/internal/vaults/aegis/aegis_test.go +++ b/internal/vaults/aegis/aegis_test.go @@ -25,15 +25,15 @@ func TestOpen(t *testing.T) { for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { v, err := Open(tt.filename, []byte(tt.password)) - if tt.fails { - if err == nil { - t.Fatal("Open() expected error, got none") + if tt.fails { + if err == nil { + t.Fatal("Open() expected error, got none") + } + if tt.wantErr != nil && !errors.Is(err, tt.wantErr) { + t.Fatalf("Open() error = %v, want %v", err, tt.wantErr) + } + return } - if tt.wantErr != nil && !errors.Is(err, tt.wantErr) { - t.Fatalf("Open() error = %v, want %v", err, tt.wantErr) - } - return - } entries := v.Entries() if len(entries) != 1 { diff --git a/internal/vaults/andotp/andotp.go b/internal/vaults/andotp/andotp.go index e27b852..d42a79e 100644 --- a/internal/vaults/andotp/andotp.go +++ b/internal/vaults/andotp/andotp.go @@ -32,13 +32,15 @@ type ( } ) +func init() { vaults.Register(vaultType, Open) } + func Open(filename string, pass []byte) (vaults.Vault, error) { b, err := os.ReadFile(filename) if err != nil { return nil, fmt.Errorf("%s: %w", vaultType, err) } - v := &andotp{entries: make([]entry, 0)} + v := andotp{entries: make([]entry, 0)} if v.IsPlain(b) { return nil, vaults.ErrIsPlain @@ -53,7 +55,7 @@ func Open(filename string, pass []byte) (vaults.Vault, error) { return nil, fmt.Errorf("%s: %w", vaultType, err) } - return v, nil + return &v, nil } func (v andotp) Entries() []vaults.Entry { diff --git a/internal/vaults/andotp/andotp_test.go b/internal/vaults/andotp/andotp_test.go index 040c24a..0c0d247 100644 --- a/internal/vaults/andotp/andotp_test.go +++ b/internal/vaults/andotp/andotp_test.go @@ -24,15 +24,15 @@ func TestOpen(t *testing.T) { for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { v, err := Open(tt.filename, []byte(tt.password)) - if tt.fails { - if err == nil { - t.Fatal("Open() expected error, got none") + if tt.fails { + if err == nil { + t.Fatal("Open() expected error, got none") + } + if tt.wantErr != nil && !errors.Is(err, tt.wantErr) { + t.Fatalf("Open() error = %v, want %v", err, tt.wantErr) + } + return } - if tt.wantErr != nil && !errors.Is(err, tt.wantErr) { - t.Fatalf("Open() error = %v, want %v", err, tt.wantErr) - } - return - } entries := v.Entries() if len(entries) != 1 { diff --git a/internal/vaults/keepass/keepass.go b/internal/vaults/keepass/keepass.go index 98b8e03..2e08b6d 100644 --- a/internal/vaults/keepass/keepass.go +++ b/internal/vaults/keepass/keepass.go @@ -18,6 +18,8 @@ var _ vaults.Vault = &keepass{} type keepass struct{ entries []gokeepasslib.Entry } +func init() { vaults.Register(vaultType, Open) } + func Open(filename string, pass []byte) (vaults.Vault, error) { v := keepass{entries: make([]gokeepasslib.Entry, 0)} db := gokeepasslib.NewDatabase() @@ -43,7 +45,7 @@ func Open(filename string, pass []byte) (vaults.Vault, error) { v.entries = append(v.entries, parseGroups(db.Content.Root.Groups)...) - return v, nil + return &v, nil } func (v keepass) Entries() []vaults.Entry { diff --git a/internal/vaults/protonpass/protonpass.go b/internal/vaults/protonpass/protonpass.go index 0f53e11..1216d5f 100644 --- a/internal/vaults/protonpass/protonpass.go +++ b/internal/vaults/protonpass/protonpass.go @@ -25,10 +25,10 @@ var ( ) type ( - envelope struct{ Vaults map[string]proton } + envelope struct{ Vaults map[string]protonvault } // protonvault only implements the essentials for reading OTP data. - proton struct { + protonvault struct { Name, Description string Items []struct { Data struct { @@ -43,14 +43,16 @@ type ( } ) +func init() { vaults.Register(vaultType, Open) } + func Open(filename string, pass []byte) (vaults.Vault, error) { b, err := read(filename) if err != nil { return nil, fmt.Errorf("%s: %s", vaultType, err) } - var e envelope - if e.IsPlain(b) { + var v envelope + if v.IsPlain(b) { return nil, vaults.ErrIsPlain } @@ -64,11 +66,11 @@ func Open(filename string, pass []byte) (vaults.Vault, error) { return nil, fmt.Errorf("%s: %s", vaultType, err) } - if err := json.Unmarshal(result.Bytes(), &e); err != nil { + if err := json.Unmarshal(result.Bytes(), &v); err != nil { return nil, fmt.Errorf("%s: %s", vaultType, err) } - return e, nil + return &v, nil } func (e envelope) Entries() []vaults.Entry { diff --git a/internal/vaults/registry.go b/internal/vaults/registry.go new file mode 100644 index 0000000..182b60e --- /dev/null +++ b/internal/vaults/registry.go @@ -0,0 +1,79 @@ +package vaults + +import ( + "context" + "fmt" + "slices" + "strings" +) + +const ( + ANDOTP VaultType = "andotp" + AEGIS VaultType = "aegis" + TWOFAS VaultType = "twofas" + STRATUM VaultType = "stratum" + KEEPASS VaultType = "keepass" + PROTON VaultType = "proton" +) + +type ( + OpenFn func(path string, pass []byte) (Vault, error) + VaultType string +) + +func (vt VaultType) String() string { return string(vt) } + +var registry = make(map[VaultType]OpenFn) + +// Register registers a vault type with the given open function. +func Register(vt VaultType, fn OpenFn) { registry[vt] = fn } + +// Open opens a vault of the given type at the given path using the given password. +func Open(ctx context.Context, path string, pass []byte, vt VaultType) (Vault, error) { + open, ok := registry[vt] + if !ok { + return nil, fmt.Errorf("vault type %q: not implemented", vt) + } + + type data struct { + v Vault + err error + } + + done := make(chan data, 1) + go func() { + defer func() { + for i := range pass { + pass[i] = 0 + } + }() + v, err := open(path, pass) + done <- data{v, err} + }() + + select { + case r := <-done: + return r.v, r.err + case <-ctx.Done(): + return nil, fmt.Errorf("open: operation timed out. wrong type?") + } +} + +// Types returns a slice of all registered vault types. +func Types() []VaultType { + var s []VaultType + for t := range registry { + s = append(s, t) + } + slices.Sort(s) + return s +} + +// StrTypes returns a comma-separated string of all registered vault types. +func StrTypes() string { + var s []string + for _, vt := range Types() { + s = append(s, vt.String()) + } + return strings.Join(s, ", ") +} diff --git a/internal/vaults/registry_test.go b/internal/vaults/registry_test.go new file mode 100644 index 0000000..47604b2 --- /dev/null +++ b/internal/vaults/registry_test.go @@ -0,0 +1,57 @@ +package vaults_test + +import ( + "context" + "reflect" + "strings" + "testing" + + _ "github.com/tjblackheart/andcli/v2/internal/vaults/aegis" + _ "github.com/tjblackheart/andcli/v2/internal/vaults/andotp" + _ "github.com/tjblackheart/andcli/v2/internal/vaults/keepass" + _ "github.com/tjblackheart/andcli/v2/internal/vaults/protonpass" + _ "github.com/tjblackheart/andcli/v2/internal/vaults/stratum" + _ "github.com/tjblackheart/andcli/v2/internal/vaults/twofas" + + "github.com/tjblackheart/andcli/v2/internal/vaults" +) + +func TestRegister(t *testing.T) { + for _, vt := range vaults.Types() { + _, err := vaults.Open(context.Background(), ".", nil, vt) + if err == nil { + t.Fatalf("%s: expected an error, got none", vt) + } + + if strings.Contains(err.Error(), "not implemented") { + t.Fatalf("%s: missing registered openFn: %s", vt, err) + } + } +} + +func TestTypes(t *testing.T) { + tests := []struct { + name string + want []vaults.VaultType + }{ + { + "returns defined types", + []vaults.VaultType{ + vaults.AEGIS, + vaults.ANDOTP, + vaults.KEEPASS, + vaults.PROTON, + vaults.STRATUM, + vaults.TWOFAS, + }, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + if got := vaults.Types(); !reflect.DeepEqual(got, tt.want) { + t.Errorf("Types() = %v, want %v", got, tt.want) + } + }) + } +} diff --git a/internal/vaults/stratum/stratum.go b/internal/vaults/stratum/stratum.go index d6bac80..6ef43a3 100644 --- a/internal/vaults/stratum/stratum.go +++ b/internal/vaults/stratum/stratum.go @@ -58,13 +58,15 @@ type ( } ) +func init() { vaults.Register(vaultType, Open) } + func Open(filename string, pass []byte) (vaults.Vault, error) { b, err := os.ReadFile(filename) if err != nil { return nil, err } - v := &stratum{Authenticators: make([]entry, 0)} + v := stratum{Authenticators: make([]entry, 0)} if v.IsPlain(b) { return nil, vaults.ErrIsPlain } @@ -91,7 +93,7 @@ func Open(filename string, pass []byte) (vaults.Vault, error) { } } - return v, nil + return &v, nil } func (v stratum) Entries() []vaults.Entry { diff --git a/internal/vaults/stratum/stratum_test.go b/internal/vaults/stratum/stratum_test.go index 5895293..bedaf57 100644 --- a/internal/vaults/stratum/stratum_test.go +++ b/internal/vaults/stratum/stratum_test.go @@ -30,15 +30,15 @@ func TestOpen(t *testing.T) { for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { v, err := Open(tt.filename, []byte(tt.password)) - if tt.fails { - if err == nil { - t.Fatal("Open() expected error, got none") - } - if tt.wantErr != nil && !errors.Is(err, tt.wantErr) { - t.Fatalf("Open() error = %v, want %v", err, tt.wantErr) + if tt.fails { + if err == nil { + t.Fatal("Open() expected error, got none") + } + if tt.wantErr != nil && !errors.Is(err, tt.wantErr) { + t.Fatalf("Open() error = %v, want %v", err, tt.wantErr) + } + return } - return - } entries := v.Entries() if len(entries) != 3 { diff --git a/internal/vaults/twofas/twofas.go b/internal/vaults/twofas/twofas.go index c96bf12..a7c041a 100644 --- a/internal/vaults/twofas/twofas.go +++ b/internal/vaults/twofas/twofas.go @@ -64,6 +64,8 @@ type ( } ) +func init() { vaults.Register(vaultType, Open) } + func Open(filename string, pass []byte) (vaults.Vault, error) { var v twofas @@ -94,7 +96,7 @@ func Open(filename string, pass []byte) (vaults.Vault, error) { return nil, fmt.Errorf("%s: %w", vaultType, err) } - return v, nil + return &v, nil } func (v twofas) Entries() []vaults.Entry { diff --git a/internal/vaults/twofas/twofas_test.go b/internal/vaults/twofas/twofas_test.go index e39ea50..c1d692f 100644 --- a/internal/vaults/twofas/twofas_test.go +++ b/internal/vaults/twofas/twofas_test.go @@ -25,15 +25,15 @@ func TestOpen(t *testing.T) { for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { v, err := Open(tt.filename, []byte(tt.password)) - if tt.fails { - if err == nil { - t.Fatal("Open() expected error, got none") + if tt.fails { + if err == nil { + t.Fatal("Open() expected error, got none") + } + if tt.wantErr != nil && !errors.Is(err, tt.wantErr) { + t.Fatalf("Open() error = %v, want %v", err, tt.wantErr) + } + return } - if tt.wantErr != nil && !errors.Is(err, tt.wantErr) { - t.Fatalf("Open() error = %v, want %v", err, tt.wantErr) - } - return - } entries := v.Entries() if len(entries) != 1 { diff --git a/internal/vaults/vault.go b/internal/vaults/vault.go index d352a56..3fb4ba3 100644 --- a/internal/vaults/vault.go +++ b/internal/vaults/vault.go @@ -2,7 +2,6 @@ package vaults import ( "errors" - "strings" ) // Vault is the basic skeleton of a vault implementation. @@ -11,39 +10,4 @@ type Vault interface { IsPlain([]byte) bool } -// Type is an implemented vault type name. -type Type string - -func (t Type) String() string { return string(t) } - -var ErrIsPlain error = errors.New("unencrypted vaults are not supported") - -const ( - ANDOTP Type = "andotp" - AEGIS Type = "aegis" - TWOFAS Type = "twofas" - STRATUM Type = "stratum" - KEEPASS Type = "keepass" - PROTON Type = "proton" -) - -// Returns a list containing the implemented types. -func Types() []Type { - return []Type{ - ANDOTP, - AEGIS, - TWOFAS, - STRATUM, - KEEPASS, - PROTON, - } -} - -// StrTypes returns a concatenated string of all defined types. -func StrTypes() string { - var s []string - for _, t := range Types() { - s = append(s, t.String()) - } - return strings.Join(s, ", ") -} +var ErrIsPlain error = errors.New("plaintext vaults are unsupported") diff --git a/internal/vaults/vault_test.go b/internal/vaults/vault_test.go deleted file mode 100644 index 3f41399..0000000 --- a/internal/vaults/vault_test.go +++ /dev/null @@ -1,35 +0,0 @@ -package vaults_test - -import ( - "reflect" - "testing" - - "github.com/tjblackheart/andcli/v2/internal/vaults" -) - -func TestTypes(t *testing.T) { - tests := []struct { - name string - want []vaults.Type - }{ - { - "returns defined types", - []vaults.Type{ - vaults.ANDOTP, - vaults.AEGIS, - vaults.TWOFAS, - vaults.STRATUM, - vaults.KEEPASS, - vaults.PROTON, - }, - }, - } - - for _, tt := range tests { - t.Run(tt.name, func(t *testing.T) { - if got := vaults.Types(); !reflect.DeepEqual(got, tt.want) { - t.Errorf("Types() = %v, want %v", got, tt.want) - } - }) - } -} From 4227e2f107f35e3d2da6a45e0eae5b133f84b719 Mon Sep 17 00:00:00 2001 From: tjblackheart Date: Mon, 3 Aug 2026 07:33:05 +0200 Subject: [PATCH 2/2] fix: vault type usage list --- internal/config/flags.go | 3 ++- 1 file changed, 2 insertions(+), 1 deletion(-) diff --git a/internal/config/flags.go b/internal/config/flags.go index b3a078c..5a81560 100644 --- a/internal/config/flags.go +++ b/internal/config/flags.go @@ -14,7 +14,7 @@ import ( var ( set = flag.NewFlagSet("default", flag.ExitOnError) vfile = set.StringP("file", "f", "", "Path to the encrypted vault (deprecated: Pass the filename directly)") - vtype = set.StringP("type", "t", "", fmt.Sprintf("Vault type (%s)", vaults.StrTypes())) + vtype = set.StringP("type", "t", "", "Vault types") cmd = set.StringP("clipboard-cmd", "c", "", "A custom clipboard command, including args (xclip, wl-copy, pbcopy etc.)") pwstdin = set.Bool("passwd-stdin", false, "Read the vault password from stdin. If set, skips the password input.") query = set.StringP("query", "q", "", "Query the vault directly and skip TUI functionality") @@ -27,6 +27,7 @@ var ( // Parses given flags into the existing config. func (cfg *Config) parseFlags() error { set.Usage = func() { usage(true) } + set.Lookup("type").Usage = fmt.Sprintf("Vault type (%s)", vaults.StrTypes()) if err := set.Parse(os.Args[1:]); err != nil { log.Printf("%s: %s", buildinfo.AppName, err)