package expo.modules.signalexpo import expo.modules.kotlin.modules.Module import expo.modules.kotlin.modules.ModuleDefinition import expo.modules.kotlin.Promise import org.signal.libsignal.protocol.IdentityKey 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.ecc.Curve 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.SessionRecord import org.signal.libsignal.protocol.state.IdentityKeyStore import org.signal.libsignal.protocol.state.PreKeyStore import org.signal.libsignal.protocol.state.SignedPreKeyStore import org.signal.libsignal.protocol.state.SessionStore import org.signal.libsignal.protocol.message.PreKeySignalMessage import org.signal.libsignal.protocol.message.SignalMessage import org.signal.libsignal.protocol.message.CiphertextMessage import org.signal.libsignal.protocol.util.KeyHelper class SignalExpoModule : Module() { // Shared state private var identityKeyPair: IdentityKeyPair? = null private var localRegistrationId: Int = 0 // Protocol stores private val sessionStore = InMemorySessionStore() private val identityStore = InMemoryIdentityStore() private val preKeyStore = InMemoryPreKeyStore() private val signedPreKeyStore = InMemorySignedPreKeyStore() override fun definition() = ModuleDefinition { Name("SignalExpo") // MARK: - Key Generation Functions Function("generateIdentityKeyPair") { val keyPair = IdentityKeyPair.generate() mapOf( "publicKey" to keyPair.publicKey.serialize(), "privateKey" to keyPair.serialize() ) } Function("generateRegistrationId") { KeyHelper.generateRegistrationId(false) } Function("generatePreKeys") { start: Int, count: Int -> val preKeys = mutableListOf>() for (i in 0 until count) { val preKeyId = start + i val preKeyPair = Curve.generateKeyPair() val preKeyRecord = PreKeyRecord(preKeyId, preKeyPair) preKeys.add(mapOf( "id" to preKeyId, "publicKey" to preKeyRecord.keyPair.publicKey.serialize(), "privateKey" to preKeyRecord.keyPair.privateKey.serialize() )) } preKeys } Function("generateSignedPreKey") { identityKeyPairData: Map, signedPreKeyId: Int -> val privateKeyBytes = identityKeyPairData["privateKey"] ?: throw IllegalArgumentException("Missing privateKey") // Reconstruct identity key pair val identityKeyPair = IdentityKeyPair(privateKeyBytes) // Generate signed pre-key val signedPreKeyPair = Curve.generateKeyPair() val timestamp = System.currentTimeMillis() // Sign the public key with the identity key val signature = Curve.calculateSignature( identityKeyPair.privateKey, signedPreKeyPair.publicKey.serialize() ) mapOf( "id" to signedPreKeyId, "publicKey" to signedPreKeyPair.publicKey.serialize(), "privateKey" to signedPreKeyPair.privateKey.serialize(), "signature" to signature, "timestamp" to timestamp ) } // MARK: - Initialization AsyncFunction("initialize") { identityKeyPairData: Map, registrationId: Int, preKeysData: List>, signedPreKeyData: Map, promise: Promise -> try { val privateKeyBytes = identityKeyPairData["privateKey"] ?: throw IllegalArgumentException("Missing privateKey") // Set identity identityKeyPair = IdentityKeyPair(privateKeyBytes) localRegistrationId = registrationId // Configure identity store identityStore.identityKeyPair = identityKeyPair identityStore.localRegistrationId = registrationId // Store pre-keys for (preKeyData in preKeysData) { val id = (preKeyData["id"] as? Number)?.toInt() ?: continue val privateKey = preKeyData["privateKey"] as? ByteArray ?: continue val preKeyPair = Curve.generateKeyPair() // We need to reconstruct from private key val preKeyRecord = PreKeyRecord(id, Curve.decodePrivatePoint(privateKey).let { org.signal.libsignal.protocol.ecc.ECKeyPair(it.publicKey, it) }) preKeyStore.storePreKey(id, preKeyRecord) } // Store signed pre-key val signedId = (signedPreKeyData["id"] as? Number)?.toInt() ?: throw IllegalArgumentException("Missing signed pre-key id") val signedPrivateKey = signedPreKeyData["privateKey"] as? ByteArray ?: throw IllegalArgumentException("Missing signed pre-key privateKey") val signature = signedPreKeyData["signature"] as? ByteArray ?: throw IllegalArgumentException("Missing signed pre-key signature") val timestamp = (signedPreKeyData["timestamp"] as? Number)?.toLong() ?: throw IllegalArgumentException("Missing signed pre-key timestamp") val signedKeyPair = Curve.decodePrivatePoint(signedPrivateKey).let { org.signal.libsignal.protocol.ecc.ECKeyPair(it.publicKey, it) } val signedPreKeyRecord = SignedPreKeyRecord(signedId, timestamp, signedKeyPair, signature) signedPreKeyStore.storeSignedPreKey(signedId, signedPreKeyRecord) promise.resolve(null) } catch (e: Exception) { promise.reject("INIT_ERROR", e.message, e) } } // MARK: - Session Management Functions AsyncFunction("createSession") { addressData: Map, bundleData: Map, promise: Promise -> try { val address = parseAddress(addressData) val bundle = parsePreKeyBundle(bundleData) val sessionBuilder = SessionBuilder(sessionStore, preKeyStore, signedPreKeyStore, identityStore, address) sessionBuilder.process(bundle) promise.resolve(null) } catch (e: Exception) { promise.reject("SESSION_ERROR", e.message, e) } } AsyncFunction("hasSession") { addressData: Map, promise: Promise -> try { val address = parseAddress(addressData) val hasSession = sessionStore.containsSession(address) promise.resolve(hasSession) } catch (e: Exception) { promise.reject("SESSION_ERROR", e.message, e) } } AsyncFunction("deleteSession") { addressData: Map, promise: Promise -> try { val address = parseAddress(addressData) sessionStore.deleteSession(address) promise.resolve(null) } catch (e: Exception) { promise.reject("SESSION_ERROR", e.message, e) } } // MARK: - Encryption/Decryption Functions AsyncFunction("encrypt") { addressData: Map, plaintext: ByteArray, promise: Promise -> try { val address = parseAddress(addressData) val sessionCipher = SessionCipher(sessionStore, preKeyStore, signedPreKeyStore, identityStore, address) val ciphertext = sessionCipher.encrypt(plaintext) promise.resolve(mapOf( "type" to getMessageTypeName(ciphertext.type), "body" to ciphertext.serialize() )) } catch (e: Exception) { promise.reject("ENCRYPT_ERROR", e.message, e) } } AsyncFunction("decrypt") { addressData: Map, ciphertextData: Map, promise: Promise -> try { val address = parseAddress(addressData) val type = ciphertextData["type"] as? String ?: throw IllegalArgumentException("Missing ciphertext type") val body = ciphertextData["body"] as? ByteArray ?: throw IllegalArgumentException("Missing ciphertext body") val sessionCipher = SessionCipher(sessionStore, preKeyStore, signedPreKeyStore, identityStore, address) val plaintext = when (type) { "preKey" -> sessionCipher.decrypt(PreKeySignalMessage(body)) "whisper" -> sessionCipher.decrypt(SignalMessage(body)) else -> throw IllegalArgumentException("Unsupported message type: $type") } promise.resolve(mapOf( "plaintext" to plaintext )) } catch (e: Exception) { promise.reject("DECRYPT_ERROR", e.message, e) } } // MARK: - Storage Accessors AsyncFunction("getIdentityPublicKey") { promise: Promise -> try { val keyPair = identityKeyPair ?: throw IllegalStateException("Not initialized") promise.resolve(keyPair.publicKey.serialize()) } catch (e: Exception) { promise.reject("STORAGE_ERROR", e.message, e) } } AsyncFunction("getLocalRegistrationId") { promise: Promise -> try { if (identityKeyPair == null) { throw IllegalStateException("Not initialized") } promise.resolve(localRegistrationId) } catch (e: Exception) { promise.reject("STORAGE_ERROR", e.message, e) } } // MARK: - Clear all stores (for user switching) AsyncFunction("clear") { promise: Promise -> try { // Clear identity identityKeyPair = null localRegistrationId = 0 identityStore.identityKeyPair = null identityStore.localRegistrationId = 0 // Reset all in-memory stores sessionStore.clear() preKeyStore.clear() signedPreKeyStore.clear() identityStore.clear() promise.resolve(null) } catch (e: Exception) { promise.reject("CLEAR_ERROR", e.message, e) } } } // MARK: - Helper Methods private fun parseAddress(data: Map): SignalProtocolAddress { val name = data["name"] as? String ?: throw IllegalArgumentException("Missing address name") val deviceId = (data["deviceId"] as? Number)?.toInt() ?: throw IllegalArgumentException("Missing address deviceId") return SignalProtocolAddress(name, deviceId) } private fun parsePreKeyBundle(data: Map): PreKeyBundle { val registrationId = (data["registrationId"] as? Number)?.toInt() ?: throw IllegalArgumentException("Missing registrationId") val deviceId = (data["deviceId"] as? Number)?.toInt() ?: throw IllegalArgumentException("Missing deviceId") val signedPreKeyId = (data["signedPreKeyId"] as? Number)?.toInt() ?: throw IllegalArgumentException("Missing signedPreKeyId") val signedPreKeyPublic = data["signedPreKeyPublic"] as? ByteArray ?: throw IllegalArgumentException("Missing signedPreKeyPublic") val signedPreKeySignature = data["signedPreKeySignature"] as? ByteArray ?: throw IllegalArgumentException("Missing signedPreKeySignature") val identityKey = data["identityKey"] as? ByteArray ?: throw IllegalArgumentException("Missing identityKey") val preKeyId = (data["preKeyId"] as? Number)?.toInt() val preKeyPublic = data["preKeyPublic"] as? ByteArray return if (preKeyId != null && preKeyPublic != null) { PreKeyBundle( registrationId, deviceId, preKeyId, Curve.decodePoint(preKeyPublic, 0), signedPreKeyId, Curve.decodePoint(signedPreKeyPublic, 0), signedPreKeySignature, IdentityKey(identityKey, 0) ) } else { PreKeyBundle( registrationId, deviceId, 0, null, signedPreKeyId, Curve.decodePoint(signedPreKeyPublic, 0), signedPreKeySignature, IdentityKey(identityKey, 0) ) } } private fun getMessageTypeName(type: Int): String { return when (type) { CiphertextMessage.WHISPER_TYPE -> "whisper" CiphertextMessage.PREKEY_TYPE -> "preKey" CiphertextMessage.SENDERKEY_TYPE -> "senderKey" CiphertextMessage.PLAINTEXT_CONTENT_TYPE -> "plaintext" else -> "unknown" } } } // MARK: - In-Memory Stores class InMemorySessionStore : SessionStore { private val sessions = mutableMapOf() override fun loadSession(address: SignalProtocolAddress): SessionRecord { return sessions[address] ?: SessionRecord() } override fun getSubDeviceSessions(name: String): List { return sessions.keys.filter { it.name == name }.map { it.deviceId } } override fun storeSession(address: SignalProtocolAddress, record: SessionRecord) { sessions[address] = record } override fun containsSession(address: SignalProtocolAddress): Boolean { return sessions.containsKey(address) && sessions[address]?.hasSenderChain() == true } override fun deleteSession(address: SignalProtocolAddress) { sessions.remove(address) } override fun deleteAllSessions(name: String) { sessions.keys.filter { it.name == name }.forEach { sessions.remove(it) } } fun clear() { sessions.clear() } } class InMemoryIdentityStore : IdentityKeyStore { var identityKeyPair: IdentityKeyPair? = null var localRegistrationId: Int = 0 private val identities = mutableMapOf() override fun getIdentityKeyPair(): IdentityKeyPair { return identityKeyPair ?: throw IllegalStateException("Not initialized") } override fun getLocalRegistrationId(): Int { return localRegistrationId } override fun saveIdentity(address: SignalProtocolAddress, identityKey: IdentityKey): Boolean { val existing = identities[address] identities[address] = identityKey return existing != null && existing != identityKey } override fun isTrustedIdentity(address: SignalProtocolAddress, identityKey: IdentityKey, direction: IdentityKeyStore.Direction): Boolean { val existingIdentity = identities[address] ?: return true // Trust on first use return existingIdentity == identityKey } override fun getIdentity(address: SignalProtocolAddress): IdentityKey? { return identities[address] } fun clear() { identities.clear() } } class InMemoryPreKeyStore : PreKeyStore { private val preKeys = mutableMapOf() override fun loadPreKey(preKeyId: Int): PreKeyRecord { return preKeys[preKeyId] ?: throw org.signal.libsignal.protocol.InvalidKeyIdException("PreKey not found: $preKeyId") } override fun storePreKey(preKeyId: Int, record: PreKeyRecord) { preKeys[preKeyId] = record } override fun containsPreKey(preKeyId: Int): Boolean { return preKeys.containsKey(preKeyId) } override fun removePreKey(preKeyId: Int) { preKeys.remove(preKeyId) } fun clear() { preKeys.clear() } } class InMemorySignedPreKeyStore : SignedPreKeyStore { private val signedPreKeys = mutableMapOf() override fun loadSignedPreKey(signedPreKeyId: Int): SignedPreKeyRecord { return signedPreKeys[signedPreKeyId] ?: throw org.signal.libsignal.protocol.InvalidKeyIdException("SignedPreKey not found: $signedPreKeyId") } override fun loadSignedPreKeys(): List { return signedPreKeys.values.toList() } override fun storeSignedPreKey(signedPreKeyId: Int, record: SignedPreKeyRecord) { signedPreKeys[signedPreKeyId] = record } override fun containsSignedPreKey(signedPreKeyId: Int): Boolean { return signedPreKeys.containsKey(signedPreKeyId) } override fun removeSignedPreKey(signedPreKeyId: Int) { signedPreKeys.remove(signedPreKeyId) } fun clear() { signedPreKeys.clear() } }