mirror of
https://github.com/priyanshujain/sanderling.git
synced 2026-10-02 11:07: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
@@ -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