Fix streaming, add thread safety, settings UI, and request indicator
Some checks are pending
Release / Build and Release APK (push) Waiting to run
Some checks are pending
Release / Build and Release APK (push) Waiting to run
Fixes: - Streaming was bypassed for all requests (system prompt auto-inject made messages.size always > 1). Stream=true now routes directly to generateStreaming before multi-turn check - LiteRTModel: add Mutex to serialize Engine calls (not thread-safe) New features: - Settings card: temperature slider, max tokens, agent system prompt toggle (all saved to SharedPreferences, applied on server start) - Active request indicator: "⚡ Processing request…" shown in status card while inference is running - onActiveRequest callback from AIApiServer → ApiServerService → MainActivity Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com>
This commit is contained in:
@@ -7,6 +7,8 @@ import com.google.ai.edge.litertlm.Engine
|
||||
import com.google.ai.edge.litertlm.EngineConfig
|
||||
import kotlinx.coroutines.Dispatchers
|
||||
import kotlinx.coroutines.flow.catch
|
||||
import kotlinx.coroutines.sync.Mutex
|
||||
import kotlinx.coroutines.sync.withLock
|
||||
import kotlinx.coroutines.withContext
|
||||
import java.io.File
|
||||
|
||||
@@ -29,42 +31,49 @@ class LiteRTModel private constructor(
|
||||
override var isReady: Boolean = true
|
||||
private set
|
||||
|
||||
// LiteRT Engine is not thread-safe — serialize all inference calls
|
||||
private val mutex = Mutex()
|
||||
|
||||
override suspend fun generate(
|
||||
prompt: String,
|
||||
maxTokens: Int,
|
||||
temperature: Float
|
||||
): String = withContext(Dispatchers.Default) {
|
||||
val conversation = engine.createConversation()
|
||||
try {
|
||||
conversation.sendMessage(prompt).toString()
|
||||
} catch (e: Exception) {
|
||||
Log.e(TAG, "LiteRT inference error", e)
|
||||
throw OnDeviceModel.InferenceException("Generation failed: ${e.message}", e)
|
||||
} finally {
|
||||
conversation.close()
|
||||
): String = mutex.withLock {
|
||||
withContext(Dispatchers.Default) {
|
||||
val conversation = engine.createConversation()
|
||||
try {
|
||||
conversation.sendMessage(prompt).toString()
|
||||
} catch (e: Exception) {
|
||||
Log.e(TAG, "LiteRT inference error", e)
|
||||
throw OnDeviceModel.InferenceException("Generation failed: ${e.message}", e)
|
||||
} finally {
|
||||
conversation.close()
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
override suspend fun generateStreaming(
|
||||
prompt: String,
|
||||
onToken: (String) -> Unit
|
||||
): String = withContext(Dispatchers.Default) {
|
||||
val conversation = engine.createConversation()
|
||||
val sb = StringBuilder()
|
||||
try {
|
||||
conversation.sendMessageAsync(prompt)
|
||||
.catch { e ->
|
||||
throw OnDeviceModel.InferenceException("Streaming failed: ${e.message}", e)
|
||||
}
|
||||
.collect { message ->
|
||||
val token = message.toString()
|
||||
sb.append(token)
|
||||
onToken(token)
|
||||
}
|
||||
} finally {
|
||||
conversation.close()
|
||||
): String = mutex.withLock {
|
||||
withContext(Dispatchers.Default) {
|
||||
val conversation = engine.createConversation()
|
||||
val sb = StringBuilder()
|
||||
try {
|
||||
conversation.sendMessageAsync(prompt)
|
||||
.catch { e ->
|
||||
throw OnDeviceModel.InferenceException("Streaming failed: ${e.message}", e)
|
||||
}
|
||||
.collect { message ->
|
||||
val token = message.toString()
|
||||
sb.append(token)
|
||||
onToken(token)
|
||||
}
|
||||
} finally {
|
||||
conversation.close()
|
||||
}
|
||||
sb.toString()
|
||||
}
|
||||
sb.toString()
|
||||
}
|
||||
|
||||
override fun close() {
|
||||
|
||||
@@ -27,9 +27,16 @@ import java.util.concurrent.atomic.AtomicLong
|
||||
* -H "Content-Type: application/json" \
|
||||
* -d '{"messages":[{"role":"user","content":"Hello!"}]}'
|
||||
*/
|
||||
data class ServerConfig(
|
||||
val defaultTemperature: Float = 0.7f,
|
||||
val defaultMaxTokens: Int = 1024,
|
||||
val autoSystemPrompt: Boolean = true
|
||||
)
|
||||
|
||||
class AIApiServer(
|
||||
port: Int,
|
||||
private val model: OnDeviceModel
|
||||
private val model: OnDeviceModel,
|
||||
private val config: ServerConfig = ServerConfig()
|
||||
) : NanoHTTPD(port) {
|
||||
|
||||
private val gson = Gson()
|
||||
@@ -37,6 +44,7 @@ class AIApiServer(
|
||||
val requestCount = AtomicLong(0)
|
||||
|
||||
var onRequestLogged: ((String) -> Unit)? = null
|
||||
var onActiveRequest: ((Boolean) -> Unit)? = null
|
||||
|
||||
override fun serve(session: IHTTPSession): Response {
|
||||
val method = session.method
|
||||
@@ -113,31 +121,48 @@ class AIApiServer(
|
||||
return errorResponse(400, "messages array is required and must not be empty")
|
||||
}
|
||||
|
||||
// Auto-inject agent system prompt if the conversation has no system message.
|
||||
// Auto-inject default tools if the request provides none.
|
||||
// This makes the server zero-config as a coding agent for any OpenAI-compatible client.
|
||||
val messages = if (raw.messages.none { it.role == "system" }) {
|
||||
// Auto-inject agent system prompt if enabled and no system message present
|
||||
val messages = if (config.autoSystemPrompt && raw.messages.none { it.role == "system" }) {
|
||||
listOf(Message(role = "system", content = AgentConfig.SYSTEM_PROMPT)) + raw.messages
|
||||
} else {
|
||||
raw.messages
|
||||
}
|
||||
val request = raw.copy(
|
||||
messages = messages,
|
||||
tools = raw.tools.takeUnless { it.isNullOrEmpty() } ?: AgentConfig.DEFAULT_TOOLS
|
||||
tools = raw.tools.takeUnless { it.isNullOrEmpty() }
|
||||
?: if (config.autoSystemPrompt) AgentConfig.DEFAULT_TOOLS else null,
|
||||
temperature = if (raw.temperature == 0.7f) config.defaultTemperature else raw.temperature,
|
||||
max_tokens = if (raw.max_tokens == 8192) config.defaultMaxTokens else raw.max_tokens
|
||||
)
|
||||
|
||||
val id = "chatcmpl-${UUID.randomUUID().toString().take(8)}"
|
||||
val hasTools = !request.tools.isNullOrEmpty()
|
||||
|
||||
log("Chat: ${request.messages.size} messages, tools=${request.tools?.size ?: 0}, stream=${request.stream}")
|
||||
log("Chat: ${request.messages.size} messages, tools=${request.tools?.size ?: 0}, stream=${request.stream}, temp=${request.temperature}")
|
||||
|
||||
// ── Streaming — always uses flat prompt + generateStreaming ────────────
|
||||
if (request.stream) {
|
||||
val prompt = buildFlatPrompt(request.messages)
|
||||
onActiveRequest?.invoke(true)
|
||||
return try {
|
||||
handleStreamingResponse(id, prompt, request)
|
||||
} finally {
|
||||
onActiveRequest?.invoke(false)
|
||||
}
|
||||
}
|
||||
|
||||
// ── Tool calling / multi-turn chat ─────────────────────────────────────
|
||||
if (hasTools || request.messages.size > 1 || request.messages.any { it.role == "system" }) {
|
||||
val convMessages = request.messages.map { it.toConvMessage() }
|
||||
val toolDefs = request.tools?.map { it.toToolDef() } ?: emptyList()
|
||||
|
||||
val result = runBlocking {
|
||||
model.chat(convMessages, toolDefs, request.max_tokens, request.temperature)
|
||||
onActiveRequest?.invoke(true)
|
||||
val result = try {
|
||||
runBlocking {
|
||||
model.chat(convMessages, toolDefs, request.max_tokens, request.temperature)
|
||||
}
|
||||
} finally {
|
||||
onActiveRequest?.invoke(false)
|
||||
}
|
||||
|
||||
if (result.toolCalls != null) {
|
||||
|
||||
@@ -34,6 +34,7 @@ class ApiServerService : Service() {
|
||||
|
||||
var onStatusChanged: ((ServerState) -> Unit)? = null
|
||||
var onLog: ((String) -> Unit)? = null
|
||||
var onActiveRequest: ((Boolean) -> Unit)? = null
|
||||
|
||||
val isRunning: Boolean get() = server != null
|
||||
val requestCount: Long get() = server?.requestCount?.get() ?: 0
|
||||
@@ -70,8 +71,15 @@ class ApiServerService : Service() {
|
||||
|
||||
// Start the HTTP server
|
||||
notifyLog("Starting API server on port $port...")
|
||||
val apiServer = AIApiServer(port, model!!)
|
||||
val prefs = getSharedPreferences("pixel10_prefs", MODE_PRIVATE)
|
||||
val serverConfig = ServerConfig(
|
||||
defaultTemperature = prefs.getFloat("temperature", 0.7f),
|
||||
defaultMaxTokens = prefs.getInt("max_tokens", 1024),
|
||||
autoSystemPrompt = prefs.getBoolean("auto_system_prompt", true)
|
||||
)
|
||||
val apiServer = AIApiServer(port, model!!, serverConfig)
|
||||
apiServer.onRequestLogged = { msg -> notifyLog(msg) }
|
||||
apiServer.onActiveRequest = { active -> onActiveRequest?.invoke(active) }
|
||||
apiServer.start()
|
||||
server = apiServer
|
||||
|
||||
|
||||
@@ -13,6 +13,7 @@ import android.os.Build
|
||||
import android.os.Bundle
|
||||
import android.os.IBinder
|
||||
import android.view.View
|
||||
import android.widget.SeekBar
|
||||
import androidx.activity.result.contract.ActivityResultContracts
|
||||
import androidx.appcompat.app.AppCompatActivity
|
||||
import androidx.core.content.ContextCompat
|
||||
@@ -53,6 +54,16 @@ class MainActivity : AppCompatActivity() {
|
||||
service?.onLog = { message ->
|
||||
runOnUiThread { appendLog(message) }
|
||||
}
|
||||
service?.onActiveRequest = { active ->
|
||||
runOnUiThread {
|
||||
if (active) {
|
||||
binding.tvActiveRequest.text = "⚡ Processing request…"
|
||||
binding.tvActiveRequest.visibility = View.VISIBLE
|
||||
} else {
|
||||
binding.tvActiveRequest.visibility = View.GONE
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
if (service?.isRunning == true) {
|
||||
updateStatus(ApiServerService.ServerState.RUNNING)
|
||||
@@ -73,8 +84,21 @@ class MainActivity : AppCompatActivity() {
|
||||
prefs = getSharedPreferences("pixel10_prefs", MODE_PRIVATE)
|
||||
requestNotificationPermission()
|
||||
|
||||
// Restore saved HF token
|
||||
// Restore saved settings
|
||||
binding.etHfToken.setText(prefs.getString("hf_token", ""))
|
||||
val savedTemp = (prefs.getFloat("temperature", 0.7f) * 100).toInt()
|
||||
binding.seekTemperature.progress = savedTemp
|
||||
binding.tvTemperatureValue.text = "%.1f".format(savedTemp / 100f)
|
||||
binding.etMaxTokens.setText(prefs.getInt("max_tokens", 1024).toString())
|
||||
binding.switchSystemPrompt.isChecked = prefs.getBoolean("auto_system_prompt", true)
|
||||
|
||||
binding.seekTemperature.setOnSeekBarChangeListener(object : SeekBar.OnSeekBarChangeListener {
|
||||
override fun onProgressChanged(seekBar: SeekBar, progress: Int, fromUser: Boolean) {
|
||||
binding.tvTemperatureValue.text = "%.1f".format(progress / 100f)
|
||||
}
|
||||
override fun onStartTrackingTouch(seekBar: SeekBar) {}
|
||||
override fun onStopTrackingTouch(seekBar: SeekBar) {}
|
||||
})
|
||||
|
||||
binding.btnToggle.setOnClickListener {
|
||||
if (service?.isRunning == true) stopServer() else startServer()
|
||||
@@ -201,7 +225,16 @@ class MainActivity : AppCompatActivity() {
|
||||
}
|
||||
}
|
||||
|
||||
private fun saveSettings() {
|
||||
prefs.edit()
|
||||
.putFloat("temperature", binding.seekTemperature.progress / 100f)
|
||||
.putInt("max_tokens", binding.etMaxTokens.text.toString().toIntOrNull() ?: 1024)
|
||||
.putBoolean("auto_system_prompt", binding.switchSystemPrompt.isChecked)
|
||||
.apply()
|
||||
}
|
||||
|
||||
private fun startServer() {
|
||||
saveSettings()
|
||||
val port = binding.etPort.text.toString().toIntOrNull() ?: 8080
|
||||
val intent = Intent(this, ApiServerService::class.java).apply {
|
||||
action = ApiServerService.ACTION_START
|
||||
@@ -227,6 +260,9 @@ class MainActivity : AppCompatActivity() {
|
||||
val canEdit = state == ApiServerService.ServerState.STOPPED ||
|
||||
state == ApiServerService.ServerState.ERROR
|
||||
binding.etPort.isEnabled = canEdit
|
||||
binding.seekTemperature.isEnabled = canEdit
|
||||
binding.etMaxTokens.isEnabled = canEdit
|
||||
binding.switchSystemPrompt.isEnabled = canEdit
|
||||
|
||||
when (state) {
|
||||
ApiServerService.ServerState.STOPPED -> {
|
||||
|
||||
@@ -102,6 +102,16 @@
|
||||
android:textColor="@color/log_text"
|
||||
android:textSize="13sp"
|
||||
android:layout_marginTop="2dp" />
|
||||
|
||||
<TextView
|
||||
android:id="@+id/tvActiveRequest"
|
||||
android:layout_width="wrap_content"
|
||||
android:layout_height="wrap_content"
|
||||
android:textColor="@color/primary"
|
||||
android:textSize="13sp"
|
||||
android:textStyle="bold"
|
||||
android:layout_marginTop="2dp"
|
||||
android:visibility="gone" />
|
||||
</LinearLayout>
|
||||
</com.google.android.material.card.MaterialCardView>
|
||||
|
||||
@@ -176,38 +186,143 @@
|
||||
</LinearLayout>
|
||||
</com.google.android.material.card.MaterialCardView>
|
||||
|
||||
<!-- Port Config -->
|
||||
<LinearLayout
|
||||
<!-- Settings Card -->
|
||||
<com.google.android.material.card.MaterialCardView
|
||||
android:id="@+id/layoutPort"
|
||||
android:layout_width="0dp"
|
||||
android:layout_height="wrap_content"
|
||||
android:orientation="horizontal"
|
||||
android:gravity="center_vertical"
|
||||
android:layout_marginTop="16dp"
|
||||
android:layout_marginTop="12dp"
|
||||
app:cardBackgroundColor="@color/surface_variant"
|
||||
app:cardCornerRadius="16dp"
|
||||
app:cardElevation="0dp"
|
||||
app:strokeWidth="0dp"
|
||||
app:layout_constraintTop_toBottomOf="@id/cardModel"
|
||||
app:layout_constraintStart_toStartOf="parent"
|
||||
app:layout_constraintEnd_toEndOf="parent">
|
||||
|
||||
<TextView
|
||||
android:layout_width="wrap_content"
|
||||
<LinearLayout
|
||||
android:layout_width="match_parent"
|
||||
android:layout_height="wrap_content"
|
||||
android:text="@string/port_label"
|
||||
android:textColor="@color/on_surface"
|
||||
android:textSize="16sp" />
|
||||
android:orientation="vertical"
|
||||
android:padding="16dp">
|
||||
|
||||
<com.google.android.material.textfield.TextInputEditText
|
||||
android:id="@+id/etPort"
|
||||
android:layout_width="100dp"
|
||||
android:layout_height="48dp"
|
||||
android:layout_marginStart="12dp"
|
||||
android:text="@string/port_default"
|
||||
android:inputType="number"
|
||||
android:textColor="@color/on_surface"
|
||||
android:backgroundTint="@color/primary"
|
||||
android:fontFamily="monospace"
|
||||
android:textSize="16sp"
|
||||
android:gravity="center" />
|
||||
</LinearLayout>
|
||||
<!-- Port row -->
|
||||
<LinearLayout
|
||||
android:layout_width="match_parent"
|
||||
android:layout_height="wrap_content"
|
||||
android:orientation="horizontal"
|
||||
android:gravity="center_vertical">
|
||||
|
||||
<TextView
|
||||
android:layout_width="0dp"
|
||||
android:layout_height="wrap_content"
|
||||
android:layout_weight="1"
|
||||
android:text="@string/port_label"
|
||||
android:textColor="@color/on_surface"
|
||||
android:textSize="14sp" />
|
||||
|
||||
<com.google.android.material.textfield.TextInputEditText
|
||||
android:id="@+id/etPort"
|
||||
android:layout_width="80dp"
|
||||
android:layout_height="40dp"
|
||||
android:text="@string/port_default"
|
||||
android:inputType="number"
|
||||
android:textColor="@color/on_surface"
|
||||
android:backgroundTint="@color/primary"
|
||||
android:fontFamily="monospace"
|
||||
android:textSize="14sp"
|
||||
android:gravity="center" />
|
||||
</LinearLayout>
|
||||
|
||||
<!-- Temperature row -->
|
||||
<LinearLayout
|
||||
android:layout_width="match_parent"
|
||||
android:layout_height="wrap_content"
|
||||
android:orientation="horizontal"
|
||||
android:gravity="center_vertical"
|
||||
android:layout_marginTop="12dp">
|
||||
|
||||
<TextView
|
||||
android:layout_width="0dp"
|
||||
android:layout_height="wrap_content"
|
||||
android:layout_weight="1"
|
||||
android:text="@string/setting_temperature"
|
||||
android:textColor="@color/on_surface"
|
||||
android:textSize="14sp" />
|
||||
|
||||
<TextView
|
||||
android:id="@+id/tvTemperatureValue"
|
||||
android:layout_width="36dp"
|
||||
android:layout_height="wrap_content"
|
||||
android:text="0.7"
|
||||
android:textColor="@color/log_text"
|
||||
android:textSize="13sp"
|
||||
android:fontFamily="monospace"
|
||||
android:gravity="end" />
|
||||
|
||||
<SeekBar
|
||||
android:id="@+id/seekTemperature"
|
||||
android:layout_width="120dp"
|
||||
android:layout_height="wrap_content"
|
||||
android:layout_marginStart="8dp"
|
||||
android:max="100"
|
||||
android:progress="70" />
|
||||
</LinearLayout>
|
||||
|
||||
<!-- Max tokens row -->
|
||||
<LinearLayout
|
||||
android:layout_width="match_parent"
|
||||
android:layout_height="wrap_content"
|
||||
android:orientation="horizontal"
|
||||
android:gravity="center_vertical"
|
||||
android:layout_marginTop="8dp">
|
||||
|
||||
<TextView
|
||||
android:layout_width="0dp"
|
||||
android:layout_height="wrap_content"
|
||||
android:layout_weight="1"
|
||||
android:text="@string/setting_max_tokens"
|
||||
android:textColor="@color/on_surface"
|
||||
android:textSize="14sp" />
|
||||
|
||||
<com.google.android.material.textfield.TextInputEditText
|
||||
android:id="@+id/etMaxTokens"
|
||||
android:layout_width="80dp"
|
||||
android:layout_height="40dp"
|
||||
android:text="1024"
|
||||
android:inputType="number"
|
||||
android:textColor="@color/on_surface"
|
||||
android:backgroundTint="@color/primary"
|
||||
android:fontFamily="monospace"
|
||||
android:textSize="14sp"
|
||||
android:gravity="center" />
|
||||
</LinearLayout>
|
||||
|
||||
<!-- System prompt toggle -->
|
||||
<LinearLayout
|
||||
android:layout_width="match_parent"
|
||||
android:layout_height="wrap_content"
|
||||
android:orientation="horizontal"
|
||||
android:gravity="center_vertical"
|
||||
android:layout_marginTop="8dp">
|
||||
|
||||
<TextView
|
||||
android:layout_width="0dp"
|
||||
android:layout_height="wrap_content"
|
||||
android:layout_weight="1"
|
||||
android:text="@string/setting_auto_system_prompt"
|
||||
android:textColor="@color/on_surface"
|
||||
android:textSize="14sp" />
|
||||
|
||||
<com.google.android.material.switchmaterial.SwitchMaterial
|
||||
android:id="@+id/switchSystemPrompt"
|
||||
android:layout_width="wrap_content"
|
||||
android:layout_height="wrap_content"
|
||||
android:checked="true" />
|
||||
</LinearLayout>
|
||||
|
||||
</LinearLayout>
|
||||
</com.google.android.material.card.MaterialCardView>
|
||||
|
||||
<!-- Start/Stop Button -->
|
||||
<com.google.android.material.button.MaterialButton
|
||||
|
||||
@@ -25,7 +25,10 @@
|
||||
<string name="model_not_downloaded">No local model. Enter HuggingFace token and download.</string>
|
||||
<string name="btn_start">Start Server</string>
|
||||
<string name="btn_stop">Stop Server</string>
|
||||
<string name="port_label">Port:</string>
|
||||
<string name="port_label">Port</string>
|
||||
<string name="setting_temperature">Temperature</string>
|
||||
<string name="setting_max_tokens">Max tokens</string>
|
||||
<string name="setting_auto_system_prompt">Agent system prompt</string>
|
||||
<string name="port_default">8080</string>
|
||||
<string name="requests_served">Requests served: %d</string>
|
||||
<string name="request_log_label">Request Log</string>
|
||||
|
||||
Reference in New Issue
Block a user