telegabber/xmpp/omemo.go
2026-07-28 23:14:32 -04:00

116 lines
3.3 KiB
Go

package xmpp
import (
"encoding/base64"
"fmt"
"dev.narayana.im/narayana/telegabber/e2ee"
"dev.narayana.im/narayana/telegabber/e2ee/omemo"
"dev.narayana.im/narayana/telegabber/xmpp/extensions"
"gosrc.io/xmpp/stanza"
)
// hasOMEMOPayload reports whether msg carries any of the three OMEMO
// <encrypted> wire shapes, regardless of whether it also has a <body> -
// used to fix HandleMessage's body-gate, since OMEMO stanzas are commonly
// body-less (a plaintext fallback body is a courtesy for non-supporting
// clients, not a wire requirement - see gateway.omemoFallbackBody).
func hasOMEMOPayload(msg stanza.Message) bool {
var legacy extensions.OMEMO0Encrypted
if msg.Get(&legacy) {
return true
}
var modern extensions.OMEMOEncrypted
return msg.Get(&modern)
}
// decodeOMEMOEnvelope extracts msg's OMEMO <encrypted> payload (whichever
// of the three namespaces is present) as an opaque e2ee.Envelope, ready for
// Backend.Decrypt. ok is false if msg carries no OMEMO payload at all - not
// an error, just nothing for the decrypt hook to do.
func decodeOMEMOEnvelope(msg stanza.Message) (env e2ee.Envelope, ok bool, err error) {
var legacy extensions.OMEMO0Encrypted
if msg.Get(&legacy) {
env, err = decodeOMEMO0Envelope(legacy)
return env, true, err
}
var modern extensions.OMEMOEncrypted
if msg.Get(&modern) {
env, err = decodeOMEMOModernEnvelope(modern)
return env, true, err
}
return e2ee.Envelope{}, false, nil
}
func decodeOMEMO0Envelope(enc extensions.OMEMO0Encrypted) (e2ee.Envelope, error) {
payload, err := base64.StdEncoding.DecodeString(enc.Payload)
if err != nil {
return e2ee.Envelope{}, fmt.Errorf("omemo: decode payload: %w", err)
}
iv, err := base64.StdEncoding.DecodeString(enc.Header.IV)
if err != nil {
return e2ee.Envelope{}, fmt.Errorf("omemo: decode iv: %w", err)
}
wireEnv := &omemo.WireEnvelope{
Version: omemo.Omemo0,
SenderSID: enc.Header.SID,
IV: iv,
Payload: payload,
}
for _, k := range enc.Header.Keys {
ct, err := base64.StdEncoding.DecodeString(k.Text)
if err != nil {
return e2ee.Envelope{}, fmt.Errorf("omemo: decode key %d: %w", k.RID, err)
}
wireEnv.Keys = append(wireEnv.Keys, omemo.WireKey{
DeviceID: k.RID,
IsPreKey: k.PreKey,
Ciphertext: ct,
})
}
raw, err := wireEnv.Encode()
if err != nil {
return e2ee.Envelope{}, err
}
return e2ee.Envelope{Backend: omemo.Name, Raw: raw}, nil
}
func decodeOMEMOModernEnvelope(enc extensions.OMEMOEncrypted) (e2ee.Envelope, error) {
version := omemo.Omemo2
if enc.Namespace() == "urn:xmpp:omemo:1" {
version = omemo.Omemo1
}
payload, err := base64.StdEncoding.DecodeString(enc.Payload)
if err != nil {
return e2ee.Envelope{}, fmt.Errorf("omemo: decode payload: %w", err)
}
wireEnv := &omemo.WireEnvelope{
Version: version,
SenderSID: enc.Header.SID,
Payload: payload,
}
for _, group := range enc.Header.Keys {
for _, k := range group.Keys {
ct, err := base64.StdEncoding.DecodeString(k.Text)
if err != nil {
return e2ee.Envelope{}, fmt.Errorf("omemo: decode key %d: %w", k.RID, err)
}
wireEnv.Keys = append(wireEnv.Keys, omemo.WireKey{
RecipientJID: group.JID,
DeviceID: k.RID,
IsPreKey: k.Kex,
Ciphertext: ct,
})
}
}
raw, err := wireEnv.Encode()
if err != nil {
return e2ee.Envelope{}, err
}
return e2ee.Envelope{Backend: omemo.Name, Raw: raw}, nil
}