Skip to content

Commit 939b191

Browse files
jatezzzclaude
andauthored
ADFA-4388 | embedding model crash (#1429)
* fix(ADFA-4388): Prevent crash when loading embedding models for chat Added multi-layer protection to detect and reject embedding models: **Native Layer (C++):** - Check pooling_type in new_context() - reject if not LLAMA_POOLING_TYPE_NONE - Added get_pooling_type() JNI function for Kotlin validation - Clear error messages explaining embedding vs generative models **Kotlin Layer:** - Validate model during load() in LLamaAndroid.kt - Catch IllegalStateException and wrap with user-friendly message - File format validation for ONNX, PyTorch, TensorFlow, etc. **UI Layer:** - Proper exception handling in AiSettingsViewModel - Display error in ModelLoadingState.Error instead of crashing - Keep bottom sheet expanded after file picker to show error/status **Infrastructure:** - Rebuilt llama.cpp AAR with updated native code (v8) - Updated LLAMA_LIB_VERSION to 8 in DynamicLibraryLoader The app now gracefully handles embedding models with clear error messages instead of crashing with SIGABRT. Co-Authored-By: Claude Sonnet 4.5 <noreply@anthropic.com> * chore(ADFA-4388): Remove explanatory comments from Kotlin files Co-Authored-By: Claude Sonnet 4.5 <noreply@anthropic.com> * chore(ADFA-4388): Remove comments and extract magic strings to constants - Removed inline comments explaining bottom sheet behavior - Extracted file extension strings to named constants (EXT_*) - Extracted keyword strings to named constants (KEYWORD_*) - Improved code maintainability and reduced duplication Co-Authored-By: Claude Sonnet 4.5 <noreply@anthropic.com> * fix(ADFA-4388): block embedding models from being loaded for chat Reject embedding/Mini models (all-MiniLM, mpnet, e5, *embed*) up front by filename instead of only warning, so they never reach the native loader. The deep C++ pooling_type check remains as a second line of defense. Introduce a ModelLoadResult sealed type (Loaded/Rejected/Failed) so callers can distinguish an unsupported/embedding model from a generic load failure; migrate AiSettingsViewModel, LocalLlmRepositoryImpl and ChatViewModel to it, and extract ChatViewModel's local-LLM setup into focused helpers. Move all user-facing model-load error messages to the resources module strings.xml. Free native model/context/batch handles on load failure paths to avoid leaks, and harden batch null/negative-token guards in the JNI layer. * fix(ADFA-4388): surface specific model-load rejection reason in UI Local model rejections/failures previously collapsed to a generic message, hiding the embedding-model guidance the load result carries. - LocalLlmRepositoryImpl.loadModel: emit AgentState.Error(result.message) for Rejected/Failed instead of the generic loaded_failure status. - ChatViewModel: propagate the rejection/failure message via localModelLoadError so retrieveAgentResponse posts the specific reason, falling back to the generic "model not loaded" text when absent. --------- Co-authored-by: Claude Sonnet 4.5 <noreply@anthropic.com>
1 parent e812858 commit 939b191

10 files changed

Lines changed: 795 additions & 109 deletions

File tree

app/src/main/java/com/itsaky/androidide/agent/fragments/AiSettingsFragment.kt

Lines changed: 18 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -16,10 +16,12 @@ import androidx.core.net.toUri
1616
import androidx.fragment.app.Fragment
1717
import androidx.fragment.app.viewModels
1818
import androidx.navigation.fragment.findNavController
19+
import com.google.android.material.bottomsheet.BottomSheetBehavior
1920
import com.google.android.material.textfield.TextInputEditText
2021
import com.google.android.material.textfield.TextInputLayout
2122
import com.google.android.material.materialswitch.MaterialSwitch
2223
import com.itsaky.androidide.R
24+
import com.itsaky.androidide.activities.editor.BaseEditorActivity
2325
import com.itsaky.androidide.agent.repository.AiBackend
2426
import com.itsaky.androidide.agent.repository.Util.getCurrentBackend
2527
import com.itsaky.androidide.agent.viewmodel.AiSettingsViewModel
@@ -28,6 +30,7 @@ import com.itsaky.androidide.agent.viewmodel.ModelLoadingState
2830
import com.itsaky.androidide.databinding.FragmentAiSettingsBinding
2931
import com.itsaky.androidide.utils.flashInfo
3032
import com.itsaky.androidide.utils.getFileName
33+
import com.itsaky.androidide.viewmodel.BottomSheetViewModel
3134
import java.text.SimpleDateFormat
3235
import java.util.Date
3336
import java.util.Locale
@@ -50,6 +53,13 @@ class AiSettingsFragment : Fragment(R.layout.fragment_ai_settings) {
5053
val uriString = it.toString()
5154
viewModel.loadModelFromUri(uriString, requireContext())
5255
flashInfo("Attempting to load selected model...")
56+
57+
view?.postDelayed({
58+
(activity as? BaseEditorActivity)?.bottomSheetViewModel?.setSheetState(
59+
sheetState = BottomSheetBehavior.STATE_EXPANDED,
60+
currentTab = BottomSheetViewModel.TAB_AGENT
61+
)
62+
}, 100)
5363
}
5464
}
5565

