mirror of
https://github.com/priyanshujain/sanderling.git
synced 2026-10-02 19:17:10 +00:00
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:
1 parent
270d4bef7e
commit
3feb7b2d7f
18 files changed
+17
-1540
No files matched your search
@@ -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"));
|
||||
|
||||
|
||||
@@ -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")
|
||||
}
|
||||
@@ -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.
|
||||
@@ -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))
|
||||
}
|
||||
}
|
||||
Reference in new issue
Block a user