From f76292769ae08b6b8e993587a2e7b812dc5f0d20 Mon Sep 17 00:00:00 2001 From: Greyson Parrelli Date: Fri, 29 May 2026 14:02:42 -0400 Subject: [PATCH] Update MessageApiV2 to use libsignal-net. --- .../securesms/jobs/IndividualSendJobV2.kt | 58 ++- .../org/signal/network/api/MessageApiV2.kt | 165 ++----- .../signal/network/service/MessageService.kt | 298 +++++++---- .../signal/network/api/MessageApiV2Test.kt | 196 -------- .../network/service/MessageServiceTest.kt | 461 +++++++++++++----- 5 files changed, 608 insertions(+), 570 deletions(-) delete mode 100644 lib/network/src/test/java/org/signal/network/api/MessageApiV2Test.kt diff --git a/app/src/main/java/org/thoughtcrime/securesms/jobs/IndividualSendJobV2.kt b/app/src/main/java/org/thoughtcrime/securesms/jobs/IndividualSendJobV2.kt index 6ffa233a32..172ad35081 100644 --- a/app/src/main/java/org/thoughtcrime/securesms/jobs/IndividualSendJobV2.kt +++ b/app/src/main/java/org/thoughtcrime/securesms/jobs/IndividualSendJobV2.kt @@ -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 diff --git a/lib/network/src/main/java/org/signal/network/api/MessageApiV2.kt b/lib/network/src/main/java/org/signal/network/api/MessageApiV2.kt index 87c89ffaec..58d5827eeb 100644 --- a/lib/network/src/main/java/org/signal/network/api/MessageApiV2.kt +++ b/lib/network/src/main/java/org/signal/network/api/MessageApiV2.kt @@ -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 { - 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, + auth: UserBasedSendAuthorization, + onlineOnly: Boolean, + urgent: Boolean + ): RequestResult { + return unauthWebSocket.runCatchingWithChatConnection { connection -> + UnauthMessagesService(connection).sendMessage(serviceId.libSignalServiceId, timestamp, contents, auth, onlineOnly, urgent) } } - @Serializable - data class SendMessageRequest( - val messages: List, - 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 = emptyList(), - val extraDevices: List = emptyList() - ) - - @Serializable - data class StaleDevices( - val staleDevices: List = 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 = 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, val retryAfter: Duration?) : SendMessageError() - data class RateLimited(val retryAfter: Duration?) : SendMessageError() - data object ServerRejected : SendMessageError() + suspend fun sendUnsealedSenderMessage( + serviceId: ServiceId, + timestamp: Long, + contents: List, + onlineOnly: Boolean, + urgent: Boolean + ): RequestResult { + return authWebSocket.runCatchingWithChatConnection { connection -> + AuthMessagesService(connection).sendMessage(serviceId.libSignalServiceId, timestamp, contents, onlineOnly, urgent) + } } } diff --git a/lib/network/src/main/java/org/signal/network/service/MessageService.kt b/lib/network/src/main/java/org/signal/network/service/MessageService.kt index f836cb10d8..6c94413263 100644 --- a/lib/network/src/main/java/org/signal/network/service/MessageService.kt +++ b/lib/network/src/main/java/org/signal/network/service/MessageService.kt @@ -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.sendSealed( + serviceId: ServiceId, + encryptedMessages: List, + 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 -> { + 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.sendUnsealed( + serviceId: ServiceId, + timestamp: Long, + encryptedMessages: List, + 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 -> { + 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.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.encryptForAllDevices( - recipient: SignalServiceAddress, + serviceId: ServiceId, envelopeContent: EnvelopeContent, sealedSenderAccess: SealedSenderAccess? ): List { - 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.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 { - val subDevices: MutableSet = (protocolStore.getSubDeviceSessions(recipient.identifier) + SignalServiceAddress.DEFAULT_DEVICE_ID).toMutableSet() + private fun targetDeviceIds(serviceId: ServiceId): List { + val subDevices: MutableSet = (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.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.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.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 ) - 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, val retryAfter: Duration?) : SendError + data class ChallengeRequired(val token: String, val options: Set, 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() } } diff --git a/lib/network/src/test/java/org/signal/network/api/MessageApiV2Test.kt b/lib/network/src/test/java/org/signal/network/api/MessageApiV2Test.kt deleted file mode 100644 index 38278c46f6..0000000000 --- a/lib/network/src/test/java/org/signal/network/api/MessageApiV2Test.kt +++ /dev/null @@ -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 - 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()) } 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()) } 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()) } 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()) } 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()) } returns response - } - - private fun assertNonSuccess(result: RequestResult<*, *>, expected: MessageApiV2.SendMessageError) { - val nonSuccess = result as RequestResult.NonSuccess - assertThat(nonSuccess.error).isEqualTo(expected) - } -} diff --git a/lib/network/src/test/java/org/signal/network/service/MessageServiceTest.kt b/lib/network/src/test/java/org/signal/network/service/MessageServiceTest.kt index d550a5e0ec..0825c7817f 100644 --- a/lib/network/src/test/java/org/signal/network/service/MessageServiceTest.kt +++ b/lib/network/src/test/java/org/signal/network/service/MessageServiceTest.kt @@ -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 { 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() + + 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() + 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) }