Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
Expand Up @@ -16,10 +16,12 @@ import androidx.core.net.toUri
import androidx.fragment.app.Fragment
import androidx.fragment.app.viewModels
import androidx.navigation.fragment.findNavController
import com.google.android.material.bottomsheet.BottomSheetBehavior
import com.google.android.material.textfield.TextInputEditText
import com.google.android.material.textfield.TextInputLayout
import com.google.android.material.materialswitch.MaterialSwitch
import com.itsaky.androidide.R
import com.itsaky.androidide.activities.editor.BaseEditorActivity
import com.itsaky.androidide.agent.repository.AiBackend
import com.itsaky.androidide.agent.repository.Util.getCurrentBackend
import com.itsaky.androidide.agent.viewmodel.AiSettingsViewModel
Expand All @@ -28,6 +30,7 @@ import com.itsaky.androidide.agent.viewmodel.ModelLoadingState
import com.itsaky.androidide.databinding.FragmentAiSettingsBinding
import com.itsaky.androidide.utils.flashInfo
import com.itsaky.androidide.utils.getFileName
import com.itsaky.androidide.viewmodel.BottomSheetViewModel
import java.text.SimpleDateFormat
import java.util.Date
import java.util.Locale
Expand All @@ -50,6 +53,13 @@ class AiSettingsFragment : Fragment(R.layout.fragment_ai_settings) {
val uriString = it.toString()
viewModel.loadModelFromUri(uriString, requireContext())
flashInfo("Attempting to load selected model...")

view?.postDelayed({
(activity as? BaseEditorActivity)?.bottomSheetViewModel?.setSheetState(
sheetState = BottomSheetBehavior.STATE_EXPANDED,
currentTab = BottomSheetViewModel.TAB_AGENT
)
}, 100)
}
}

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

