[PATCH net-next v2 1/7] psp: support rx rekey operation

From: Daniel Zahka

Date: Fri Oct 09 2026 - 16:48:08 EST


Support updating the rx PSP state used by a socket. This is most
useful to do after a device key rotation has occurred and the key
being used by the peer will become stale after the next device key
rotation.

Create a new psp_assoc for the socket to use that contains the same tx
state, but with new rx state. Copy the previous rx state to use as a
fallback while the peer is in the process of switching to the new key.

skbs matching the current rx spi and generation or the previous rx spi
and generation will be accepted. It might be reasonable to stop
accepting skbs on the previous state after seeing skbs with the new
state and after waiting a grace period, but that is not implemented
for now.

Signed-off-by: Daniel Zahka <daniel.zahka@xxxxxxxxx>
---
v2:
- copy old rx state by value instead of pointer to prev
- place non-datapath fields at end of struct psp_assoc
- disallow rx rekey when socket is not in full psp state, or
dev/version doesn't match.
- require the new rx spi to have the opposite phase bit from the previous
one
---
include/net/psp/functions.h | 10 ++++++----
include/net/psp/types.h | 14 ++++++++++++++
net/psp/psp.h | 3 ++-
net/psp/psp_sock.c | 46 +++++++++++++++++++++++++++++++++++++++++----
4 files changed, 64 insertions(+), 9 deletions(-)

diff --git a/include/net/psp/functions.h b/include/net/psp/functions.h
index b23c30898389..fb80e15e4368 100644
--- a/include/net/psp/functions.h
+++ b/include/net/psp/functions.h
@@ -77,10 +77,12 @@ psp_is_allowed_nondata(struct sk_buff *skb, struct psp_assoc *pas)
static inline bool
psp_pse_matches_pas(struct psp_skb_ext *pse, struct psp_assoc *pas)
{
- return pse && pas->rx.spi == pse->spi &&
- pas->generation == pse->generation &&
- pas->version == pse->version &&
- pas->dev_id == pse->dev_id;
+ return pse && pas->version == pse->version &&
+ pas->dev_id == pse->dev_id &&
+ ((pas->rx.spi == pse->spi &&
+ pas->generation == pse->generation) ||
+ (pas->prev_spi && pas->prev_spi == pse->spi &&
+ pas->prev_generation == pse->generation));
}

static inline enum skb_drop_reason
diff --git a/include/net/psp/types.h b/include/net/psp/types.h
index b8905efbd604..52c78cc72f1f 100644
--- a/include/net/psp/types.h
+++ b/include/net/psp/types.h
@@ -3,6 +3,7 @@
#ifndef __NET_PSP_H
#define __NET_PSP_H

+#include <linux/bits.h>
#include <linux/mutex.h>
#include <linux/refcount.h>
#include <net/net_trackers.h>
@@ -152,6 +153,15 @@ struct psp_key_parsed {
u8 key[PSP_MAX_KEY];
};

