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.
This commit is contained in:
pj committed 2026-04-17 22:48:10 +07:00
1 parent ca92bb1492
commit 2cc7617590
3 files changed
+309

No files matched your search

+1
View File
@@ -24,4 +24,5 @@ android {
dependencies { dependencies {
testImplementation("junit:junit:4.13.2") testImplementation("junit:junit:4.13.2")
testImplementation("org.json:json:20240303")
} }
@@ -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<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, 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?>): 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
}
@@ -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<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 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)
}
}