diff --git a/sdk/android/src/main/kotlin/dev/uatu/sdk/LocalAbstractTransport.kt b/sdk/android/src/main/kotlin/dev/uatu/sdk/LocalAbstractTransport.kt new file mode 100644 index 0000000..c945c82 --- /dev/null +++ b/sdk/android/src/main/kotlin/dev/uatu/sdk/LocalAbstractTransport.kt @@ -0,0 +1,26 @@ +package dev.uatu.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() + } + } +} diff --git a/sdk/android/src/main/kotlin/dev/uatu/sdk/SocketClient.kt b/sdk/android/src/main/kotlin/dev/uatu/sdk/SocketClient.kt new file mode 100644 index 0000000..4a0d556 --- /dev/null +++ b/sdk/android/src/main/kotlin/dev/uatu/sdk/SocketClient.kt @@ -0,0 +1,108 @@ +package dev.uatu.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, "uatu-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) + } + } +} diff --git a/sdk/android/src/test/kotlin/dev/uatu/sdk/SocketClientTest.kt b/sdk/android/src/test/kotlin/dev/uatu/sdk/SocketClientTest.kt new file mode 100644 index 0000000..36d3d34 --- /dev/null +++ b/sdk/android/src/test/kotlin/dev/uatu/sdk/SocketClientTest.kt @@ -0,0 +1,190 @@ +package dev.uatu.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() + 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() + val messages = CopyOnWriteArrayList() + @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() + 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)) + } +}