telegabber/e2ee/store/badgerstore/store_test.go
2026-07-28 20:32:56 -04:00

259 lines
7.1 KiB
Go

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 != 1 {
t.Fatalf("expected first block to start at 1 (XEP-0384 requires positive ids), got %d", start1)
}
start2, err := store.NextPreKeyIDs(5)
if err != nil {
t.Fatalf("NextPreKeyIDs: %v", err)
}
if start2 != 11 {
t.Fatalf("expected second block to start at 11, 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("<devices/>")); err != nil {
t.Fatalf("SaveRemoteDeviceListCache: %v", err)
}
doc, found, err := store.RemoteDeviceListCache("bob@example.com")
if err != nil || !found || string(doc) != "<devices/>" {
t.Fatalf("RemoteDeviceListCache: doc=%q found=%v err=%v", doc, found, err)
}
}