Update MessageApiV2 to use libsignal-net.

This commit is contained in:
Greyson Parrelli
2026-06-03 13:55:36 -04:00
committed by Michelle Tang
parent 51c4afe5f5
commit f76292769a
5 changed files with 608 additions and 570 deletions
@@ -17,6 +17,7 @@ import okio.utf8Size
import org.signal.core.models.ServiceId
import org.signal.core.util.logging.Log
import org.signal.core.util.orNull
import org.signal.libsignal.net.ChallengeOption
import org.signal.libsignal.protocol.SignalProtocolAddress
import org.signal.network.service.MessageService
import org.thoughtcrime.securesms.BuildConfig
@@ -201,7 +202,7 @@ class IndividualSendJobV2 private constructor(parameters: Parameters, private va
val syntheticResult = SendMessageResult.success(
SignalServiceAddress(recipient.requireServiceId(), recipient.e164.orNull()),
success.devices,
success.sentUnidentified,
success.sentSealedSender,
false,
0L,
Optional.of(content)
@@ -220,7 +221,7 @@ class IndividualSendJobV2 private constructor(parameters: Parameters, private va
SignalDatabase.pendingPniSignatureMessages.insertIfNecessary(recipient.id, message.sentTimeMillis, syntheticResult)
}
SignalDatabase.messages.markAsSent(messageId, success.sentUnidentified)
SignalDatabase.messages.markAsSent(messageId, success.sentSealedSender)
PushSendJob.markAttachmentsUploaded(messageId, message)
SignalDatabase.threads.updateSilently(threadId, false)
@@ -232,11 +233,11 @@ class IndividualSendJobV2 private constructor(parameters: Parameters, private va
}
val accessMode = recipient.sealedSenderAccessMode
if (success.sentUnidentified && accessMode == SealedSenderAccessMode.UNKNOWN && recipient.profileKey == null) {
if (success.sentSealedSender && accessMode == SealedSenderAccessMode.UNKNOWN && recipient.profileKey == null) {
SignalDatabase.recipients.setSealedSenderAccessMode(recipient.id, SealedSenderAccessMode.UNRESTRICTED)
} else if (success.sentUnidentified && accessMode == SealedSenderAccessMode.UNKNOWN) {
} else if (success.sentSealedSender && accessMode == SealedSenderAccessMode.UNKNOWN) {
SignalDatabase.recipients.setSealedSenderAccessMode(recipient.id, SealedSenderAccessMode.ENABLED)
} else if (!success.sentUnidentified && accessMode != SealedSenderAccessMode.DISABLED) {
} else if (!success.sentSealedSender && accessMode != SealedSenderAccessMode.DISABLED) {
SignalDatabase.recipients.setSealedSenderAccessMode(recipient.id, SealedSenderAccessMode.DISABLED)
}
@@ -259,36 +260,41 @@ class IndividualSendJobV2 private constructor(parameters: Parameters, private va
ifLeft = { error ->
when (error) {
is MessageService.SendError.IdentityMismatch -> {
Log.w(TAG, "${logPrefix(message.sentTimeMillis)} Identity mismatch for ${error.recipient.identifier}", error.cause)
val externalRecipient = Recipient.external(error.recipient.identifier)
Log.w(TAG, "${logPrefix(message.sentTimeMillis)} Identity mismatch for ${error.serviceId}", error.exception)
val externalRecipient = Recipient.external(error.serviceId.toString())
if (externalRecipient == null) {
Log.w(TAG, "${logPrefix(message.sentTimeMillis)} Failed to create a Recipient for the identifier!")
} else {
SignalDatabase.messages.addMismatchedIdentity(messageId, externalRecipient.id, error.cause.untrustedIdentity)
SignalDatabase.messages.addMismatchedIdentity(messageId, externalRecipient.id, error.exception.untrustedIdentity)
SignalDatabase.messages.markAsSentFailed(messageId)
RetrieveProfileJob.enqueue(externalRecipient.id, true)
}
Result.success()
}
MessageService.SendError.NotRegistered -> {
Log.w(TAG, "${logPrefix(message.sentTimeMillis)} Recipient not registered")
is MessageService.SendError.NotRegistered -> {
Log.w(TAG, "${logPrefix(message.sentTimeMillis)} Recipient not registered", error)
SignalDatabase.messages.markAsSentFailed(messageId)
PushSendJob.notifyMediaMessageDeliveryFailed(context, messageId)
AppDependencies.jobManager.add(DirectoryRefreshJob(false))
Result.success()
}
MessageService.SendError.Unauthorized -> {
Log.w(TAG, "${logPrefix(message.sentTimeMillis)} Unauthorized send")
is MessageService.SendError.Unauthorized -> {
Log.w(TAG, "${logPrefix(message.sentTimeMillis)} Unauthorized send", error)
Result.failure()
}
is MessageService.SendError.ChallengeRequired -> {
Log.w(TAG, "${logPrefix(message.sentTimeMillis)} Challenge required (options=${error.options})")
Log.w(TAG, "${logPrefix(message.sentTimeMillis)} Challenge required (options=${error.options})", error)
val proofResponse = ProofRequiredResponse().apply {
token = error.token
options = error.options
options = error.options.map {
when (it) {
ChallengeOption.PUSH_CHALLENGE -> "pushChallenge"
ChallengeOption.CAPTCHA -> "captcha"
}
}
}
val proofException = ProofRequiredException(proofResponse, error.retryAfter?.inWholeSeconds ?: 0L)
val threadRecipient = SignalDatabase.threads.getRecipientForThreadId(threadId)
@@ -299,23 +305,23 @@ class IndividualSendJobV2 private constructor(parameters: Parameters, private va
}
}
MessageService.SendError.ServerRejected -> {
Log.w(TAG, "${logPrefix(message.sentTimeMillis)} Server rejected the send")
is MessageService.SendError.ServerRejected -> {
Log.w(TAG, "${logPrefix(message.sentTimeMillis)} Server rejected the send", error)
Result.failure()
}
is MessageService.SendError.ContentTooLarge -> {
Log.w(TAG, "${logPrefix(message.sentTimeMillis)} Content too large (${error.size} > ${error.maxAllowed} bytes); failing.")
Log.w(TAG, "${logPrefix(message.sentTimeMillis)} Content too large (${error.size} > ${error.maxAllowed} bytes). Failing.", error)
Result.failure()
}
MessageService.SendError.SessionAttemptsExhausted -> {
Log.w(TAG, "${logPrefix(message.sentTimeMillis)} Exhausted device-resolution attempts; retrying")
is MessageService.SendError.SessionAttemptsExhausted -> {
Log.w(TAG, "${logPrefix(message.sentTimeMillis)} Exhausted device-resolution attempts. Retrying", error)
Result.retry(nextRunAttemptBackoff(runAttempt + 1))
}
is MessageService.SendError.PreKeyUnavailable -> {
Log.w(TAG, "${logPrefix(message.sentTimeMillis)} Prekey unavailable: ${error.reason}")
Log.w(TAG, "${logPrefix(message.sentTimeMillis)} Prekey unavailable: ${error.reason}", error)
Result.retry(nextRunAttemptBackoff(runAttempt + 1))
}
@@ -323,16 +329,16 @@ class IndividualSendJobV2 private constructor(parameters: Parameters, private va
val defaultBackoff = nextRunAttemptBackoff(runAttempt + 1)
val serverBackoff = error.retryAfter?.inWholeMilliseconds ?: 0L
val backoff = maxOf(defaultBackoff, serverBackoff)
Log.w(TAG, "${logPrefix(message.sentTimeMillis)} Rate limited, retryAfter=${error.retryAfter}, using backoff=${backoff}ms")
Log.w(TAG, "${logPrefix(message.sentTimeMillis)} Rate limited, retryAfter=${error.retryAfter}, using backoff=${backoff}ms", error)
Result.retry(backoff)
}
is MessageService.SendError.NetworkError -> {
Log.w(TAG, "${logPrefix(message.sentTimeMillis)} Network error", error.cause)
Log.w(TAG, "${logPrefix(message.sentTimeMillis)} Network error", error.exception)
Result.retry(nextRunAttemptBackoff(runAttempt + 1))
}
is MessageService.SendError.ApplicationError -> when (val cause = error.cause) {
is MessageService.SendError.ApplicationError -> when (val cause = error.exception) {
is RuntimeException -> {
Log.e(TAG, "${logPrefix(message.sentTimeMillis)} Encountered a fatal application error. Crash imminent.", cause)
Result.fatalFailure(cause)
@@ -397,7 +403,7 @@ class IndividualSendJobV2 private constructor(parameters: Parameters, private va
}
return AppDependencies.messageService.sendMessage(
recipient = SignalServiceAddress(recipient.requireServiceId(), recipient.e164.orNull()),
serviceId = recipient.requireServiceId(),
envelopeContent = envelopeContent,
timestamp = dataMessage.timestamp!!,
sealedSenderAccess = SealedSenderAccessUtil.getSealedSenderAccessFor(recipient),
@@ -434,7 +440,7 @@ class IndividualSendJobV2 private constructor(parameters: Parameters, private va
unidentifiedStatus = listOf(
SyncMessage.Sent.UnidentifiedDeliveryStatus(
destinationServiceIdBinary = recipientServiceId.toByteString(),
unidentified = primaryResult.sentUnidentified,
unidentified = primaryResult.sentSealedSender,
destinationPniIdentityKey = pniIdentityKey
)
)
@@ -444,7 +450,7 @@ class IndividualSendJobV2 private constructor(parameters: Parameters, private va
val syncEnvelope = EnvelopeContent.encrypted(syncContent, ContentHint.IMPLICIT, Optional.empty())
return AppDependencies.messageService.sendMessage(
recipient = SignalServiceAddress(SignalStore.account.requireAci()),
serviceId = SignalStore.account.requireAci(),
envelopeContent = syncEnvelope,
timestamp = timestamp,
sealedSenderAccess = null, // We don't use sealed sender for sync messages
@@ -5,18 +5,16 @@
package org.signal.network.api
import arrow.core.getOrElse
import kotlinx.serialization.Serializable
import kotlinx.serialization.Transient
import org.signal.core.util.serialization.SignalJson
import org.signal.libsignal.net.BadRequestError
import org.signal.core.models.ServiceId
import org.signal.libsignal.net.AuthMessagesService
import org.signal.libsignal.net.RequestResult
import org.signal.network.websocket.WebSocketRequestMessage
import org.signal.network.websocket.put
import org.whispersystems.signalservice.api.crypto.SealedSenderAccess
import org.signal.libsignal.net.SealedSendFailure
import org.signal.libsignal.net.SingleOutboundSealedSenderMessage
import org.signal.libsignal.net.SingleOutboundUnsealedMessage
import org.signal.libsignal.net.UnauthMessagesService
import org.signal.libsignal.net.UnsealedSendFailure
import org.signal.libsignal.net.UserBasedSendAuthorization
import org.whispersystems.signalservice.api.websocket.SignalWebSocket
import java.io.IOException
import kotlin.time.Duration
/**
* Collection of message-related endpoints.
@@ -25,136 +23,29 @@ class MessageApiV2(
private val authWebSocket: SignalWebSocket.AuthenticatedWebSocket,
private val unauthWebSocket: SignalWebSocket.UnauthenticatedWebSocket
) {
/**
* Sends a message to a single recipient. Uses the unauthenticated websocket if [sealedSenderAccess] is provided,
* and the authenticated websocket otherwise.
*
* PUT /v1/messages/[destination]?story=[story]
* - 200: Success
* - 401: Authorization or [sealedSenderAccess] is missing or incorrect
* - 404: Recipient is not a registered Signal user
* - 409: Mismatched devices for the recipient
* - 410: Stale devices for some recipient devices
* - 428: Sender must complete a challenge before proceeding
* - 508: Server rejected the message
*/
suspend fun sendMessage(
destination: String,
messageList: SendMessageRequest,
sealedSenderAccess: SealedSenderAccess?,
story: Boolean
): RequestResult<SendMessageResponse, SendMessageError> {
val requestBody = SignalJson.encode(SendMessageRequest.serializer(), messageList).getOrElse { return RequestResult.ApplicationError(it.cause) }
val request = WebSocketRequestMessage.put("/v1/messages/$destination?story=$story", requestBody)
return try {
val response = if (sealedSenderAccess == null) {
authWebSocket.requestSuspend(request)
} else {
unauthWebSocket.requestSuspend(request, sealedSenderAccess)
}
when (response.status) {
200 -> {
SignalJson
.decode(SendMessageResponse.serializer(), response.body)
.map { it.copy(sentUnidentified = response.isUnidentified) }
.fold(
ifLeft = { RequestResult.ApplicationError(it.cause) },
ifRight = { RequestResult.Success(it) }
)
}
401 -> {
RequestResult.NonSuccess(SendMessageError.Unauthorized)
}
404 -> {
RequestResult.NonSuccess(SendMessageError.NotRegistered)
}
409 -> {
SignalJson
.decode(MismatchedDevices.serializer(), response.body)
.fold(
ifLeft = { RequestResult.ApplicationError(it.cause) },
ifRight = { RequestResult.NonSuccess(SendMessageError.MismatchedDevicesError(it)) }
)
}
410 -> {
SignalJson
.decode(StaleDevices.serializer(), response.body)
.fold(
ifLeft = { RequestResult.ApplicationError(it.cause) },
ifRight = { RequestResult.NonSuccess(SendMessageError.StaleDevicesError(it)) }
)
}
428 -> {
SignalJson
.decode(ProofRequiredResponseBody.serializer(), response.body)
.fold(
ifLeft = { RequestResult.ApplicationError(it.cause) },
ifRight = { RequestResult.NonSuccess(SendMessageError.ChallengeRequired(it.token, it.options, response.retryAfter())) }
)
}
429 -> RequestResult.NonSuccess(SendMessageError.RateLimited(response.retryAfter()))
508 -> RequestResult.NonSuccess(SendMessageError.ServerRejected)
else -> RequestResult.ApplicationError(IllegalStateException("Unexpected response code: ${response.status}"))
}
} catch (e: IOException) {
RequestResult.RetryableNetworkError(e)
} catch (e: Throwable) {
RequestResult.ApplicationError(e)
suspend fun sendSealedSenderMessage(
serviceId: ServiceId,
timestamp: Long,
contents: List<SingleOutboundSealedSenderMessage>,
auth: UserBasedSendAuthorization,
onlineOnly: Boolean,
urgent: Boolean
): RequestResult<Unit, SealedSendFailure> {
return unauthWebSocket.runCatchingWithChatConnection { connection ->
UnauthMessagesService(connection).sendMessage(serviceId.libSignalServiceId, timestamp, contents, auth, onlineOnly, urgent)
}
}
@Serializable
data class SendMessageRequest(
val messages: List<Message>,
val timestamp: Long,
val online: Boolean = false,
val urgent: Boolean = true
)
@Serializable
data class Message(
val type: Int,
val destinationDeviceId: Int,
val destinationRegistrationId: Int,
val content: String
)
@Serializable
data class SendMessageResponse(
val needsSync: Boolean = false,
@Transient val sentUnidentified: Boolean = false
)
@Serializable
data class MismatchedDevices(
val missingDevices: List<Int> = emptyList(),
val extraDevices: List<Int> = emptyList()
)
@Serializable
data class StaleDevices(
val staleDevices: List<Int> = emptyList()
)
/**
* Body of a 428 response. [token] is the proof-required challenge token; [options] is the
* list of supported challenge mechanisms (e.g. "captcha", "pushChallenge").
*/
@Serializable
private data class ProofRequiredResponseBody(
val token: String,
val options: List<String> = emptyList()
)
sealed class SendMessageError : BadRequestError {
data object Unauthorized : SendMessageError()
data object NotRegistered : SendMessageError()
data class MismatchedDevicesError(val devices: MismatchedDevices) : SendMessageError()
data class StaleDevicesError(val devices: StaleDevices) : SendMessageError()
data class ChallengeRequired(val token: String, val options: List<String>, val retryAfter: Duration?) : SendMessageError()
data class RateLimited(val retryAfter: Duration?) : SendMessageError()
data object ServerRejected : SendMessageError()
suspend fun sendUnsealedSenderMessage(
serviceId: ServiceId,
timestamp: Long,
contents: List<SingleOutboundUnsealedMessage>,
onlineOnly: Boolean,
urgent: Boolean
): RequestResult<Unit, UnsealedSendFailure> {
return authWebSocket.runCatchingWithChatConnection { connection ->
AuthMessagesService(connection).sendMessage(serviceId.libSignalServiceId, timestamp, contents, onlineOnly, urgent)
}
}
}
@@ -11,8 +11,21 @@ import arrow.core.raise.either
import kotlinx.coroutines.Dispatchers
import kotlinx.coroutines.withContext
import org.jetbrains.annotations.VisibleForTesting
import org.signal.core.models.ServiceId
import org.signal.core.util.Base64.decodeBase64OrThrow
import org.signal.core.util.logging.Log
import org.signal.libsignal.net.ChallengeOption
import org.signal.libsignal.net.MismatchedDeviceException
import org.signal.libsignal.net.RateLimitChallengeException
import org.signal.libsignal.net.RequestResult
import org.signal.libsignal.net.RequestUnauthorizedException
import org.signal.libsignal.net.SealedSendFailure
import org.signal.libsignal.net.ServiceIdNotFoundException
import org.signal.libsignal.net.SingleOutboundSealedSenderMessage
import org.signal.libsignal.net.SingleOutboundUnsealedMessage
import org.signal.libsignal.net.UnsealedSendFailure
import org.signal.libsignal.net.UserBasedAuthorization
import org.signal.libsignal.net.UserBasedSendAuthorization
import org.signal.libsignal.protocol.IdentityKey
import org.signal.libsignal.protocol.InvalidKeyException
import org.signal.libsignal.protocol.SessionBuilder
@@ -20,6 +33,7 @@ import org.signal.libsignal.protocol.SignalProtocolAddress
import org.signal.libsignal.protocol.UntrustedIdentityException
import org.signal.libsignal.protocol.ecc.ECPublicKey
import org.signal.libsignal.protocol.kem.KEMPublicKey
import org.signal.libsignal.protocol.message.SignalMessage
import org.signal.libsignal.protocol.state.PreKeyBundle
import org.signal.network.api.KeysApiV2
import org.signal.network.api.MessageApiV2
@@ -33,6 +47,7 @@ import org.whispersystems.signalservice.api.push.SignalServiceAddress
import org.whispersystems.signalservice.internal.push.OutgoingPushMessage
import java.io.IOException
import kotlin.time.Duration
import kotlin.time.toKotlinDuration
/**
* Sends an [EnvelopeContent] to a single recipient, driving the full one-to-one flow:
@@ -69,10 +84,10 @@ open class MessageService(
private val localProtocolAddress: SignalProtocolAddress = SignalProtocolAddress(localAddress.identifier, localDeviceId)
/**
* Sends [envelopeContent] to [recipient]. Handles things like establishing sessions with newly-discovered linked devices.
* Sends [envelopeContent] to [serviceId]. Handles things like establishing sessions with newly-discovered linked devices.
*/
suspend fun sendMessage(
recipient: SignalServiceAddress,
serviceId: ServiceId,
envelopeContent: EnvelopeContent,
timestamp: Long,
sealedSenderAccess: SealedSenderAccess?,
@@ -89,89 +104,222 @@ open class MessageService(
}
var encryptedReported = false
var activeSealedSenderAccess = sealedSenderAccess
var sealedSender = sealedSenderAccess != null || story
// Certain errors self-resolve by mutating external state, like creating new sessions.
// Trying several times in a loop lets us re-read that external state and use it in the next attempt.
for (attempt in 0 until MAX_DEVICE_RECOVERY_ATTEMPTS) {
val encrypted = encryptForAllDevices(recipient, envelopeContent, sealedSenderAccess)
Log.d(TAG, "Starting message send attempt ${attempt + 1} to $serviceId")
val encryptedMessages = encryptForAllDevices(serviceId, envelopeContent, activeSealedSenderAccess)
if (!encryptedReported) {
onEncrypted?.invoke()
encryptedReported = true
}
val request = MessageApiV2.SendMessageRequest(
messages = encrypted.map { it.toWireMessage() },
timestamp = timestamp,
online = isOnline,
urgent = urgent
)
when (val result = messageApi.sendMessage(recipient.identifier, request, sealedSenderAccess, story)) {
is RequestResult.Success -> {
val response = result.result
val devices = encrypted.map { it.destinationDeviceId }
return@either SendSuccess(envelopeContent = envelopeContent, sentUnidentified = response.sentUnidentified, devices = devices)
}
is RequestResult.NonSuccess -> when (val err = result.error) {
is MessageApiV2.SendMessageError.MismatchedDevicesError -> {
handleMismatched(recipient, err.devices, sealedSenderAccess)
if (sealedSender) {
val result = sendSealed(serviceId, encryptedMessages, timestamp, isOnline, urgent, activeSealedSenderAccess, story)
when (result) {
SealedSendResult.Success -> {}
SealedSendResult.MismatchedDevices -> { continue }
SealedSendResult.InvalidAccessKey -> {
Log.w(TAG, "Sealed sender access was rejected for $serviceId. Falling back to an unsealed send.")
activeSealedSenderAccess = null
sealedSender = story
continue
}
is MessageApiV2.SendMessageError.StaleDevicesError -> {
for (deviceId in err.devices.staleDevices) {
protocolStore.archiveSession(SignalProtocolAddress(recipient.identifier, deviceId))
}
}
MessageApiV2.SendMessageError.Unauthorized -> raise(SendError.Unauthorized)
MessageApiV2.SendMessageError.NotRegistered -> raise(SendError.NotRegistered)
is MessageApiV2.SendMessageError.ChallengeRequired -> raise(SendError.ChallengeRequired(err.token, err.options, err.retryAfter))
MessageApiV2.SendMessageError.ServerRejected -> raise(SendError.ServerRejected)
is MessageApiV2.SendMessageError.RateLimited -> raise(SendError.RateLimited(err.retryAfter))
}
is RequestResult.RetryableNetworkError -> raise(SendError.NetworkError(result.networkError))
is RequestResult.ApplicationError -> raise(SendError.ApplicationError(result.cause))
} else {
val result = sendUnsealed(serviceId, timestamp, encryptedMessages, isOnline, urgent)
when (result) {
UnsealedSendResult.Success -> {}
UnsealedSendResult.MismatchedDevices -> { continue }
}
}
val devices = encryptedMessages.map { it.destinationDeviceId }
Log.d(TAG, "Successfully sent ${if (sealedSender) "a sealed" else "an unsealed"} message to $serviceId, devices: $devices")
return@either SendSuccess(
envelopeContent = envelopeContent,
sentSealedSender = sealedSender,
devices = devices
)
}
Log.w(TAG, "Exhausted device-recovery attempts for ${recipient.identifier}")
raise(SendError.SessionAttemptsExhausted)
Log.w(TAG, "Exhausted device-recovery attempts for $serviceId")
raise(SendError.SessionAttemptsExhausted())
}
}
private suspend fun Raise<SendError>.sendSealed(
serviceId: ServiceId,
encryptedMessages: List<OutgoingPushMessage>,
timestamp: Long,
online: Boolean,
urgent: Boolean,
sealedSenderAccess: SealedSenderAccess?,
story: Boolean
): SealedSendResult {
val auth = if (story) {
UserBasedSendAuthorization.Story
} else if (sealedSenderAccess is SealedSenderAccess.IndividualUnidentifiedAccessFirst) {
if (sealedSenderAccess.unidentifiedAccess.unidentifiedAccessKey.all { it == 0.toByte() }) {
UserBasedAuthorization.UnrestrictedUnauthenticatedAccess
} else {
UserBasedAuthorization.AccessKey(sealedSenderAccess.unidentifiedAccess.unidentifiedAccessKey)
}
} else {
raise(SendError.ApplicationError(IllegalArgumentException("Bad sealed sender access!")))
}
val result = messageApi.sendSealedSenderMessage(
serviceId = serviceId,
timestamp = timestamp,
contents = encryptedMessages.map {
SingleOutboundSealedSenderMessage(
deviceId = it.destinationDeviceId,
registrationId = it.destinationRegistrationId,
message = it.content.decodeBase64OrThrow()
)
},
auth = auth,
onlineOnly = online,
urgent = urgent
)
return when (result) {
is RequestResult.Success -> {
SealedSendResult.Success
}
is RequestResult.RetryableNetworkError -> {
raise(SendError.NetworkError(result.networkError))
}
is RequestResult.ApplicationError -> {
raise(SendError.ApplicationError(result.cause))
}
is RequestResult.NonSuccess<SealedSendFailure> -> {
when (val error = result.error) {
is MismatchedDeviceException -> {
handleMismatched(error, sealedSenderAccess)
SealedSendResult.MismatchedDevices
}
is RequestUnauthorizedException -> {
SealedSendResult.InvalidAccessKey
}
is ServiceIdNotFoundException -> {
raise(SendError.NotRegistered())
}
}
}
}
}
private suspend fun Raise<SendError>.sendUnsealed(
serviceId: ServiceId,
timestamp: Long,
encryptedMessages: List<OutgoingPushMessage>,
online: Boolean,
urgent: Boolean
): UnsealedSendResult {
val result = messageApi.sendUnsealedSenderMessage(
serviceId = serviceId,
timestamp = timestamp,
contents = encryptedMessages.map {
SingleOutboundUnsealedMessage(
deviceId = it.destinationDeviceId,
registrationId = it.destinationRegistrationId,
message = SignalMessage(it.content.decodeBase64OrThrow())
)
},
onlineOnly = online,
urgent = urgent
)
return when (result) {
is RequestResult.Success -> {
UnsealedSendResult.Success
}
is RequestResult.RetryableNetworkError -> {
raise(SendError.NetworkError(result.networkError))
}
is RequestResult.ApplicationError -> {
raise(SendError.ApplicationError(result.cause))
}
is RequestResult.NonSuccess<UnsealedSendFailure> -> {
when (val error = result.error) {
is MismatchedDeviceException -> {
handleMismatched(error, sealedSenderAccess = null)
UnsealedSendResult.MismatchedDevices
}
is ServiceIdNotFoundException -> {
raise(SendError.NotRegistered())
}
is RateLimitChallengeException -> {
raise(SendError.ChallengeRequired(error.token, error.options, error.retryLater?.toKotlinDuration()))
}
}
}
}
}
suspend fun Raise<SendError>.handleMismatched(error: MismatchedDeviceException, sealedSenderAccess: SealedSenderAccess?) {
Log.w(TAG, "Handling mismatched devices: ${error.entries}")
for (entry in error.entries) {
for (staleDeviceId in entry.staleDevices) {
Log.w(TAG, "Archiving stale session: (${entry.account}, $staleDeviceId)")
protocolStore.archiveSession(SignalProtocolAddress(entry.account, staleDeviceId))
}
for (extraDeviceId in entry.extraDevices) {
Log.w(TAG, "Archiving extra session: (${entry.account}, $extraDeviceId)")
protocolStore.archiveSession(SignalProtocolAddress(entry.account, extraDeviceId))
}
for (missingDeviceId in entry.missingDevices) {
Log.w(TAG, "Initializing session for missing device: (${entry.account}, $missingDeviceId)")
val address = SignalProtocolAddress(entry.account, missingDeviceId)
initializeSession(ServiceId.fromLibSignal(entry.account), address, sealedSenderAccess)
}
}
}
private fun Raise<SendError>.encryptForAllDevices(
recipient: SignalServiceAddress,
serviceId: ServiceId,
envelopeContent: EnvelopeContent,
sealedSenderAccess: SealedSenderAccess?
): List<OutgoingPushMessage> {
return targetDeviceIds(recipient).map { deviceId ->
val address = SignalProtocolAddress(recipient.identifier, deviceId)
encryptContent(recipient, address, envelopeContent, sealedSenderAccess)
return targetDeviceIds(serviceId).map { deviceId ->
val address = SignalProtocolAddress(serviceId.libSignalServiceId, deviceId)
encryptContent(serviceId, address, envelopeContent, sealedSenderAccess)
}
}
private fun Raise<SendError>.encryptContent(
recipient: SignalServiceAddress,
serviceId: ServiceId,
address: SignalProtocolAddress,
envelopeContent: EnvelopeContent,
sealedSenderAccess: SealedSenderAccess?
): OutgoingPushMessage = try {
cipher.encrypt(address, sealedSenderAccess, envelopeContent)
} catch (e: UntrustedIdentityException) {
raise(SendError.IdentityMismatch(recipient, e))
raise(SendError.IdentityMismatch(serviceId, e))
} catch (e: InvalidKeyException) {
raise(SendError.ApplicationError(e))
}
private fun targetDeviceIds(recipient: SignalServiceAddress): List<Int> {
val subDevices: MutableSet<Int> = (protocolStore.getSubDeviceSessions(recipient.identifier) + SignalServiceAddress.DEFAULT_DEVICE_ID).toMutableSet()
private fun targetDeviceIds(serviceId: ServiceId): List<Int> {
val subDevices: MutableSet<Int> = (protocolStore.getSubDeviceSessions(serviceId.toString()) + SignalServiceAddress.DEFAULT_DEVICE_ID).toMutableSet()
// When sending to self, skip our own device.
if (recipient.matches(localAddress)) {
if (serviceId == localAddress.serviceId) {
subDevices -= localDeviceId
}
return subDevices
.filter { it == SignalServiceAddress.DEFAULT_DEVICE_ID || protocolStore.containsSession(SignalProtocolAddress(recipient.identifier, it)) }
.filter { it == SignalServiceAddress.DEFAULT_DEVICE_ID || protocolStore.containsSession(SignalProtocolAddress(serviceId.libSignalServiceId, it)) }
.sorted()
.toList()
}
@@ -180,7 +328,7 @@ open class MessageService(
*/
@VisibleForTesting
internal open suspend fun Raise<SendError>.initializeSession(
recipient: SignalServiceAddress,
serviceId: ServiceId,
address: SignalProtocolAddress,
sealedSenderAccess: SealedSenderAccess?
) {
@@ -188,7 +336,7 @@ open class MessageService(
is RequestResult.Success -> result.result
is RequestResult.NonSuccess -> {
when (val e = result.error) {
KeysApiV2.GetPreKeysError.Unauthorized -> raise(SendError.Unauthorized)
KeysApiV2.GetPreKeysError.Unauthorized -> raise(SendError.Unauthorized())
KeysApiV2.GetPreKeysError.NotFound -> raise(SendError.PreKeyUnavailable("No prekeys found for $address"))
is KeysApiV2.GetPreKeysError.RateLimited -> raise(SendError.RateLimited(e.retryAfter))
}
@@ -205,34 +353,12 @@ open class MessageService(
try {
SignalSessionBuilder(sessionLock, SessionBuilder(protocolStore, address, localProtocolAddress)).process(bundle)
} catch (e: UntrustedIdentityException) {
raise(SendError.IdentityMismatch(recipient, e))
raise(SendError.IdentityMismatch(serviceId, e))
} catch (e: InvalidKeyException) {
raise(SendError.ApplicationError(e))
}
}
private suspend fun Raise<SendError>.handleMismatched(
recipient: SignalServiceAddress,
mismatched: MessageApiV2.MismatchedDevices,
sealedSenderAccess: SealedSenderAccess?
) {
for (extra in mismatched.extraDevices) {
protocolStore.archiveSession(SignalProtocolAddress(recipient.identifier, extra))
}
for (missing in mismatched.missingDevices) {
val address = SignalProtocolAddress(recipient.identifier, missing)
initializeSession(recipient, address, sealedSenderAccess)
}
}
private fun OutgoingPushMessage.toWireMessage(): MessageApiV2.Message = MessageApiV2.Message(
type = type,
destinationDeviceId = destinationDeviceId,
destinationRegistrationId = destinationRegistrationId,
content = content
)
private fun Raise<SendError>.buildPreKeyBundle(
identityKey: ByteArray,
item: KeysApiV2.PreKeyResponseItem,
@@ -269,52 +395,60 @@ open class MessageService(
*/
data class SendSuccess(
val envelopeContent: EnvelopeContent,
val sentUnidentified: Boolean,
val sentSealedSender: Boolean,
val devices: List<Int>
)
sealed interface SendError {
private enum class SealedSendResult {
Success, InvalidAccessKey, MismatchedDevices
}
private enum class UnsealedSendResult {
Success, MismatchedDevices
}
sealed class SendError : Exception() {
/** You discovered a safety number change during sending. */
data class IdentityMismatch(val recipient: SignalServiceAddress, val cause: UntrustedIdentityException) : SendError
data class IdentityMismatch(val serviceId: ServiceId, val exception: UntrustedIdentityException) : SendError()
/** The recipient is no longer registered. */
data object NotRegistered : SendError
class NotRegistered : SendError()
/** Invalid credentials. You are likely no longer registered. */
data object Unauthorized : SendError
class Unauthorized : SendError()
/**
* The server wants you to complete a push challenge/captcha before continuing.
* [token] is the challenge token; [options] enumerates the supported challenge mechanisms
* (e.g. "captcha", "pushChallenge"). [retryAfter] is the Retry-After hint, if provided.
*/
data class ChallengeRequired(val token: String, val options: List<String>, val retryAfter: Duration?) : SendError
data class ChallengeRequired(val token: String, val options: Set<ChallengeOption>, val retryAfter: Duration?) : SendError()
/** The server has fully rejected your request. This usually only happens during times of turmoil. Fail and require user action to resend. */
data object ServerRejected : SendError
class ServerRejected : SendError()
/**
* The encoded content exceeded the configured size cap. Permanent failure for this message —
* retrying with the same content won't help.
*/
data class ContentTooLarge(val size: Long, val maxAllowed: Long) : SendError
data class ContentTooLarge(val size: Long, val maxAllowed: Long) : SendError()
/**
* Each send attempt may result in us having to establish sessions with linked devices and such. This indicates that we hit our max attempt count while
* trying to handle these situations. It should be safe to retry with normal backoff.
*/
data object SessionAttemptsExhausted : SendError
class SessionAttemptsExhausted : SendError()
/** We needed to establish a session, but the server was missing either a signed or kyber prekey for the user. */
data class PreKeyUnavailable(val reason: String) : SendError
data class PreKeyUnavailable(val reason: String) : SendError()
/** You're rate-limited. Use the [retryAfter] for your backoff. */
data class RateLimited(val retryAfter: Duration?) : SendError
data class RateLimited(val retryAfter: Duration?) : SendError()
/** A generic, retryable network error. */
data class NetworkError(val cause: IOException) : SendError
data class NetworkError(val exception: IOException) : SendError()
/** An unexpected error. You should likely crash. */
data class ApplicationError(val cause: Throwable) : SendError
data class ApplicationError(val exception: Throwable) : SendError()
}
}
@@ -1,196 +0,0 @@
/*
* Copyright 2026 Signal Messenger, LLC
* SPDX-License-Identifier: AGPL-3.0-only
*/
package org.signal.network.api
import assertk.assertThat
import assertk.assertions.isEqualTo
import assertk.assertions.isInstanceOf
import assertk.assertions.isSameInstanceAs
import io.mockk.coEvery
import io.mockk.every
import io.mockk.mockk
import kotlinx.coroutines.test.runTest
import org.junit.Test
import org.signal.libsignal.net.RequestResult
import org.signal.network.websocket.WebSocketRequestMessage
import org.signal.network.websocket.WebsocketResponse
import org.whispersystems.signalservice.api.crypto.SealedSenderAccess
import org.whispersystems.signalservice.api.websocket.SignalWebSocket
import java.io.IOException
import kotlin.time.Duration.Companion.seconds
class MessageApiV2Test {
private val authSocket: SignalWebSocket.AuthenticatedWebSocket = mockk()
private val unauthSocket: SignalWebSocket.UnauthenticatedWebSocket = mockk()
private val api = MessageApiV2(authSocket, unauthSocket)
private val request = MessageApiV2.SendMessageRequest(
messages = listOf(MessageApiV2.Message(type = 1, destinationDeviceId = 1, destinationRegistrationId = 42, content = "abc")),
timestamp = 1_700_000_000L
)
@Test
fun `200 parses SendMessageResponse and flags sentUnidentified from response`() = runTest {
stubAuth(status = 200, body = """{"needsSync": true}""", unidentified = true)
val result = api.sendMessage("destination-id", request, sealedSenderAccess = null, story = false)
assertThat(result).isInstanceOf(RequestResult.Success::class)
val success = result as RequestResult.Success<MessageApiV2.SendMessageResponse>
assertThat(success.result.needsSync).isEqualTo(true)
assertThat(success.result.sentUnidentified).isEqualTo(true)
}
@Test
fun `401 maps to Unauthorized`() = runTest {
stubAuth(status = 401)
val result = api.sendMessage("destination-id", request, sealedSenderAccess = null, story = false)
assertNonSuccess(result, MessageApiV2.SendMessageError.Unauthorized)
}
@Test
fun `404 maps to NotRegistered`() = runTest {
stubAuth(status = 404)
val result = api.sendMessage("destination-id", request, sealedSenderAccess = null, story = false)
assertNonSuccess(result, MessageApiV2.SendMessageError.NotRegistered)
}
@Test
fun `409 parses MismatchedDevices body`() = runTest {
stubAuth(status = 409, body = """{"missingDevices": [2, 3], "extraDevices": [5]}""")
val result = api.sendMessage("destination-id", request, sealedSenderAccess = null, story = false)
val nonSuccess = result as RequestResult.NonSuccess
val err = nonSuccess.error as MessageApiV2.SendMessageError.MismatchedDevicesError
assertThat(err.devices.missingDevices).isEqualTo(listOf(2, 3))
assertThat(err.devices.extraDevices).isEqualTo(listOf(5))
}
@Test
fun `410 parses StaleDevices body`() = runTest {
stubAuth(status = 410, body = """{"staleDevices": [2]}""")
val result = api.sendMessage("destination-id", request, sealedSenderAccess = null, story = false)
val nonSuccess = result as RequestResult.NonSuccess
val err = nonSuccess.error as MessageApiV2.SendMessageError.StaleDevicesError
assertThat(err.devices.staleDevices).isEqualTo(listOf(2))
}
@Test
fun `428 parses ProofRequired body and Retry-After header`() = runTest {
val response: WebsocketResponse = mockk {
every { status } returns 428
every { body } returns """{"token": "abc123", "options": ["captcha", "pushChallenge"]}"""
every { isUnidentified } returns false
every { getHeader("retry-after") } returns "120"
}
coEvery { authSocket.requestSuspend(any<WebSocketRequestMessage>()) } returns response
val result = api.sendMessage("destination-id", request, sealedSenderAccess = null, story = false)
val err = (result as RequestResult.NonSuccess).error as MessageApiV2.SendMessageError.ChallengeRequired
assertThat(err.token).isEqualTo("abc123")
assertThat(err.options).isEqualTo(listOf("captcha", "pushChallenge"))
assertThat(err.retryAfter).isEqualTo(120.seconds)
}
@Test
fun `429 with retry-after header maps to RateLimited with Duration`() = runTest {
val response: WebsocketResponse = mockk {
every { status } returns 429
every { body } returns "{}"
every { isUnidentified } returns false
every { getHeader("retry-after") } returns "42"
}
coEvery { authSocket.requestSuspend(any<WebSocketRequestMessage>()) } returns response
val result = api.sendMessage("destination-id", request, sealedSenderAccess = null, story = false)
val err = (result as RequestResult.NonSuccess).error as MessageApiV2.SendMessageError.RateLimited
assertThat(err.retryAfter).isEqualTo(42.seconds)
}
@Test
fun `429 without retry-after header maps to RateLimited with null Duration`() = runTest {
val response: WebsocketResponse = mockk {
every { status } returns 429
every { body } returns "{}"
every { isUnidentified } returns false
every { getHeader("retry-after") } returns null
}
coEvery { authSocket.requestSuspend(any<WebSocketRequestMessage>()) } returns response
val result = api.sendMessage("destination-id", request, sealedSenderAccess = null, story = false)
val err = (result as RequestResult.NonSuccess).error as MessageApiV2.SendMessageError.RateLimited
assertThat(err.retryAfter).isEqualTo(null)
}
@Test
fun `508 maps to ServerRejected`() = runTest {
stubAuth(status = 508)
val result = api.sendMessage("destination-id", request, sealedSenderAccess = null, story = false)
assertNonSuccess(result, MessageApiV2.SendMessageError.ServerRejected)
}
@Test
fun `unexpected status maps to ApplicationError`() = runTest {
stubAuth(status = 418)
val result = api.sendMessage("destination-id", request, sealedSenderAccess = null, story = false)
assertThat(result).isInstanceOf(RequestResult.ApplicationError::class)
}
@Test
fun `IOException from socket becomes RetryableNetworkError`() = runTest {
val ioError = IOException("socket closed")
coEvery { authSocket.requestSuspend(any<WebSocketRequestMessage>()) } throws ioError
val result = api.sendMessage("destination-id", request, sealedSenderAccess = null, story = false)
val retry = result as RequestResult.RetryableNetworkError
assertThat(retry.networkError).isSameInstanceAs(ioError)
}
@Test
fun `sealedSenderAccess routes to unauthenticated socket`() = runTest {
val sealed: SealedSenderAccess = mockk()
val response: WebsocketResponse = mockk {
every { status } returns 200
every { body } returns """{"needsSync": false}"""
every { isUnidentified } returns true
}
coEvery { unauthSocket.requestSuspend(any(), sealed) } returns response
val result = api.sendMessage("destination-id", request, sealedSenderAccess = sealed, story = false)
assertThat(result).isInstanceOf(RequestResult.Success::class)
}
private fun stubAuth(status: Int, body: String = "{}", unidentified: Boolean = false) {
val response: WebsocketResponse = mockk {
every { this@mockk.status } returns status
every { this@mockk.body } returns body
every { isUnidentified } returns unidentified
}
coEvery { authSocket.requestSuspend(any<WebSocketRequestMessage>()) } returns response
}
private fun assertNonSuccess(result: RequestResult<*, *>, expected: MessageApiV2.SendMessageError) {
val nonSuccess = result as RequestResult.NonSuccess
assertThat(nonSuccess.error).isEqualTo(expected)
}
}
@@ -12,6 +12,7 @@ import assertk.assertions.isEqualTo
import assertk.assertions.isInstanceOf
import io.mockk.coEvery
import io.mockk.coVerify
import io.mockk.coVerifyOrder
import io.mockk.every
import io.mockk.mockk
import io.mockk.spyk
@@ -19,9 +20,28 @@ import io.mockk.verify
import kotlinx.coroutines.test.runTest
import org.junit.Test
import org.signal.core.models.ServiceId
import org.signal.core.util.Base64
import org.signal.libsignal.net.MismatchedDeviceException
import org.signal.libsignal.net.RequestResult
import org.signal.libsignal.net.RequestUnauthorizedException
import org.signal.libsignal.net.ServiceIdNotFoundException
import org.signal.libsignal.net.UserBasedAuthorization
import org.signal.libsignal.net.UserBasedSendAuthorization
import org.signal.libsignal.protocol.IdentityKeyPair
import org.signal.libsignal.protocol.SessionBuilder
import org.signal.libsignal.protocol.SessionCipher
import org.signal.libsignal.protocol.SignalProtocolAddress
import org.signal.libsignal.protocol.UntrustedIdentityException
import org.signal.libsignal.protocol.ecc.ECKeyPair
import org.signal.libsignal.protocol.kem.KEMKeyPair
import org.signal.libsignal.protocol.kem.KEMKeyType
import org.signal.libsignal.protocol.message.CiphertextMessage
import org.signal.libsignal.protocol.message.PreKeySignalMessage
import org.signal.libsignal.protocol.state.KyberPreKeyRecord
import org.signal.libsignal.protocol.state.PreKeyBundle
import org.signal.libsignal.protocol.state.PreKeyRecord
import org.signal.libsignal.protocol.state.SignedPreKeyRecord
import org.signal.libsignal.protocol.state.impl.InMemorySignalProtocolStore
import org.signal.network.api.KeysApiV2
import org.signal.network.api.MessageApiV2
import org.whispersystems.signalservice.api.SignalServiceAccountDataStore
@@ -29,6 +49,7 @@ import org.whispersystems.signalservice.api.SignalSessionLock
import org.whispersystems.signalservice.api.crypto.EnvelopeContent
import org.whispersystems.signalservice.api.crypto.SealedSenderAccess
import org.whispersystems.signalservice.api.crypto.SignalServiceCipher
import org.whispersystems.signalservice.api.crypto.UnidentifiedAccess
import org.whispersystems.signalservice.api.push.SignalServiceAddress
import org.whispersystems.signalservice.internal.push.OutgoingPushMessage
import java.io.IOException
@@ -57,230 +78,411 @@ class MessageServiceTest {
@Test
fun `happy path with existing session returns Success`() = runTest {
val service = newService()
every { protocolStore.getSubDeviceSessions(recipient.identifier) } returns emptyList()
every { protocolStore.containsSession(SignalProtocolAddress(recipient.identifier, 1)) } returns true
every { cipher.encrypt(any(), any(), any()) } returns OutgoingPushMessage(1, 1, 100, "payload")
every { protocolStore.getSubDeviceSessions(recipientAci.toString()) } returns emptyList()
every { cipher.encrypt(any(), any(), any()) } returns OutgoingPushMessage(1, 1, 100, "AAAA")
coEvery { messageApi.sendSealedSenderMessage(eq(recipientAci), any(), any(), any(), any(), any()) } returns
RequestResult.Success(Unit)
coEvery { messageApi.sendMessage(recipient.identifier, any(), null, false) } returns
RequestResult.Success(MessageApiV2.SendMessageResponse(sentUnidentified = true))
val result = service.sendMessage(recipient, envelopeContent, timestamp, sealedSenderAccess = null, story = false, isOnline = false)
val result = service.sendMessage(recipientAci, envelopeContent, timestamp, sealedSenderAccess = null, story = true, isOnline = false)
val success = (result as Either.Right).value
assertThat(success.sentUnidentified).isEqualTo(true)
assertThat(success.sentSealedSender).isEqualTo(true)
assertThat(success.devices).isEqualTo(listOf(1))
}
@Test
fun `isOnline true is forwarded to the send request`() = runTest {
val service = newService()
every { protocolStore.getSubDeviceSessions(recipient.identifier) } returns emptyList()
every { protocolStore.containsSession(SignalProtocolAddress(recipient.identifier, 1)) } returns true
every { cipher.encrypt(any(), any(), any()) } returns OutgoingPushMessage(1, 1, 100, "payload")
coEvery { messageApi.sendMessage(any(), any(), any(), any()) } returns
RequestResult.Success(MessageApiV2.SendMessageResponse())
every { protocolStore.getSubDeviceSessions(recipientAci.toString()) } returns emptyList()
every { cipher.encrypt(any(), any(), any()) } returns OutgoingPushMessage(1, 1, 100, "AAAA")
coEvery { messageApi.sendSealedSenderMessage(any(), any(), any(), any(), any(), any()) } returns
RequestResult.Success(Unit)
service.sendMessage(recipient, envelopeContent, timestamp, sealedSenderAccess = null, story = false, isOnline = true)
service.sendMessage(recipientAci, envelopeContent, timestamp, sealedSenderAccess = null, story = true, isOnline = true)
coVerify {
messageApi.sendMessage(
recipient.identifier,
match<MessageApiV2.SendMessageRequest> { it.online },
null,
false
)
messageApi.sendSealedSenderMessage(eq(recipientAci), any(), any(), any(), eq(true), any())
}
}
@Test
fun `urgent false is forwarded to the send request`() = runTest {
val service = newService()
every { protocolStore.getSubDeviceSessions(recipientAci.toString()) } returns emptyList()
every { cipher.encrypt(any(), any(), any()) } returns OutgoingPushMessage(1, 1, 100, "AAAA")
coEvery { messageApi.sendSealedSenderMessage(any(), any(), any(), any(), any(), any()) } returns
RequestResult.Success(Unit)
service.sendMessage(recipientAci, envelopeContent, timestamp, sealedSenderAccess = null, story = true, isOnline = false, urgent = false)
coVerify {
messageApi.sendSealedSenderMessage(eq(recipientAci), any(), any(), any(), any(), eq(false))
}
}
@Test
fun `story uses Story auth`() = runTest {
val service = newService()
every { protocolStore.getSubDeviceSessions(recipientAci.toString()) } returns emptyList()
every { cipher.encrypt(any(), any(), any()) } returns OutgoingPushMessage(1, 1, 100, "AAAA")
coEvery { messageApi.sendSealedSenderMessage(any(), any(), any(), any(), any(), any()) } returns
RequestResult.Success(Unit)
service.sendMessage(recipientAci, envelopeContent, timestamp, sealedSenderAccess = null, story = true, isOnline = false)
coVerify {
messageApi.sendSealedSenderMessage(any(), any(), any(), eq(UserBasedSendAuthorization.Story), any(), any())
}
}
@Test
fun `non-story sealed send with non-zero access key uses AccessKey auth`() = runTest {
val service = newService()
every { protocolStore.getSubDeviceSessions(recipientAci.toString()) } returns emptyList()
every { cipher.encrypt(any(), any(), any()) } returns OutgoingPushMessage(1, 1, 100, "AAAA")
coEvery { messageApi.sendSealedSenderMessage(any(), any(), any(), any(), any(), any()) } returns
RequestResult.Success(Unit)
val accessKey = ByteArray(16) { 1 }
val sealed = individualUnidentifiedAccessFirst(accessKey)
service.sendMessage(recipientAci, envelopeContent, timestamp, sealedSenderAccess = sealed, story = false, isOnline = false)
coVerify {
messageApi.sendSealedSenderMessage(any(), any(), any(), eq(UserBasedAuthorization.AccessKey(accessKey)), any(), any())
}
}
@Test
fun `non-story sealed send with zero access key uses UnrestrictedUnauthenticatedAccess auth`() = runTest {
val service = newService()
every { protocolStore.getSubDeviceSessions(recipientAci.toString()) } returns emptyList()
every { cipher.encrypt(any(), any(), any()) } returns OutgoingPushMessage(1, 1, 100, "AAAA")
coEvery { messageApi.sendSealedSenderMessage(any(), any(), any(), any(), any(), any()) } returns
RequestResult.Success(Unit)
val sealed = individualUnidentifiedAccessFirst(ByteArray(16))
service.sendMessage(recipientAci, envelopeContent, timestamp, sealedSenderAccess = sealed, story = false, isOnline = false)
coVerify {
messageApi.sendSealedSenderMessage(any(), any(), any(), eq(UserBasedAuthorization.UnrestrictedUnauthenticatedAccess), any(), any())
}
}
@Test
fun `non-story sealed send with unsupported access type raises ApplicationError`() = runTest {
val service = newService()
every { protocolStore.getSubDeviceSessions(recipientAci.toString()) } returns emptyList()
every { cipher.encrypt(any(), any(), any()) } returns OutgoingPushMessage(1, 1, 100, "AAAA")
val unsupportedAccess = mockk<SealedSenderAccess.IndividualGroupSendTokenFirst>()
val result = service.sendMessage(recipientAci, envelopeContent, timestamp, sealedSenderAccess = unsupportedAccess, story = false, isOnline = false)
val app = (result as Either.Left).value as MessageService.SendError.ApplicationError
assertThat(app.exception).isInstanceOf(IllegalArgumentException::class)
}
@Test
fun `sub-device without session is excluded from target devices`() = runTest {
val service = newService()
every { protocolStore.getSubDeviceSessions(recipient.identifier) } returns listOf(2, 3)
every { protocolStore.containsSession(SignalProtocolAddress(recipient.identifier, 2)) } returns true
every { protocolStore.containsSession(SignalProtocolAddress(recipient.identifier, 3)) } returns false
every { cipher.encrypt(any(), any(), any()) } returns OutgoingPushMessage(1, 1, 100, "payload")
every { protocolStore.getSubDeviceSessions(recipientAci.toString()) } returns listOf(2, 3)
every { protocolStore.containsSession(SignalProtocolAddress(recipientAci.libSignalServiceId, 2)) } returns true
every { protocolStore.containsSession(SignalProtocolAddress(recipientAci.libSignalServiceId, 3)) } returns false
every { cipher.encrypt(any(), any(), any()) } returns OutgoingPushMessage(1, 1, 100, "AAAA")
coEvery { messageApi.sendMessage(recipient.identifier, any(), null, false) } returns
RequestResult.Success(MessageApiV2.SendMessageResponse())
coEvery { messageApi.sendSealedSenderMessage(any(), any(), any(), any(), any(), any()) } returns
RequestResult.Success(Unit)
val result = service.sendMessage(recipient, envelopeContent, timestamp, sealedSenderAccess = null, story = false, isOnline = false)
val result = service.sendMessage(recipientAci, envelopeContent, timestamp, sealedSenderAccess = null, story = true, isOnline = false)
assertThat(result).isInstanceOf(Either.Right::class)
verify { cipher.encrypt(SignalProtocolAddress(recipient.identifier, 1), any(), any()) }
verify { cipher.encrypt(SignalProtocolAddress(recipient.identifier, 2), any(), any()) }
verify(exactly = 0) { cipher.encrypt(SignalProtocolAddress(recipient.identifier, 3), any(), any()) }
verify { cipher.encrypt(SignalProtocolAddress(recipientAci.libSignalServiceId, 1), any(), any()) }
verify { cipher.encrypt(SignalProtocolAddress(recipientAci.libSignalServiceId, 2), any(), any()) }
verify(exactly = 0) { cipher.encrypt(SignalProtocolAddress(recipientAci.libSignalServiceId, 3), any(), any()) }
}
@Test
fun `409 MismatchedDevices archives extras, fetches missing prekeys, and retries`() = runTest {
fun `MismatchedDeviceException archives extras and fetches missing prekeys`() = runTest {
val service = newService()
every { protocolStore.getSubDeviceSessions(recipient.identifier) } returns emptyList()
every { protocolStore.containsSession(any()) } returns true
every { cipher.encrypt(any(), any(), any()) } returns OutgoingPushMessage(1, 1, 100, "payload")
every { protocolStore.getSubDeviceSessions(recipientAci.toString()) } returns emptyList()
every { cipher.encrypt(any(), any(), any()) } returns OutgoingPushMessage(1, 1, 100, "AAAA")
val mismatched = MessageApiV2.MismatchedDevices(missingDevices = listOf(2), extraDevices = listOf(5))
coEvery { messageApi.sendMessage(recipient.identifier, any(), null, false) } returnsMany listOf(
RequestResult.NonSuccess(MessageApiV2.SendMessageError.MismatchedDevicesError(mismatched)),
RequestResult.Success(MessageApiV2.SendMessageResponse())
)
coEvery { keysApi.getPreKey(recipient.identifier, 2, null) } returns
val mismatch = mismatchedException(missing = intArrayOf(2), extra = intArrayOf(5))
coEvery { messageApi.sendSealedSenderMessage(any(), any(), any(), any(), any(), any()) } returns
RequestResult.NonSuccess(mismatch)
coEvery { keysApi.getPreKey(recipientAci.toString(), 2, null) } returns
RequestResult.Success(KeysApiV2.PreKeyResponse(identityKey = ByteArray(0), devices = emptyList()))
val result = service.sendMessage(recipient, envelopeContent, timestamp, sealedSenderAccess = null, story = false, isOnline = false)
service.sendMessage(recipientAci, envelopeContent, timestamp, sealedSenderAccess = null, story = true, isOnline = false)
assertThat(result).isInstanceOf(Either.Right::class)
verify { protocolStore.archiveSession(SignalProtocolAddress(recipient.identifier, 5)) }
coVerify { keysApi.getPreKey(recipient.identifier, 2, null) }
verify { protocolStore.archiveSession(SignalProtocolAddress(recipientAci.libSignalServiceId, 5)) }
coVerify { keysApi.getPreKey(recipientAci.toString(), 2, null) }
}
@Test
fun `410 StaleDevices archives stales and retries`() = runTest {
fun `MismatchedDeviceException archives stale devices`() = runTest {
val service = newService()
every { protocolStore.getSubDeviceSessions(recipient.identifier) } returns emptyList()
every { protocolStore.containsSession(any()) } returns true
every { cipher.encrypt(any(), any(), any()) } returns OutgoingPushMessage(1, 1, 100, "payload")
every { protocolStore.getSubDeviceSessions(recipientAci.toString()) } returns emptyList()
every { cipher.encrypt(any(), any(), any()) } returns OutgoingPushMessage(1, 1, 100, "AAAA")
val stale = MessageApiV2.StaleDevices(staleDevices = listOf(3))
coEvery { messageApi.sendMessage(recipient.identifier, any(), null, false) } returnsMany listOf(
RequestResult.NonSuccess(MessageApiV2.SendMessageError.StaleDevicesError(stale)),
RequestResult.Success(MessageApiV2.SendMessageResponse())
)
val mismatch = mismatchedException(stale = intArrayOf(3))
coEvery { messageApi.sendSealedSenderMessage(any(), any(), any(), any(), any(), any()) } returns
RequestResult.NonSuccess(mismatch)
val result = service.sendMessage(recipient, envelopeContent, timestamp, sealedSenderAccess = null, story = false, isOnline = false)
service.sendMessage(recipientAci, envelopeContent, timestamp, sealedSenderAccess = null, story = true, isOnline = false)
assertThat(result).isInstanceOf(Either.Right::class)
verify { protocolStore.archiveSession(SignalProtocolAddress(recipient.identifier, 3)) }
verify { protocolStore.archiveSession(SignalProtocolAddress(recipientAci.libSignalServiceId, 3)) }
}
@Test
fun `repeated device conflicts exhaust retries`() = runTest {
fun `single mismatched device recovers and the next send attempt succeeds`() = runTest {
val service = newService()
every { protocolStore.getSubDeviceSessions(recipient.identifier) } returns emptyList()
every { protocolStore.containsSession(any()) } returns true
every { cipher.encrypt(any(), any(), any()) } returns OutgoingPushMessage(1, 1, 100, "payload")
every { protocolStore.getSubDeviceSessions(recipientAci.toString()) } returns emptyList()
every { cipher.encrypt(any(), any(), any()) } returns OutgoingPushMessage(1, 1, 100, "AAAA")
val stale = MessageApiV2.StaleDevices(staleDevices = listOf(4))
coEvery { messageApi.sendMessage(recipient.identifier, any(), null, false) } returns
RequestResult.NonSuccess(MessageApiV2.SendMessageError.StaleDevicesError(stale))
coEvery { messageApi.sendSealedSenderMessage(any(), any(), any(), any(), any(), any()) } returns
RequestResult.NonSuccess(mismatchedException(missing = intArrayOf(2))) andThen
RequestResult.Success(Unit)
coEvery { keysApi.getPreKey(recipientAci.toString(), 2, null) } returns
RequestResult.Success(KeysApiV2.PreKeyResponse(identityKey = ByteArray(0), devices = emptyList()))
val result = service.sendMessage(recipient, envelopeContent, timestamp, sealedSenderAccess = null, story = false, isOnline = false)
val result = service.sendMessage(recipientAci, envelopeContent, timestamp, sealedSenderAccess = null, story = true, isOnline = false)
assertThat(result).isEqualTo(Either.Left(MessageService.SendError.SessionAttemptsExhausted))
val success = (result as Either.Right).value
assertThat(success.devices).isEqualTo(listOf(1))
coVerify(exactly = 2) { messageApi.sendSealedSenderMessage(eq(recipientAci), any(), any(), any(), any(), any()) }
coVerify { keysApi.getPreKey(recipientAci.toString(), 2, null) }
}
@Test
fun `401 maps to Unauthorized`() = runTest {
fun `ServiceIdNotFoundException maps to NotRegistered`() = runTest {
val service = newService()
every { protocolStore.getSubDeviceSessions(recipient.identifier) } returns emptyList()
every { protocolStore.containsSession(any()) } returns true
every { cipher.encrypt(any(), any(), any()) } returns OutgoingPushMessage(1, 1, 100, "payload")
coEvery { messageApi.sendMessage(any(), any(), any(), any()) } returns
RequestResult.NonSuccess(MessageApiV2.SendMessageError.Unauthorized)
every { protocolStore.getSubDeviceSessions(recipientAci.toString()) } returns emptyList()
every { cipher.encrypt(any(), any(), any()) } returns OutgoingPushMessage(1, 1, 100, "AAAA")
coEvery { messageApi.sendSealedSenderMessage(any(), any(), any(), any(), any(), any()) } returns
RequestResult.NonSuccess(ServiceIdNotFoundException("not registered"))
val result = service.sendMessage(recipient, envelopeContent, timestamp, sealedSenderAccess = null, story = false, isOnline = false)
val result = service.sendMessage(recipientAci, envelopeContent, timestamp, sealedSenderAccess = null, story = true, isOnline = false)
assertThat(result).isEqualTo(Either.Left(MessageService.SendError.Unauthorized))
val left = (result as Either.Left).value
assertThat(left).isInstanceOf(MessageService.SendError.NotRegistered::class)
}
@Test
fun `404 maps to NotRegistered`() = runTest {
fun `RequestUnauthorizedException from sealed send is retried and exhausts attempts`() = runTest {
val service = newService()
every { protocolStore.getSubDeviceSessions(recipient.identifier) } returns emptyList()
every { protocolStore.containsSession(any()) } returns true
every { cipher.encrypt(any(), any(), any()) } returns OutgoingPushMessage(1, 1, 100, "payload")
coEvery { messageApi.sendMessage(any(), any(), any(), any()) } returns
RequestResult.NonSuccess(MessageApiV2.SendMessageError.NotRegistered)
every { protocolStore.getSubDeviceSessions(recipientAci.toString()) } returns emptyList()
every { cipher.encrypt(any(), any(), any()) } returns OutgoingPushMessage(1, 1, 100, "AAAA")
coEvery { messageApi.sendSealedSenderMessage(any(), any(), any(), any(), any(), any()) } returns
RequestResult.NonSuccess(RequestUnauthorizedException("bad access"))
val result = service.sendMessage(recipient, envelopeContent, timestamp, sealedSenderAccess = null, story = false, isOnline = false)
val result = service.sendMessage(recipientAci, envelopeContent, timestamp, sealedSenderAccess = null, story = true, isOnline = false)
assertThat(result).isEqualTo(Either.Left(MessageService.SendError.NotRegistered))
val left = (result as Either.Left).value
assertThat(left).isInstanceOf(MessageService.SendError.SessionAttemptsExhausted::class)
coVerify(exactly = 3) { messageApi.sendSealedSenderMessage(eq(recipientAci), any(), any(), any(), any(), any()) }
}
@Test
fun `send 429 propagates retry-after duration via SendResult RateLimited`() = runTest {
fun `RequestUnauthorizedException from sealed send falls back to an unsealed send`() = runTest {
val service = newService()
every { protocolStore.getSubDeviceSessions(recipient.identifier) } returns emptyList()
every { protocolStore.containsSession(any()) } returns true
every { cipher.encrypt(any(), any(), any()) } returns OutgoingPushMessage(1, 1, 100, "payload")
coEvery { messageApi.sendMessage(any(), any(), any(), any()) } returns
RequestResult.NonSuccess(MessageApiV2.SendMessageError.RateLimited(retryAfter = 30.seconds))
every { protocolStore.getSubDeviceSessions(recipientAci.toString()) } returns emptyList()
every { cipher.encrypt(any(), any(), any()) } returns OutgoingPushMessage(1, 1, 100, validSerializedSignalMessageBase64())
coEvery { messageApi.sendSealedSenderMessage(any(), any(), any(), any(), any(), any()) } returns
RequestResult.NonSuccess(RequestUnauthorizedException("bad access"))
coEvery { messageApi.sendUnsealedSenderMessage(any(), any(), any(), any(), any()) } returns
RequestResult.Success(Unit)
val result = service.sendMessage(recipient, envelopeContent, timestamp, sealedSenderAccess = null, story = false, isOnline = false)
val sealed = individualUnidentifiedAccessFirst(ByteArray(16) { 1 })
assertThat(result).isEqualTo(Either.Left(MessageService.SendError.RateLimited(retryAfter = 30.seconds)))
val result = service.sendMessage(recipientAci, envelopeContent, timestamp, sealedSenderAccess = sealed, story = false, isOnline = false)
val success = (result as Either.Right).value
assertThat(success.sentSealedSender).isEqualTo(false)
coVerifyOrder {
messageApi.sendSealedSenderMessage(eq(recipientAci), any(), any(), any(), any(), any())
messageApi.sendUnsealedSenderMessage(eq(recipientAci), any(), any(), any(), any())
}
}
@Test
fun `prekey 429 during mismatched-device recovery propagates retry-after as RateLimited`() = runTest {
fun `RetryableNetworkError maps to NetworkError`() = runTest {
val service = newService()
every { protocolStore.getSubDeviceSessions(recipient.identifier) } returns emptyList()
every { protocolStore.containsSession(any()) } returns true
every { cipher.encrypt(any(), any(), any()) } returns OutgoingPushMessage(1, 1, 100, "payload")
val mismatched = MessageApiV2.MismatchedDevices(missingDevices = listOf(2), extraDevices = emptyList())
coEvery { messageApi.sendMessage(recipient.identifier, any(), null, false) } returns
RequestResult.NonSuccess(MessageApiV2.SendMessageError.MismatchedDevicesError(mismatched))
coEvery { keysApi.getPreKey(recipient.identifier, 2, null) } returns
RequestResult.NonSuccess(KeysApiV2.GetPreKeysError.RateLimited(retryAfter = 60.seconds))
val result = service.sendMessage(recipient, envelopeContent, timestamp, sealedSenderAccess = null, story = false, isOnline = false)
assertThat(result).isEqualTo(Either.Left(MessageService.SendError.RateLimited(retryAfter = 60.seconds)))
}
@Test
fun `IOException from send maps to NetworkError`() = runTest {
val service = newService()
every { protocolStore.getSubDeviceSessions(recipient.identifier) } returns emptyList()
every { protocolStore.containsSession(any()) } returns true
every { cipher.encrypt(any(), any(), any()) } returns OutgoingPushMessage(1, 1, 100, "payload")
every { protocolStore.getSubDeviceSessions(recipientAci.toString()) } returns emptyList()
every { cipher.encrypt(any(), any(), any()) } returns OutgoingPushMessage(1, 1, 100, "AAAA")
val ioError = IOException("down")
coEvery { messageApi.sendMessage(any(), any(), any(), any()) } returns RequestResult.RetryableNetworkError(ioError)
coEvery { messageApi.sendSealedSenderMessage(any(), any(), any(), any(), any(), any()) } returns
RequestResult.RetryableNetworkError(ioError)
val result = service.sendMessage(recipient, envelopeContent, timestamp, sealedSenderAccess = null, story = false, isOnline = false)
val result = service.sendMessage(recipientAci, envelopeContent, timestamp, sealedSenderAccess = null, story = true, isOnline = false)
val network = (result as Either.Left).value as MessageService.SendError.NetworkError
assertThat(network.cause).isEqualTo(ioError)
assertThat(network.exception).isEqualTo(ioError)
}
@Test
fun `ApplicationError from send is propagated`() = runTest {
val service = newService()
every { protocolStore.getSubDeviceSessions(recipientAci.toString()) } returns emptyList()
every { cipher.encrypt(any(), any(), any()) } returns OutgoingPushMessage(1, 1, 100, "AAAA")
val cause = IllegalStateException("boom")
coEvery { messageApi.sendSealedSenderMessage(any(), any(), any(), any(), any(), any()) } returns
RequestResult.ApplicationError(cause)
val result = service.sendMessage(recipientAci, envelopeContent, timestamp, sealedSenderAccess = null, story = true, isOnline = false)
val app = (result as Either.Left).value as MessageService.SendError.ApplicationError
assertThat(app.exception).isEqualTo(cause)
}
@Test
fun `UntrustedIdentityException during encryption maps to IdentityMismatch`() = runTest {
val service = newService()
every { protocolStore.getSubDeviceSessions(recipient.identifier) } returns emptyList()
every { protocolStore.containsSession(any()) } returns true
every { protocolStore.getSubDeviceSessions(recipientAci.toString()) } returns emptyList()
val untrusted = UntrustedIdentityException(recipient.identifier)
every { cipher.encrypt(any(), any(), any()) } throws untrusted
val result = service.sendMessage(recipient, envelopeContent, timestamp, sealedSenderAccess = null, story = false, isOnline = false)
val result = service.sendMessage(recipientAci, envelopeContent, timestamp, sealedSenderAccess = null, story = true, isOnline = false)
val mismatch = (result as Either.Left).value as MessageService.SendError.IdentityMismatch
assertThat(mismatch.cause).isEqualTo(untrusted)
assertThat(mismatch.exception).isEqualTo(untrusted)
}
@Test
fun `content larger than max returns ContentTooLarge before encryption`() = runTest {
val largeContent: EnvelopeContent = mockk {
every { size() } returns 1024
}
val service = newService(maxContentSizeBytes = 100)
val result = service.sendMessage(recipientAci, largeContent, timestamp, sealedSenderAccess = null, story = true, isOnline = false)
val tooLarge = (result as Either.Left).value as MessageService.SendError.ContentTooLarge
assertThat(tooLarge.size).isEqualTo(1024)
assertThat(tooLarge.maxAllowed).isEqualTo(100)
verify(exactly = 0) { cipher.encrypt(any(), any(), any()) }
}
@Test
fun `prekey fetch 404 during mismatched-device recovery propagates as PreKeyUnavailable`() = runTest {
val service = newService()
every { protocolStore.getSubDeviceSessions(recipient.identifier) } returns emptyList()
every { protocolStore.containsSession(any()) } returns true
every { cipher.encrypt(any(), any(), any()) } returns OutgoingPushMessage(1, 1, 100, "payload")
every { protocolStore.getSubDeviceSessions(recipientAci.toString()) } returns emptyList()
every { cipher.encrypt(any(), any(), any()) } returns OutgoingPushMessage(1, 1, 100, "AAAA")
val mismatched = MessageApiV2.MismatchedDevices(missingDevices = listOf(2), extraDevices = emptyList())
coEvery { messageApi.sendMessage(recipient.identifier, any(), null, false) } returns
RequestResult.NonSuccess(MessageApiV2.SendMessageError.MismatchedDevicesError(mismatched))
coEvery { keysApi.getPreKey(recipient.identifier, 2, null) } returns
coEvery { messageApi.sendSealedSenderMessage(any(), any(), any(), any(), any(), any()) } returns
RequestResult.NonSuccess(mismatchedException(missing = intArrayOf(2)))
coEvery { keysApi.getPreKey(recipientAci.toString(), 2, null) } returns
RequestResult.NonSuccess(KeysApiV2.GetPreKeysError.NotFound)
val result = service.sendMessage(recipient, envelopeContent, timestamp, sealedSenderAccess = null, story = false, isOnline = false)
val result = service.sendMessage(recipientAci, envelopeContent, timestamp, sealedSenderAccess = null, story = true, isOnline = false)
val left = (result as Either.Left).value
assertThat(left).isInstanceOf(MessageService.SendError.PreKeyUnavailable::class)
}
@Test
fun `prekey 429 during mismatched-device recovery propagates retry-after as RateLimited`() = runTest {
val service = newService()
every { protocolStore.getSubDeviceSessions(recipientAci.toString()) } returns emptyList()
every { cipher.encrypt(any(), any(), any()) } returns OutgoingPushMessage(1, 1, 100, "AAAA")
coEvery { messageApi.sendSealedSenderMessage(any(), any(), any(), any(), any(), any()) } returns
RequestResult.NonSuccess(mismatchedException(missing = intArrayOf(2)))
coEvery { keysApi.getPreKey(recipientAci.toString(), 2, null) } returns
RequestResult.NonSuccess(KeysApiV2.GetPreKeysError.RateLimited(retryAfter = 60.seconds))
val result = service.sendMessage(recipientAci, envelopeContent, timestamp, sealedSenderAccess = null, story = true, isOnline = false)
assertThat(result).isEqualTo(Either.Left(MessageService.SendError.RateLimited(retryAfter = 60.seconds)))
}
private fun individualUnidentifiedAccessFirst(accessKey: ByteArray): SealedSenderAccess.IndividualUnidentifiedAccessFirst {
val ua = mockk<UnidentifiedAccess>()
every { ua.unidentifiedAccessKey } returns accessKey
return SealedSenderAccess.IndividualUnidentifiedAccessFirst(ua)
}
private fun mismatchedException(
missing: IntArray = intArrayOf(),
extra: IntArray = intArrayOf(),
stale: IntArray = intArrayOf()
): MismatchedDeviceException {
val entry = MismatchedDeviceException.Entry(
account = recipientAci.libSignalServiceId,
missingDevices = missing,
extraDevices = extra,
staleDevices = stale
)
return MismatchedDeviceException("mismatched", arrayOf(entry))
}
/**
* Produces a genuinely valid, native-parseable serialized [org.signal.libsignal.protocol.message.SignalMessage]
* (a WHISPER_TYPE message) encoded as base64. The unsealed send path wraps the encrypted bytes in
* `SignalMessage(bytes)`, whose native constructor rejects arbitrary input — so the fallback test needs real
* ciphertext here rather than a placeholder like "AAAA". We get one by establishing a real session between two
* in-memory stores and capturing a post-handshake reply (the initiator's first message is a PreKeySignalMessage,
* so we round-trip once to reach a plain SignalMessage).
*/
private fun validSerializedSignalMessageBase64(): String {
val aliceAddress = SignalProtocolAddress("alice", 1)
val bobAddress = SignalProtocolAddress("bob", 1)
val aliceStore = InMemorySignalProtocolStore(IdentityKeyPair.generate(), 1)
val bobStore = InMemorySignalProtocolStore(IdentityKeyPair.generate(), 2)
SessionBuilder(aliceStore, bobAddress, aliceAddress).process(createPreKeyBundle(bobStore, deviceId = bobAddress.deviceId))
val aliceCipher = SessionCipher(aliceStore, aliceAddress, bobAddress)
val preKeyMessage = aliceCipher.encrypt("hello".toByteArray())
val bobCipher = SessionCipher(bobStore, bobAddress, aliceAddress)
bobCipher.decrypt(PreKeySignalMessage(preKeyMessage.serialize()))
val reply = bobCipher.encrypt("reply".toByteArray())
check(reply.type == CiphertextMessage.WHISPER_TYPE) { "Expected a WHISPER_TYPE SignalMessage but got ${reply.type}" }
return Base64.encodeWithPadding(reply.serialize())
}
private fun createPreKeyBundle(store: InMemorySignalProtocolStore, deviceId: Int): PreKeyBundle {
val preKeyPair = ECKeyPair.generate()
val signedPreKeyPair = ECKeyPair.generate()
val signedPreKeySignature = store.identityKeyPair.privateKey.calculateSignature(signedPreKeyPair.publicKey.serialize())
val kyberPreKeyPair = KEMKeyPair.generate(KEMKeyType.KYBER_1024)
val kyberPreKeySignature = store.identityKeyPair.privateKey.calculateSignature(kyberPreKeyPair.publicKey.serialize())
val preKeyId = 1
val signedPreKeyId = 2
val kyberPreKeyId = 3
store.storePreKey(preKeyId, PreKeyRecord(preKeyId, preKeyPair))
store.storeSignedPreKey(signedPreKeyId, SignedPreKeyRecord(signedPreKeyId, 1L, signedPreKeyPair, signedPreKeySignature))
store.storeKyberPreKey(kyberPreKeyId, KyberPreKeyRecord(kyberPreKeyId, 1L, kyberPreKeyPair, kyberPreKeySignature))
return PreKeyBundle(
store.localRegistrationId,
deviceId,
preKeyId,
preKeyPair.publicKey,
signedPreKeyId,
signedPreKeyPair.publicKey,
signedPreKeySignature,
store.identityKeyPair.publicKey,
kyberPreKeyId,
kyberPreKeyPair.publicKey,
kyberPreKeySignature
)
}
/**
* Spy with `initializeSession` stubbed so tests don't exercise real crypto / native session building.
* The stub still invokes [KeysApiV2.getPreKey] and forwards non-success [RequestResult]s as the real
* implementation would; happy path is a no-op.
*/
private fun newService(): MessageService {
private fun newService(maxContentSizeBytes: Long = 0L): MessageService {
val spy: MessageService = spyk(
MessageService(
localAddress = localAddress,
@@ -289,7 +491,8 @@ class MessageServiceTest {
keysApi = keysApi,
protocolStore = protocolStore,
sessionLock = sessionLock,
cipher = cipher
cipher = cipher,
maxContentSizeBytes = maxContentSizeBytes
)
)
coEvery {
@@ -304,7 +507,7 @@ class MessageServiceTest {
is RequestResult.Success -> Unit
is RequestResult.NonSuccess -> raiseArg.raise(
when (val e = r.error) {
KeysApiV2.GetPreKeysError.Unauthorized -> MessageService.SendError.Unauthorized
KeysApiV2.GetPreKeysError.Unauthorized -> MessageService.SendError.Unauthorized()
KeysApiV2.GetPreKeysError.NotFound -> MessageService.SendError.PreKeyUnavailable("No prekeys found for $addressArg")
is KeysApiV2.GetPreKeysError.RateLimited -> MessageService.SendError.RateLimited(e.retryAfter)
}