From 2cc7617590f5451fc25fea0a65185df47133c8fe Mon Sep 17 00:00:00 2001 From: PJ Date: Fri, 17 Apr 2026 22:48:10 +0700 Subject: [PATCH] feat(sdk-android): Kotlin wire protocol matching the Go encoder Uses org.json (already on Android; org.json:json test dep for JVM unit tests). Field names and framing mirror internal/agent/protocol.go so Go-encoded frames decode on the SDK without a shared schema file. --- sdk/android/build.gradle.kts | 1 + .../src/main/kotlin/dev/uatu/sdk/Protocol.kt | 145 ++++++++++++++++ .../test/kotlin/dev/uatu/sdk/ProtocolTest.kt | 163 ++++++++++++++++++ 3 files changed, 309 insertions(+) create mode 100644 sdk/android/src/main/kotlin/dev/uatu/sdk/Protocol.kt create mode 100644 sdk/android/src/test/kotlin/dev/uatu/sdk/ProtocolTest.kt diff --git a/sdk/android/build.gradle.kts b/sdk/android/build.gradle.kts index 7d340ad..e13ac72 100644 --- a/sdk/android/build.gradle.kts +++ b/sdk/android/build.gradle.kts @@ -24,4 +24,5 @@ android { dependencies { testImplementation("junit:junit:4.13.2") + testImplementation("org.json:json:20240303") } diff --git a/sdk/android/src/main/kotlin/dev/uatu/sdk/Protocol.kt b/sdk/android/src/main/kotlin/dev/uatu/sdk/Protocol.kt new file mode 100644 index 0000000..21ffe48 --- /dev/null +++ b/sdk/android/src/main/kotlin/dev/uatu/sdk/Protocol.kt @@ -0,0 +1,145 @@ +package dev.uatu.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 version: String? = null, + val platform: String? = null, + val appPackage: String? = null, + val snapshots: Map? = 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, 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): Message = + Message(MessageType.STATE, id = id, snapshots = snapshots) + + 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 + + @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) + 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.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), + 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)) } + }, + 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 +} diff --git a/sdk/android/src/test/kotlin/dev/uatu/sdk/ProtocolTest.kt b/sdk/android/src/test/kotlin/dev/uatu/sdk/ProtocolTest.kt new file mode 100644 index 0000000..78ec50d --- /dev/null +++ b/sdk/android/src/test/kotlin/dev/uatu/sdk/ProtocolTest.kt @@ -0,0 +1,163 @@ +package dev.uatu.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("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( + "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 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("\"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) + } +}