Expand Down Expand Up @@ -204,6 +214,13 @@ class AiSettingsFragment : Fragment(R.layout.fragment_ai_settings) {
}
if (hasPermission) {
viewModel.loadModelFromUri(savedUri, requireContext())

view?.postDelayed({
(activity as? BaseEditorActivity)?.bottomSheetViewModel?.setSheetState(
sheetState = BottomSheetBehavior.STATE_EXPANDED,
currentTab = BottomSheetViewModel.TAB_AGENT
)
}, 100)
} else {
requireActivity().getSharedPreferences(PREFS_NAME, Context.MODE_PRIVATE).edit {
remove(SAVED_MODEL_URI_KEY)
Expand Down
Original file line number Diff line number Diff line change
@@ -0,0 +1,16 @@
package com.itsaky.androidide.agent.model

sealed interface ModelLoadResult {
data class Loaded(
val modelName: String
) : ModelLoadResult

data class Rejected(
val message: String
) : ModelLoadResult

data class Failed(
val message: String,
val cause: Throwable? = null
) : ModelLoadResult
}
Original file line number Diff line number Diff line change
Expand Up @@ -2,7 +2,9 @@ package com.itsaky.androidide.agent.repository

import android.content.Context
import androidx.core.net.toUri
import com.itsaky.androidide.agent.model.ModelLoadResult
import com.itsaky.androidide.llamacpp.api.ILlamaController
import com.itsaky.androidide.resources.R
import com.itsaky.androidide.utils.DynamicLibraryLoader
import kotlinx.coroutines.CoroutineDispatcher
import kotlinx.coroutines.Dispatchers
Expand All @@ -14,6 +16,7 @@ import kotlinx.coroutines.withContext
import org.slf4j.LoggerFactory
import java.io.File
import java.io.FileOutputStream
import kotlin.coroutines.cancellation.CancellationException

/**
* A wrapper class for the LLamaAndroid library, loaded dynamically.
Expand Down Expand Up @@ -63,6 +66,29 @@ class LlmInferenceEngine(
private const val CONTEXT_SIZE_MID_MEM = 2048
private const val CONTEXT_SIZE_HIGH_MEM = 3072
private const val CONTEXT_SIZE_MAX = 4096

private const val EXT_ONNX = ".onnx"
private const val EXT_PT = ".pt"
private const val EXT_PTH = ".pth"
private const val EXT_BIN = ".bin"
private const val EXT_SAFETENSORS = ".safetensors"
private const val EXT_PB = ".pb"
private const val EXT_TFLITE = ".tflite"
private const val EXT_GGML = ".ggml"
private const val EXT_GGUF = ".gguf"

private const val KEYWORD_TENSORFLOW = "tensorflow"
private const val KEYWORD_ALL_MINI = "all-mini"
private const val KEYWORD_ALL_MPNET = "all-mpnet"
private const val KEYWORD_E5 = "e5-"
private const val KEYWORD_EMBED = "embed"
private const val KEYWORD_LLAMA = "llama"
private const val KEYWORD_H2O = "h2o"
private const val KEYWORD_DANUBE = "danube"
private const val KEYWORD_QWEN = "qwen"
private const val KEYWORD_GEMMA3 = "gemma3"
private const val KEYWORD_GEMMA_3 = "gemma-3"
private const val KEYWORD_GEMMA = "gemma"
}

/**
Expand Down Expand Up @@ -267,10 +293,16 @@ class LlmInferenceEngine(
context: Context,
modelUriString: String,
expectedSha256: String? = null
): Boolean {
): ModelLoadResult {
Comment thread
coderabbitai[bot] marked this conversation as resolved.
return modelLoadMutex.withLock {
if (!ensureInitialized(context)) return@withLock false
if (!ensureInitialized(context)) {
return@withLock ModelLoadResult.Failed(
message = context.getString(R.string.model_error_engine_init)
)
}

if (isModelLoaded) unloadModel()

withContext(ioDispatcher) {
loadModelFromUri(context, modelUriString, expectedSha256)
}
Expand All @@ -290,33 +322,71 @@ class LlmInferenceEngine(
context: Context,
modelUriString: String,
expectedSha256: String?
): Boolean {
): ModelLoadResult {
val modelUri = modelUriString.toUri()
val displayName = resolveModelDisplayName(context, modelUri)

return try {
Comment thread
coderabbitai[bot] marked this conversation as resolved.
val modelUri = modelUriString.toUri()
val displayName = resolveModelDisplayName(context, modelUri)
validateModelFormat(context, displayName)

val destinationFile = File(context.cacheDir, "local_model.gguf")

if (!copyModelToCache(context, modelUri, destinationFile)) {
return false
return ModelLoadResult.Failed(
message = context.getString(R.string.model_error_copy_failed)
)
}
log.info("Model copied to cache at {}", destinationFile.path)

if (!verifyModelHash(destinationFile, expectedSha256)) {
return false
return ModelLoadResult.Failed(
message = context.getString(R.string.model_error_verification_failed)
)
}

llamaController?.load(destinationFile.path)

isModelLoaded = true
loadedModelPath = destinationFile.path
loadedModelSourceUri = modelUriString
loadedModelName = displayName
currentModelFamily = detectModelFamily(displayName)
log.info("Successfully loaded local model: {}", loadedModelName)
true

ModelLoadResult.Loaded(displayName)
} catch (e: CancellationException) {
resetLoadedModelState()
throw e
} catch (e: IllegalStateException) {
resetLoadedModelState()

if (e.message?.contains("embedding model", ignoreCase = true) == true) {
log.error("Cannot use embedding model for chat: {}", displayName, e)

ModelLoadResult.Rejected(
context.getString(R.string.model_error_embedding, displayName)
)
} else {
log.error("Failed to load model: {}", displayName, e)

ModelLoadResult.Failed(
message = context.getString(R.string.model_error_load_failed, displayName),
cause = e
)
}
} catch (e: IllegalArgumentException) {
log.error("Model validation failed: {}", displayName, e)
resetLoadedModelState()

ModelLoadResult.Rejected(
e.message ?: context.getString(R.string.model_error_format_unsupported)
)
} catch (e: Exception) {
log.error("Failed to initialize or load model from file", e)
resetLoadedModelState()
false

ModelLoadResult.Failed(
message = context.getString(R.string.model_error_load_failed, displayName),
cause = e
)
}
}

Expand Down Expand Up @@ -458,14 +528,58 @@ class LlmInferenceEngine(
}
}

/**
* Validates that the model file format is supported.
* This app uses llama.cpp which only supports GGUF format.
*
* @throws IllegalArgumentException if the model format is not supported
*/
private fun validateModelFormat(context: Context, filename: String) {
val lowerName = filename.lowercase()

when {
lowerName.endsWith(EXT_ONNX) -> {
throw IllegalArgumentException(context.getString(R.string.model_error_format_onnx))
}
lowerName.endsWith(EXT_PT) || lowerName.endsWith(EXT_PTH) || lowerName.endsWith(EXT_BIN) -> {
throw IllegalArgumentException(context.getString(R.string.model_error_format_pytorch))
}
lowerName.endsWith(EXT_SAFETENSORS) -> {
throw IllegalArgumentException(context.getString(R.string.model_error_format_safetensors))
}
lowerName.endsWith(EXT_PB) || lowerName.contains(KEYWORD_TENSORFLOW) -> {
throw IllegalArgumentException(context.getString(R.string.model_error_format_tensorflow))
}
lowerName.endsWith(EXT_TFLITE) -> {
throw IllegalArgumentException(context.getString(R.string.model_error_format_tflite))
}
lowerName.endsWith(EXT_GGML) -> {
throw IllegalArgumentException(context.getString(R.string.model_error_format_ggml))
}
!lowerName.endsWith(EXT_GGUF) -> {
log.warn("Model file '{}' doesn't have $EXT_GGUF extension. May fail to load.", filename)
}
}

if (lowerName.contains(KEYWORD_ALL_MINI) ||
lowerName.contains(KEYWORD_ALL_MPNET) ||
lowerName.contains(KEYWORD_E5) ||
(lowerName.contains(KEYWORD_EMBED) && !lowerName.contains(KEYWORD_LLAMA))) {
log.error("Rejecting embedding model based on filename: {}", filename)
throw IllegalArgumentException(
context.getString(R.string.model_error_embedding, filename)
)
}
}

private fun detectModelFamily(path: String): ModelFamily {
val lowerPath = path.lowercase()
return when {
lowerPath.contains("h2o") || lowerPath.contains("danube") -> ModelFamily.H2O
lowerPath.contains("qwen") -> ModelFamily.QWEN
lowerPath.contains("gemma-3") || lowerPath.contains("gemma3") -> ModelFamily.GEMMA3
lowerPath.contains("gemma") -> ModelFamily.GEMMA2
lowerPath.contains("llama") -> ModelFamily.LLAMA3
lowerPath.contains(KEYWORD_H2O) || lowerPath.contains(KEYWORD_DANUBE) -> ModelFamily.H2O
lowerPath.contains(KEYWORD_QWEN) -> ModelFamily.QWEN
lowerPath.contains(KEYWORD_GEMMA_3) || lowerPath.contains(KEYWORD_GEMMA3) -> ModelFamily.GEMMA3
lowerPath.contains(KEYWORD_GEMMA) -> ModelFamily.GEMMA2
lowerPath.contains(KEYWORD_LLAMA) -> ModelFamily.LLAMA3
else -> ModelFamily.UNKNOWN
}
}
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -8,6 +8,7 @@ import com.itsaky.androidide.agent.AgentState
import com.itsaky.androidide.agent.ChatMessage
import com.itsaky.androidide.agent.Sender
import com.itsaky.androidide.agent.ToolExecutionTracker
import com.itsaky.androidide.agent.model.ModelLoadResult
import com.itsaky.androidide.resources.R
import kotlinx.coroutines.Dispatchers
import kotlinx.coroutines.flow.MutableStateFlow
Expand Down Expand Up @@ -70,16 +71,28 @@ class LocalLlmRepositoryImpl(

suspend fun loadModel(modelUriString: String): Boolean {
onStateUpdate?.invoke(AgentState.Processing("Loading local model..."))
val success = engine.initModelFromFile(context, modelUriString)
val status =
if (success) {
context.getString(R.string.agent_local_model_loaded_success)
} else {
context.getString(R.string.agent_local_model_loaded_failure)

return when (val result = engine.initModelFromFile(context, modelUriString)) {
is ModelLoadResult.Loaded -> {
onStateUpdate?.invoke(
AgentState.Processing(context.getString(R.string.agent_local_model_loaded_success))
)
onStateUpdate?.invoke(AgentState.Idle)
true
}
onStateUpdate?.invoke(AgentState.Processing(status))
onStateUpdate?.invoke(AgentState.Idle)
return success

is ModelLoadResult.Rejected -> {
log.warn("Model rejected: {}", result.message)
onStateUpdate?.invoke(AgentState.Error(result.message))
false
}

is ModelLoadResult.Failed -> {
log.error(result.message, result.cause)
onStateUpdate?.invoke(AgentState.Error(result.message))
false
}
}
}

private val tools: Map<String, Tool> = listOf(
Expand Down
Loading
Loading