branch: elpa/jabber
commit 66a5ed51e3a0a794d02b6c1cace4b5aef43602cf
Author: Thanos Apollo <[email protected]>
Commit: Thanos Apollo <[email protected]>

    omemo: Make skipped-key consumption transactional
---
 src/jabber-omemo-core.c           | 53 +++++++++++++++++++++++++++++++--------
 src/picomemo/omemo.c              | 11 ++++++++
 src/picomemo/omemo.h              |  3 +++
 tests/jabber-test-omemo-module.el | 46 +++++++++++++++++++++++++++++++++
 4 files changed, 103 insertions(+), 10 deletions(-)

diff --git a/src/jabber-omemo-core.c b/src/jabber-omemo-core.c
index 27703f2a17..c126332558 100644
--- a/src/jabber-omemo-core.c
+++ b/src/jabber-omemo-core.c
@@ -65,6 +65,14 @@ struct session_skipped {
 static struct session_skipped *g_skipped;
 static size_t g_skipped_count, g_skipped_cap;
 
+static void
+skipped_clear(void *ptr, size_t size)
+{
+    volatile unsigned char *p = ptr;
+    while (size--)
+        *p++ = 0;
+}
+
 static struct session_skipped *
 skipped_find(struct omemoSession *s, int create)
 {
@@ -94,8 +102,8 @@ skipped_drop(struct omemoSession *s)
     for (size_t i = 0; i < g_skipped_count; i++) {
         if (g_skipped[i].session == s) {
             if (g_skipped[i].keys) {
-                memset(g_skipped[i].keys, 0,
-                       g_skipped[i].cap * sizeof(struct skipped_key));
+                skipped_clear(g_skipped[i].keys,
+                              g_skipped[i].cap * sizeof(struct skipped_key));
                 free(g_skipped[i].keys);
             }
             g_skipped[i] = g_skipped[--g_skipped_count];
@@ -135,14 +143,29 @@ int omemoLoadMessageKey(struct omemoSession *s, struct 
omemoMessageKey *k)
         struct skipped_key *sk = &e->keys[i];
         if (sk->nr == k->nr && !memcmp(sk->dh, k->dh, 32)) {
             memcpy(k->mk, sk->mk, 32);
-            /* Single use: replace with the last entry and zero it. */
+            return 0;
+        }
+    }
+    return 1; /* not found */
+}
+
+int omemoRemoveMessageKey(struct omemoSession *s,
+                          const struct omemoMessageKey *k)
+{
+    struct session_skipped *e = skipped_find(s, 0);
+    if (!e)
+        return OMEMO_ESTORE;
+    for (size_t i = 0; i < e->count; i++) {
+        struct skipped_key *sk = &e->keys[i];
+        if (sk->nr == k->nr && !memcmp(sk->dh, k->dh, 32)) {
             e->keys[i] = e->keys[e->count - 1];
-            memset(&e->keys[e->count - 1], 0, sizeof(struct skipped_key));
+            skipped_clear(&e->keys[e->count - 1],
+                          sizeof(struct skipped_key));
             e->count--;
             return 0;
         }
     }
-    return 1; /* not found */
+    return OMEMO_ESTORE;
 }
 
 int omemoStoreMessageKey(struct omemoSession *s,
@@ -835,14 +858,22 @@ F_session_skipped_keys(emacs_env *env, ptrdiff_t nargs, 
emacs_value *args,
     if (env->non_local_exit_check(env))
         return Qnil_v;
 
+    struct session_skipped *e = skipped_find(session, 0);
+    if (!e || !e->count)
+        return Qnil_v;
+    size_t count = e->count;
+    struct skipped_key *snapshot = malloc(count * sizeof *snapshot);
+    if (!snapshot && count) {
+        signal_error(env, OMEMO_ESTORE, "cannot snapshot skipped keys");
+        return Qnil_v;
+    }
+    memcpy(snapshot, e->keys, count * sizeof *snapshot);
+
     emacs_value Qlist = env->intern(env, "list");
     emacs_value Qcons = env->intern(env, "cons");
     emacs_value result = Qnil_v;
-    struct session_skipped *e = skipped_find(session, 0);
-    if (!e)
-        return result;
-    for (size_t i = e->count; i > 0; i--) {
-        struct skipped_key *sk = &e->keys[i - 1];
+    for (size_t i = count; i > 0; i--) {
+        struct skipped_key *sk = &snapshot[i - 1];
         emacs_value entry_args[] = {
             env->make_integer(env, sk->nr),
             make_unibyte(env, sk->dh, 32),
@@ -852,6 +883,8 @@ F_session_skipped_keys(emacs_env *env, ptrdiff_t nargs, 
emacs_value *args,
         emacs_value cons_args[] = { entry, result };
         result = env->funcall(env, Qcons, 2, cons_args);
     }
+    skipped_clear(snapshot, count * sizeof *snapshot);
+    free(snapshot);
     return result;
 }
 
diff --git a/src/picomemo/omemo.c b/src/picomemo/omemo.c
index d76a01a78d..44f63b5844 100644
--- a/src/picomemo/omemo.c
+++ b/src/picomemo/omemo.c
@@ -104,6 +104,13 @@ int WEAK omemoLoadMessageKey(struct omemoSession *s,
   return 1;
 }
 
+int WEAK omemoRemoveMessageKey(struct omemoSession *s,
+                               const struct omemoMessageKey *sk) {
+  (void)s;
+  (void)sk;
+  return 0;
+}
+
 int WEAK omemoStoreMessageKey(struct omemoSession *s,
                               const struct omemoMessageKey *sk,
                               uint64_t n) {
@@ -877,11 +884,13 @@ static int DecryptKeyImpl(struct omemoSession *session,
 
   omemoKey mk;
   struct omemoMessageKey mkey = {0};
+  bool loadedmkey = false;
   memcpy(mkey.dh, headerdh, 32);
   mkey.nr = headern;
   int r;
   if (!(r = omemoLoadMessageKey(session, &mkey))) {
     memcpy(mk, mkey.mk, 32);
+    loadedmkey = true;
   } else if (r < 0) {
     return r;
   } else {
@@ -918,6 +927,8 @@ static int DecryptKeyImpl(struct omemoSession *session,
   uint8_t pad = tmp[encn - 1];
   if (pad > 16 || pad > encn || encn - pad > *keyn)
     return OMEMO_ECORRUPT;
+  if (loadedmkey && omemoRemoveMessageKey(session, &mkey))
+    return OMEMO_ESTORE;
   memcpy(key, tmp, encn - pad);
   *keyn = encn - pad;
   session->init = SESSION_READY;
diff --git a/src/picomemo/omemo.h b/src/picomemo/omemo.h
index 4ecd2f662a..41d0f0b080 100644
--- a/src/picomemo/omemo.h
+++ b/src/picomemo/omemo.h
@@ -137,6 +137,9 @@ typedef int (*omemoRandomCallback)(void *p, size_t n);
 int omemoLoadMessageKey(struct omemoSession *s,
                         struct omemoMessageKey *sk);
 
+int omemoRemoveMessageKey(struct omemoSession *s,
+                          const struct omemoMessageKey *sk);
+
 int omemoStoreMessageKey(struct omemoSession *s,
                          const struct omemoMessageKey *sk,
                          uint64_t n);
diff --git a/tests/jabber-test-omemo-module.el 
b/tests/jabber-test-omemo-module.el
index 5d1440eac4..d8a9d088e0 100644
--- a/tests/jabber-test-omemo-module.el
+++ b/tests/jabber-test-omemo-module.el
@@ -537,6 +537,24 @@ Alice has initiated a session towards Bob's bundle."
     (jabber-omemo--session-set-skipped-keys session nil)
     (should (null (jabber-omemo--session-skipped-keys session)))))
 
+(ert-deftest jabber-test-omemo-module-skipped-keys-survive-finalizers ()
+  "Enumerating skipped keys is safe while other sessions are finalized."
+  (let* ((session (jabber-omemo--make-session))
+         (keys (list (list 3 (make-string 32 ?d) (make-string 32 ?m))
+                     (list 7 (make-string 32 ?e) (make-string 32 ?n)))))
+    (jabber-omemo--session-set-skipped-keys session keys)
+    (dotimes (i 200)
+      (let ((disposable (jabber-omemo--make-session)))
+        (jabber-omemo--session-set-skipped-keys
+         disposable
+         (list (list i
+                     (make-string 32 (+ ?a (% i 26)))
+                     (make-string 32 (+ ?A (% i 26))))))))
+    (let ((gc-cons-threshold 1))
+      (should (equal keys (jabber-omemo--session-skipped-keys session))))
+    (garbage-collect)
+    (should (equal keys (jabber-omemo--session-skipped-keys session)))))
+
 (ert-deftest jabber-test-omemo-module-out-of-order-decrypt ()
   "A message skipped over in the ratchet still decrypts afterwards."
   (jabber-test-omemo-module--with-session-pair
@@ -561,6 +579,34 @@ Alice has initiated a session towards Bob's bundle."
                            (plist-get m2 :pre-key-p) (plist-get m2 :data))))
       (should (null (jabber-omemo--session-skipped-keys bob-session))))))
 
+(ert-deftest jabber-test-omemo-module-failed-decrypt-keeps-skipped-key ()
+  "A corrupted late message does not consume its skipped key."
+  (jabber-test-omemo-module--with-session-pair
+    (let* ((k1 (make-string 32 ?1))
+           (k2 (make-string 32 ?2))
+           (m1 (jabber-omemo--encrypt-key alice-session k1))
+           (m2 (jabber-omemo--encrypt-key alice-session k2))
+           (bob-session (jabber-omemo--make-session)))
+      (jabber-omemo--decrypt-key
+       bob-session bob (plist-get m2 :pre-key-p) (plist-get m2 :data))
+      (let* ((corrupt (copy-sequence (plist-get m1 :data)))
+             (last (1- (length corrupt))))
+        (aset corrupt last (logxor 1 (aref corrupt last)))
+        (should-error
+         (jabber-omemo--decrypt-key
+          bob-session bob (plist-get m1 :pre-key-p) corrupt)
+         :type 'jabber-omemo-error))
+      (should (= 1 (length (jabber-omemo--session-skipped-keys bob-session))))
+      (should (string= k1 (jabber-omemo--decrypt-key
+                           bob-session bob
+                           (plist-get m1 :pre-key-p)
+                           (plist-get m1 :data))))
+      (should (null (jabber-omemo--session-skipped-keys bob-session)))
+      (should-error
+       (jabber-omemo--decrypt-key
+        bob-session bob (plist-get m1 :pre-key-p) (plist-get m1 :data))
+       :type 'jabber-omemo-error))))
+
 (ert-deftest jabber-test-omemo-module-skipped-keys-survive-reserialization ()
   "Skipped keys carried over to a reloaded session still decrypt."
   (jabber-test-omemo-module--with-session-pair

Reply via email to