package badgerstore_test import ( "testing" "dev.narayana.im/narayana/telegabber/e2ee/omemo/libsignal" "dev.narayana.im/narayana/telegabber/e2ee/store/badgerstore" ) func openTestStore(t *testing.T) *badgerstore.Store { t.Helper() db, err := badgerstore.Open(t.TempDir(), nil) if err != nil { t.Fatalf("Open: %v", err) } t.Cleanup(func() { db.Close() }) ctx, err := libsignal.NewContext() if err != nil { t.Fatalf("NewContext: %v", err) } t.Cleanup(ctx.Close) return db.Store(ctx, "testlogin", "1234@example.com") } func TestIdentityKeyPairRoundTrip(t *testing.T) { store := openTestStore(t) if has, err := store.HasIdentityKeyPair(); err != nil || has { t.Fatalf("expected no identity key pair yet, has=%v err=%v", has, err) } // A fresh context is fine for generation - the record format doesn't // depend on which *Context instance produced it. ctx, err := libsignal.NewContext() if err != nil { t.Fatalf("NewContext: %v", err) } defer ctx.Close() idKeyPair, err := libsignal.GenerateIdentityKeyPair(ctx) if err != nil { t.Fatalf("GenerateIdentityKeyPair: %v", err) } wantPublic, _, err := libsignal.SplitIdentityKeyPair(ctx, idKeyPair.Record) if err != nil { t.Fatalf("SplitIdentityKeyPair: %v", err) } if err := store.SaveIdentityKeyPair(idKeyPair.Record); err != nil { t.Fatalf("SaveIdentityKeyPair: %v", err) } if has, err := store.HasIdentityKeyPair(); err != nil || !has { t.Fatalf("expected identity key pair to be saved, has=%v err=%v", has, err) } gotPublic, _, err := store.GetIdentityKeyPair() if err != nil { t.Fatalf("GetIdentityKeyPair: %v", err) } if string(gotPublic) != string(wantPublic) { t.Fatalf("public key mismatch after round trip") } } func TestRegistrationIDRoundTrip(t *testing.T) { store := openTestStore(t) if err := store.SaveRegistrationID(4242); err != nil { t.Fatalf("SaveRegistrationID: %v", err) } got, err := store.GetLocalRegistrationID() if err != nil { t.Fatalf("GetLocalRegistrationID: %v", err) } if got != 4242 { t.Fatalf("got %d, want 4242", got) } } func TestPreKeyRoundTrip(t *testing.T) { store := openTestStore(t) if store.ContainsPreKey(7) { t.Fatalf("expected prekey 7 to not exist yet") } if err := store.StorePreKey(7, []byte("prekey-record")); err != nil { t.Fatalf("StorePreKey: %v", err) } if !store.ContainsPreKey(7) { t.Fatalf("expected prekey 7 to exist") } record, found, err := store.LoadPreKey(7) if err != nil || !found { t.Fatalf("LoadPreKey: found=%v err=%v", found, err) } if string(record) != "prekey-record" { t.Fatalf("got %q", record) } if err := store.RemovePreKey(7); err != nil { t.Fatalf("RemovePreKey: %v", err) } if store.ContainsPreKey(7) { t.Fatalf("expected prekey 7 to be removed") } } func TestSignedPreKeyRoundTrip(t *testing.T) { store := openTestStore(t) if err := store.StoreSignedPreKey(3, []byte("signed-record")); err != nil { t.Fatalf("StoreSignedPreKey: %v", err) } if !store.ContainsSignedPreKey(3) { t.Fatalf("expected signed prekey 3 to exist") } record, found, err := store.LoadSignedPreKey(3) if err != nil || !found || string(record) != "signed-record" { t.Fatalf("LoadSignedPreKey: record=%q found=%v err=%v", record, found, err) } if err := store.RemoveSignedPreKey(3); err != nil { t.Fatalf("RemoveSignedPreKey: %v", err) } if store.ContainsSignedPreKey(3) { t.Fatalf("expected signed prekey 3 to be removed") } } func TestNextPreKeyIDs(t *testing.T) { store := openTestStore(t) start1, err := store.NextPreKeyIDs(10) if err != nil { t.Fatalf("NextPreKeyIDs: %v", err) } if start1 != 0 { t.Fatalf("expected first block to start at 0, got %d", start1) } start2, err := store.NextPreKeyIDs(5) if err != nil { t.Fatalf("NextPreKeyIDs: %v", err) } if start2 != 10 { t.Fatalf("expected second block to start at 10, got %d", start2) } } func TestSessionRoundTrip(t *testing.T) { store := openTestStore(t) addr := libsignal.Address{Name: "bob@example.com", DeviceID: 1} if store.ContainsSession(addr) { t.Fatalf("expected no session yet") } if err := store.StoreSession(addr, []byte("session-record")); err != nil { t.Fatalf("StoreSession: %v", err) } if !store.ContainsSession(addr) { t.Fatalf("expected session to exist") } record, found, err := store.LoadSession(addr) if err != nil || !found || string(record) != "session-record" { t.Fatalf("LoadSession: record=%q found=%v err=%v", record, found, err) } otherAddr := libsignal.Address{Name: "bob@example.com", DeviceID: 2} if err := store.StoreSession(otherAddr, []byte("session-record-2")); err != nil { t.Fatalf("StoreSession (device 2): %v", err) } ids, err := store.GetSubDeviceSessions("bob@example.com") if err != nil { t.Fatalf("GetSubDeviceSessions: %v", err) } if len(ids) != 2 { t.Fatalf("expected 2 device sessions, got %d (%v)", len(ids), ids) } count, err := store.DeleteAllSessions("bob@example.com") if err != nil { t.Fatalf("DeleteAllSessions: %v", err) } if count != 2 { t.Fatalf("expected to delete 2 sessions, deleted %d", count) } if store.ContainsSession(addr) || store.ContainsSession(otherAddr) { t.Fatalf("expected both sessions to be gone") } } func TestIsTrustedIdentityTOFU(t *testing.T) { store := openTestStore(t) addr := libsignal.Address{Name: "bob@example.com", DeviceID: 1} key1 := []byte("identity-key-v1") trusted, err := store.IsTrustedIdentity(addr, key1) if err != nil { t.Fatalf("IsTrustedIdentity (unseen): %v", err) } if !trusted { t.Fatalf("expected an unseen identity to be trusted (TOFU)") } if err := store.SaveIdentity(addr, key1); err != nil { t.Fatalf("SaveIdentity: %v", err) } trusted, err = store.IsTrustedIdentity(addr, key1) if err != nil || !trusted { t.Fatalf("expected the same key to remain trusted: trusted=%v err=%v", trusted, err) } key2 := []byte("identity-key-v2-different") trusted, err = store.IsTrustedIdentity(addr, key2) if err != nil { t.Fatalf("IsTrustedIdentity (mismatched): %v", err) } if trusted { t.Fatalf("expected a mismatched key to be untrusted, not silently re-trusted") } } func TestEnabledFlag(t *testing.T) { store := openTestStore(t) enabled, err := store.Enabled() if err != nil { t.Fatalf("Enabled (default): %v", err) } if enabled { t.Fatalf("expected OMEMO to be off by default") } if err := store.SetEnabled(true); err != nil { t.Fatalf("SetEnabled: %v", err) } enabled, err = store.Enabled() if err != nil || !enabled { t.Fatalf("expected OMEMO to be enabled after SetEnabled(true): enabled=%v err=%v", enabled, err) } } func TestRemoteDeviceListCache(t *testing.T) { store := openTestStore(t) if _, found, err := store.RemoteDeviceListCache("bob@example.com"); err != nil || found { t.Fatalf("expected no cached device list yet: found=%v err=%v", found, err) } if err := store.SaveRemoteDeviceListCache("bob@example.com", []byte("")); err != nil { t.Fatalf("SaveRemoteDeviceListCache: %v", err) } doc, found, err := store.RemoteDeviceListCache("bob@example.com") if err != nil || !found || string(doc) != "" { t.Fatalf("RemoteDeviceListCache: doc=%q found=%v err=%v", doc, found, err) } }