package provider import ( "context" "errors" "strings" "testing" "cloud.campbellwireless.net/git/barkstack/treatvault/internal/vault" ) const ( testVaultID = "0011223344556677" testRevision = "00112233445566778899aabbccddeeff" ) func serviceSpecJSON(names string, secrets string) string { return `{"Name":"barkstack-pawsql","Labels":{"io.barkstack.treatvault.secrets":"true","io.barkstack.treatvault.names":"` + names + `"},"TaskTemplate":{"ContainerSpec":{"Secrets":[` + secrets + `]}}}` } func TestSyncCreatesVersionedDockerSecret(t *testing.T) { runner := &fakeRunner{responses: []fakeResponse{{}, {}, {}}} provider := Docker{Runner: runner} snapshot := vault.Snapshot{ Version: 1, VaultID: testVaultID, Secrets: map[string]vault.Record{"database_password": {Revision: testRevision, Value: []byte("super-secret")}}, } states, err := provider.Sync(context.Background(), snapshot) if err != nil { t.Fatal(err) } const physical = "barkstack_tv_" + testVaultID + "_" + testRevision if len(states) != 1 || states[0].PhysicalName != physical { t.Fatalf("states = %#v", states) } create := runner.callContaining(t, "secret create") for _, want := range []string{ "--label io.barkstack.treatvault=true", "--label io.barkstack.treatvault.name=database_password", physical + " -", } { if !strings.Contains(create.args, want) { t.Errorf("create args = %q, missing %q", create.args, want) } } if string(create.input) != "super-secret" || strings.Contains(create.args, "super-secret") { t.Fatalf("secret was not passed exclusively on stdin: %#v", create) } } func TestSyncRotatesMountedSecretWithoutChangingTarget(t *testing.T) { const ( oldRevision = "aaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaa" oldPhysical = "barkstack_tv_" + testVaultID + "_" + oldRevision newPhysical = "barkstack_tv_" + testVaultID + "_" + testRevision ) runner := &fakeRunner{responses: []fakeResponse{ {output: "old-id\n"}, // secret ls {output: secretSpecJSON("database_password", oldRevision, oldPhysical, testVaultID)}, // secret inspect {}, // secret create (new revision) {output: "service-id\n"}, // service ls {output: serviceSpecJSON("database_password", mountJSON(oldPhysical, "barkstack_database_password"))}, // service inspect {}, // service update {err: errors.New("secret is in use by old task")}, // secret rm }} provider := Docker{Runner: runner} states, err := provider.Sync(context.Background(), vault.Snapshot{ Version: 1, VaultID: testVaultID, Secrets: map[string]vault.Record{"database_password": {Revision: testRevision, Value: []byte("rotated")}}, }) if err != nil { t.Fatal(err) } if len(states) != 1 || !states[0].InUse { t.Fatalf("states = %#v", states) } update := runner.callContaining(t, "service update") for _, want := range []string{ "--secret-rm " + oldPhysical, "--secret-add source=" + newPhysical + ",target=barkstack_database_password", "barkstack-pawsql", } { if !strings.Contains(update.args, want) { t.Errorf("update args = %q, missing %q", update.args, want) } } } func TestSyncMountsMissingDesiredSecretOnConsumer(t *testing.T) { const physical = "barkstack_tv_" + testVaultID + "_" + testRevision runner := &fakeRunner{responses: []fakeResponse{ {}, // secret ls (empty vault listing) {}, // secret create {output: "service-id\n"}, {output: serviceSpecJSON("database_password", "")}, {}, // service update }} provider := Docker{Runner: runner} states, err := provider.Sync(context.Background(), vault.Snapshot{ Version: 1, VaultID: testVaultID, Secrets: map[string]vault.Record{"database_password": {Revision: testRevision, Value: []byte("value")}}, }) if err != nil { t.Fatal(err) } update := runner.callContaining(t, "service update") if want := "--secret-add source=" + physical + ",target=barkstack_database_password"; !strings.Contains(update.args, want) { t.Errorf("update args = %q, missing %q", update.args, want) } if !states[0].InUse { t.Fatalf("states = %#v, want InUse", states) } } func TestSyncReportsConfiguredSecretMissingFromVault(t *testing.T) { runner := &fakeRunner{responses: []fakeResponse{ {}, // secret ls (empty) {output: "service-id\n"}, {output: serviceSpecJSON("database_password", "")}, }} _, err := (Docker{Runner: runner}).Sync(context.Background(), vault.Snapshot{Version: 1, VaultID: testVaultID, Secrets: map[string]vault.Record{}}) if !errors.Is(err, ErrSecretInUse) { t.Fatalf("Sync() error = %v", err) } } func TestSyncKeepsOrphanedMountAndReportsGap(t *testing.T) { const oldPhysical = "barkstack_tv_" + testVaultID + "_aaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaa" runner := &fakeRunner{responses: []fakeResponse{ {output: "old-id\n"}, {output: secretSpecJSON("removed_secret", "aaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaa", oldPhysical, testVaultID)}, {output: "service-id\n"}, {output: serviceSpecJSON("database_password", mountJSON(oldPhysical, "barkstack_database_password"))}, {err: errors.New("secret is in use")}, }} _, err := (Docker{Runner: runner}).Sync(context.Background(), vault.Snapshot{Version: 1, VaultID: testVaultID, Secrets: map[string]vault.Record{}}) if !errors.Is(err, ErrSecretInUse) { t.Fatalf("Sync() error = %v", err) } } func secretSpecJSON(name, revision, physical, vaultID string) string { return `{"Name":"` + physical + `","Labels":{"io.barkstack.treatvault":"true","io.barkstack.treatvault.name":"` + name + `","io.barkstack.treatvault.vault":"` + vaultID + `","io.barkstack.treatvault.revision":"` + revision + `"}}` } func mountJSON(source, target string) string { return `{"SecretName":"` + source + `","File":{"Name":"` + target + `"}}` } type fakeResponse struct { output string err error } type fakeCall struct { args string input []byte } type fakeRunner struct { responses []fakeResponse calls []fakeCall } func (r *fakeRunner) Run(_ context.Context, input []byte, args ...string) (string, error) { r.calls = append(r.calls, fakeCall{args: strings.Join(args, " "), input: append([]byte(nil), input...)}) if len(r.responses) == 0 { return "", errors.New("unexpected Docker command") } response := r.responses[0] r.responses = r.responses[1:] return response.output, response.err } func (r *fakeRunner) callContaining(t *testing.T, prefix string) fakeCall { t.Helper() for _, call := range r.calls { if strings.HasPrefix(call.args, prefix) { return call } } t.Fatalf("calls = %#v, missing %q", r.calls, prefix) return fakeCall{} }