diff --git a/Android/src/app/eval/AndroidManifest.xml b/Android/src/app/eval/AndroidManifest.xml new file mode 100644 index 000000000..c8a08d994 --- /dev/null +++ b/Android/src/app/eval/AndroidManifest.xml @@ -0,0 +1,28 @@ + + + + + + + + + + + diff --git a/Android/src/app/eval/java/com/google/ai/edge/gallery/eval/MiniHttpServer.kt b/Android/src/app/eval/java/com/google/ai/edge/gallery/eval/MiniHttpServer.kt new file mode 100644 index 000000000..5dbcfa5b7 --- /dev/null +++ b/Android/src/app/eval/java/com/google/ai/edge/gallery/eval/MiniHttpServer.kt @@ -0,0 +1,153 @@ +/* + * Copyright 2024 The Android Open Source Project + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ +package com.google.ai.edge.gallery.eval + +import android.util.Log +import java.io.BufferedInputStream +import java.io.OutputStream +import java.net.ServerSocket +import java.net.Socket +import kotlinx.coroutines.CoroutineScope +import kotlinx.coroutines.Dispatchers +import kotlinx.coroutines.cancel +import kotlinx.coroutines.launch + +class MiniHttpServer(val port: Int, val handler: (Request) -> Response) { + private var serverSocket: ServerSocket? = null + private val scope = CoroutineScope(Dispatchers.IO) + @Volatile private var running = false + + fun start() { + running = true + serverSocket = ServerSocket(port) + scope.launch { + while (running) { + try { + val socket = serverSocket?.accept() ?: break + scope.launch { handleConnection(socket) } + } catch (e: Exception) { + if (running) Log.e(TAG, "Error accepting connection", e) + } + } + } + } + + fun stop() { + running = false + serverSocket?.close() + serverSocket = null + scope.cancel() + } + + private fun handleConnection(socket: Socket) { + try { + val input = BufferedInputStream(socket.inputStream) + val output = socket.outputStream + + val headerBuilder = StringBuilder() + var currentByte = input.read() + + // Read headers until \r\n\r\n + while (currentByte != -1) { + headerBuilder.append(currentByte.toChar()) + if (headerBuilder.endsWith("\r\n\r\n")) break + currentByte = input.read() + } + + val headerString = headerBuilder.toString() + val headerLines = headerString.lines() + if (headerLines.isEmpty()) return + + val requestLineParts = headerLines[0].split(" ") + if (requestLineParts.size < 3) return + val method = requestLineParts[0] + val path = requestLineParts[1] + + // Parse headers + val headers = mutableMapOf() + var contentLength = 0 + for (i in 1 until headerLines.size) { + val line = headerLines[i] + if (line.isEmpty()) continue + val parts = line.split(":", limit = 2) + if (parts.size == 2) { + val key = parts[0].trim().lowercase() + val value = parts[1].trim() + headers[key] = value + if (key == "content-length") { + contentLength = value.toIntOrNull() ?: 0 + } + } + } + + // Safely read exact number of BYTES for the body + val body = + if (contentLength > 0) { + val bodyBytes = ByteArray(contentLength) + var bytesRead = 0 + while (bytesRead < contentLength) { + val read = input.read(bodyBytes, bytesRead, contentLength - bytesRead) + if (read == -1) break + bytesRead += read + } + String(bodyBytes, Charsets.UTF_8) // Decode to string safely here! + } else "" + + val request = Request(method, path, headers, body) + val response = handler(request) + writeResponse(output, response) + } catch (e: Exception) { + Log.e(TAG, "Error handling connection", e) + } finally { + socket.close() + } + } + + private fun writeResponse(output: OutputStream, response: Response) { + // Use the byte size of the encoded UTF-8 body to compute Content-Length. + // Using string character length (response.body.length) would cause response + // truncation on the client side for responses containing multi-byte characters. + val bodyBytes = response.body.toByteArray() + val statusLine = "HTTP/1.1 ${response.status.code} ${response.status.message}\r\n" + output.write(statusLine.toByteArray()) + output.write("Content-Type: ${response.contentType}\r\n".toByteArray()) + output.write("Content-Length: ${bodyBytes.size}\r\n".toByteArray()) + output.write("Connection: close\r\n".toByteArray()) + output.write("\r\n".toByteArray()) + output.write(bodyBytes) + output.flush() + } + + data class Request( + val method: String, + val path: String, + val headers: Map, + val body: String, + ) + + data class Response(val status: Status, val contentType: String, val body: String) + + enum class Status(val code: Int, val message: String) { + OK(200, "OK"), + BAD_REQUEST(400, "Bad Request"), + NOT_FOUND(404, "Not Found"), + INTERNAL_ERROR(500, "Internal Server Error"), + } + + companion object { + private const val TAG = "MiniHttpServer" + } +} diff --git a/Android/src/app/eval/javatests/com/google/ai/edge/gallery/eval/EvalAppTestSuite.kt b/Android/src/app/eval/javatests/com/google/ai/edge/gallery/eval/EvalAppTestSuite.kt new file mode 100644 index 000000000..4c7fe9b70 --- /dev/null +++ b/Android/src/app/eval/javatests/com/google/ai/edge/gallery/eval/EvalAppTestSuite.kt @@ -0,0 +1,22 @@ +/* + * Copyright 2024 The Android Open Source Project + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ +package com.google.ai.edge.gallery.eval + +import org.junit.runner.RunWith +import org.junit.runners.Suite + +/** Test suite to run all local unit tests for the on-device evaluation app. */ +@RunWith(Suite::class) @Suite.SuiteClasses(MiniHttpServerTest::class) class EvalAppTestSuite diff --git a/Android/src/app/eval/javatests/com/google/ai/edge/gallery/eval/MiniHttpServerTest.kt b/Android/src/app/eval/javatests/com/google/ai/edge/gallery/eval/MiniHttpServerTest.kt new file mode 100644 index 000000000..f1437409d --- /dev/null +++ b/Android/src/app/eval/javatests/com/google/ai/edge/gallery/eval/MiniHttpServerTest.kt @@ -0,0 +1,105 @@ +/* + * Copyright 2024 The Android Open Source Project + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ +package com.google.ai.edge.gallery.eval + +import com.google.common.truth.Truth.assertThat +import java.io.BufferedReader +import java.io.InputStreamReader +import java.net.HttpURLConnection +import java.net.URL +import org.junit.After +import org.junit.Before +import org.junit.Test +import org.junit.runner.RunWith +import org.robolectric.RobolectricTestRunner + +/** + * Unit tests for [MiniHttpServer] to verify socket handling, request parsing, and routing logic. + */ +@RunWith(RobolectricTestRunner::class) +class MiniHttpServerTest { + + private lateinit var server: MiniHttpServer + private val port = 8081 + + @Before + fun setUp() { + server = + MiniHttpServer(port) { request -> + if (request.path == "/test") { + MiniHttpServer.Response(MiniHttpServer.Status.OK, "text/plain", "Test OK") + } else { + MiniHttpServer.Response(MiniHttpServer.Status.NOT_FOUND, "text/plain", "Not Found") + } + } + server.start() + } + + @After + fun tearDown() { + server.stop() + } + + @Test + fun handleRequest_validRoute_returnsOk() { + val url = URL("http://localhost:$port/test") + val connection = url.openConnection() as HttpURLConnection + connection.requestMethod = "GET" + + assertThat(connection.responseCode).isEqualTo(200) + + val response = BufferedReader(InputStreamReader(connection.inputStream)).readText() + assertThat(response).isEqualTo("Test OK") + } + + @Test + fun handleRequest_invalidRoute_returnsNotFound() { + val url = URL("http://localhost:$port/unknown") + val connection = url.openConnection() as HttpURLConnection + connection.requestMethod = "GET" + + assertThat(connection.responseCode).isEqualTo(404) + } + + @Test + fun handleRequest_emptyRequest_returnsEarly() { + val socket = java.net.Socket("localhost", port) + socket.getOutputStream().close() // Immediately close + socket.close() // Should return early without throwing + } + + @Test + fun handleRequest_malformedRequest_returnsEarly() { + val socket = java.net.Socket("localhost", port) + val out = socket.getOutputStream() + out.write("GET\r\n\r\n".toByteArray()) + out.flush() + socket.close() + } + + @Test + fun handleRequest_headersWithEmptyLinesAndContentLength_doesNotCrash() { + val socket = java.net.Socket("localhost", port) + val out = socket.getOutputStream() + val req = "POST /test HTTP/1.1\r\nContent-Length: 4\r\n\r\n\r\nbody" + out.write(req.toByteArray()) + out.flush() + val response = + java.io.BufferedReader(java.io.InputStreamReader(socket.getInputStream())).readText() + assertThat(response).contains("200 OK") + socket.close() + } +}