+/**
+ * enum psp_assoc_flags - flags of struct psp_assoc
+ * @PSP_ASSOC_SKIP_TX_KEY_DEL: Do not delete Tx key from psp_dev. It was
+ * copied to a newer psp_assoc during an Rx rekey.
+ */
+enum psp_assoc_flags {
+ PSP_ASSOC_SKIP_TX_KEY_DEL = BIT(0),
+};
+
struct psp_assoc {
struct psp_dev *psd;

@@ -159,6 +169,10 @@ struct psp_assoc {
u8 generation;
u8 version;
u8 peer_tx;
+ u8 prev_generation;
+ u8 flags; /* Slow path, protected by psd->lock */
+
+ __be32 prev_spi;

u32 upgrade_seq;

diff --git a/net/psp/psp.h b/net/psp/psp.h
index b123c2427905..92d7b92acaed 100644
--- a/net/psp/psp.h
+++ b/net/psp/psp.h
@@ -63,7 +63,8 @@ static inline bool psp_dev_has_sadb(struct psp_dev *psd)
static inline bool psp_assoc_needs_tx_key_del(struct psp_assoc *pas)
{
lockdep_assert_held(&pas->psd->lock);
- return psp_dev_has_sadb(pas->psd) && pas->tx.spi;
+ return psp_dev_has_sadb(pas->psd) && pas->tx.spi &&
+ !(pas->flags & PSP_ASSOC_SKIP_TX_KEY_DEL);
}

#endif /* __PSP_PSP_H */
diff --git a/net/psp/psp_sock.c b/net/psp/psp_sock.c
index a6b1c42dd626..b5887171c84e 100644
--- a/net/psp/psp_sock.c
+++ b/net/psp/psp_sock.c
@@ -149,10 +149,42 @@ void psp_sk_assoc_free(struct sock *sk)
psp_assoc_put(pas);
}

+static int psp_sock_rx_rekey(struct psp_assoc *pas, struct psp_assoc *prev,
+ struct netlink_ext_ack *extack)
+{
+ if (pas->psd != prev->psd) {
+ NL_SET_ERR_MSG(extack, "PSP device mismatch with existing state");
+ return -EINVAL;
+ }
+ if (pas->version != prev->version) {
+ NL_SET_ERR_MSG(extack, "PSP version mismatch with existing state");
+ return -EINVAL;
+ }
+ if (!prev->tx.spi || !prev->peer_tx) {
+ NL_SET_ERR_MSG(extack, "Socket PSP state is not fully established");
+ return -EBUSY;
+ }
+ if (!((pas->rx.spi ^ prev->rx.spi) & cpu_to_be32(PSP_SPI_KEY_PHASE))) {
+ NL_SET_ERR_MSG(extack, "New and prev SPI have same phase bit");
+ return -EINVAL;
+ }
+
+ pas->peer_tx = 1;
+ pas->prev_spi = prev->rx.spi;
+ pas->prev_generation = prev->generation;
+
+ memcpy(&pas->tx, &prev->tx, sizeof(pas->tx));
+ memcpy(pas->drv_data, prev->drv_data, pas->psd->caps->assoc_drv_spc);
+ prev->flags |= PSP_ASSOC_SKIP_TX_KEY_DEL;
+
+ return 0;
+}
+
int psp_sock_assoc_set_rx(struct sock *sk, struct psp_assoc *pas,
struct psp_key_parsed *key,
struct netlink_ext_ack *extack)
{
+ struct psp_assoc *prev;
int err;

memcpy(&pas->rx, key, sizeof(*key));
@@ -165,10 +197,11 @@ int psp_sock_assoc_set_rx(struct sock *sk, struct psp_assoc *pas,
goto exit_unlock;
}

- if (psp_sk_assoc(sk)) {
- NL_SET_ERR_MSG(extack, "Socket already has PSP state");
- err = -EBUSY;
- goto exit_unlock;
+ prev = psp_sk_assoc(sk);
+ if (prev) {
+ err = psp_sock_rx_rekey(pas, prev, extack);
+ if (err)
+ goto exit_unlock;
} else if (sk_has_decrypt_user(sk)) {
NL_SET_ERR_MSG(extack, "Socket has incompatible state");
err = -EINVAL;
@@ -177,6 +210,7 @@ int psp_sock_assoc_set_rx(struct sock *sk, struct psp_assoc *pas,

refcount_inc(&pas->refcnt);
rcu_assign_pointer(sk->psp_assoc, pas);
+ psp_assoc_put(prev);
err = 0;

exit_unlock:
@@ -304,6 +338,10 @@ void psp_assocs_key_rotated(struct psp_dev *psd)
pas->generation |= ~PSP_GEN_VALID_MASK;
psd->stats.stales++;
}
+
+ list_for_each_entry(pas, &psd->active_assocs, assocs_list)
+ pas->prev_generation |= ~PSP_GEN_VALID_MASK;
+
list_splice_init(&psd->prev_assocs, &psd->stale_assocs);
list_splice_init(&psd->active_assocs, &psd->prev_assocs);


--
2.52.0