fix(folio): detect screens from unique element presence, not id: selectors

testTag() in Compose is not exposed as resource-id without testTagsAsResourceId.
Use desc: selectors for elements unique to each screen instead of id: path queries.
This commit is contained in:
pj committed 2026-04-25 19:45:06 +07:00
1 parent 270d4bef7e
commit 3feb7b2d7f
18 files changed
+17 -1540

No files matched your search

+17 -15
View File
@@ -36,15 +36,17 @@ function parseCents(desc: string | null | undefined): number {
return Number(parts[1]) || 0;
}
const loggedIn = extract(s => s.ax.find("id:LoginScreen") == null);
// Screen detection via unique element presence
const loggedIn = extract(s => s.ax.find("desc:login_submit") == null);
const route = extract<string | null>(s => {
if (s.ax.find("id:LoginScreen")) return "login";
if (s.ax.find("id:HomeScreen")) return "home";
if (s.ax.find("id:AddAccountScreen")) return "add-account";
if (s.ax.find("id:LedgerScreen")) return "ledger";
if (s.ax.find("id:AddTransactionScreen")) return "add-transaction";
if (s.ax.find("desc:login_submit")) return "login";
if (s.ax.find("desc:add_account_button")) return "home";
if (s.ax.find("desc:account_name_field")) return "add-account";
if (s.ax.find("descPrefix:active_account:")) return "ledger";
if (s.ax.find("desc:txn_amount")) return "add-transaction";
return null;
});
const accounts = extract(s => s.ax.findAll("descPrefix:account:")
.map(el => parseAccount(el.desc)));
const ledgerRows = extract(s => s.ax.findAll("descPrefix:ledger_row:")
@@ -56,15 +58,15 @@ const activeAccountId = extract(s =>
const focusedInput = extract(s =>
s.ax.find("descPrefix:focused_input:")?.desc?.split(":")[1] ?? null);
const loginEmailField = extract(s => s.ax.find("id:LoginScreen > desc:login_email"));
const loginPasswordField = extract(s => s.ax.find("id:LoginScreen > desc:login_password"));
const loginSubmit = extract(s => s.ax.find("id:LoginScreen > desc:login_submit"));
const addAccountButton = extract(s => s.ax.find("id:HomeScreen > desc:add_account_button"));
const accountNameField = extract(s => s.ax.find("id:AddAccountScreen > desc:account_name_field"));
const addAccountSubmit = extract(s => s.ax.find("id:AddAccountScreen > desc:add_account_submit"));
const addTxnButton = extract(s => s.ax.find("id:LedgerScreen > desc:add_txn_button"));
const txnAmountField = extract(s => s.ax.find("id:AddTransactionScreen > desc:txn_amount"));
const txnSubmit = extract(s => s.ax.find("id:AddTransactionScreen > desc:txn_submit"));
const loginEmailField = extract(s => s.ax.find("desc:login_email"));
const loginPasswordField = extract(s => s.ax.find("desc:login_password"));
const loginSubmit = extract(s => s.ax.find("desc:login_submit"));
const addAccountButton = extract(s => s.ax.find("desc:add_account_button"));
const accountNameField = extract(s => s.ax.find("desc:account_name_field"));
const addAccountSubmit = extract(s => s.ax.find("desc:add_account_submit"));
const addTxnButton = extract(s => s.ax.find("desc:add_txn_button"));
const txnAmountField = extract(s => s.ax.find("desc:txn_amount"));
const txnSubmit = extract(s => s.ax.find("desc:txn_submit"));
const accountCards = extract(s => s.ax.findAll("descPrefix:account:"));
const backButton = extract(s => s.ax.find("desc:Back"));
-98
View File
@@ -1,98 +0,0 @@
import com.vanniktech.maven.publish.AndroidSingleVariantLibrary
import com.vanniktech.maven.publish.JavadocJar
import com.vanniktech.maven.publish.SourcesJar
plugins {
id("com.android.library") version "8.13.0"
kotlin("android") version "2.1.21"
id("com.vanniktech.maven.publish") version "0.36.0"
id("org.jetbrains.dokka") version "2.2.0"
id("org.jetbrains.dokka-javadoc") version "2.2.0"
}
version = findProperty("sanderling.version") as String? ?: "0.0.0-dev"
group = "io.github.priyanshujain.sanderling"
android {
namespace = "dev.sanderling.sdk"
compileSdk = 35
defaultConfig {
minSdk = 24
consumerProguardFiles("consumer-rules.pro")
}
compileOptions {
sourceCompatibility = JavaVersion.VERSION_17
targetCompatibility = JavaVersion.VERSION_17
}
kotlinOptions {
jvmTarget = "17"
}
testOptions {
unitTests.isReturnDefaultValues = true
}
}
mavenPublishing {
publishToMavenCentral(automaticRelease = true)
// Sign only when a release-signing key is provided (env or Gradle
// property). Unsigned runs are useful for `publishToMavenLocal` dry-runs;
// CI always has the key set so the actual Central push is always signed.
if (findProperty("signingInMemoryKey") != null) {
signAllPublications()
}
configure(
AndroidSingleVariantLibrary(
javadocJar = JavadocJar.Dokka("dokkaGeneratePublicationJavadoc"),
sourcesJar = SourcesJar.Sources(),
variant = "release",
),
)
coordinates(
groupId = "io.github.priyanshujain.sanderling",
artifactId = "sdk-android",
version = version.toString(),
)
pom {
name.set("sanderling sdk-android")
description.set(
"Android runtime SDK for sanderling, a property-based UI fuzzer for mobile apps. " +
"Exposes a content-provider accessibility bridge consumed by the sanderling CLI at test time.",
)
url.set("https://github.com/priyanshujain/sanderling")
licenses {
license {
name.set("Apache License, Version 2.0")
url.set("https://www.apache.org/licenses/LICENSE-2.0.txt")
distribution.set("repo")
}
}
developers {
developer {
id.set("priyanshujain")
name.set("Priyanshu Jain")
url.set("https://github.com/priyanshujain")
}
}
scm {
url.set("https://github.com/priyanshujain/sanderling")
connection.set("scm:git:git://github.com/priyanshujain/sanderling.git")
developerConnection.set("scm:git:ssh://[email protected]/priyanshujain/sanderling.git")
}
}
}
dependencies {
testImplementation("junit:junit:4.13.2")
testImplementation("org.json:json:20240303")
}
-2
View File
@@ -1,2 +0,0 @@
# ProGuard rules shipped to consumers of dev.sanderling:sdk-android.
# None needed for v0.1. Socket and semaphore APIs are all reflection-free.
-2
View File
@@ -1,2 +0,0 @@
<?xml version="1.0" encoding="utf-8"?>
<manifest xmlns:android="http://schemas.android.com/apk/res/android" />
@@ -1,15 +0,0 @@
package dev.sanderling.sdk
import android.os.Handler
import android.os.Looper
import android.view.Choreographer
class ChoreographerPoster : FrameCallbackPoster {
private val mainHandler = Handler(Looper.getMainLooper())
override fun postFrameCallback(callback: () -> Unit) {
mainHandler.post {
Choreographer.getInstance().postFrameCallback { callback() }
}
}
}
@@ -1,62 +0,0 @@
package dev.sanderling.sdk
import java.io.PrintWriter
import java.io.StringWriter
internal class ExceptionRecorder(private val capacity: Int = DEFAULT_CAPACITY) {
data class Entry(
val className: String,
val message: String,
val stackTrace: String,
val unixMillis: Long,
)
private val buffer: ArrayDeque<Entry> = ArrayDeque()
private var chainedHandler: Thread.UncaughtExceptionHandler? = null
@Volatile private var installed: Boolean = false
@Synchronized
fun install() {
if (installed) return
chainedHandler = Thread.getDefaultUncaughtExceptionHandler()
Thread.setDefaultUncaughtExceptionHandler { thread, throwable ->
record(throwable)
chainedHandler?.uncaughtException(thread, throwable)
}
installed = true
}
@Synchronized
fun uninstall() {
if (!installed) return
Thread.setDefaultUncaughtExceptionHandler(chainedHandler)
chainedHandler = null
installed = false
}
@Synchronized
fun record(throwable: Throwable, now: Long = System.currentTimeMillis()) {
val stackTrace = StringWriter().also { throwable.printStackTrace(PrintWriter(it)) }.toString()
val entry = Entry(
className = throwable.javaClass.name,
message = throwable.message ?: "",
stackTrace = stackTrace,
unixMillis = now,
)
if (buffer.size >= capacity) {
buffer.removeFirst()
}
buffer.addLast(entry)
}
@Synchronized
fun drain(): List<Entry> {
val snapshot = buffer.toList()
buffer.clear()
return snapshot
}
companion object {
const val DEFAULT_CAPACITY: Int = 50
}
}
@@ -1,26 +0,0 @@
package dev.sanderling.sdk
import android.net.LocalSocket
import android.net.LocalSocketAddress
import java.io.IOException
import java.io.InputStream
import java.io.OutputStream
class LocalAbstractTransport(private val socketName: String) : AgentTransport {
@Throws(IOException::class)
override fun connect(): AgentConnection {
val socket = LocalSocket()
socket.connect(LocalSocketAddress(socketName, LocalSocketAddress.Namespace.ABSTRACT))
return LocalSocketConnection(socket)
}
private class LocalSocketConnection(private val socket: LocalSocket) : AgentConnection {
override val input: InputStream = socket.inputStream
override val output: OutputStream = socket.outputStream
override fun close() {
try { socket.shutdownInput() } catch (_: IOException) {}
try { socket.shutdownOutput() } catch (_: IOException) {}
socket.close()
}
}
}
@@ -1,51 +0,0 @@
package dev.sanderling.sdk
import java.util.concurrent.CountDownLatch
import java.util.concurrent.Semaphore
import java.util.concurrent.TimeUnit
import java.util.concurrent.TimeoutException
import java.util.concurrent.atomic.AtomicReference
fun interface FrameCallbackPoster {
fun postFrameCallback(callback: () -> Unit)
}
class Pauser(
private val poster: FrameCallbackPoster,
private val pauseTimeoutMillis: Long = 5_000L,
) {
@Volatile private var currentGate: Semaphore? = null
/**
* Schedules extractors to run on the frame-callback thread (the SDK's
* "main thread" analogue) and blocks that thread after they complete
* until release() is called or pauseTimeoutMillis elapses.
* Returns the extractor output. Must be called from a worker thread.
*/
@Throws(TimeoutException::class)
fun pauseAndSnapshot(extractors: () -> Map<String, Any?>): Map<String, Any?> {
val gate = Semaphore(0)
val ready = CountDownLatch(1)
val captured = AtomicReference<Result<Map<String, Any?>>>()
poster.postFrameCallback {
captured.set(runCatching { extractors() })
ready.countDown()
try {
gate.tryAcquire(pauseTimeoutMillis, TimeUnit.MILLISECONDS)
} catch (_: InterruptedException) {
Thread.currentThread().interrupt()
}
}
currentGate = gate
if (!ready.await(pauseTimeoutMillis, TimeUnit.MILLISECONDS)) {
throw TimeoutException("extractors did not run within ${pauseTimeoutMillis}ms")
}
return captured.get().getOrThrow()
}
fun release() {
currentGate?.release()
}
}
@@ -1,183 +0,0 @@
package dev.sanderling.sdk
import java.io.DataInputStream
import java.io.DataOutputStream
import java.io.IOException
import java.io.InputStream
import java.io.OutputStream
import org.json.JSONArray
import org.json.JSONObject
enum class MessageType(val wire: String) {
HELLO("HELLO"),
PAUSE("PAUSE"),
RESUME("RESUME"),
STATE("STATE"),
EXTRACT_RESULT("EXTRACT_RESULT"),
GOODBYE("GOODBYE");
companion object {
fun fromWire(wire: String): MessageType =
values().firstOrNull { it.wire == wire }
?: throw IOException("unknown message type: $wire")
}
}
data class Message(
val type: MessageType,
val id: Long = 0L,
val protocolVersion: Int = 0,
val version: String? = null,
val platform: String? = null,
val appPackage: String? = null,
val snapshots: Map<String, Any?>? = null,
val exceptions: List<Map<String, Any?>>? = null,
val extractor: String? = null,
val result: Any? = null,
val error: String? = null,
val reason: String? = null,
) {
companion object {
fun hello(version: String, platform: String, appPackage: String): Message =
Message(
MessageType.HELLO,
protocolVersion = Protocol.PROTOCOL_VERSION,
version = version,
platform = platform,
appPackage = appPackage,
)
fun pause(id: Long): Message = Message(MessageType.PAUSE, id = id)
fun resume(id: Long): Message = Message(MessageType.RESUME, id = id)
fun state(
id: Long,
snapshots: Map<String, Any?>,
exceptions: List<Map<String, Any?>>? = null,
): Message = Message(
MessageType.STATE,
id = id,
snapshots = snapshots,
exceptions = exceptions,
)
fun extractResult(id: Long, extractor: String, result: Any?, error: String? = null): Message =
Message(MessageType.EXTRACT_RESULT, id = id, extractor = extractor, result = result, error = error)
fun goodbye(reason: String): Message = Message(MessageType.GOODBYE, reason = reason)
}
}
object Protocol {
const val MAX_FRAME_SIZE: Int = 16 * 1024 * 1024
// Wire-format version. Must match agent.ProtocolVersion on the Go side.
const val PROTOCOL_VERSION: Int = 1
@Throws(IOException::class)
fun write(output: OutputStream, message: Message) {
val bytes = toJson(message).toString().toByteArray(Charsets.UTF_8)
if (bytes.size > MAX_FRAME_SIZE) {
throw IOException("frame of ${bytes.size} bytes exceeds maximum $MAX_FRAME_SIZE")
}
DataOutputStream(output).apply {
writeInt(bytes.size)
write(bytes)
flush()
}
}
@Throws(IOException::class)
fun read(input: InputStream): Message {
val dataInput = DataInputStream(input)
val length = dataInput.readInt()
if (length < 0 || length > MAX_FRAME_SIZE) {
throw IOException("frame of $length bytes exceeds maximum $MAX_FRAME_SIZE")
}
val bytes = ByteArray(length)
dataInput.readFully(bytes)
return fromJson(JSONObject(String(bytes, Charsets.UTF_8)))
}
private fun toJson(message: Message): JSONObject {
val json = JSONObject()
json.put("type", message.type.wire)
if (message.id != 0L) json.put("id", message.id)
if (message.protocolVersion != 0) json.put("protocol_version", message.protocolVersion)
message.version?.let { json.put("version", it) }
message.platform?.let { json.put("platform", it) }
message.appPackage?.let { json.put("app_package", it) }
message.snapshots?.let { snapshots ->
val snapshotsJson = JSONObject()
for ((key, value) in snapshots) {
snapshotsJson.put(key, wrap(value))
}
json.put("snapshots", snapshotsJson)
}
message.exceptions?.let { exceptions ->
val array = JSONArray()
for (entry in exceptions) {
val entryJson = JSONObject()
for ((key, value) in entry) entryJson.put(key, wrap(value))
array.put(entryJson)
}
json.put("exceptions", array)
}
message.extractor?.let { json.put("extractor", it) }
message.result?.let { json.put("result", wrap(it)) }
message.error?.let { json.put("error", it) }
message.reason?.let { json.put("reason", it) }
return json
}
private fun fromJson(json: JSONObject): Message {
val typeString = json.optString("type", "")
if (typeString.isEmpty()) throw IOException("missing type")
return Message(
type = MessageType.fromWire(typeString),
id = json.optLong("id", 0L),
protocolVersion = json.optInt("protocol_version", 0),
version = json.optStringOrNull("version"),
platform = json.optStringOrNull("platform"),
appPackage = json.optStringOrNull("app_package"),
snapshots = json.optJSONObject("snapshots")?.let { snapshotsJson ->
snapshotsJson.keys().asSequence().associateWith { unwrap(snapshotsJson.get(it)) }
},
exceptions = json.optJSONArray("exceptions")?.let { array ->
buildList {
for (index in 0 until array.length()) {
val item = array.optJSONObject(index) ?: continue
add(item.keys().asSequence().associateWith { unwrap(item.get(it)) })
}
}
},
extractor = json.optStringOrNull("extractor"),
result = if (json.has("result") && !json.isNull("result")) unwrap(json.get("result")) else null,
error = json.optStringOrNull("error"),
reason = json.optStringOrNull("reason"),
)
}
private fun wrap(value: Any?): Any = when (value) {
null -> JSONObject.NULL
is Number, is Boolean, is String -> value
is Map<*, *> -> JSONObject().also { json ->
for ((key, nested) in value) json.put(key.toString(), wrap(nested))
}
is List<*> -> JSONArray().also { array ->
for (item in value) array.put(wrap(item))
}
else -> value.toString()
}
private fun unwrap(value: Any?): Any? = when (value) {
JSONObject.NULL, null -> null
is JSONObject -> value.keys().asSequence().associateWith { unwrap(value.get(it)) }
is JSONArray -> buildList { for (index in 0 until value.length()) add(unwrap(value.get(index))) }
else -> value
}
private fun JSONObject.optStringOrNull(key: String): String? =
if (has(key) && !isNull(key)) getString(key) else null
}
@@ -1,72 +0,0 @@
package dev.sanderling.sdk
import android.app.Application
import android.util.Log
import kotlin.properties.ReadOnlyProperty
import kotlin.reflect.KProperty
internal fun String.camelToSnakeCase(): String = buildString {
for ((i, c) in this@camelToSnakeCase.withIndex()) {
if (c.isUpperCase() && i > 0) append('_')
append(c.lowercaseChar())
}
}
data class Configuration(
val socketName: String = "sanderling-agent",
val pauseTimeoutMillis: Long = 5_000L,
)
object Sanderling {
const val VERSION: String = "0.0.1"
private const val LOG_TAG = "Sanderling"
@Volatile private var runtime: SanderlingRuntime? = null
@JvmOverloads
@Synchronized
fun start(application: Application, configuration: Configuration = Configuration()) {
if (runtime != null) return
val newRuntime = SanderlingRuntime(
transport = LocalAbstractTransport(configuration.socketName),
pauser = Pauser(ChoreographerPoster(), configuration.pauseTimeoutMillis),
version = VERSION,
platform = "android",
appPackage = application.packageName,
)
newRuntime.start()
runtime = newRuntime
Log.i(LOG_TAG, "SDK started (package=${application.packageName} socket=${configuration.socketName})")
}
fun extract(name: String, function: () -> Any?) {
val activeRuntime = runtime
?: throw IllegalStateException("Sanderling.start must be called before registering extractors")
activeRuntime.register(name, function)
}
fun <T> snapshot(function: () -> T): SnapshotDelegate<T> = SnapshotDelegate(function)
/**
* Records a caught [Throwable] so it surfaces in the next STATE message's
* exceptions field. Useful for coroutine CoroutineExceptionHandler,
* OkHttp interceptors, or anywhere else the host app catches errors it
* still wants verified against properties like noUncaughtExceptions.
*/
fun reportError(throwable: Throwable) {
runtime?.reportError(throwable)
}
@Synchronized
internal fun stopForTest() {
runtime?.stop()
runtime = null
}
}
class SnapshotDelegate<T>(private val function: () -> T) {
operator fun provideDelegate(thisRef: Any?, prop: KProperty<*>): ReadOnlyProperty<Any?, T> {
Sanderling.extract(prop.name.camelToSnakeCase(), function as () -> Any?)
return ReadOnlyProperty { _, _ -> function() }
}
}
@@ -1,97 +0,0 @@
package dev.sanderling.sdk
import android.util.Log
internal class SanderlingRuntime(
transport: AgentTransport,
private val pauser: Pauser,
private val version: String,
private val platform: String,
private val appPackage: String,
private val exceptionRecorder: ExceptionRecorder = ExceptionRecorder(),
) {
private val extractors = LinkedHashMap<String, () -> Any?>()
@Volatile private var sender: SocketClient.MessageSender? = null
private val socketClient = SocketClient(transport, AgentHandler())
fun start() {
exceptionRecorder.install()
socketClient.start()
}
fun stop() {
socketClient.stop()
exceptionRecorder.uninstall()
}
fun register(name: String, extractor: () -> Any?) {
synchronized(extractors) { extractors[name] = extractor }
}
fun reportError(throwable: Throwable) {
exceptionRecorder.record(throwable)
}
internal fun snapshot(): Map<String, Any?> {
val drained = synchronized(extractors) { LinkedHashMap(extractors) }
val result = LinkedHashMap<String, Any?>(drained.size)
for ((name, extractor) in drained) {
result[name] = runCatching { extractor() }
.onFailure { cause -> Log.w(LOG_TAG, "extractor $name threw: $cause") }
.getOrNull()
}
return result
}
private inner class AgentHandler : SocketClient.Handler {
override fun onConnected(sender: SocketClient.MessageSender) {
this@SanderlingRuntime.sender = sender
try {
sender.send(Message.hello(version, platform, appPackage))
} catch (cause: Exception) {
Log.w(LOG_TAG, "failed to send HELLO: $cause")
}
}
override fun onMessage(message: Message) {
when (message.type) {
MessageType.PAUSE -> handlePause(message.id)
MessageType.RESUME -> pauser.release()
MessageType.GOODBYE -> socketClient.stop()
else -> Log.w(LOG_TAG, "unexpected message type ${message.type} from host")
}
}
override fun onDisconnected(cause: Throwable?) {
sender = null
pauser.release()
}
private fun handlePause(id: Long) {
val snapshots = try {
pauser.pauseAndSnapshot { snapshot() }
} catch (cause: Exception) {
Log.w(LOG_TAG, "snapshot failed: $cause")
emptyMap()
}
val exceptions = exceptionRecorder.drain().map { entry ->
mapOf(
"class" to entry.className,
"message" to entry.message,
"stack_trace" to entry.stackTrace,
"unix_millis" to entry.unixMillis,
)
}.takeIf { it.isNotEmpty() }
val activeSender = sender ?: return
try {
activeSender.send(Message.state(id, snapshots, exceptions))
} catch (cause: Exception) {
Log.w(LOG_TAG, "failed to send STATE: $cause")
}
}
}
companion object {
private const val LOG_TAG = "Sanderling"
}
}
@@ -1,108 +0,0 @@
package dev.sanderling.sdk
import java.io.IOException
import java.io.InputStream
import java.io.OutputStream
import java.util.concurrent.atomic.AtomicBoolean
interface AgentTransport {
@Throws(IOException::class)
fun connect(): AgentConnection
}
interface AgentConnection {
val input: InputStream
val output: OutputStream
fun close()
}
data class Backoff(
val initialDelayMillis: Long = 500L,
val maxDelayMillis: Long = 10_000L,
val multiplier: Double = 2.0,
) {
fun next(previousDelayMillis: Long): Long =
if (previousDelayMillis <= 0L) initialDelayMillis
else minOf((previousDelayMillis * multiplier).toLong(), maxDelayMillis)
}
class SocketClient(
private val transport: AgentTransport,
private val handler: Handler,
private val backoff: Backoff = Backoff(),
private val threadFactory: (Runnable) -> Thread = { runnable -> Thread(runnable, "sanderling-agent-reader") },
private val sleeper: (Long) -> Unit = { millis -> if (millis > 0L) Thread.sleep(millis) },
) {
interface Handler {
fun onConnected(sender: MessageSender)
fun onMessage(message: Message)
fun onDisconnected(cause: Throwable?)
}
fun interface MessageSender {
@Throws(IOException::class)
fun send(message: Message)
}
private val running = AtomicBoolean(false)
@Volatile private var workerThread: Thread? = null
@Volatile private var connection: AgentConnection? = null
fun start() {
if (!running.compareAndSet(false, true)) return
val thread = threadFactory { runLoop() }
thread.isDaemon = true
workerThread = thread
thread.start()
}
fun stop() {
if (!running.compareAndSet(true, false)) return
try { connection?.close() } catch (_: IOException) {}
workerThread?.interrupt()
try { workerThread?.join(1_000L) } catch (_: InterruptedException) { Thread.currentThread().interrupt() }
}
private fun runLoop() {
var delayMillis = 0L
while (running.get()) {
val connection = try {
transport.connect()
} catch (e: IOException) {
handler.onDisconnected(e)
if (!running.get()) return
delayMillis = backoff.next(delayMillis)
try { sleeper(delayMillis) } catch (_: InterruptedException) { return }
continue
}
this.connection = connection
delayMillis = 0L
serve(connection)
if (!running.get()) return
delayMillis = backoff.next(delayMillis)
try { sleeper(delayMillis) } catch (_: InterruptedException) { return }
}
}
private fun serve(connection: AgentConnection) {
val sender = MessageSender { message ->
synchronized(connection.output) {
Protocol.write(connection.output, message)
}
}
handler.onConnected(sender)
var disconnectCause: Throwable? = null
try {
while (running.get()) {
val message = Protocol.read(connection.input)
handler.onMessage(message)
}
} catch (e: IOException) {
disconnectCause = e
} finally {
try { connection.close() } catch (_: IOException) {}
this.connection = null
handler.onDisconnected(disconnectCause)
}
}
}
@@ -1,64 +0,0 @@
package dev.sanderling.sdk
import org.junit.Assert.assertEquals
import org.junit.Assert.assertTrue
import org.junit.Test
class ExceptionRecorderTest {
@Test fun recordsClassMessageAndStackTrace() {
val recorder = ExceptionRecorder()
recorder.record(RuntimeException("boom"))
val drained = recorder.drain()
assertEquals(1, drained.size)
val entry = drained[0]
assertEquals("java.lang.RuntimeException", entry.className)
assertEquals("boom", entry.message)
assertTrue(
"stackTrace should include the class name, got: ${entry.stackTrace}",
entry.stackTrace.contains("RuntimeException"),
)
}
@Test fun drainClearsBuffer() {
val recorder = ExceptionRecorder()
recorder.record(RuntimeException("first"))
recorder.record(RuntimeException("second"))
assertEquals(2, recorder.drain().size)
assertEquals(0, recorder.drain().size)
}
@Test fun dropsOldestWhenOverCapacity() {
val recorder = ExceptionRecorder(capacity = 2)
recorder.record(RuntimeException("a"))
recorder.record(RuntimeException("b"))
recorder.record(RuntimeException("c"))
val drained = recorder.drain()
assertEquals(2, drained.size)
assertEquals("b", drained[0].message)
assertEquals("c", drained[1].message)
}
@Test fun installChainsExistingHandler() {
val recorder = ExceptionRecorder()
val original = Thread.getDefaultUncaughtExceptionHandler()
var chainedInvoked = false
Thread.setDefaultUncaughtExceptionHandler { _, _ -> chainedInvoked = true }
try {
recorder.install()
// Simulate an uncaught exception by invoking the installed handler
// directly — we don't need to actually terminate a thread.
Thread.getDefaultUncaughtExceptionHandler()!!.uncaughtException(
Thread.currentThread(),
IllegalStateException("chain me"),
)
assertTrue("chained handler should have fired", chainedInvoked)
assertEquals(1, recorder.drain().size)
} finally {
recorder.uninstall()
Thread.setDefaultUncaughtExceptionHandler(original)
}
}
}
@@ -1,133 +0,0 @@
package dev.sanderling.sdk
import java.util.concurrent.CountDownLatch
import java.util.concurrent.Executors
import java.util.concurrent.TimeUnit
import java.util.concurrent.TimeoutException
import java.util.concurrent.atomic.AtomicBoolean
import java.util.concurrent.atomic.AtomicReference
import org.junit.After
import org.junit.Assert.assertEquals
import org.junit.Assert.assertFalse
import org.junit.Assert.assertNotNull
import org.junit.Assert.assertTrue
import org.junit.Assert.fail
import org.junit.Test
class PauserTest {
// Runs posted callbacks on a dedicated "main thread" executor, modelling
// Choreographer's contract that callbacks fire off-thread from the caller.
class FakeFrameThread : FrameCallbackPoster {
val executor = Executors.newSingleThreadExecutor { runnable -> Thread(runnable, "fake-frame-thread") }
val threadRef = AtomicReference<Thread>()
override fun postFrameCallback(callback: () -> Unit) {
executor.submit {
threadRef.compareAndSet(null, Thread.currentThread())
callback()
}
}
fun shutdown() { executor.shutdownNow() }
}
private lateinit var frameThread: FakeFrameThread
@After fun tearDown() {
if (::frameThread.isInitialized) frameThread.shutdown()
}
@Test fun extractorsRunOnFrameThreadAndSnapshotReturns() {
frameThread = FakeFrameThread()
val pauser = Pauser(frameThread, pauseTimeoutMillis = 2_000L)
val extractorThread = AtomicReference<Thread>()
// Run on a worker thread so we can observe the separation.
val snapshot = runOnWorker {
pauser.pauseAndSnapshot {
extractorThread.set(Thread.currentThread())
mapOf("screen" to "home", "count" to 3)
}.also { snapshot ->
// Immediately release so the frame thread can exit the callback.
pauser.release()
snapshot
}
}
assertEquals("home", snapshot["screen"])
assertEquals(3, snapshot["count"])
assertNotNull("extractor must have run", extractorThread.get())
assertEquals("fake-frame-thread", extractorThread.get().name)
}
@Test fun frameThreadStaysBlockedUntilRelease() {
frameThread = FakeFrameThread()
val pauser = Pauser(frameThread, pauseTimeoutMillis = 2_000L)
val releasedMarker = AtomicBoolean(false)
val latch = CountDownLatch(1)
val worker = Thread {
pauser.pauseAndSnapshot { emptyMap() }
// Now the frame thread is blocked inside the callback. Verify by
// posting another callback and checking it does NOT run until we release.
val secondCallbackRan = CountDownLatch(1)
frameThread.postFrameCallback { secondCallbackRan.countDown() }
assertFalse("second callback should be queued, not run",
secondCallbackRan.await(200, TimeUnit.MILLISECONDS))
pauser.release()
assertTrue("second callback should run after release",
secondCallbackRan.await(2, TimeUnit.SECONDS))
releasedMarker.set(true)
latch.countDown()
}
worker.start()
assertTrue("worker must finish", latch.await(5, TimeUnit.SECONDS))
assertTrue("release must have happened", releasedMarker.get())
}
@Test fun timeoutPropagatesWhenFrameThreadNeverRuns() {
val stuck = FrameCallbackPoster { /* never invokes callback */ }
val pauser = Pauser(stuck, pauseTimeoutMillis = 150L)
try {
pauser.pauseAndSnapshot { emptyMap() }
fail("expected TimeoutException")
} catch (_: TimeoutException) {
// pass
}
}
@Test fun extractorExceptionBubbles() {
frameThread = FakeFrameThread()
val pauser = Pauser(frameThread, pauseTimeoutMillis = 2_000L)
try {
runOnWorker {
pauser.pauseAndSnapshot {
pauser.release() // drop the lock before throwing so main can unwind
throw IllegalStateException("extractor boom")
}
}
fail("expected IllegalStateException")
} catch (e: IllegalStateException) {
assertEquals("extractor boom", e.message)
}
}
@Test fun releaseWithoutActivePauseIsNoOp() {
frameThread = FakeFrameThread()
val pauser = Pauser(frameThread, pauseTimeoutMillis = 2_000L)
pauser.release() // should not throw
}
private fun <T> runOnWorker(block: () -> T): T {
val result = AtomicReference<Result<T>>()
val thread = Thread { result.set(runCatching { block() }) }
thread.start()
thread.join(5_000L)
return result.get().getOrThrow()
}
}
@@ -1,181 +0,0 @@
package dev.sanderling.sdk
import java.io.ByteArrayInputStream
import java.io.ByteArrayOutputStream
import java.io.DataInputStream
import java.io.EOFException
import java.io.IOException
import org.junit.Assert.assertArrayEquals
import org.junit.Assert.assertEquals
import org.junit.Assert.assertNotNull
import org.junit.Assert.assertNull
import org.junit.Assert.assertThrows
import org.junit.Assert.assertTrue
import org.junit.Test
class ProtocolTest {
private fun roundTrip(message: Message): Message {
val output = ByteArrayOutputStream()
Protocol.write(output, message)
return Protocol.read(ByteArrayInputStream(output.toByteArray()))
}
@Test fun roundTripHello() {
val got = roundTrip(Message.hello("0.0.1", "android", "in.okcredit.merchant"))
assertEquals(MessageType.HELLO, got.type)
assertEquals(Protocol.PROTOCOL_VERSION, got.protocolVersion)
assertEquals("0.0.1", got.version)
assertEquals("android", got.platform)
assertEquals("in.okcredit.merchant", got.appPackage)
}
@Test fun roundTripPauseResume() {
val pause = roundTrip(Message.pause(42))
assertEquals(MessageType.PAUSE, pause.type)
assertEquals(42L, pause.id)
val resume = roundTrip(Message.resume(43))
assertEquals(MessageType.RESUME, resume.type)
assertEquals(43L, resume.id)
}
@Test fun roundTripState() {
val snapshots = mapOf<String, Any?>(
"screen" to "customer_ledger",
"ledger.balance" to 1500,
"is_signed_in" to true,
)
val got = roundTrip(Message.state(7, snapshots))
assertEquals(MessageType.STATE, got.type)
assertEquals(7L, got.id)
assertNotNull(got.snapshots)
assertEquals("customer_ledger", got.snapshots!!["screen"])
assertEquals(1500, got.snapshots["ledger.balance"])
assertEquals(true, got.snapshots["is_signed_in"])
}
@Test fun roundTripExtractResult() {
val ok = roundTrip(Message.extractResult(1, "ledger.balance", 2500))
assertEquals("ledger.balance", ok.extractor)
assertEquals(2500, ok.result)
assertNull(ok.error)
val failed = roundTrip(Message.extractResult(2, "ledger.balance", null, "no active customer"))
assertEquals("no active customer", failed.error)
assertNull(failed.result)
}
@Test fun roundTripGoodbye() {
val got = roundTrip(Message.goodbye("app terminated"))
assertEquals(MessageType.GOODBYE, got.type)
assertEquals("app terminated", got.reason)
}
@Test fun frameFormatIsBigEndianLengthPlusJson() {
val output = ByteArrayOutputStream()
Protocol.write(output, Message.pause(99))
val raw = output.toByteArray()
assertTrue("frame must have 4-byte header plus body", raw.size > 4)
val length = DataInputStream(ByteArrayInputStream(raw.copyOfRange(0, 4))).readInt()
assertEquals(raw.size - 4, length)
val payload = String(raw.copyOfRange(4, raw.size), Charsets.UTF_8)
assertTrue("payload should contain PAUSE type, got $payload", payload.contains("\"type\":\"PAUSE\""))
}
@Test fun emptyReaderThrowsEof() {
assertThrows(EOFException::class.java) {
Protocol.read(ByteArrayInputStream(ByteArray(0)))
}
}
@Test fun oversizedFrameRejected() {
val header = ByteArray(4)
val tooBig = Protocol.MAX_FRAME_SIZE + 1
header[0] = (tooBig ushr 24).toByte()
header[1] = (tooBig ushr 16).toByte()
header[2] = (tooBig ushr 8).toByte()
header[3] = tooBig.toByte()
val error = assertThrows(IOException::class.java) {
Protocol.read(ByteArrayInputStream(header))
}
assertTrue("expected size error, got: ${error.message}", error.message!!.contains("exceeds maximum"))
}
@Test fun missingTypeRejected() {
val payload = "{\"id\":1}".toByteArray(Charsets.UTF_8)
val header = ByteArray(4)
header[0] = (payload.size ushr 24).toByte()
header[1] = (payload.size ushr 16).toByte()
header[2] = (payload.size ushr 8).toByte()
header[3] = payload.size.toByte()
val frame = header + payload
val error = assertThrows(IOException::class.java) {
Protocol.read(ByteArrayInputStream(frame))
}
assertTrue("expected missing-type error, got: ${error.message}", error.message!!.contains("missing type"))
}
@Test fun roundTripStateWithExceptions() {
val exceptions = listOf(
mapOf<String, Any?>(
"class" to "java.lang.RuntimeException",
"message" to "boom",
"stack_trace" to "at Foo.bar(Foo.kt:42)",
"unix_millis" to 1_700_000_000_000L,
),
)
val got = roundTrip(Message.state(3, mapOf("screen" to "home"), exceptions))
assertNotNull(got.exceptions)
assertEquals(1, got.exceptions!!.size)
assertEquals("java.lang.RuntimeException", got.exceptions[0]["class"])
assertEquals("boom", got.exceptions[0]["message"])
}
@Test fun streamsMultipleFrames() {
val messages = listOf(
Message.hello("v", "android", "com.x"),
Message.pause(1),
Message.state(1, mapOf("x" to 42)),
Message.resume(1),
Message.goodbye("done"),
)
val output = ByteArrayOutputStream()
for (message in messages) Protocol.write(output, message)
val input = ByteArrayInputStream(output.toByteArray())
for (want in messages) {
val got = Protocol.read(input)
assertEquals(want.type, got.type)
}
}
@Test fun sharedWireFormatMatchesGoEncoder() {
// Fixture encoded by the Go side (see internal/agent/protocol.go).
// Ensures both encoders agree on field names and ordering conventions.
val output = ByteArrayOutputStream()
Protocol.write(output, Message.hello("0.0.1", "android", "com.x"))
val payload = String(output.toByteArray().copyOfRange(4, output.size()), Charsets.UTF_8)
assertTrue(payload.contains("\"type\":\"HELLO\""))
assertTrue(payload.contains("\"protocol_version\":1"))
assertTrue(payload.contains("\"version\":\"0.0.1\""))
assertTrue(payload.contains("\"platform\":\"android\""))
assertTrue(payload.contains("\"app_package\":\"com.x\""))
}
@Test fun bytesAreConsumedInOrder() {
// Regression guard: a second read shouldn't see stale bytes.
val output = ByteArrayOutputStream()
Protocol.write(output, Message.pause(1))
Protocol.write(output, Message.pause(2))
val bytes = output.toByteArray()
val input = ByteArrayInputStream(bytes)
assertEquals(1L, Protocol.read(input).id)
assertEquals(2L, Protocol.read(input).id)
// ByteArrayInputStream should be drained.
val leftover = ByteArray(bytes.size)
val remaining = input.read(leftover)
assertEquals(-1, remaining)
assertArrayEquals(ByteArray(bytes.size), leftover)
}
}
@@ -1,146 +0,0 @@
package dev.sanderling.sdk
import java.util.concurrent.CopyOnWriteArrayList
import java.util.concurrent.CountDownLatch
import java.util.concurrent.Executors
import java.util.concurrent.TimeUnit
import java.util.concurrent.atomic.AtomicInteger
import org.junit.After
import org.junit.Assert.assertArrayEquals
import org.junit.Assert.assertEquals
import org.junit.Assert.assertNotNull
import org.junit.Assert.assertTrue
import org.junit.Test
class SanderlingRuntimeTest {
private lateinit var transport: SocketClientTest.FakeTransport
private lateinit var frameThread: PauserTest.FakeFrameThread
private lateinit var runtime: SanderlingRuntime
private fun newRuntime(): SanderlingRuntime {
transport = SocketClientTest.FakeTransport()
frameThread = PauserTest.FakeFrameThread()
val pauser = Pauser(frameThread, pauseTimeoutMillis = 2_000L)
return SanderlingRuntime(
transport = transport,
pauser = pauser,
version = "0.0.1",
platform = "android",
appPackage = "com.example.sanderling_test",
).also { runtime = it }
}
@After fun tearDown() {
if (::runtime.isInitialized) runtime.stop()
if (::frameThread.isInitialized) frameThread.shutdown()
}
@Test fun startSendsHelloWithSdkMetadata() {
newRuntime().start()
val server = transport.nextServerEndpoint()
val hello = Protocol.read(server.input)
assertEquals(MessageType.HELLO, hello.type)
assertEquals("0.0.1", hello.version)
assertEquals("android", hello.platform)
assertEquals("com.example.sanderling_test", hello.appPackage)
}
@Test fun pauseTriggersExtractorsAndReturnsState() {
newRuntime().start()
val server = transport.nextServerEndpoint()
Protocol.read(server.input) // drain HELLO
runtime.register("screen") { "customer_ledger" }
runtime.register("ledger.balance") { 1500 }
Protocol.write(server.output, Message.pause(7))
val state = Protocol.read(server.input)
assertEquals(MessageType.STATE, state.type)
assertEquals(7L, state.id)
val snapshots = state.snapshots ?: error("state.snapshots must not be null")
assertEquals("customer_ledger", snapshots["screen"])
assertEquals(1500, snapshots["ledger.balance"])
Protocol.write(server.output, Message.resume(7))
// After resume, the frame thread can accept subsequent callbacks.
Protocol.write(server.output, Message.pause(8))
val nextState = Protocol.read(server.input)
assertEquals(MessageType.STATE, nextState.type)
assertEquals(8L, nextState.id)
}
@Test fun extractorInvocationOrderMatchesRegistration() {
val runtime = newRuntime()
val order = CopyOnWriteArrayList<String>()
runtime.register("first") { order += "first"; 1 }
runtime.register("second") { order += "second"; 2 }
runtime.register("third") { order += "third"; 3 }
val snapshot = runtime.snapshot()
assertEquals(listOf("first", "second", "third"), order)
assertEquals(listOf("first", "second", "third"), snapshot.keys.toList())
assertArrayEquals(arrayOf(1, 2, 3), snapshot.values.toList().toTypedArray())
}
@Test fun extractorThrowIsIsolatedAndReportsNull() {
val runtime = newRuntime()
runtime.register("ok") { "value" }
runtime.register("boom") { throw IllegalStateException("oops") }
runtime.register("later") { 42 }
val snapshot = runtime.snapshot()
assertEquals("value", snapshot["ok"])
assertEquals(null, snapshot["boom"])
assertEquals(42, snapshot["later"])
}
@Test fun concurrentExtractorRegistrationIsSafe() {
val runtime = newRuntime()
val registrations = 500
val pool = Executors.newFixedThreadPool(8)
val latch = CountDownLatch(registrations)
val index = AtomicInteger(0)
repeat(registrations) {
pool.submit {
val id = index.getAndIncrement()
runtime.register("ext-$id") { id }
latch.countDown()
}
}
assertTrue(latch.await(5, TimeUnit.SECONDS))
pool.shutdown()
val snapshot = runtime.snapshot()
assertEquals(registrations, snapshot.size)
}
@Test fun helloIsSentOnReconnect() {
val runtime = newRuntime()
val shortBackoff = Backoff(initialDelayMillis = 10L, maxDelayMillis = 10L, multiplier = 1.0)
// Swap in a client with short backoff by re-creating the runtime's client indirectly.
// We'll re-use the existing runtime instead and just close the first connection.
runtime.start()
val firstServer = transport.nextServerEndpoint()
Protocol.read(firstServer.input) // drain first HELLO
firstServer.close()
val secondServer = transport.nextServerEndpoint(timeoutMillis = 3_000L)
val secondHello = Protocol.read(secondServer.input)
assertEquals(MessageType.HELLO, secondHello.type)
}
@Test fun registerBeforeStartQueuesCorrectlyOnceStarted() {
val runtime = newRuntime()
runtime.register("early") { 1 }
runtime.register("middle") { 2 }
runtime.start()
runtime.register("late") { 3 }
val snapshot = runtime.snapshot()
assertEquals(3, snapshot.size)
assertEquals(1, snapshot["early"])
assertEquals(2, snapshot["middle"])
assertEquals(3, snapshot["late"])
}
}
@@ -1,95 +0,0 @@
package dev.sanderling.sdk
import org.junit.After
import org.junit.Assert.assertEquals
import org.junit.Test
class SnapshotDelegateTest {
@After fun tearDown() = Sanderling.stopForTest()
@Test fun camelToSnakeCaseConvertsKnownNames() {
assertEquals("logged_in", "loggedIn".camelToSnakeCase())
assertEquals("total_balance", "totalBalance".camelToSnakeCase())
assertEquals("txn_form_type", "txnFormType".camelToSnakeCase())
assertEquals("active_account_id", "activeAccountId".camelToSnakeCase())
assertEquals("add_account_error", "addAccountError".camelToSnakeCase())
assertEquals("ledger_balance", "ledgerBalance".camelToSnakeCase())
assertEquals("ledger_rows", "ledgerRows".camelToSnakeCase())
assertEquals("auth_status", "authStatus".camelToSnakeCase())
assertEquals("login_error", "loginError".camelToSnakeCase())
assertEquals("txn_error", "txnError".camelToSnakeCase())
assertEquals("account_count", "accountCount".camelToSnakeCase())
assertEquals("focused_input", "focusedInput".camelToSnakeCase())
assertEquals("txn_form_account_id", "txnFormAccountId".camelToSnakeCase())
}
@Test fun camelToSnakeCaseLeavesAlreadyLowercase() {
assertEquals("screen", "screen".camelToSnakeCase())
assertEquals("accounts", "accounts".camelToSnakeCase())
}
@Test fun snapshotDelegateRegistersWithDerivedKey() {
val transport = SocketClientTest.FakeTransport()
val frameThread = PauserTest.FakeFrameThread()
val runtime = SanderlingRuntime(
transport = transport,
pauser = Pauser(frameThread, pauseTimeoutMillis = 2_000L),
version = "0.0.1",
platform = "android",
appPackage = "com.example.test",
)
runtime.start()
// Inject runtime into Sanderling via reflection so snapshot() can register
val runtimeField = Sanderling::class.java.getDeclaredField("runtime")
runtimeField.isAccessible = true
runtimeField.set(Sanderling, runtime)
var callCount = 0
val obj = object {
val loggedIn by Sanderling.snapshot { callCount++; true }
}
val snapshot = runtime.snapshot()
assertEquals(true, snapshot["logged_in"])
assertEquals(1, callCount)
// getValue delegates back to lambda
callCount = 0
val value = obj.loggedIn
assertEquals(true, value)
assertEquals(1, callCount)
runtime.stop()
frameThread.shutdown()
}
@Test fun snapshotDelegateGetValueReturnsFreshResult() {
val transport = SocketClientTest.FakeTransport()
val frameThread = PauserTest.FakeFrameThread()
val runtime = SanderlingRuntime(
transport = transport,
pauser = Pauser(frameThread, pauseTimeoutMillis = 2_000L),
version = "0.0.1",
platform = "android",
appPackage = "com.example.test",
)
runtime.start()
val runtimeField = Sanderling::class.java.getDeclaredField("runtime")
runtimeField.isAccessible = true
runtimeField.set(Sanderling, runtime)
var counter = 0
val obj = object {
val accountCount by Sanderling.snapshot { counter }
}
counter = 5
assertEquals(5, obj.accountCount)
counter = 10
assertEquals(10, obj.accountCount)
runtime.stop()
frameThread.shutdown()
}
}
@@ -1,190 +0,0 @@
package dev.sanderling.sdk
import java.io.IOException
import java.io.PipedInputStream
import java.io.PipedOutputStream
import java.util.concurrent.CopyOnWriteArrayList
import java.util.concurrent.CountDownLatch
import java.util.concurrent.LinkedBlockingQueue
import java.util.concurrent.TimeUnit
import java.util.concurrent.atomic.AtomicInteger
import org.junit.After
import org.junit.Assert.assertEquals
import org.junit.Assert.assertNotNull
import org.junit.Assert.assertSame
import org.junit.Assert.assertTrue
import org.junit.Test
class SocketClientTest {
/** In-memory transport that pairs each connect() call with a matching
* server endpoint a test can read/write on. */
class FakeTransport : AgentTransport {
data class Endpoints(val client: AgentConnection, val server: AgentConnection)
private val pending = LinkedBlockingQueue<Endpoints>()
val connectCount = AtomicInteger(0)
@Volatile var failNextConnect: Throwable? = null
override fun connect(): AgentConnection {
connectCount.incrementAndGet()
failNextConnect?.let { cause ->
failNextConnect = null
throw IOException("connect failed", cause)
}
val toClient = PipedOutputStream()
val fromServer = PipedInputStream(toClient, 64 * 1024)
val toServer = PipedOutputStream()
val fromClient = PipedInputStream(toServer, 64 * 1024)
val client = pipedConnection(fromServer, toServer)
val server = pipedConnection(fromClient, toClient)
pending.offer(Endpoints(client, server))
return client
}
fun nextServerEndpoint(timeoutMillis: Long = 2_000L): AgentConnection {
val endpoints = pending.poll(timeoutMillis, TimeUnit.MILLISECONDS)
?: error("no server endpoint emerged within $timeoutMillis ms")
return endpoints.server
}
private fun pipedConnection(input: PipedInputStream, output: PipedOutputStream): AgentConnection {
return object : AgentConnection {
override val input = input
override val output = output
override fun close() {
try { input.close() } catch (_: IOException) {}
try { output.close() } catch (_: IOException) {}
}
}
}
}
class RecordingHandler : SocketClient.Handler {
val connected = CountDownLatch(1)
val disconnected = LinkedBlockingQueue<Throwable?>()
val messages = CopyOnWriteArrayList<Message>()
@Volatile var sender: SocketClient.MessageSender? = null
override fun onConnected(sender: SocketClient.MessageSender) {
this.sender = sender
connected.countDown()
}
override fun onMessage(message: Message) {
messages += message
}
override fun onDisconnected(cause: Throwable?) {
disconnected.offer(cause ?: SentinelDisconnect)
}
fun waitForMessages(expected: Int, timeoutMillis: Long = 2_000L) {
val deadline = System.currentTimeMillis() + timeoutMillis
while (messages.size < expected && System.currentTimeMillis() < deadline) Thread.sleep(5L)
assertEquals("expected $expected messages, got ${messages.size}", expected, messages.size)
}
}
private object SentinelDisconnect : Throwable()
private lateinit var client: SocketClient
@After fun tearDown() {
if (::client.isInitialized) client.stop()
}
@Test fun connectsAndDeliversIncomingMessages() {
val transport = FakeTransport()
val handler = RecordingHandler()
client = SocketClient(transport, handler).also { it.start() }
val server = transport.nextServerEndpoint()
Protocol.write(server.output, Message.pause(1))
Protocol.write(server.output, Message.resume(1))
handler.waitForMessages(2)
assertEquals(MessageType.PAUSE, handler.messages[0].type)
assertEquals(1L, handler.messages[0].id)
assertEquals(MessageType.RESUME, handler.messages[1].type)
}
@Test fun sendGoesThroughOutputStream() {
val transport = FakeTransport()
val handler = RecordingHandler()
client = SocketClient(transport, handler).also { it.start() }
assertTrue("connected latch must fire", handler.connected.await(2, TimeUnit.SECONDS))
val sender = handler.sender ?: error("sender not set after onConnected")
sender.send(Message.hello("0.0.1", "android", "com.x"))
val server = transport.nextServerEndpoint()
val received = Protocol.read(server.input)
assertEquals(MessageType.HELLO, received.type)
assertEquals("com.x", received.appPackage)
}
@Test fun reconnectsWithBackoffAfterConnectFailure() {
val transport = FakeTransport()
transport.failNextConnect = RuntimeException("no socket yet")
val handler = RecordingHandler()
val observedSleeps = CopyOnWriteArrayList<Long>()
val shortBackoff = Backoff(initialDelayMillis = 25L, maxDelayMillis = 25L, multiplier = 1.0)
client = SocketClient(
transport,
handler,
backoff = shortBackoff,
sleeper = { millis -> observedSleeps += millis },
).also { it.start() }
assertTrue("connected latch must fire after retry", handler.connected.await(2, TimeUnit.SECONDS))
assertTrue("should have retried at least once", transport.connectCount.get() >= 2)
assertTrue("sleeper should have been called with backoff delay, got $observedSleeps",
observedSleeps.any { it > 0L })
}
@Test fun reconnectsAfterServerClosesConnection() {
val transport = FakeTransport()
val handler = RecordingHandler()
val shortBackoff = Backoff(initialDelayMillis = 10L, maxDelayMillis = 10L, multiplier = 1.0)
client = SocketClient(transport, handler, backoff = shortBackoff).also { it.start() }
val firstServer = transport.nextServerEndpoint()
assertTrue(handler.connected.await(2, TimeUnit.SECONDS))
firstServer.close()
// Handler.onDisconnected should fire.
val cause = handler.disconnected.poll(2, TimeUnit.SECONDS)
assertNotNull("expected disconnect notification", cause)
// A second connect() should happen — fetch the new server endpoint to prove it.
val secondServer = transport.nextServerEndpoint(timeoutMillis = 2_000L)
assertSame("second endpoint must exist", secondServer, secondServer)
assertTrue("connect count should be >= 2, got ${transport.connectCount.get()}",
transport.connectCount.get() >= 2)
}
@Test fun stopClosesConnection() {
val transport = FakeTransport()
val handler = RecordingHandler()
client = SocketClient(transport, handler).also { it.start() }
transport.nextServerEndpoint()
assertTrue(handler.connected.await(2, TimeUnit.SECONDS))
client.stop()
val cause = handler.disconnected.poll(2, TimeUnit.SECONDS)
assertNotNull("stop() should trigger onDisconnected", cause)
}
@Test fun backoffGrowsExponentially() {
val backoff = Backoff(initialDelayMillis = 100L, maxDelayMillis = 800L, multiplier = 2.0)
assertEquals(100L, backoff.next(0L))
assertEquals(200L, backoff.next(100L))
assertEquals(400L, backoff.next(200L))
assertEquals(800L, backoff.next(400L))
assertEquals(800L, backoff.next(800L))
}
}