diff --git a/Android/src/app/build.gradle.kts b/Android/src/app/build.gradle.kts index f5eaf41878..d872a938b9 100644 --- a/Android/src/app/build.gradle.kts +++ b/Android/src/app/build.gradle.kts @@ -126,6 +126,11 @@ dependencies { implementation(libs.mcp.kotlin.sdk) implementation(libs.ktor.client.android) implementation(libs.ktor.client.core) + implementation(libs.ktor.server.core) + implementation(libs.ktor.server.cio) + implementation(libs.ktor.server.content.negotiation) + implementation(libs.ktor.server.cors) + implementation(libs.ktor.serialization.kotlinx.json) } protobuf { diff --git a/Android/src/app/src/main/AndroidManifest.xml b/Android/src/app/src/main/AndroidManifest.xml index b78b08a3f4..a3ea35e5eb 100644 --- a/Android/src/app/src/main/AndroidManifest.xml +++ b/Android/src/app/src/main/AndroidManifest.xml @@ -31,7 +31,12 @@ + + + + + @@ -58,6 +63,7 @@ + + + + ) @@ -62,6 +77,14 @@ interface DataStoreRepository { */ fun readFirebaseAnalytics(): Boolean + fun saveOpenAiApiServerPreferences(preferences: OpenAiApiServerPreferences) + + fun readOpenAiApiServerPreferences(): OpenAiApiServerPreferences + + fun saveFrpPreferences(preferences: FrpPreferences) + + fun readFrpPreferences(): FrpPreferences + fun saveSecret(key: String, value: String) fun readSecret(key: String): String? @@ -193,6 +216,60 @@ class DefaultDataStoreRepository( } } + override fun saveOpenAiApiServerPreferences(preferences: OpenAiApiServerPreferences) { + runBlocking { + dataStore.updateData { settings -> + settings + .toBuilder() + .setOpenaiApiServerEnabled(preferences.enabled) + .setOpenaiApiServerPort(preferences.port) + .setOpenaiApiServerModel(preferences.modelName) + .build() + } + } + } + + override fun readOpenAiApiServerPreferences(): OpenAiApiServerPreferences { + return runBlocking { + val settings = dataStore.data.first() + OpenAiApiServerPreferences( + enabled = settings.openaiApiServerEnabled, + port = settings.openaiApiServerPort.takeIf { it in 1024..65535 } ?: 8080, + modelName = settings.openaiApiServerModel, + ) + } + } + + override fun saveFrpPreferences(preferences: FrpPreferences) { + runBlocking { + dataStore.updateData { settings -> + settings + .toBuilder() + .setFrpEnabled(preferences.enabled) + .setFrpServerAddress(preferences.serverAddress) + .setFrpServerPort(preferences.serverPort) + .setFrpToken(preferences.token) + .setFrpRemotePort(preferences.remotePort) + .setFrpCustomDomain(preferences.customDomain) + .build() + } + } + } + + override fun readFrpPreferences(): FrpPreferences { + return runBlocking { + val settings = dataStore.data.first() + FrpPreferences( + enabled = settings.frpEnabled, + serverAddress = settings.frpServerAddress, + serverPort = settings.frpServerPort, + token = settings.frpToken, + remotePort = settings.frpRemotePort, + customDomain = settings.frpCustomDomain, + ) + } + } + override fun saveSecret(key: String, value: String) { runBlocking { userDataDataStore.updateData { userData -> diff --git a/Android/src/app/src/main/java/com/google/ai/edge/gallery/server/FrpManager.kt b/Android/src/app/src/main/java/com/google/ai/edge/gallery/server/FrpManager.kt new file mode 100644 index 0000000000..2497c5afe9 --- /dev/null +++ b/Android/src/app/src/main/java/com/google/ai/edge/gallery/server/FrpManager.kt @@ -0,0 +1,163 @@ +/* + * Copyright 2026 Google LLC + * + * 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.server + +import android.content.Context +import android.net.Uri +import android.util.Log +import dagger.hilt.android.qualifiers.ApplicationContext +import java.io.File +import javax.inject.Inject +import javax.inject.Singleton +import kotlinx.coroutines.CoroutineScope +import kotlinx.coroutines.Dispatchers +import kotlinx.coroutines.SupervisorJob +import kotlinx.coroutines.flow.MutableStateFlow +import kotlinx.coroutines.flow.asStateFlow +import kotlinx.coroutines.launch + +private const val TAG = "AGFrpManager" + +@Singleton +class FrpManager @Inject constructor(@ApplicationContext private val context: Context) { + private var process: Process? = null + private val frpScope = CoroutineScope(SupervisorJob() + Dispatchers.IO) + private val _isRunning = MutableStateFlow(false) + val isRunning = _isRunning.asStateFlow() + + fun isBinaryAvailable(): Boolean { + val binary = getFrpcBinary() + return binary != null && binary.exists() + } + + fun importBinary(uri: Uri): Boolean { + return try { + val targetFile = File(context.filesDir, "frpc") + context.contentResolver.openInputStream(uri)?.use { input -> + targetFile.outputStream().use { output -> + input.copyTo(output) + } + } + targetFile.setExecutable(true) + Log.i(TAG, "frpc binary imported successfully to ${targetFile.absolutePath}") + true + } catch (e: Exception) { + Log.e(TAG, "Failed to import frpc binary", e) + false + } + } + + fun deleteBinary(): Boolean { + val file = File(context.filesDir, "frpc") + return if (file.exists()) { + file.delete() + } else { + true + } + } + + fun start(serverAddr: String, serverPort: Int, token: String, localPort: Int, remotePort: Int, customDomain: String = "") { + if (_isRunning.value) stop() + + val frpcBinary = getFrpcBinary() + if (frpcBinary == null || !frpcBinary.exists()) { + val msg = "frpc binary not found. Please run: adb push frpc /data/data/com.google.aiedge.gallery/files/frpc" + Log.e(TAG, msg) + return + } + + if (!frpcBinary.canExecute()) { + frpcBinary.setExecutable(true) + } + + val configFile = File(context.filesDir, "frpc.toml") + val proxyConfig = if (customDomain.isNotEmpty()) { + """ + type = "http" + localPort = $localPort + customDomains = ["$customDomain"] + """ + } else { + """ + type = "tcp" + localPort = $localPort + remotePort = $remotePort + """ + } + + configFile.writeText(""" + serverAddr = "$serverAddr" + serverPort = $serverPort + auth.token = "$token" + + [[proxies]] + name = "openai-api" + $proxyConfig + """.trimIndent()) + + frpScope.launch { + try { + Log.i(TAG, "Starting frpc...") + val builder = ProcessBuilder(frpcBinary.absolutePath, "-c", configFile.absolutePath) + .directory(context.filesDir) + .redirectErrorStream(true) + + val p = builder.start() + process = p + _isRunning.value = true + + p.inputStream.bufferedReader().use { reader -> + var line: String? + while (reader.readLine().also { line = it } != null) { + Log.d(TAG, "frpc: $line") + } + } + } catch (e: Exception) { + Log.e(TAG, "Error running frpc", e) + } finally { + _isRunning.value = false + process = null + } + } + } + + fun stop() { + process?.destroy() + process = null + _isRunning.value = false + } + + private fun getFrpcBinary(): File? { + // Expected binary name based on architecture + val arch = android.os.Build.SUPPORTED_ABIS.firstOrNull() ?: return null + val binaryName = when { + arch.contains("arm64") -> "frpc_android_arm64" + arch.contains("armeabi") -> "frpc_android_arm" + arch.contains("x86_64") -> "frpc_android_amd64" + else -> "frpc" + } + + // Check in files directory + val file = File(context.filesDir, "frpc") + if (file.exists()) return file + + val archFile = File(context.filesDir, binaryName) + if (archFile.exists()) return archFile + + return null + } +} diff --git a/Android/src/app/src/main/java/com/google/ai/edge/gallery/server/OpenAiApiModels.kt b/Android/src/app/src/main/java/com/google/ai/edge/gallery/server/OpenAiApiModels.kt new file mode 100644 index 0000000000..1be5306662 --- /dev/null +++ b/Android/src/app/src/main/java/com/google/ai/edge/gallery/server/OpenAiApiModels.kt @@ -0,0 +1,146 @@ +/* + * Copyright 2026 Google LLC + * + * 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.server + +import kotlinx.serialization.SerialName +import kotlinx.serialization.Serializable +import kotlinx.serialization.json.JsonArray +import kotlinx.serialization.json.JsonElement +import kotlinx.serialization.json.JsonObject +import kotlinx.serialization.json.JsonPrimitive +import kotlinx.serialization.json.contentOrNull +import kotlinx.serialization.json.jsonPrimitive + +@Serializable +data class OpenAiChatCompletionRequest( + val model: String = "", + val messages: List, + val stream: Boolean = false, + val temperature: Double? = null, + @SerialName("top_p") val topP: Double? = null, + @SerialName("max_tokens") val maxTokens: Int? = null, +) + +@Serializable data class OpenAiChatMessage(val role: String, val content: JsonElement) + +@Serializable +data class OpenAiModelList( + @SerialName("object") val objectType: String = "list", + val data: List, +) + +@Serializable +data class OpenAiModel( + val id: String, + @SerialName("object") val objectType: String = "model", + val created: Long, + @SerialName("owned_by") val ownedBy: String = "edge-gallery", +) + +@Serializable +data class OpenAiChatCompletionResponse( + val id: String, + @SerialName("object") val objectType: String = "chat.completion", + val created: Long, + val model: String, + val choices: List, + val usage: OpenAiUsage, +) + +@Serializable +data class OpenAiChatChoice( + val index: Int = 0, + val message: OpenAiAssistantMessage, + @SerialName("finish_reason") val finishReason: String = "stop", +) + +@Serializable +data class OpenAiAssistantMessage(val role: String = "assistant", val content: String) + +@Serializable +data class OpenAiUsage( + @SerialName("prompt_tokens") val promptTokens: Int, + @SerialName("completion_tokens") val completionTokens: Int, + @SerialName("total_tokens") val totalTokens: Int, +) + +@Serializable +data class OpenAiChatCompletionChunk( + val id: String, + @SerialName("object") val objectType: String = "chat.completion.chunk", + val created: Long, + val model: String, + val choices: List, +) + +@Serializable +data class OpenAiChunkChoice( + val index: Int = 0, + val delta: OpenAiDelta, + @SerialName("finish_reason") val finishReason: String? = null, +) + +@Serializable data class OpenAiDelta(val role: String? = null, val content: String? = null) + +@Serializable data class OpenAiErrorEnvelope(val error: OpenAiError) + +@Serializable +data class OpenAiError( + val message: String, + val type: String = "invalid_request_error", + val param: String? = null, + val code: String? = null, +) + +internal fun OpenAiChatCompletionRequest.toInferencePrompt(): String { + return messages + .mapNotNull { message -> + val text = message.content.asText().trim() + if (text.isEmpty()) null + else { + val role = + when (message.role.lowercase()) { + "system" -> "System" + "developer" -> "Developer" + "assistant" -> "Assistant" + "user" -> "User" + else -> message.role.replaceFirstChar { it.uppercase() } + } + "$role: $text" + } + } + .joinToString(separator = "\n\n", postfix = "\n\nAssistant:") +} + +private fun JsonElement.asText(): String { + return when (this) { + is JsonPrimitive -> contentOrNull.orEmpty() + is JsonArray -> + mapNotNull { part -> + val obj = part as? JsonObject ?: return@mapNotNull null + if (obj["type"]?.jsonPrimitive?.contentOrNull == "text") { + obj["text"]?.jsonPrimitive?.contentOrNull + } else { + null + } + } + .joinToString("\n") + else -> "" + } +} + +internal fun estimateTokens(text: String): Int = (text.length / 4).coerceAtLeast(if (text.isEmpty()) 0 else 1) diff --git a/Android/src/app/src/main/java/com/google/ai/edge/gallery/server/OpenAiApiServer.kt b/Android/src/app/src/main/java/com/google/ai/edge/gallery/server/OpenAiApiServer.kt new file mode 100644 index 0000000000..a8dc70ec14 --- /dev/null +++ b/Android/src/app/src/main/java/com/google/ai/edge/gallery/server/OpenAiApiServer.kt @@ -0,0 +1,388 @@ +/* + * Copyright 2026 Google LLC + * + * 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.server + +import android.content.Context +import android.util.Log +import com.google.ai.edge.gallery.data.BuiltInTaskId +import com.google.ai.edge.gallery.data.Model +import com.google.ai.edge.gallery.runtime.runtimeHelper +import dagger.hilt.android.qualifiers.ApplicationContext +import io.ktor.http.ContentType +import io.ktor.http.HttpHeaders +import io.ktor.http.HttpMethod +import io.ktor.http.HttpStatusCode +import io.ktor.serialization.kotlinx.json.json +import io.ktor.server.application.Application +import io.ktor.server.application.ApplicationCall +import io.ktor.server.application.ApplicationCallPipeline +import io.ktor.server.application.call +import io.ktor.server.application.install +import io.ktor.server.cio.CIO +import io.ktor.server.engine.EmbeddedServer +import io.ktor.server.engine.embeddedServer +import io.ktor.server.plugins.contentnegotiation.ContentNegotiation +import io.ktor.server.plugins.cors.routing.CORS +import io.ktor.server.request.header +import io.ktor.server.request.httpMethod +import io.ktor.server.request.receive +import io.ktor.server.request.uri +import io.ktor.server.response.respond +import io.ktor.server.response.respondTextWriter +import io.ktor.server.routing.get +import io.ktor.server.routing.post +import io.ktor.server.routing.routing +import java.net.Inet4Address +import java.net.NetworkInterface +import java.security.MessageDigest +import java.time.Instant +import java.util.UUID +import java.util.concurrent.atomic.AtomicBoolean +import javax.inject.Inject +import javax.inject.Singleton +import kotlin.coroutines.resume +import kotlin.coroutines.resumeWithException +import kotlinx.coroutines.CoroutineScope +import kotlinx.coroutines.Dispatchers +import kotlinx.coroutines.SupervisorJob +import kotlinx.coroutines.channels.awaitClose +import kotlinx.coroutines.flow.Flow +import kotlinx.coroutines.flow.MutableStateFlow +import kotlinx.coroutines.flow.asStateFlow +import kotlinx.coroutines.flow.callbackFlow +import kotlinx.coroutines.flow.collect +import kotlinx.coroutines.launch +import kotlinx.coroutines.sync.Mutex +import kotlinx.coroutines.sync.withLock +import kotlinx.coroutines.suspendCancellableCoroutine +import kotlinx.serialization.encodeToString +import kotlinx.serialization.json.Json + +private const val TAG = "AGOpenAiApiServer" + +enum class OpenAiApiServerStatus { + STOPPED, + STARTING, + RUNNING, + ERROR, +} + +data class OpenAiApiServerState( + val status: OpenAiApiServerStatus = OpenAiApiServerStatus.STOPPED, + val port: Int = 8080, + val endpoint: String = "", + val activeModel: String = "", + val requestCount: Long = 0, + val error: String = "", +) + +@Singleton +class OpenAiApiServer +@Inject +constructor(@ApplicationContext private val context: Context) { + private val json = Json { + ignoreUnknownKeys = true + explicitNulls = false + encodeDefaults = true + } + private val serverScope = CoroutineScope(SupervisorJob() + Dispatchers.Default) + private val inferenceMutex = Mutex() + private val _state = MutableStateFlow(OpenAiApiServerState()) + val state = _state.asStateFlow() + + @Volatile private var availableModels: Map = emptyMap() + private var activeRuntimeModel: Model? = null + private var activePublicModelName: String = "" + private var engine: EmbeddedServer<*, *>? = null + + fun updateModels(models: List) { + availableModels = models.associateBy { it.name } + } + + @Synchronized + fun start(port: Int, apiKey: String, defaultModel: String) { + Log.i(TAG, "Starting API server on port $port...") + if (engine != null && _state.value.port == port) { + Log.i(TAG, "Server already running on port $port") + return + } + stop() + if (port !in 1024..65535) { + val error = "Port must be between 1024 and 65535." + Log.e(TAG, error) + _state.value = OpenAiApiServerState(status = OpenAiApiServerStatus.ERROR, error = error) + return + } + if (apiKey.isBlank()) { + val error = "API key is missing." + Log.e(TAG, error) + _state.value = OpenAiApiServerState(status = OpenAiApiServerStatus.ERROR, error = error) + return + } + if (availableModels.isEmpty()) { + val error = "Download an LLM before starting the API server." + Log.e(TAG, error) + _state.value = OpenAiApiServerState(status = OpenAiApiServerStatus.ERROR, error = error) + return + } + + _state.value = OpenAiApiServerState(status = OpenAiApiServerStatus.STARTING, port = port) + try { + val address = findLanAddress() + val host = address ?: "0.0.0.0" + Log.i(TAG, "Creating engine on $host:$port...") + + val newEngine = embeddedServer(CIO, host = host, port = port) { + configureApi(apiKey = apiKey, defaultModel = defaultModel) + } + Log.i(TAG, "Engine created, starting...") + newEngine.start(wait = false) + engine = newEngine + val endpoint = "http://${address ?: "127.0.0.1"}:$port/v1" + Log.i(TAG, "API server is running at $endpoint") + _state.value = + OpenAiApiServerState( + status = OpenAiApiServerStatus.RUNNING, + port = port, + endpoint = endpoint, + activeModel = defaultModel, + ) + } catch (e: Exception) { + Log.e(TAG, "Failed to start API server", e) + engine = null + _state.value = OpenAiApiServerState(status = OpenAiApiServerStatus.ERROR, port = port, error = e.message ?: "Unable to start server.") + } + } + + @Synchronized + fun stop() { + try { + engine?.stop(500, 2_000) + } catch (e: Exception) { + Log.w(TAG, "Error while stopping API server", e) + } + engine = null + val model = activeRuntimeModel + activeRuntimeModel = null + activePublicModelName = "" + if (model != null) { + serverScope.launch { model.runtimeHelper.cleanUp(model) {} } + } + _state.value = OpenAiApiServerState(status = OpenAiApiServerStatus.STOPPED, port = _state.value.port) + } + + private fun Application.configureApi(apiKey: String, defaultModel: String) { + install(ContentNegotiation) { json(json) } + install(CORS) { + anyHost() + allowMethod(HttpMethod.Get) + allowMethod(HttpMethod.Post) + allowHeader(HttpHeaders.Authorization) + allowHeader(HttpHeaders.ContentType) + } + routing { + intercept(ApplicationCallPipeline.Plugins) { + val method = call.request.httpMethod.value + val uri = call.request.uri + Log.d(TAG, "Incoming request: $method $uri") + try { + proceed() + Log.d(TAG, "Response finished for: $method $uri") + } catch (e: Exception) { + Log.e(TAG, "Error processing $method $uri", e) + throw e + } + } + get("/") { + call.respond(mapOf("status" to "ok", "api" to "OpenAI compatible", "base_url" to "/v1")) + } + get("/health") { call.respond(mapOf("status" to "ok")) } + get("/v1/models") { + if (!call.authorize(apiKey)) return@get + val created = Instant.now().epochSecond + call.respond(OpenAiModelList(data = availableModels.keys.sorted().map { OpenAiModel(id = it, created = created) })) + } + post("/v1/chat/completions") { + if (!call.authorize(apiKey)) return@post + try { + val request = call.receive() + if (request.messages.isEmpty()) { + call.respondError(HttpStatusCode.BadRequest, "messages must not be empty", "messages") + return@post + } + val modelName = request.model.ifBlank { defaultModel.ifBlank { availableModels.keys.firstOrNull().orEmpty() } } + if (!availableModels.containsKey(modelName)) { + call.respondError(HttpStatusCode.NotFound, "Model '$modelName' is not available on this device.", "model", "model_not_found") + return@post + } + val prompt = request.toInferencePrompt() + if (prompt.isBlank()) { + call.respondError(HttpStatusCode.BadRequest, "No supported text content was found in messages.", "messages") + return@post + } + val id = "chatcmpl-${UUID.randomUUID().toString().replace("-", "")}" + val created = Instant.now().epochSecond + if (request.stream) { + call.respondTextWriter(contentType = ContentType.Text.EventStream) { + writeSse(OpenAiChatCompletionChunk(id = id, created = created, model = modelName, choices = listOf(OpenAiChunkChoice(delta = OpenAiDelta(role = "assistant"))))) + complete(modelName = modelName, prompt = prompt) { token -> + if (token.isNotEmpty()) { + writeSse(OpenAiChatCompletionChunk(id = id, created = created, model = modelName, choices = listOf(OpenAiChunkChoice(delta = OpenAiDelta(content = token))))) + } + } + writeSse(OpenAiChatCompletionChunk(id = id, created = created, model = modelName, choices = listOf(OpenAiChunkChoice(delta = OpenAiDelta(), finishReason = "stop")))) + write("data: [DONE]\n\n") + flush() + } + } else { + val responseText = complete(modelName = modelName, prompt = prompt) + val promptTokens = estimateTokens(prompt) + val completionTokens = estimateTokens(responseText) + call.respond( + OpenAiChatCompletionResponse( + id = id, + created = created, + model = modelName, + choices = listOf(OpenAiChatChoice(message = OpenAiAssistantMessage(content = responseText))), + usage = OpenAiUsage(promptTokens, completionTokens, promptTokens + completionTokens), + ) + ) + } + _state.value = _state.value.copy(requestCount = _state.value.requestCount + 1, activeModel = modelName, error = "") + } catch (e: Exception) { + Log.e(TAG, "Completion failed", e) + _state.value = _state.value.copy(error = e.message ?: "Inference failed") + if (!call.response.isCommitted) { + call.respondError(HttpStatusCode.InternalServerError, e.message ?: "Inference failed", code = "server_error") + } + } + } + } + } + + private suspend fun complete(modelName: String, prompt: String, onToken: suspend (String) -> Unit = {}): String { + return inferenceMutex.withLock { + val model = ensureModel(modelName) + model.runtimeHelper.resetConversation(model = model) + val response = StringBuilder() + inferenceFlow(model, prompt).collect { token -> + response.append(token) + onToken(token) + } + response.toString() + } + } + + private suspend fun ensureModel(publicModelName: String): Model { + if (activeRuntimeModel != null && activePublicModelName == publicModelName) return activeRuntimeModel!! + activeRuntimeModel?.let { previous -> + suspendCancellableCoroutine { continuation -> + previous.runtimeHelper.cleanUp(previous) { if (continuation.isActive) continuation.resume(Unit) } + } + } + val source = availableModels[publicModelName] ?: error("Model '$publicModelName' is unavailable.") + val runtimeModel = + source.copy( + name = "${source.name}__openai_api", + localModelFilePathOverride = source.getPath(context), + instance = null, + initializing = false, + cleanUpAfterInit = false, + ) + activeRuntimeModel = runtimeModel + activePublicModelName = publicModelName + suspendCancellableCoroutine { continuation -> + val completed = AtomicBoolean(false) + runtimeModel.runtimeHelper.initialize( + context = context, + model = runtimeModel, + taskId = BuiltInTaskId.LLM_CHAT, + supportImage = false, + supportAudio = false, + coroutineScope = serverScope, + onDone = { message -> + if (runtimeModel.instance != null && completed.compareAndSet(false, true)) { + if (continuation.isActive) continuation.resume(Unit) + } else if ( + runtimeModel.instance == null && + (message.contains("failed", ignoreCase = true) || message.contains("unavailable", ignoreCase = true)) && + completed.compareAndSet(false, true) + ) { + if (continuation.isActive) continuation.resumeWithException(IllegalStateException(message)) + } + }, + ) + continuation.invokeOnCancellation { runtimeModel.runtimeHelper.stopResponse(runtimeModel) } + } + return runtimeModel + } + + private fun inferenceFlow(model: Model, prompt: String): Flow = callbackFlow { + val finished = AtomicBoolean(false) + model.runtimeHelper.runInference( + model = model, + input = prompt, + coroutineScope = serverScope, + resultListener = { token, done, _ -> + if (token.isNotEmpty()) trySend(token) + if (done && finished.compareAndSet(false, true)) close() + }, + cleanUpListener = { if (finished.compareAndSet(false, true)) close() }, + onError = { message -> if (finished.compareAndSet(false, true)) close(IllegalStateException(message)) }, + ) + awaitClose { + if (!finished.get()) model.runtimeHelper.stopResponse(model) + } + } + + private suspend fun ApplicationCall.authorize(apiKey: String): Boolean { + val supplied = request.header(HttpHeaders.Authorization)?.removePrefix("Bearer ")?.trim().orEmpty() + if (supplied.isNotEmpty() && MessageDigest.isEqual(supplied.toByteArray(), apiKey.toByteArray())) return true + respondError(HttpStatusCode.Unauthorized, "Invalid or missing bearer token.", code = "invalid_api_key") + return false + } + + private suspend fun ApplicationCall.respondError(status: HttpStatusCode, message: String, param: String? = null, code: String? = null) { + respond(status, OpenAiErrorEnvelope(OpenAiError(message = message, param = param, code = code))) + } + + private suspend fun java.io.Writer.writeSse(chunk: OpenAiChatCompletionChunk) { + write("data: ${json.encodeToString(chunk)}\n\n") + flush() + } + + private fun findLanAddress(): String? { + return runCatching { + NetworkInterface.getNetworkInterfaces().toList() + .flatMap { it.inetAddresses.toList() } + .filterIsInstance() + .firstOrNull { address -> + !address.isLoopbackAddress && !address.isLinkLocalAddress && address.hostAddress?.let(::isPrivateIpv4) == true + } + ?.hostAddress + } + .getOrNull() + } + + private fun isPrivateIpv4(value: String): Boolean { + val octets = value.split('.').mapNotNull { it.toIntOrNull() } + if (octets.size != 4) return false + return octets[0] == 10 || + (octets[0] == 172 && octets[1] in 16..31) || + (octets[0] == 192 && octets[1] == 168) + } +} diff --git a/Android/src/app/src/main/java/com/google/ai/edge/gallery/server/OpenAiApiServerScreen.kt b/Android/src/app/src/main/java/com/google/ai/edge/gallery/server/OpenAiApiServerScreen.kt new file mode 100644 index 0000000000..866b01754f --- /dev/null +++ b/Android/src/app/src/main/java/com/google/ai/edge/gallery/server/OpenAiApiServerScreen.kt @@ -0,0 +1,414 @@ +/* + * Copyright 2026 Google LLC + * + * 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.server + +import android.content.ClipData +import android.content.ClipboardManager +import android.content.Context +import androidx.activity.compose.rememberLauncherForActivityResult +import androidx.activity.result.contract.ActivityResultContracts +import androidx.compose.foundation.layout.Arrangement +import androidx.compose.foundation.layout.Column +import androidx.compose.foundation.layout.Row +import androidx.compose.foundation.layout.Spacer +import androidx.compose.foundation.layout.fillMaxSize +import androidx.compose.foundation.layout.fillMaxWidth +import androidx.compose.foundation.layout.height +import androidx.compose.foundation.layout.padding +import androidx.compose.foundation.layout.size +import androidx.compose.foundation.layout.width +import androidx.compose.foundation.rememberScrollState +import androidx.compose.foundation.text.KeyboardOptions +import androidx.compose.foundation.verticalScroll +import androidx.compose.material.icons.Icons +import androidx.compose.material.icons.automirrored.rounded.ArrowBack +import androidx.compose.material.icons.rounded.ContentCopy +import androidx.compose.material.icons.rounded.Refresh +import androidx.compose.material3.Button +import androidx.compose.material3.Card +import androidx.compose.material3.DropdownMenuItem +import androidx.compose.material3.ExperimentalMaterial3Api +import androidx.compose.material3.ExposedDropdownMenuBox +import androidx.compose.material3.ExposedDropdownMenuDefaults +import androidx.compose.material3.Icon +import androidx.compose.material3.IconButton +import androidx.compose.material3.MaterialTheme +import androidx.compose.material3.OutlinedButton +import androidx.compose.material3.OutlinedTextField +import androidx.compose.material3.Scaffold +import androidx.compose.material3.Switch +import androidx.compose.material3.Text +import androidx.compose.material3.TextButton +import androidx.compose.material3.TopAppBar +import androidx.compose.runtime.Composable +import androidx.compose.runtime.LaunchedEffect +import androidx.compose.runtime.collectAsState +import androidx.compose.runtime.getValue +import androidx.compose.runtime.mutableStateOf +import androidx.compose.runtime.remember +import androidx.compose.runtime.setValue +import androidx.compose.ui.Alignment +import androidx.compose.ui.Modifier +import androidx.compose.ui.platform.LocalContext +import androidx.compose.ui.res.stringResource +import androidx.compose.ui.text.font.FontFamily +import androidx.compose.ui.text.font.FontWeight +import androidx.compose.ui.text.input.KeyboardType +import androidx.compose.ui.unit.dp +import androidx.hilt.lifecycle.viewmodel.compose.hiltViewModel +import com.google.ai.edge.gallery.R +import com.google.ai.edge.gallery.ui.common.ClickableLink +import com.google.ai.edge.gallery.ui.modelmanager.ModelManagerViewModel + +@OptIn(ExperimentalMaterial3Api::class) +@Composable +fun OpenAiApiServerScreen( + modelManagerViewModel: ModelManagerViewModel, + viewModel: OpenAiApiServerViewModel = hiltViewModel(), + onBackClicked: () -> Unit, +) { + val uiState by viewModel.uiState.collectAsState() + val context = LocalContext.current + val scrollState = rememberScrollState() + + val frpPickerLauncher = rememberLauncherForActivityResult( + contract = ActivityResultContracts.GetContent() + ) { uri -> + uri?.let { viewModel.importFrpBinary(it) } + } + + LaunchedEffect(Unit) { + viewModel.updateModels(modelManagerViewModel.getAllDownloadedModels()) + } + + Scaffold( + topBar = { + TopAppBar( + title = { Text(stringResource(R.string.openai_api_server_title)) }, + navigationIcon = { + IconButton(onClick = onBackClicked) { + Icon(Icons.AutoMirrored.Rounded.ArrowBack, contentDescription = stringResource(R.string.cd_navigate_back_icon)) + } + } + ) + } + ) { innerPadding -> + Column( + modifier = Modifier + .padding(innerPadding) + .padding(16.dp) + .fillMaxSize() + .verticalScroll(scrollState), + verticalArrangement = Arrangement.spacedBy(24.dp) + ) { + Text( + stringResource(R.string.openai_api_server_description), + style = MaterialTheme.typography.bodyMedium, + color = MaterialTheme.colorScheme.onSurfaceVariant + ) + + // Server Toggle + Card(modifier = Modifier.fillMaxWidth()) { + Row( + modifier = Modifier + .padding(16.dp) + .fillMaxWidth(), + verticalAlignment = Alignment.CenterVertically, + horizontalArrangement = Arrangement.SpaceBetween + ) { + Column(modifier = Modifier.weight(1f)) { + Text( + stringResource(R.string.openai_api_server_enabled), + style = MaterialTheme.typography.titleMedium + ) + Text( + uiState.status.name, + style = MaterialTheme.typography.bodySmall, + color = if (uiState.status == OpenAiApiServerStatus.ERROR) MaterialTheme.colorScheme.error else MaterialTheme.colorScheme.primary + ) + } + Switch( + checked = uiState.enabled, + onCheckedChange = { + if (it) viewModel.start(uiState.port, uiState.selectedModel) + else viewModel.stop() + } + ) + } + } + + if (uiState.error.isNotEmpty()) { + Text( + uiState.error, + color = MaterialTheme.colorScheme.error, + style = MaterialTheme.typography.bodySmall, + modifier = Modifier.padding(horizontal = 8.dp) + ) + } + + // Configuration + Column(verticalArrangement = Arrangement.spacedBy(16.dp)) { + // Port + var portText by remember(uiState.port) { mutableStateOf(uiState.port.toString()) } + OutlinedTextField( + value = portText, + onValueChange = { + portText = it + it.toIntOrNull()?.let { port -> + if (port in 1024..65535) { + if (uiState.enabled) viewModel.start(port, uiState.selectedModel) + } + } + }, + label = { Text(stringResource(R.string.openai_api_server_port)) }, + modifier = Modifier.fillMaxWidth(), + keyboardOptions = KeyboardOptions(keyboardType = KeyboardType.Number), + singleLine = true + ) + + // Model Selector + var expanded by remember { mutableStateOf(false) } + ExposedDropdownMenuBox( + expanded = expanded, + onExpandedChange = { expanded = !expanded }, + modifier = Modifier.fillMaxWidth() + ) { + OutlinedTextField( + value = uiState.selectedModel, + onValueChange = {}, + readOnly = true, + label = { Text(stringResource(R.string.openai_api_server_model)) }, + trailingIcon = { ExposedDropdownMenuDefaults.TrailingIcon(expanded = expanded) }, + colors = ExposedDropdownMenuDefaults.outlinedTextFieldColors(), + modifier = Modifier + .menuAnchor() + .fillMaxWidth() + ) + ExposedDropdownMenu( + expanded = expanded, + onDismissRequest = { expanded = false } + ) { + uiState.availableModels.forEach { modelName -> + DropdownMenuItem( + text = { Text(modelName) }, + onClick = { + expanded = false + if (uiState.enabled) viewModel.start(uiState.port, modelName) + else viewModel.updateModels(modelManagerViewModel.getAllDownloadedModels(), preferredModel = modelName) + } + ) + } + } + } + } + + // FRP Configuration + Text( + stringResource(R.string.openai_api_server_frp_title), + style = MaterialTheme.typography.titleLarge, + modifier = Modifier.padding(top = 8.dp) + ) + + Card(modifier = Modifier.fillMaxWidth()) { + Column(modifier = Modifier.fillMaxWidth()) { + Row( + modifier = Modifier + .padding(16.dp) + .fillMaxWidth(), + verticalAlignment = Alignment.CenterVertically, + horizontalArrangement = Arrangement.SpaceBetween + ) { + Column(modifier = Modifier.weight(1f)) { + Text(stringResource(R.string.openai_api_server_frp_enable), style = MaterialTheme.typography.titleMedium) + Text( + if (uiState.frpRunning) stringResource(R.string.openai_api_server_frp_running) else stringResource(R.string.openai_api_server_frp_stopped), + style = MaterialTheme.typography.bodySmall, + color = if (uiState.frpRunning) MaterialTheme.colorScheme.primary else MaterialTheme.colorScheme.onSurfaceVariant + ) + } + Switch( + checked = uiState.frpEnabled, + onCheckedChange = { + viewModel.updateFrpConfig(it, uiState.frpServerAddr, uiState.frpServerPort, uiState.frpToken, uiState.frpRemotePort, uiState.frpCustomDomain) + } + ) + } + if (!uiState.frpBinaryMissing) { + TextButton( + onClick = { viewModel.deleteFrpBinary() }, + modifier = Modifier.padding(start = 8.dp, bottom = 8.dp) + ) { + Text("Reset/Delete frpc binary", style = MaterialTheme.typography.labelSmall, color = MaterialTheme.colorScheme.error) + } + } + if (uiState.frpBinaryMissing) { + Column(modifier = Modifier.padding(start = 16.dp, end = 16.dp, bottom = 16.dp), verticalArrangement = Arrangement.spacedBy(8.dp)) { + Text( + "frpc binary missing! You can download it from GitHub (arm64-v8a) and import it here.", + style = MaterialTheme.typography.bodySmall, + color = MaterialTheme.colorScheme.error, + ) + Row(verticalAlignment = Alignment.CenterVertically, horizontalArrangement = Arrangement.spacedBy(8.dp)) { + OutlinedButton( + onClick = { frpPickerLauncher.launch("*/*") }, + modifier = Modifier.height(32.dp), + contentPadding = ExposedDropdownMenuDefaults.ItemContentPadding + ) { + Text("Import frpc", style = MaterialTheme.typography.labelMedium) + } + ClickableLink( + url = "https://github.com/fatedier/frp/releases", + linkText = "Download from GitHub" + ) + } + } + } + } + } + + Column(verticalArrangement = Arrangement.spacedBy(16.dp)) { + OutlinedTextField( + value = uiState.frpServerAddr, + onValueChange = { + viewModel.updateFrpConfig(uiState.frpEnabled, it, uiState.frpServerPort, uiState.frpToken, uiState.frpRemotePort, uiState.frpCustomDomain) + }, + label = { Text(stringResource(R.string.openai_api_server_frp_address)) }, + modifier = Modifier.fillMaxWidth(), + singleLine = true + ) + + Row(horizontalArrangement = Arrangement.spacedBy(16.dp)) { + OutlinedTextField( + value = uiState.frpServerPort.toString(), + onValueChange = { + it.toIntOrNull()?.let { port -> + viewModel.updateFrpConfig(uiState.frpEnabled, uiState.frpServerAddr, port, uiState.frpToken, uiState.frpRemotePort, uiState.frpCustomDomain) + } + }, + label = { Text(stringResource(R.string.openai_api_server_frp_port)) }, + modifier = Modifier.weight(1f), + keyboardOptions = KeyboardOptions(keyboardType = KeyboardType.Number), + singleLine = true + ) + OutlinedTextField( + value = uiState.frpRemotePort.toString(), + onValueChange = { + it.toIntOrNull()?.let { port -> + viewModel.updateFrpConfig(uiState.frpEnabled, uiState.frpServerAddr, uiState.frpServerPort, uiState.frpToken, port, uiState.frpCustomDomain) + } + }, + label = { Text(stringResource(R.string.openai_api_server_frp_remote_port)) }, + modifier = Modifier.weight(1f), + keyboardOptions = KeyboardOptions(keyboardType = KeyboardType.Number), + singleLine = true + ) + } + + OutlinedTextField( + value = uiState.frpCustomDomain, + onValueChange = { + viewModel.updateFrpConfig(uiState.frpEnabled, uiState.frpServerAddr, uiState.frpServerPort, uiState.frpToken, uiState.frpRemotePort, it) + }, + label = { Text(stringResource(R.string.openai_api_server_frp_custom_domain)) }, + modifier = Modifier.fillMaxWidth(), + singleLine = true + ) + + OutlinedTextField( + value = uiState.frpToken, + onValueChange = { + viewModel.updateFrpConfig(uiState.frpEnabled, uiState.frpServerAddr, uiState.frpServerPort, it, uiState.frpRemotePort, uiState.frpCustomDomain) + }, + label = { Text(stringResource(R.string.openai_api_server_frp_token)) }, + modifier = Modifier.fillMaxWidth(), + singleLine = true + ) + } + + // Connection Details + if (uiState.status == OpenAiApiServerStatus.RUNNING) { + Column(verticalArrangement = Arrangement.spacedBy(16.dp)) { + DetailItem( + label = stringResource(R.string.openai_api_server_endpoint), + value = uiState.endpoint, + onCopy = { copyToClipboard(context, "Endpoint", uiState.endpoint) } + ) + + DetailItem( + label = stringResource(R.string.openai_api_server_api_key), + value = uiState.apiKey, + onCopy = { copyToClipboard(context, "API Key", uiState.apiKey) }, + onRegenerate = { viewModel.regenerateApiKey() } + ) + + Row( + modifier = Modifier.fillMaxWidth(), + horizontalArrangement = Arrangement.SpaceBetween + ) { + Text( + stringResource(R.string.openai_api_server_requests), + style = MaterialTheme.typography.titleSmall + ) + Text( + uiState.requestCount.toString(), + style = MaterialTheme.typography.bodyMedium, + fontWeight = FontWeight.Bold + ) + } + } + } + + Spacer(modifier = Modifier.height(32.dp)) + } + } +} + +@Composable +private fun DetailItem( + label: String, + value: String, + onCopy: () -> Unit, + onRegenerate: (() -> Unit)? = null, +) { + Column(verticalArrangement = Arrangement.spacedBy(4.dp)) { + Text(label, style = MaterialTheme.typography.titleSmall) + Row( + modifier = Modifier.fillMaxWidth(), + verticalAlignment = Alignment.CenterVertically + ) { + Text( + value, + style = MaterialTheme.typography.bodyMedium.copy(fontFamily = FontFamily.Monospace), + modifier = Modifier.weight(1f) + ) + IconButton(onClick = onCopy) { + Icon(Icons.Rounded.ContentCopy, contentDescription = "Copy", modifier = Modifier.size(20.dp)) + } + if (onRegenerate != null) { + IconButton(onClick = onRegenerate) { + Icon(Icons.Rounded.Refresh, contentDescription = "Regenerate", modifier = Modifier.size(20.dp)) + } + } + } + } +} + +private fun copyToClipboard(context: Context, label: String, text: String) { + val clipboard = context.getSystemService(Context.CLIPBOARD_SERVICE) as ClipboardManager + val clip = ClipData.newPlainText(label, text) + clipboard.setPrimaryClip(clip) +} diff --git a/Android/src/app/src/main/java/com/google/ai/edge/gallery/server/OpenAiApiServerService.kt b/Android/src/app/src/main/java/com/google/ai/edge/gallery/server/OpenAiApiServerService.kt new file mode 100644 index 0000000000..989056475e --- /dev/null +++ b/Android/src/app/src/main/java/com/google/ai/edge/gallery/server/OpenAiApiServerService.kt @@ -0,0 +1,154 @@ +/* + * Copyright 2026 Google LLC + * + * 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.server + +import android.app.Notification +import android.app.NotificationChannel +import android.app.NotificationManager +import android.app.PendingIntent +import android.app.Service +import android.content.Intent +import android.content.pm.ServiceInfo +import android.net.wifi.WifiManager +import android.os.Build +import android.os.IBinder +import android.os.PowerManager +import androidx.core.app.NotificationCompat +import com.google.ai.edge.gallery.MainActivity +import com.google.ai.edge.gallery.R +import dagger.hilt.android.AndroidEntryPoint +import javax.inject.Inject + +@AndroidEntryPoint +class OpenAiApiServerService : Service() { + @Inject lateinit var apiServer: OpenAiApiServer + private var wakeLock: PowerManager.WakeLock? = null + private var wifiLock: WifiManager.WifiLock? = null + + override fun onCreate() { + super.onCreate() + createNotificationChannel() + } + + override fun onStartCommand(intent: Intent?, flags: Int, startId: Int): Int { + if (intent?.action == ACTION_STOP) { + apiServer.stop() + releaseLocks() + stopForeground(STOP_FOREGROUND_REMOVE) + stopSelf() + return START_NOT_STICKY + } + + val port = intent?.getIntExtra(EXTRA_PORT, 8080) ?: 8080 + val apiKey = intent?.getStringExtra(EXTRA_API_KEY).orEmpty() + val model = intent?.getStringExtra(EXTRA_MODEL).orEmpty() + val notification = createNotification(port) + + acquireLocks() + + if (Build.VERSION.SDK_INT >= Build.VERSION_CODES.UPSIDE_DOWN_CAKE) { + startForeground(NOTIFICATION_ID, notification, ServiceInfo.FOREGROUND_SERVICE_TYPE_SPECIAL_USE or ServiceInfo.FOREGROUND_SERVICE_TYPE_CONNECTED_DEVICE) + } else { + startForeground(NOTIFICATION_ID, notification) + } + apiServer.start(port = port, apiKey = apiKey, defaultModel = model) + if (apiServer.state.value.status == OpenAiApiServerStatus.ERROR) { + releaseLocks() + stopForeground(STOP_FOREGROUND_REMOVE) + stopSelf() + return START_NOT_STICKY + } + return START_NOT_STICKY + } + + override fun onDestroy() { + apiServer.stop() + releaseLocks() + super.onDestroy() + } + + private fun acquireLocks() { + if (wakeLock == null) { + val powerManager = getSystemService(POWER_SERVICE) as PowerManager + wakeLock = powerManager.newWakeLock(PowerManager.PARTIAL_WAKE_LOCK, "OpenAiApiServer::WakeLock").apply { + acquire() + } + } + if (wifiLock == null) { + val wifiManager = getSystemService(WIFI_SERVICE) as WifiManager + wifiLock = wifiManager.createWifiLock(WifiManager.WIFI_MODE_FULL_HIGH_PERF, "OpenAiApiServer::WifiLock").apply { + acquire() + } + } + } + + private fun releaseLocks() { + wakeLock?.let { if (it.isHeld) it.release() } + wifiLock?.let { if (it.isHeld) it.release() } + wakeLock = null + wifiLock = null + } + + override fun onBind(intent: Intent?): IBinder? = null + + private fun createNotificationChannel() { + val manager = getSystemService(NotificationManager::class.java) + manager.createNotificationChannel( + NotificationChannel( + CHANNEL_ID, + getString(R.string.openai_api_notification_channel), + NotificationManager.IMPORTANCE_LOW, + ) + ) + } + + private fun createNotification(port: Int): Notification { + val openIntent = + PendingIntent.getActivity( + this, + 0, + Intent(this, MainActivity::class.java), + PendingIntent.FLAG_IMMUTABLE or PendingIntent.FLAG_UPDATE_CURRENT, + ) + val stopIntent = + PendingIntent.getService( + this, + 1, + Intent(this, OpenAiApiServerService::class.java).setAction(ACTION_STOP), + PendingIntent.FLAG_IMMUTABLE or PendingIntent.FLAG_UPDATE_CURRENT, + ) + return NotificationCompat.Builder(this, CHANNEL_ID) + .setSmallIcon(R.mipmap.ic_launcher_monochrome) + .setContentTitle(getString(R.string.openai_api_notification_title)) + .setContentText(getString(R.string.openai_api_notification_text, port)) + .setContentIntent(openIntent) + .setOngoing(true) + .setCategory(NotificationCompat.CATEGORY_SERVICE) + .addAction(0, getString(R.string.stop), stopIntent) + .build() + } + + companion object { + const val ACTION_START = "com.google.ai.edge.gallery.server.START" + const val ACTION_STOP = "com.google.ai.edge.gallery.server.STOP" + const val EXTRA_PORT = "port" + const val EXTRA_API_KEY = "api_key" + const val EXTRA_MODEL = "model" + private const val CHANNEL_ID = "openai_api_server" + private const val NOTIFICATION_ID = 8402 + } +} diff --git a/Android/src/app/src/main/java/com/google/ai/edge/gallery/server/OpenAiApiServerViewModel.kt b/Android/src/app/src/main/java/com/google/ai/edge/gallery/server/OpenAiApiServerViewModel.kt new file mode 100644 index 0000000000..bd3eb38d37 --- /dev/null +++ b/Android/src/app/src/main/java/com/google/ai/edge/gallery/server/OpenAiApiServerViewModel.kt @@ -0,0 +1,235 @@ +/* + * Copyright 2026 Google LLC + * + * 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.server + +import android.content.Context +import android.content.Intent +import android.net.Uri +import android.util.Base64 +import androidx.core.content.ContextCompat +import androidx.lifecycle.ViewModel +import androidx.lifecycle.viewModelScope +import com.google.ai.edge.gallery.data.DataStoreRepository +import com.google.ai.edge.gallery.data.FrpPreferences +import com.google.ai.edge.gallery.data.Model +import com.google.ai.edge.gallery.data.OpenAiApiServerPreferences +import dagger.hilt.android.lifecycle.HiltViewModel +import dagger.hilt.android.qualifiers.ApplicationContext +import java.security.SecureRandom +import javax.inject.Inject +import kotlinx.coroutines.flow.MutableStateFlow +import kotlinx.coroutines.flow.asStateFlow +import kotlinx.coroutines.flow.collect +import kotlinx.coroutines.flow.update +import kotlinx.coroutines.launch + +data class OpenAiApiServerUiState( + val status: OpenAiApiServerStatus = OpenAiApiServerStatus.STOPPED, + val enabled: Boolean = false, + val port: Int = 8080, + val selectedModel: String = "", + val availableModels: List = emptyList(), + val apiKey: String = "", + val endpoint: String = "", + val requestCount: Long = 0, + val error: String = "", + // FRP state + val frpEnabled: Boolean = false, + val frpServerAddr: String = "", + val frpServerPort: Int = 7000, + val frpToken: String = "", + val frpRemotePort: Int = 8080, + val frpCustomDomain: String = "", + val frpRunning: Boolean = false, + val frpBinaryMissing: Boolean = false, +) + +@HiltViewModel +class OpenAiApiServerViewModel +@Inject +constructor( + @ApplicationContext private val context: Context, + private val dataStoreRepository: DataStoreRepository, + private val apiServer: OpenAiApiServer, + private val frpManager: FrpManager, +) : ViewModel() { + private val preferences = dataStoreRepository.readOpenAiApiServerPreferences() + private val frpPrefs = dataStoreRepository.readFrpPreferences() + private val _uiState = + MutableStateFlow( + OpenAiApiServerUiState( + enabled = preferences.enabled, + port = preferences.port, + selectedModel = preferences.modelName, + apiKey = loadOrCreateApiKey(), + frpEnabled = frpPrefs.enabled, + frpServerAddr = frpPrefs.serverAddress, + frpServerPort = frpPrefs.serverPort, + frpToken = frpPrefs.token, + frpRemotePort = frpPrefs.remotePort, + frpCustomDomain = frpPrefs.customDomain, + frpBinaryMissing = !frpManager.isBinaryAvailable(), + ) + ) + val uiState = _uiState.asStateFlow() + + init { + viewModelScope.launch { + apiServer.state.collect { serverState -> + _uiState.update { + it.copy( + status = serverState.status, + endpoint = serverState.endpoint, + requestCount = serverState.requestCount, + error = serverState.error, + ) + } + } + } + viewModelScope.launch { + frpManager.isRunning.collect { running -> + _uiState.update { it.copy(frpRunning = running) } + } + } + } + + fun updateModels(models: List, preferredModel: String = "") { + apiServer.updateModels(models) + val names = models.map { it.name }.distinct().sorted() + val current = _uiState.value.selectedModel + val selected = + when { + current in names -> current + preferredModel in names -> preferredModel + else -> names.firstOrNull().orEmpty() + } + _uiState.update { it.copy(availableModels = names, selectedModel = selected) } + if (_uiState.value.enabled && names.isNotEmpty() && apiServer.state.value.status == OpenAiApiServerStatus.STOPPED) { + start(port = _uiState.value.port, modelName = selected) + } + } + + fun start(port: Int, modelName: String) { + if (modelName.isBlank()) { + _uiState.update { it.copy(error = "Download an LLM before starting the server.") } + return + } + if (port !in 1024..65535) { + _uiState.update { it.copy(error = "Port must be between 1024 and 65535.") } + return + } + val prefs = OpenAiApiServerPreferences(enabled = true, port = port, modelName = modelName) + dataStoreRepository.saveOpenAiApiServerPreferences(prefs) + _uiState.update { it.copy(enabled = true, port = port, selectedModel = modelName, error = "") } + ContextCompat.startForegroundService( + context, + Intent(context, OpenAiApiServerService::class.java) + .setAction(OpenAiApiServerService.ACTION_START) + .putExtra(OpenAiApiServerService.EXTRA_PORT, port) + .putExtra(OpenAiApiServerService.EXTRA_API_KEY, _uiState.value.apiKey) + .putExtra(OpenAiApiServerService.EXTRA_MODEL, modelName), + ) + if (_uiState.value.frpEnabled) { + frpManager.start( + _uiState.value.frpServerAddr, + _uiState.value.frpServerPort, + _uiState.value.frpToken, + port, + _uiState.value.frpRemotePort, + _uiState.value.frpCustomDomain + ) + } + } + + fun stop() { + dataStoreRepository.saveOpenAiApiServerPreferences( + OpenAiApiServerPreferences(enabled = false, port = _uiState.value.port, modelName = _uiState.value.selectedModel) + ) + _uiState.update { it.copy(enabled = false) } + context.startService(Intent(context, OpenAiApiServerService::class.java).setAction(OpenAiApiServerService.ACTION_STOP)) + frpManager.stop() + } + + fun updateFrpConfig(enabled: Boolean, serverAddr: String, serverPort: Int, token: String, remotePort: Int, customDomain: String) { + val prefs = FrpPreferences( + enabled = enabled, + serverAddress = serverAddr, + serverPort = serverPort, + token = token, + remotePort = remotePort, + customDomain = customDomain + ) + dataStoreRepository.saveFrpPreferences(prefs) + _uiState.update { + it.copy( + frpEnabled = enabled, + frpServerAddr = serverAddr, + frpServerPort = serverPort, + frpToken = token, + frpRemotePort = remotePort, + frpCustomDomain = customDomain, + frpBinaryMissing = !frpManager.isBinaryAvailable() + ) + } + + if (enabled && _uiState.value.enabled && _uiState.value.status == OpenAiApiServerStatus.RUNNING) { + frpManager.start(serverAddr, serverPort, token, _uiState.value.port, remotePort, customDomain) + } else { + frpManager.stop() + } + } + + fun importFrpBinary(uri: Uri) { + if (frpManager.importBinary(uri)) { + _uiState.update { it.copy(frpBinaryMissing = false) } + } + } + + fun deleteFrpBinary() { + if (frpManager.deleteBinary()) { + _uiState.update { it.copy(frpBinaryMissing = true) } + } + } + + fun reportPermissionDenied() { + _uiState.update { it.copy(error = "Local network permission is required to accept LAN connections.") } + } + + fun regenerateApiKey() { + val apiKey = generateApiKey() + dataStoreRepository.saveSecret(API_KEY_SECRET, apiKey) + _uiState.update { it.copy(apiKey = apiKey) } + if (_uiState.value.status == OpenAiApiServerStatus.RUNNING) { + start(_uiState.value.port, _uiState.value.selectedModel) + } + } + + private fun loadOrCreateApiKey(): String { + return dataStoreRepository.readSecret(API_KEY_SECRET)?.takeIf { it.isNotBlank() } + ?: generateApiKey().also { dataStoreRepository.saveSecret(API_KEY_SECRET, it) } + } + + private fun generateApiKey(): String { + val bytes = ByteArray(24) + SecureRandom().nextBytes(bytes) + return "sk-local-${Base64.encodeToString(bytes, Base64.NO_WRAP or Base64.URL_SAFE or Base64.NO_PADDING)}" + } + + companion object { + private const val API_KEY_SECRET = "openai_api_server_key" + } +} diff --git a/Android/src/app/src/main/java/com/google/ai/edge/gallery/ui/home/HomeScreen.kt b/Android/src/app/src/main/java/com/google/ai/edge/gallery/ui/home/HomeScreen.kt index d3d30e5335..e139f12309 100644 --- a/Android/src/app/src/main/java/com/google/ai/edge/gallery/ui/home/HomeScreen.kt +++ b/Android/src/app/src/main/java/com/google/ai/edge/gallery/ui/home/HomeScreen.kt @@ -57,6 +57,7 @@ import androidx.compose.foundation.text.TextAutoSize import androidx.compose.foundation.verticalScroll import androidx.compose.material.icons.Icons import androidx.compose.material.icons.automirrored.rounded.ListAlt +import androidx.compose.material.icons.rounded.Dns import androidx.compose.material.icons.rounded.Error import androidx.compose.material.icons.rounded.Flag import androidx.compose.material.icons.rounded.Notifications @@ -165,6 +166,7 @@ fun HomeScreen( navigateToTaskScreen: (Task) -> Unit, onModelsClicked: () -> Unit, onNotificationsClicked: () -> Unit, + onOpenAiServerClicked: () -> Unit, enableAnimation: Boolean, modifier: Modifier = Modifier, gm4: Boolean = false, @@ -330,6 +332,28 @@ fun HomeScreen( } Spacer(modifier = Modifier.height(16.dp)) Row(modifier = Modifier.fillMaxWidth()) { + SquareDrawerItem( + label = stringResource(R.string.openai_api_server_title), + description = "Configure OpenAI compatible API server", + icon = Icons.Rounded.Dns, + onClick = { + scope.launch { drawerState.close() } + scope.launch { + delay(50) + onOpenAiServerClicked() + } + }, + modifier = Modifier.weight(1f), + iconBrush = + linearGradient( + colors = + listOf( + MaterialTheme.customColors.taskBgGradientColors[0][0], + MaterialTheme.customColors.taskBgGradientColors[0][1], + ) + ), + ) + Spacer(modifier = Modifier.weight(1f)) } } } diff --git a/Android/src/app/src/main/java/com/google/ai/edge/gallery/ui/navigation/GalleryNavGraph.kt b/Android/src/app/src/main/java/com/google/ai/edge/gallery/ui/navigation/GalleryNavGraph.kt index 6d330bb078..fda44a02fc 100644 --- a/Android/src/app/src/main/java/com/google/ai/edge/gallery/ui/navigation/GalleryNavGraph.kt +++ b/Android/src/app/src/main/java/com/google/ai/edge/gallery/ui/navigation/GalleryNavGraph.kt @@ -88,6 +88,7 @@ import com.google.ai.edge.gallery.ui.modelmanager.ModelInitializationStatusType import com.google.ai.edge.gallery.ui.modelmanager.ModelManager import com.google.ai.edge.gallery.ui.modelmanager.ModelManagerViewModel import com.google.ai.edge.gallery.ui.notifications.NotificationsScreen +import com.google.ai.edge.gallery.server.OpenAiApiServerScreen import kotlinx.coroutines.Dispatchers import kotlinx.coroutines.delay import kotlinx.coroutines.launch @@ -99,6 +100,7 @@ private const val ROUTE_MODEL = "route_model" private const val ROUTE_BENCHMARK = "benchmark" private const val ROUTE_MODEL_MANAGER = "model_manager" private const val ROUTE_NOTIFICATIONS = "notifications" +private const val ROUTE_OPENAI_SERVER = "openai_server" private const val ENTER_ANIMATION_DURATION_MS = 500 private val ENTER_ANIMATION_EASING = EaseOutExpo private const val ENTER_ANIMATION_DELAY_MS = 100 @@ -212,6 +214,7 @@ fun GalleryNavHost( }, onModelsClicked = { navController.navigate(ROUTE_MODEL_MANAGER) }, onNotificationsClicked = { navController.navigate(ROUTE_NOTIFICATIONS) }, + onOpenAiServerClicked = { navController.navigate(ROUTE_OPENAI_SERVER) }, gm4 = true, ) } @@ -432,6 +435,18 @@ fun GalleryNavHost( NotificationsScreen(navigateUp = { navController.navigateUp() }) } + // OpenAI Server page. + composable( + route = ROUTE_OPENAI_SERVER, + enterTransition = { slideUpEnter() }, + exitTransition = { slideDownExit() }, + ) { + OpenAiApiServerScreen( + modelManagerViewModel = modelManagerViewModel, + onBackClicked = { navController.navigateUp() } + ) + } + // Benchmark creation page. composable( route = "$ROUTE_BENCHMARK/{modelName}", diff --git a/Android/src/app/src/main/proto/settings.proto b/Android/src/app/src/main/proto/settings.proto index 8844c6d4c2..1f8d0909db 100644 --- a/Android/src/app/src/main/proto/settings.proto +++ b/Android/src/app/src/main/proto/settings.proto @@ -80,6 +80,19 @@ message Settings { repeated string viewed_promo_id = 10; bool disable_firebase_analytics = 11; + + // Local OpenAI-compatible API server preferences. The API key is kept in UserData.secrets. + bool openai_api_server_enabled = 12; + int32 openai_api_server_port = 13; + string openai_api_server_model = 14; + + // FRP configuration + bool frp_enabled = 15; + string frp_server_address = 16; + int32 frp_server_port = 17; + string frp_token = 18; + int32 frp_remote_port = 19; + string frp_custom_domain = 20; } message McpAuth { diff --git a/Android/src/app/src/main/res/values/strings.xml b/Android/src/app/src/main/res/values/strings.xml index 15a6bc13f1..5f65474d0c 100644 --- a/Android/src/app/src/main/res/values/strings.xml +++ b/Android/src/app/src/main/res/values/strings.xml @@ -229,6 +229,32 @@ Failed to save to photo album Copied to clipboard + + OpenAI API Server + OpenAI API Server is running + Port: %1$d + Stop + OpenAI API Server + Turn this device into an OpenAI-compatible API server. Any downloaded model can be used as an endpoint. + Server Enabled + Server Port + Default Model + API Key + Endpoint + Total Requests + Regenerate Key + Copy Endpoint + Copy API Key + Internet Access (FRP) + Enable FRP Tunnel + Tunnel Running + Tunnel Stopped + FRP Server Address + FRP Server Port + Remote Port + FRP Token + Custom Domain (Optional) + Parameter Parameters diff --git a/Android/src/gradle/gradle-daemon-jvm.properties b/Android/src/gradle/gradle-daemon-jvm.properties new file mode 100644 index 0000000000..6c1139ec06 --- /dev/null +++ b/Android/src/gradle/gradle-daemon-jvm.properties @@ -0,0 +1,12 @@ +#This file is generated by updateDaemonJvm +toolchainUrl.FREE_BSD.AARCH64=https\://api.foojay.io/disco/v3.0/ids/ec7520a1e057cd116f9544c42142a16b/redirect +toolchainUrl.FREE_BSD.X86_64=https\://api.foojay.io/disco/v3.0/ids/4c4f879899012ff0a8b2e2117df03b0e/redirect +toolchainUrl.LINUX.AARCH64=https\://api.foojay.io/disco/v3.0/ids/ec7520a1e057cd116f9544c42142a16b/redirect +toolchainUrl.LINUX.X86_64=https\://api.foojay.io/disco/v3.0/ids/4c4f879899012ff0a8b2e2117df03b0e/redirect +toolchainUrl.MAC_OS.AARCH64=https\://api.foojay.io/disco/v3.0/ids/73bcfb608d1fde9fb62e462f834a3299/redirect +toolchainUrl.MAC_OS.X86_64=https\://api.foojay.io/disco/v3.0/ids/846ee0d876d26a26f37aa1ce8de73224/redirect +toolchainUrl.UNIX.AARCH64=https\://api.foojay.io/disco/v3.0/ids/ec7520a1e057cd116f9544c42142a16b/redirect +toolchainUrl.UNIX.X86_64=https\://api.foojay.io/disco/v3.0/ids/4c4f879899012ff0a8b2e2117df03b0e/redirect +toolchainUrl.WINDOWS.AARCH64=https\://api.foojay.io/disco/v3.0/ids/9482ddec596298c84656d31d16652665/redirect +toolchainUrl.WINDOWS.X86_64=https\://api.foojay.io/disco/v3.0/ids/39701d92e1756bb2f141eb67cd4c660e/redirect +toolchainVersion=21 diff --git a/Android/src/gradle/libs.versions.toml b/Android/src/gradle/libs.versions.toml index 3a6e42f6f7..af0cbb8551 100644 --- a/Android/src/gradle/libs.versions.toml +++ b/Android/src/gradle/libs.versions.toml @@ -98,6 +98,11 @@ mlkit-genai-prompt = { group = "com.google.mlkit", name = "genai-prompt", versio mcp-kotlin-sdk = { group = "io.modelcontextprotocol", name = "kotlin-sdk", version.ref = "mcp" } ktor-client-android = { group = "io.ktor", name = "ktor-client-android", version.ref = "ktor" } ktor-client-core = { group = "io.ktor", name = "ktor-client-core", version.ref = "ktor" } +ktor-server-core = { group = "io.ktor", name = "ktor-server-core", version.ref = "ktor" } +ktor-server-cio = { group = "io.ktor", name = "ktor-server-cio", version.ref = "ktor" } +ktor-server-content-negotiation = { group = "io.ktor", name = "ktor-server-content-negotiation", version.ref = "ktor" } +ktor-server-cors = { group = "io.ktor", name = "ktor-server-cors", version.ref = "ktor" } +ktor-serialization-kotlinx-json = { group = "io.ktor", name = "ktor-serialization-kotlinx-json", version.ref = "ktor" } [plugins] android-application = { id = "com.android.application", version.ref = "agp" } @@ -108,4 +113,4 @@ protobuf = {id = "com.google.protobuf", version.ref = "protobuf"} hilt-application = { id = "com.google.dagger.hilt.android", version.ref = "hilt" } oss-licenses = {id = "com.google.android.gms.oss-licenses-plugin", version.ref = "ossLicenses"} google-services = { id = "com.google.gms.google-services", version.ref = "googleService" } -ksp = { id = "com.google.devtools.ksp", version.ref = "ksp" } \ No newline at end of file +ksp = { id = "com.google.devtools.ksp", version.ref = "ksp" }