diff --git a/client/internal/pqkem/capability_test.go b/client/internal/pqkem/capability_test.go index b1c29d582..93cf1546b 100644 --- a/client/internal/pqkem/capability_test.go +++ b/client/internal/pqkem/capability_test.go @@ -64,3 +64,35 @@ func TestManager_EstablishedPeerNotDowngraded(t *testing.T) { assert.True(t, ok, "an established peer must keep its PSK despite a stray zero") assert.NotEqual(t, PSK{}, psk) } + +// TestManager_ErrorMarkerIsBenign verifies finding F's wire half: a responder failure is +// signalled with an error marker (not an empty answer), so the initiator does not read it +// as "peer has no KEM". The marker must be benign — it must not disturb the in-flight +// exchange, which still converges when the real answer arrives. +func TestManager_ErrorMarkerIsBenign(t *testing.T) { + dA, dB, wgA, wgB, _ := pair(t) // dB ("bbbb") initiator, dA ("aaaa") responder + defer dA.Stop() + defer dB.Stop() + + offer, err := dB.SignalOffer("aaaa") + require.NoError(t, err) + require.NotNil(t, offer) + + _, decoded, err := Decode(offer) + require.NoError(t, err) + offerID := decoded.(*OfferMsg).ExchangeID + + // An error marker for the in-flight offer must be accepted without error and must not + // tear the exchange down (unlike an empty answer, which signals non-capability). + marker := (&ErrorMsg{ExchangeID: offerID}).Encode() + require.NoError(t, dB.SignalOnAnswer("aaaa", marker)) + + // The real answer still converges both sides on the same PSK. + answer, err := dA.SignalOnOffer("bbbb", offer) + require.NoError(t, err) + require.NotNil(t, answer) + require.NoError(t, dB.SignalOnAnswer("aaaa", answer)) + + assert.Equal(t, wgB.psk("aaaa"), wgA.psk("bbbb"), "the exchange must still converge after a benign error marker") + assert.NotEqual(t, PSK{}, wgB.psk("aaaa")) +} diff --git a/client/internal/pqkem/manager.go b/client/internal/pqkem/manager.go index 8912a27c5..429d38716 100644 --- a/client/internal/pqkem/manager.go +++ b/client/internal/pqkem/manager.go @@ -365,7 +365,15 @@ func (m *Manager) SignalOnOffer(remoteID RemoteID, offer []byte) ([]byte, error) if typ != MsgOffer { return nil, fmt.Errorf("expected offer from %s, got type %d", remoteID, typ) } - return m.processOffer(remoteID, msg.(*OfferMsg), viaSignalLabel) + o := msg.(*OfferMsg) + answer, err := m.processOffer(remoteID, o, viaSignalLabel) + if err != nil { + // Reply with an error marker rather than an empty answer: the initiator must not + // read our transient failure as "peer does not run the KEM" and mark us + // non-capable. The exchange times out on the initiator and re-bootstraps. + return (&ErrorMsg{ExchangeID: o.ExchangeID}).Encode(), err + } + return answer, nil } // SignalOnAnswer processes a KEM answer the host extracted from an incoming answer. @@ -375,6 +383,12 @@ func (m *Manager) SignalOnAnswer(remoteID RemoteID, answer []byte) error { if err != nil { return fmt.Errorf("decode signal answer from %s: %w", remoteID, err) } + if typ == MsgError { + // The responder runs the KEM but failed to answer this offer. It is capable, so + // leave the exchange to time out and re-bootstrap; do not mark it non-capable. + m.trace("pqkem: peer reported an error answering our offer", "peer", remoteID) + return nil + } if typ != MsgAnswer { return fmt.Errorf("expected answer from %s, got type %d", remoteID, typ) } @@ -416,6 +430,9 @@ func (m *Manager) OnDataPathMessage(remoteID RemoteID, raw []byte) error { return m.pushDataPath(remoteID, answer) case MsgAnswer: return m.processAnswer(remoteID, msg.(*AnswerMsg), viaDataPathLabel) + case MsgError: + m.trace("pqkem: peer reported an error over the data path", "peer", remoteID) + return nil default: return fmt.Errorf("unhandled data-path message type %d from %s", typ, remoteID) } diff --git a/client/internal/pqkem/message.go b/client/internal/pqkem/message.go index e483e9d72..1e1fd6e85 100644 --- a/client/internal/pqkem/message.go +++ b/client/internal/pqkem/message.go @@ -28,12 +28,17 @@ const ( headerSize = 1 + 1 + ExchangeIDSize ) -// MsgType tags the two message kinds of the exchange. +// MsgType tags the message kinds of the exchange. type MsgType uint8 const ( MsgOffer MsgType = iota + 1 MsgAnswer + // MsgError is a responder's "I run the KEM but could not answer this offer" reply. + // It is a non-empty, payload-less marker so the initiator does not mistake a + // transient failure for an empty "peer does not run the KEM" answer and mark the + // peer non-capable. The exchange still times out and re-bootstraps. + MsgError ) // ExchangeID is the per-exchange correlator. The zero value means "none" (an offer @@ -76,6 +81,17 @@ func (m *AnswerMsg) Encode() ([]byte, error) { return frame(MsgAnswer, m.ExchangeID, m.KEMAnswer), nil } +// ErrorMsg is the payload-less marker a responder returns when it runs the KEM but +// could not answer a given offer, naming the offer's exchange. +type ErrorMsg struct { + ExchangeID ExchangeID +} + +// Encode serialises the error marker (header only, no payload). +func (m *ErrorMsg) Encode() []byte { + return frame(MsgError, m.ExchangeID, nil) +} + // Decode parses a framed message into one of *OfferMsg / *AnswerMsg. func Decode(buf []byte) (MsgType, any, error) { if len(buf) < headerSize { @@ -103,6 +119,11 @@ func Decode(buf []byte) (MsgType, any, error) { return typ, nil, fmt.Errorf("answer payload: got %d, want %d", len(payload), AnswerSize) } return typ, &AnswerMsg{ExchangeID: id, KEMAnswer: payload}, nil + case MsgError: + if len(payload) != 0 { + return typ, nil, fmt.Errorf("error marker payload: got %d, want 0", len(payload)) + } + return typ, &ErrorMsg{ExchangeID: id}, nil default: return typ, nil, fmt.Errorf("unknown message type %d", typ) } diff --git a/client/internal/pqkem/message_test.go b/client/internal/pqkem/message_test.go index 524125e47..0ce9be69d 100644 --- a/client/internal/pqkem/message_test.go +++ b/client/internal/pqkem/message_test.go @@ -30,6 +30,15 @@ func TestMessageRoundTrip(t *testing.T) { require.NoError(t, err) require.Equal(t, MsgAnswer, typ) require.Equal(t, answer, decoded.(*AnswerMsg).KEMAnswer) + + // The error marker is a non-empty, payload-less message carrying the exchange id, so + // the initiator can tell a responder failure from an empty "no KEM" answer. + errBytes := (&ErrorMsg{ExchangeID: id}).Encode() + require.NotEmpty(t, errBytes) + typ, decoded, err = Decode(errBytes) + require.NoError(t, err) + require.Equal(t, MsgError, typ) + require.Equal(t, id, decoded.(*ErrorMsg).ExchangeID) } func TestDecodeRejects(t *testing.T) {