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