@@ -111,7 +121,7 @@ class AiSettingsFragment : Fragment(R.layout.fragment_ai_settings) {
111121
val browseButton = view.findViewById<Button>(R.id.btn_browse_model)
112122
val loadSavedButton = view.findViewById<Button>(R.id.loadSavedButton)
113123
val modelStatusTextView = view.findViewById<TextView>(R.id.model_status_text_view)
114-
val engineStatusTextView = view.findViewById<TextView>(R.id.engine_status_text) // <-- NEW: Get reference to the new TextView
124+
val engineStatusTextView = view.findViewById<TextView>(R.id.engine_status_text)
115125
val simplePromptSwitch = view.findViewById<MaterialSwitch>(R.id.switch_simple_local_prompt)
116126
val shaInput = view.findViewById<TextInputEditText>(R.id.local_model_sha_input)
117127

@@ -204,6 +214,13 @@ class AiSettingsFragment : Fragment(R.layout.fragment_ai_settings) {
204214
}
205215
if (hasPermission) {
206216
viewModel.loadModelFromUri(savedUri, requireContext())
217+
218+
view?.postDelayed({
219+
(activity as? BaseEditorActivity)?.bottomSheetViewModel?.setSheetState(
220+
sheetState = BottomSheetBehavior.STATE_EXPANDED,
221+
currentTab = BottomSheetViewModel.TAB_AGENT
222+
)
223+
}, 100)
207224
} else {
208225
requireActivity().getSharedPreferences(PREFS_NAME, Context.MODE_PRIVATE).edit {
209226
remove(SAVED_MODEL_URI_KEY)
Lines changed: 16 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,16 @@
1+
package com.itsaky.androidide.agent.model
2+
3+
sealed interface ModelLoadResult {
4+
data class Loaded(
5+
val modelName: String
6+
) : ModelLoadResult
7+
8+
data class Rejected(
9+
val message: String
10+
) : ModelLoadResult
11+
12+
data class Failed(
13+
val message: String,
14+
val cause: Throwable? = null
15+
) : ModelLoadResult
16+
}

app/src/main/java/com/itsaky/androidide/agent/repository/LlmInferenceEngine.kt

Lines changed: 130 additions & 16 deletions
Original file line numberDiff line numberDiff line change
@@ -2,7 +2,9 @@ package com.itsaky.androidide.agent.repository
22

33
import android.content.Context
44
import androidx.core.net.toUri
5+
import com.itsaky.androidide.agent.model.ModelLoadResult
56
import com.itsaky.androidide.llamacpp.api.ILlamaController
7+
import com.itsaky.androidide.resources.R
68
import com.itsaky.androidide.utils.DynamicLibraryLoader
79
import kotlinx.coroutines.CoroutineDispatcher
810
import kotlinx.coroutines.Dispatchers
@@ -14,6 +16,7 @@ import kotlinx.coroutines.withContext
1416
import org.slf4j.LoggerFactory
1517
import java.io.File
1618
import java.io.FileOutputStream
19+
import kotlin.coroutines.cancellation.CancellationException
1720

1821
/**
1922
* A wrapper class for the LLamaAndroid library, loaded dynamically.
@@ -63,6 +66,29 @@ class LlmInferenceEngine(
6366
private const val CONTEXT_SIZE_MID_MEM = 2048
6467
private const val CONTEXT_SIZE_HIGH_MEM = 3072
6568
private const val CONTEXT_SIZE_MAX = 4096
69+
70+
private const val EXT_ONNX = ".onnx"
71+
private const val EXT_PT = ".pt"
72+
private const val EXT_PTH = ".pth"
73+
private const val EXT_BIN = ".bin"
74+
private const val EXT_SAFETENSORS = ".safetensors"
75+
private const val EXT_PB = ".pb"
76+
private const val EXT_TFLITE = ".tflite"
77+
private const val EXT_GGML = ".ggml"
78+
private const val EXT_GGUF = ".gguf"
79+
80+
private const val KEYWORD_TENSORFLOW = "tensorflow"
81+
private const val KEYWORD_ALL_MINI = "all-mini"
82+
private const val KEYWORD_ALL_MPNET = "all-mpnet"
83+
private const val KEYWORD_E5 = "e5-"
84+
private const val KEYWORD_EMBED = "embed"
85+
private const val KEYWORD_LLAMA = "llama"
86+
private const val KEYWORD_H2O = "h2o"
87+
private const val KEYWORD_DANUBE = "danube"
88+
private const val KEYWORD_QWEN = "qwen"
89+
private const val KEYWORD_GEMMA3 = "gemma3"
90+
private const val KEYWORD_GEMMA_3 = "gemma-3"
91+
private const val KEYWORD_GEMMA = "gemma"
6692
}
6793

6894
/**
@@ -267,10 +293,16 @@ class LlmInferenceEngine(
267293
context: Context,
268294
modelUriString: String,
269295
expectedSha256: String? = null
270-
): Boolean {
296+
): ModelLoadResult {
271297
return modelLoadMutex.withLock {
272-
if (!ensureInitialized(context)) return@withLock false
298+
if (!ensureInitialized(context)) {
299+
return@withLock ModelLoadResult.Failed(
300+
message = context.getString(R.string.model_error_engine_init)
301+
)
302+
}
303+
273304
if (isModelLoaded) unloadModel()
305+
274306
withContext(ioDispatcher) {
275307
loadModelFromUri(context, modelUriString, expectedSha256)
276308
}
@@ -290,33 +322,71 @@ class LlmInferenceEngine(
290322
context: Context,
291323
modelUriString: String,
292324
expectedSha256: String?
293-
): Boolean {
325+
): ModelLoadResult {
326+
val modelUri = modelUriString.toUri()
327+
val displayName = resolveModelDisplayName(context, modelUri)
328+
294329
return try {
295-
val modelUri = modelUriString.toUri()
296-
val displayName = resolveModelDisplayName(context, modelUri)
330+
validateModelFormat(context, displayName)
331+
297332
val destinationFile = File(context.cacheDir, "local_model.gguf")
298333

299334
if (!copyModelToCache(context, modelUri, destinationFile)) {
300-
return false
335+
return ModelLoadResult.Failed(
336+
message = context.getString(R.string.model_error_copy_failed)
337+
)
301338
}
302-
log.info("Model copied to cache at {}", destinationFile.path)
303339

304340
if (!verifyModelHash(destinationFile, expectedSha256)) {
305-
return false
341+
return ModelLoadResult.Failed(
342+
message = context.getString(R.string.model_error_verification_failed)
343+
)
306344
}
307345

308346
llamaController?.load(destinationFile.path)
347+
309348
isModelLoaded = true
310349
loadedModelPath = destinationFile.path
311350
loadedModelSourceUri = modelUriString
312351
loadedModelName = displayName
313352
currentModelFamily = detectModelFamily(displayName)
314-
log.info("Successfully loaded local model: {}", loadedModelName)
315-
true
353+
354+
ModelLoadResult.Loaded(displayName)
355+
} catch (e: CancellationException) {
356+
resetLoadedModelState()
357+
throw e
358+
} catch (e: IllegalStateException) {
359+
resetLoadedModelState()
360+
361+
if (e.message?.contains("embedding model", ignoreCase = true) == true) {
362+
log.error("Cannot use embedding model for chat: {}", displayName, e)
363+
364+
ModelLoadResult.Rejected(
365+
context.getString(R.string.model_error_embedding, displayName)
366+
)
367+
} else {
368+
log.error("Failed to load model: {}", displayName, e)
369+
370+
ModelLoadResult.Failed(
371+
message = context.getString(R.string.model_error_load_failed, displayName),
372+
cause = e
373+
)
374+
}
375+
} catch (e: IllegalArgumentException) {
376+
log.error("Model validation failed: {}", displayName, e)
377+
resetLoadedModelState()
378+
379+
ModelLoadResult.Rejected(
380+
e.message ?: context.getString(R.string.model_error_format_unsupported)
381+
)
316382
} catch (e: Exception) {
317383
log.error("Failed to initialize or load model from file", e)
318384
resetLoadedModelState()
319-
false
385+
386+
ModelLoadResult.Failed(
387+
message = context.getString(R.string.model_error_load_failed, displayName),
388+
cause = e
389+
)
320390
}
321391
}
322392

@@ -458,14 +528,58 @@ class LlmInferenceEngine(
458528
}
459529
}
460530

531+
/**
532+
* Validates that the model file format is supported.
533+
* This app uses llama.cpp which only supports GGUF format.
534+
*
535+
* @throws IllegalArgumentException if the model format is not supported
536+
*/
537+
private fun validateModelFormat(context: Context, filename: String) {
538+
val lowerName = filename.lowercase()
539+
540+
when {
541+
lowerName.endsWith(EXT_ONNX) -> {
542+
throw IllegalArgumentException(context.getString(R.string.model_error_format_onnx))
543+
}
544+
lowerName.endsWith(EXT_PT) || lowerName.endsWith(EXT_PTH) || lowerName.endsWith(EXT_BIN) -> {
545+
throw IllegalArgumentException(context.getString(R.string.model_error_format_pytorch))
546+
}
547+
lowerName.endsWith(EXT_SAFETENSORS) -> {
548+
throw IllegalArgumentException(context.getString(R.string.model_error_format_safetensors))
549+
}
550+
lowerName.endsWith(EXT_PB) || lowerName.contains(KEYWORD_TENSORFLOW) -> {
551+
throw IllegalArgumentException(context.getString(R.string.model_error_format_tensorflow))
552+
}
553+
lowerName.endsWith(EXT_TFLITE) -> {
554+
throw IllegalArgumentException(context.getString(R.string.model_error_format_tflite))
555+
}
556+
lowerName.endsWith(EXT_GGML) -> {
557+
throw IllegalArgumentException(context.getString(R.string.model_error_format_ggml))
558+
}
559+
!lowerName.endsWith(EXT_GGUF) -> {
560+
log.warn("Model file '{}' doesn't have $EXT_GGUF extension. May fail to load.", filename)
561+
}
562+
}
563+
564+
if (lowerName.contains(KEYWORD_ALL_MINI) ||
565+
lowerName.contains(KEYWORD_ALL_MPNET) ||
566+
lowerName.contains(KEYWORD_E5) ||
567+
(lowerName.contains(KEYWORD_EMBED) && !lowerName.contains(KEYWORD_LLAMA))) {
568+
log.error("Rejecting embedding model based on filename: {}", filename)
569+
throw IllegalArgumentException(
570+
context.getString(R.string.model_error_embedding, filename)
571+
)
572+
}
573+
}
574+
461575
private fun detectModelFamily(path: String): ModelFamily {
462576
val lowerPath = path.lowercase()
463577
return when {
464-
lowerPath.contains("h2o") || lowerPath.contains("danube") -> ModelFamily.H2O
465-
lowerPath.contains("qwen") -> ModelFamily.QWEN
466-
lowerPath.contains("gemma-3") || lowerPath.contains("gemma3") -> ModelFamily.GEMMA3
467-
lowerPath.contains("gemma") -> ModelFamily.GEMMA2
468-
lowerPath.contains("llama") -> ModelFamily.LLAMA3
578+
lowerPath.contains(KEYWORD_H2O) || lowerPath.contains(KEYWORD_DANUBE) -> ModelFamily.H2O
579+
lowerPath.contains(KEYWORD_QWEN) -> ModelFamily.QWEN
580+
lowerPath.contains(KEYWORD_GEMMA_3) || lowerPath.contains(KEYWORD_GEMMA3) -> ModelFamily.GEMMA3
581+
lowerPath.contains(KEYWORD_GEMMA) -> ModelFamily.GEMMA2
582+
lowerPath.contains(KEYWORD_LLAMA) -> ModelFamily.LLAMA3
469583
else -> ModelFamily.UNKNOWN
470584
}
471585
}

app/src/main/java/com/itsaky/androidide/agent/repository/LocalLlmRepositoryImpl.kt

Lines changed: 22 additions & 9 deletions
Original file line numberDiff line numberDiff line change
@@ -8,6 +8,7 @@ import com.itsaky.androidide.agent.AgentState
88
import com.itsaky.androidide.agent.ChatMessage
99
import com.itsaky.androidide.agent.Sender
1010
import com.itsaky.androidide.agent.ToolExecutionTracker
11+
import com.itsaky.androidide.agent.model.ModelLoadResult
1112
import com.itsaky.androidide.resources.R
1213
import kotlinx.coroutines.Dispatchers
1314
import kotlinx.coroutines.flow.MutableStateFlow
@@ -70,16 +71,28 @@ class LocalLlmRepositoryImpl(
7071

7172
suspend fun loadModel(modelUriString: String): Boolean {
7273
onStateUpdate?.invoke(AgentState.Processing("Loading local model..."))
73-
val success = engine.initModelFromFile(context, modelUriString)
74-
val status =
75-
if (success) {
76-
context.getString(R.string.agent_local_model_loaded_success)
77-
} else {
78-
context.getString(R.string.agent_local_model_loaded_failure)
74+
75+
return when (val result = engine.initModelFromFile(context, modelUriString)) {
76+
is ModelLoadResult.Loaded -> {
77+
onStateUpdate?.invoke(
78+
AgentState.Processing(context.getString(R.string.agent_local_model_loaded_success))
79+
)
80+
onStateUpdate?.invoke(AgentState.Idle)
81+
true
7982
}
80-
onStateUpdate?.invoke(AgentState.Processing(status))
81-
onStateUpdate?.invoke(AgentState.Idle)
82-
return success
83+
84+
is ModelLoadResult.Rejected -> {
85+
log.warn("Model rejected: {}", result.message)
86+
onStateUpdate?.invoke(AgentState.Error(result.message))
87+
false
88+
}
89+
90+
is ModelLoadResult.Failed -> {
91+
log.error(result.message, result.cause)
92+
onStateUpdate?.invoke(AgentState.Error(result.message))
93+
false
94+
}
95+
}
8396
}
8497

8598
private val tools: Map<String, Tool> = listOf(

0 commit comments

Comments
 (0)