package backup import ( "bytes" "context" "errors" "io" "os" "path/filepath" "testing" "time" "go.mongodb.org/mongo-driver/v2/bson" "go.mongodb.org/mongo-driver/v2/mongo" ) func seed(t *testing.T, client *mongo.Client, dbName string) { t.Helper() ctx := context.Background() db := client.Database(dbName) if _, err := db.Collection("servers").InsertMany(ctx, []any{ bson.M{"_id": bson.NewObjectID(), "name": "alpha", "instance_id": "i1"}, bson.M{"_id": bson.NewObjectID(), "name": "beta", "instance_id": "i1"}, }); err != nil { t.Fatalf("insert servers: %v", err) } if _, err := db.Collection("audit_logs").InsertOne(ctx, bson.M{"action": "login"}); err != nil { t.Fatalf("insert audit_logs: %v", err) } } func dumpToFile(t *testing.T, opt DumpOptions) (string, Manifest) { t.Helper() path := filepath.Join(t.TempDir(), "out.tar.gz") f, err := os.Create(path) if err != nil { t.Fatalf("create: %v", err) } opt.Out = f m, err := Dump(context.Background(), opt) if cerr := f.Close(); cerr != nil { t.Fatalf("close: %v", cerr) } if err != nil { t.Fatalf("Dump: %v", err) } return path, m } func TestDumpEnumeratesEveryCollection(t *testing.T) { client, dbName := testDB(t) seed(t, client, dbName) _, m := dumpToFile(t, DumpOptions{ Client: client, Database: dbName, KeyHex: validKeyHex, VantageVersion: "test", }) if _, ok := m.Collection("servers"); !ok { t.Fatal("servers missing from the manifest") } if _, ok := m.Collection("audit_logs"); !ok { t.Fatal("audit_logs missing; enumeration must not filter by a hardcoded list") } servers, _ := m.Collection("servers") if servers.Documents != 2 { t.Fatalf("servers documents %d, want 2", servers.Documents) } if m.MongoDB != dbName { t.Fatalf("manifest database %q, want %q", m.MongoDB, dbName) } if m.MongoServerVersion == "" { t.Fatal("manifest records no MongoDB server version") } if m.Hostname == "" { t.Fatal("manifest records no hostname") } } func TestDumpRecordsKeyFingerprint(t *testing.T) { client, dbName := testDB(t) seed(t, client, dbName) _, m := dumpToFile(t, DumpOptions{Client: client, Database: dbName, KeyHex: validKeyHex}) want, err := FingerprintHex(validKeyHex) if err != nil { t.Fatalf("FingerprintHex: %v", err) } if m.KeyFingerprint == nil || *m.KeyFingerprint != want { t.Fatalf("fingerprint %v, want %s", m.KeyFingerprint, want) } } func TestDumpRefusesWithoutAKey(t *testing.T) { client, dbName := testDB(t) seed(t, client, dbName) var buf bytes.Buffer _, err := Dump(context.Background(), DumpOptions{ Client: client, Database: dbName, Out: &buf, }) if !errors.Is(err, ErrNoKey) { t.Fatalf("got %v, want ErrNoKey", err) } if buf.Len() != 0 { t.Fatal("refusal must happen before anything is written") } } func TestDumpAllowNoKeyStampsNull(t *testing.T) { client, dbName := testDB(t) seed(t, client, dbName) _, m := dumpToFile(t, DumpOptions{Client: client, Database: dbName, AllowNoKey: true}) if m.KeyFingerprint != nil { t.Fatalf("want a null fingerprint, got %v", *m.KeyFingerprint) } } func TestDumpRejectsMalformedKey(t *testing.T) { client, dbName := testDB(t) var buf bytes.Buffer _, err := Dump(context.Background(), DumpOptions{ Client: client, Database: dbName, KeyHex: "nonsense", Out: &buf, }) if !errors.Is(err, ErrBadKey) { t.Fatalf("got %v, want ErrBadKey", err) } } func TestDumpExcludeIsRecordedAndOmitted(t *testing.T) { client, dbName := testDB(t) seed(t, client, dbName) path, m := dumpToFile(t, DumpOptions{ Client: client, Database: dbName, KeyHex: validKeyHex, Exclude: []string{"audit_logs"}, }) if _, ok := m.Collection("audit_logs"); ok { t.Fatal("excluded collection is in the manifest's collection list") } if len(m.Excluded) != 1 || m.Excluded[0] != "audit_logs" { t.Fatalf("excluded recorded as %v", m.Excluded) } r, err := Open(path) if err != nil { t.Fatalf("Open: %v", err) } defer r.Close() if _, err := r.OpenCollection("audit_logs"); err == nil { t.Fatal("excluded collection is present in the archive") } } func TestDumpPreservesAwkwardBSONTypes(t *testing.T) { client, dbName := testDB(t) ctx := context.Background() dec, err := bson.ParseDecimal128("1234.5678") if err != nil { t.Fatalf("ParseDecimal128: %v", err) } doc := bson.M{ "_id": bson.NewObjectID(), "decimal": dec, "when": bson.NewDateTimeFromTime(mustTime(t)), "binary": bson.Binary{Subtype: 0x00, Data: []byte{0x01, 0x02, 0x03}}, "nothing": nil, "nested": bson.A{bson.M{"deep": bson.A{1, 2, 3}}}, } if _, err := client.Database(dbName).Collection("odd").InsertOne(ctx, doc); err != nil { t.Fatalf("insert: %v", err) } path, _ := dumpToFile(t, DumpOptions{Client: client, Database: dbName, KeyHex: validKeyHex}) original, err := client.Database(dbName).Collection("odd").FindOne(ctx, bson.M{}).Raw() if err != nil { t.Fatalf("read back: %v", err) } r, err := Open(path) if err != nil { t.Fatalf("Open: %v", err) } defer r.Close() rc, err := r.OpenCollection("odd") if err != nil { t.Fatalf("OpenCollection: %v", err) } defer rc.Close() archived, err := io.ReadAll(rc) if err != nil { t.Fatalf("read: %v", err) } if !bytes.Equal(archived, []byte(original)) { t.Fatal("archived BSON differs from what the driver returned") } } func mustTime(t *testing.T) time.Time { t.Helper() return time.Date(2026, 9, 7, 12, 0, 0, 0, time.UTC) }