diff --git a/Cargo.lock b/Cargo.lock index 0e50dc2..61fa42e 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -147,25 +147,22 @@ checksum = "55248b47b0caf0546f7988906588779981c43bb1bc9d0c44087278f80cdb44ba" [[package]] name = "bindgen" -version = "0.69.5" +version = "0.72.1" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "271383c67ccabffb7381723dea0672a673f292304fcb45c01cc648c7a8d58088" +checksum = "993776b509cfb49c750f11b8f07a46fa23e0a1386ffc01fb1e7d343efc387895" dependencies = [ "bitflags", "cexpr", "clang-sys", "itertools", - "lazy_static", - "lazycell", "log", "prettyplease", "proc-macro2", "quote", "regex", - "rustc-hash 1.1.0", + "rustc-hash", "shlex", "syn", - "which", ] [[package]] @@ -679,15 +676,6 @@ version = "0.5.2" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "fc0fef456e4baa96da950455cd02c081ca953b141298e41db3fc7e36b1da849c" -[[package]] -name = "home" -version = "0.5.12" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "cc627f471c528ff0c4a49e1d5e60450c8f6461dd6d10ba9dcd3a61d3dff7728d" -dependencies = [ - "windows-sys 0.61.2", -] - [[package]] name = "hound" version = "3.5.1" @@ -1004,18 +992,6 @@ dependencies = [ "wasm-bindgen", ] -[[package]] -name = "lazy_static" -version = "1.5.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "bbd2bcb4c963f2ddae06a2efc7e9f3591312473c50c6685e1f298068316e66fe" - -[[package]] -name = "lazycell" -version = "1.3.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "830d08ce1d1d941e6b30645f1a0eb5643013d835ce3779a5fc208261dbe10f55" - [[package]] name = "libc" version = "0.2.177" @@ -1043,12 +1019,6 @@ dependencies = [ "redox_syscall", ] -[[package]] -name = "linux-raw-sys" -version = "0.4.15" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "d26c52dbd32dccf2d10cac7725f8eae5296885fb5703b261f7d0a0739ec807ab" - [[package]] name = "linux-raw-sys" version = "0.11.0" @@ -1184,6 +1154,7 @@ dependencies = [ "once_cell", "ort", "serde_json", + "sha2", "transcribe-rs", ] @@ -1395,7 +1366,7 @@ dependencies = [ "pin-project-lite", "quinn-proto", "quinn-udp", - "rustc-hash 2.1.1", + "rustc-hash", "rustls", "socket2", "thiserror 2.0.16", @@ -1415,7 +1386,7 @@ dependencies = [ "lru-slab", "rand 0.9.2", "ring", - "rustc-hash 2.1.1", + "rustc-hash", "rustls", "rustls-pki-types", "slab", @@ -1635,31 +1606,12 @@ version = "0.1.27" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "b50b8869d9fc858ce7266cce0194bd74df58b9d0e3f6df3a9fc8eb470d95c09d" -[[package]] -name = "rustc-hash" -version = "1.1.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "08d43f7aa6b08d49f382cde6a7982047c3426db949b1424bc4b7ec9ae12c6ce2" - [[package]] name = "rustc-hash" version = "2.1.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "357703d41365b4b27c590e3ed91eabb1b663f07c4c084095e60cbed4362dff0d" -[[package]] -name = "rustix" -version = "0.38.44" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "fdb5bc1ae2baa591800df16c9ca78619bf65c0488b41b96ccec5d11220d8c154" -dependencies = [ - "bitflags", - "errno", - "libc", - "linux-raw-sys 0.4.15", - "windows-sys 0.59.0", -] - [[package]] name = "rustix" version = "1.1.2" @@ -1669,7 +1621,7 @@ dependencies = [ "bitflags", "errno", "libc", - "linux-raw-sys 0.11.0", + "linux-raw-sys", "windows-sys 0.61.2", ] @@ -1796,6 +1748,12 @@ dependencies = [ "libc", ] +[[package]] +name = "semver" +version = "1.0.28" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "8a7852d02fc848982e0c167ef163aaff9cd91dc640ba85e263cb1ce46fae51cd" + [[package]] name = "serde" version = "1.0.228" @@ -1982,7 +1940,7 @@ dependencies = [ "fastrand", "getrandom 0.3.4", "once_cell", - "rustix 1.1.2", + "rustix", "windows-sys 0.61.2", ] @@ -2448,37 +2406,27 @@ dependencies = [ "rustls-pki-types", ] -[[package]] -name = "which" -version = "4.4.2" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "87ba24419a2078cd2b0f2ede2691b6c66d8e47836da3b6db8265ebad47afbfc7" -dependencies = [ - "either", - "home", - "once_cell", - "rustix 0.38.44", -] - [[package]] name = "whisper-rs" -version = "0.13.2" +version = "0.16.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "40b6fc553156b521663bfa8e713e7ad58c7ca262d46de9998cd7f2e4de5ba0d9" +checksum = "2088172d00f936c348d6a72f488dc2660ab3f507263a195df308a3c2383229f6" dependencies = [ + "libc", "whisper-rs-sys", ] [[package]] name = "whisper-rs-sys" -version = "0.11.1" +version = "0.15.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "76bab42b2c319e3a1e0280137c59368072348d3277873c7588b6466a127dca58" +checksum = "6986c0fe081241d391f09b9a071fbcbb59720c3563628c3c829057cf69f2a56f" dependencies = [ "bindgen", "cfg-if", "cmake", "fs_extra", + "semver", ] [[package]] @@ -2768,7 +2716,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "32e45ad4206f6d2479085147f02bc2ef834ac85886624a23575ae137c8aa8156" dependencies = [ "libc", - "rustix 1.1.2", + "rustix", ] [[package]] diff --git a/Cargo.toml b/Cargo.toml index 524c699..b401aa9 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -11,13 +11,17 @@ crate-type = ["cdylib", "staticlib"] default = ["native-transcription"] native-transcription = ["dep:ort", "dep:transcribe-rs"] ios-onnx = ["native-transcription"] +ios-dual-backend = ["ios-onnx", "transcribe-rs/whisper"] +android-dual-backend = ["native-transcription", "transcribe-rs/whisper"] +whisper-mobile-spike = ["dep:sha2", "dep:transcribe-rs", "transcribe-rs/whisper"] [dependencies] log = "0.4" once_cell = "1.19" serde_json = "1" -transcribe-rs = { path = "transcribe-rs", optional = true } +transcribe-rs = { path = "transcribe-rs", optional = true, default-features = false, features = ["parakeet"] } ort = { version = "=2.0.0-rc.10", optional = true } +sha2 = { version = "0.10", optional = true } [target.'cfg(target_os = "android")'.dependencies] android_logger = "0.13" diff --git a/app/build.gradle.kts b/app/build.gradle.kts index efe8816..cdcfe8c 100644 --- a/app/build.gradle.kts +++ b/app/build.gradle.kts @@ -15,6 +15,9 @@ val hasReleaseSigning = listOf( releaseKeyAlias, releaseKeyPassword, ).all { !it.isNullOrBlank() } +val cargoExecutable = System.getenv("CARGO") + ?: File(System.getProperty("user.home"), ".cargo/bin/cargo").takeIf(File::isFile)?.absolutePath + ?: "cargo" android { namespace = "me.maxistar.voiceinbox" @@ -122,6 +125,27 @@ val cargoNdkBuild by tasks.registering(Exec::class) { ?: System.getenv("ANDROID_NDK") ?: android.ndkDirectory.absolutePath environment("ANDROID_NDK_HOME", ndkDir) + val ndkPrebuiltDir = File(ndkDir, "toolchains/llvm/prebuilt").listFiles() + ?.firstOrNull(File::isDirectory) ?: throw GradleException("NDK LLVM prebuilt directory was not found") + val whisperCompatDir = layout.buildDirectory.dir("whisper-android-compat").get().asFile.apply { mkdirs() } + val whisperCompatArchive = File(whisperCompatDir, "libggml-blas.a") + if (!whisperCompatArchive.exists()) { + exec { commandLine(File(ndkPrebuiltDir, "bin/llvm-ar").absolutePath, "crs", whisperCompatArchive.absolutePath) } + } + val cmake = File(android.sdkDirectory, "cmake").listFiles() + ?.sortedByDescending(File::getName) + ?.map { File(it, "bin/cmake") } + ?.firstOrNull(File::isFile) + ?: throw GradleException("Android SDK CMake is required for the Whisper build") + environment("CMAKE", cmake.absolutePath) + environment("CMAKE_GENERATOR", "Ninja") + environment("CMAKE_MAKE_PROGRAM", File(cmake.parentFile, "ninja").absolutePath) + environment("CMAKE_ANDROID_ARCH_ABI", "arm64-v8a") + environment("RUSTFLAGS", "-Lnative=${whisperCompatDir.absolutePath}") + environment( + "CMAKE_TOOLCHAIN_FILE", + rootProject.file("scripts/whisper-android-cmake/android.toolchain.cmake").absolutePath, + ) val extractDir = layout.buildDirectory.dir("ort-extracted").get().asFile environment("ORT_LIB_LOCATION", File(extractDir, "jni/arm64-v8a").absolutePath) @@ -129,10 +153,11 @@ val cargoNdkBuild by tasks.registering(Exec::class) { val jniLibsDir = project.file("src/main/jniLibs") commandLine( - "cargo", "ndk", + cargoExecutable, "ndk", "-t", "arm64-v8a", "-o", jniLibsDir.absolutePath, "build", "--release", + "--features", "android-dual-backend", ) doLast { diff --git a/app/src/androidTest/java/me/maxistar/voiceinbox/EndToEndTranscriptionInstrumentedTest.kt b/app/src/androidTest/java/me/maxistar/voiceinbox/EndToEndTranscriptionInstrumentedTest.kt index 6eb3838..0199582 100644 --- a/app/src/androidTest/java/me/maxistar/voiceinbox/EndToEndTranscriptionInstrumentedTest.kt +++ b/app/src/androidTest/java/me/maxistar/voiceinbox/EndToEndTranscriptionInstrumentedTest.kt @@ -25,7 +25,10 @@ class EndToEndTranscriptionInstrumentedTest { @Test fun wavAndM4aBatchAppendSeparateTranscriptEntries() { - val model = SpeechModelRepository(targetContext.noBackupFilesDir.resolve("models")).inspect() + val model = SpeechModelRepository( + targetContext.noBackupFilesDir.resolve("models"), + SpeechModelCatalog.defaultModel.manifest, + ).inspect() assumeTrue(model is InstalledSpeechModelState.Ready) targetContext.deleteDatabase(AndroidSqlDelightAudioCatalogFactory.DATABASE_NAME) diff --git a/app/src/androidTest/java/me/maxistar/voiceinbox/MainActivityInstrumentedTest.kt b/app/src/androidTest/java/me/maxistar/voiceinbox/MainActivityInstrumentedTest.kt index 52e992d..b17282f 100644 --- a/app/src/androidTest/java/me/maxistar/voiceinbox/MainActivityInstrumentedTest.kt +++ b/app/src/androidTest/java/me/maxistar/voiceinbox/MainActivityInstrumentedTest.kt @@ -61,6 +61,7 @@ class MainActivityInstrumentedTest { openActionBarOverflowOrOptionsMenu(InstrumentationRegistry.getInstrumentation().targetContext) onView(withText(R.string.menu_settings)).check(matches(isDisplayed())) + onView(withText(R.string.menu_documentation)).check(matches(isDisplayed())) } } diff --git a/app/src/androidTest/java/me/maxistar/voiceinbox/SpeechModelLocalImporterInstrumentedTest.kt b/app/src/androidTest/java/me/maxistar/voiceinbox/SpeechModelLocalImporterInstrumentedTest.kt index 3050909..c17c58a 100644 --- a/app/src/androidTest/java/me/maxistar/voiceinbox/SpeechModelLocalImporterInstrumentedTest.kt +++ b/app/src/androidTest/java/me/maxistar/voiceinbox/SpeechModelLocalImporterInstrumentedTest.kt @@ -29,7 +29,16 @@ class SpeechModelLocalImporterInstrumentedTest { ) val progress = mutableListOf() - val installed = SpeechModelLocalImporter(context.contentResolver, repository) + val installed = SpeechModelLocalImporter( + repository = repository, + requiredDocuments = { _, _ -> + mapOf( + "recording.wav" to DocumentsContract.buildDocumentUriUsingTree(tree, "wav").toString(), + "recording.m4a" to DocumentsContract.buildDocumentUriUsingTree(tree, "m4a").toString(), + ) + }, + openInputStream = { context.contentResolver.openInputStream(Uri.parse(it)) }, + ) .import(tree.toString()) { progress += it } .getOrThrow() diff --git a/app/src/androidTest/java/me/maxistar/voiceinbox/TranscriptionWorkerInstrumentedTest.kt b/app/src/androidTest/java/me/maxistar/voiceinbox/TranscriptionWorkerInstrumentedTest.kt index 8e0902a..da6eb25 100644 --- a/app/src/androidTest/java/me/maxistar/voiceinbox/TranscriptionWorkerInstrumentedTest.kt +++ b/app/src/androidTest/java/me/maxistar/voiceinbox/TranscriptionWorkerInstrumentedTest.kt @@ -241,7 +241,10 @@ class TranscriptionWorkerInstrumentedTest { private fun requireModel() { assumeTrue( - SpeechModelRepository(targetContext.noBackupFilesDir.resolve("models")).inspect() + SpeechModelRepository( + targetContext.noBackupFilesDir.resolve("models"), + SpeechModelCatalog.defaultModel.manifest, + ).inspect() is InstalledSpeechModelState.Ready, ) } diff --git a/app/src/main/java/me/maxistar/voiceinbox/AndroidMainScreenStateHost.kt b/app/src/main/java/me/maxistar/voiceinbox/AndroidMainScreenStateHost.kt index 7d7a316..4649bd7 100644 --- a/app/src/main/java/me/maxistar/voiceinbox/AndroidMainScreenStateHost.kt +++ b/app/src/main/java/me/maxistar/voiceinbox/AndroidMainScreenStateHost.kt @@ -69,6 +69,7 @@ data class AndroidMainScreenState( val importEnabled: Boolean, val folderSync: AndroidFolderSyncPresentation, val onboardingHint: AndroidOnboardingHintPresentation, + val transcriptionActive: Boolean, ) { val refreshFolderVisible: Boolean get() = folderSync.visible val refreshFolderEnabled: Boolean get() = folderSync.enabled @@ -120,6 +121,7 @@ object AndroidTaskListSnapshotMapper { entriesById = input.entries.associateBy(AudioCatalogEntry::id), importEnabled = input.importEnabled, folderSync = input.folderSync, + transcriptionActive = input.transcription.active, onboardingHint = AndroidOnboardingHintPresenter.present( lifecycle = input.onboardingLifecycle, filter = input.filter, @@ -137,6 +139,10 @@ object AndroidTaskListSnapshotMapper { } +internal object AndroidTaskListAnimationPolicy { + fun suppressStructuralAnimations(transcriptionActive: Boolean): Boolean = transcriptionActive +} + class AndroidMainScreenStateHost( private val savedStateHandle: SavedStateHandle, ) : ViewModel() { diff --git a/app/src/main/java/me/maxistar/voiceinbox/MainActivity.kt b/app/src/main/java/me/maxistar/voiceinbox/MainActivity.kt index fb1faa8..c23c5af 100644 --- a/app/src/main/java/me/maxistar/voiceinbox/MainActivity.kt +++ b/app/src/main/java/me/maxistar/voiceinbox/MainActivity.kt @@ -37,6 +37,7 @@ import java.util.concurrent.Executors import java.util.UUID class MainActivity : AppCompatActivity(), StartupProcessingDialogFragment.Listener { + private val speechModel = SpeechModelCatalog.defaultModel private lateinit var importAudio: FloatingActionButton private lateinit var newTab: MaterialButton private lateinit var processedTab: MaterialButton @@ -44,6 +45,8 @@ class MainActivity : AppCompatActivity(), StartupProcessingDialogFragment.Listen private lateinit var taskFilters: MaterialButtonToggleGroup private lateinit var taskList: RecyclerView private lateinit var taskAdapter: TaskListAdapter + private var taskListItemAnimator: RecyclerView.ItemAnimator? = null + private var taskListAnimatorSuppressed = false private lateinit var taskActionRouter: AndroidTaskActionRouter private val taskStateHost: AndroidMainScreenStateHost by viewModels() @@ -82,6 +85,7 @@ class MainActivity : AppCompatActivity(), StartupProcessingDialogFragment.Listen private var modelDownloadAvailable = false private var modelDownloadProgress: Int? = null private var modelInstallCanCancel = false + private var selectedDownloadModel = SpeechModelCatalog.defaultModel private var scanMessage: String? = null private var transcriptionFinished = false private var transcriptionPhase: String? = null @@ -123,6 +127,8 @@ class MainActivity : AppCompatActivity(), StartupProcessingDialogFragment.Listen private val stopRefreshIndicator = Runnable(::stopRefreshIndicatorAnimation) private val queuedImportUris = mutableListOf() private var onboardingHintLifecycle = AndroidOnboardingHintLifecycle.DISMISSED + private var pendingModelPackageUri: Uri? = null + private var pendingModelPackageCatalogId: String? = null private val outputPicker = registerForActivityResult( ActivityResultContracts.OpenDocument(), @@ -173,7 +179,9 @@ class MainActivity : AppCompatActivity(), StartupProcessingDialogFragment.Listen documentAccess = DocumentAccess(contentResolver) folderScanner = AudioFolderScanner(contentResolver) catalog = AndroidSqlDelightAudioCatalogFactory(this).create() - modelReadiness = getSharedModelReadiness(SpeechModelRepository(noBackupFilesDir.resolve("models"))) + modelReadiness = getSharedModelReadiness { + SpeechModelRepository.forActive(noBackupFilesDir.resolve("models")) + } startupPolicyStore = StartupProcessingPolicyStore( getSharedPreferences(StartupProcessingPolicyStore.PREFERENCES_NAME, MODE_PRIVATE), ) @@ -184,6 +192,8 @@ class MainActivity : AppCompatActivity(), StartupProcessingDialogFragment.Listen startupCoordinator = StartupProcessingCoordinator.restore( savedInstanceState?.getString(STATE_STARTUP_PROCESSING_STAGE), ) + pendingModelPackageUri = savedInstanceState?.getString(STATE_PENDING_MODEL_URI)?.let(Uri::parse) + pendingModelPackageCatalogId = savedInstanceState?.getString(STATE_PENDING_MODEL_CATALOG_ID) restoreRetainedPresentation() taskActionRouter = AndroidTaskActionRouter( currentState = { taskStateHost.state.value }, @@ -222,6 +232,11 @@ class MainActivity : AppCompatActivity(), StartupProcessingDialogFragment.Listen cleanupImportedAudio() handleShareIntent(intent) refreshModel() + restorePendingModelConfirmation() + if (intent.getBooleanExtra(EXTRA_OPEN_MODEL_FOLDER_PICKER, false)) { + intent.removeExtra(EXTRA_OPEN_MODEL_FOLDER_PICKER) + modelFolderPicker.launch(null) + } } override fun onStart() { @@ -241,6 +256,8 @@ class MainActivity : AppCompatActivity(), StartupProcessingDialogFragment.Listen override fun onSaveInstanceState(outState: Bundle) { outState.putString(STATE_STARTUP_PROCESSING_STAGE, startupCoordinator.savedStage()) + outState.putString(STATE_PENDING_MODEL_URI, pendingModelPackageUri?.toString()) + outState.putString(STATE_PENDING_MODEL_CATALOG_ID, pendingModelPackageCatalogId) super.onSaveInstanceState(outState) } @@ -272,6 +289,10 @@ class MainActivity : AppCompatActivity(), StartupProcessingDialogFragment.Listen startActivity(Intent(this, SettingsActivity::class.java)) true } + R.id.menuDocumentation -> { + openDocumentation() + true + } else -> super.onOptionsItemSelected(item) } @@ -297,6 +318,7 @@ class MainActivity : AppCompatActivity(), StartupProcessingDialogFragment.Listen taskList.layoutManager = LinearLayoutManager(this) taskList.adapter = taskAdapter (taskList.itemAnimator as? SimpleItemAnimator)?.supportsChangeAnimations = false + taskListItemAnimator = taskList.itemAnimator } private fun restoreRetainedPresentation() { @@ -642,21 +664,15 @@ class MainActivity : AppCompatActivity(), StartupProcessingDialogFragment.Listen return } if (!alreadyPersisted) SpeechModelImportPermission.recordOwned(this, uri) - modelPresentationKnown = true - modelReady = false - modelSetupState = ModelSetupSnapshotState.INSTALLING - setModelUi("Checking local speech model", canDownload = false) - updateControls() importExecutor.execute { runCatching { - SpeechModelDirectoryReader(contentResolver).requiredDocuments( - uri, - EmbeddedSpeechModel.manifest, - ) - }.onSuccess { + SpeechModelDirectoryReader(contentResolver).inspectPackage(uri) + }.onSuccess { source -> runOnUiThread { if (activityDestroyed) return@runOnUiThread - SpeechModelImportWorker.enqueue(this, uri) + pendingModelPackageUri = uri + pendingModelPackageCatalogId = source.descriptor.catalogId + showModelPackageConfirmation(source.descriptor) } }.onFailure { error -> SpeechModelImportPermission.releaseOwnedIfUnused(this) @@ -675,6 +691,46 @@ class MainActivity : AppCompatActivity(), StartupProcessingDialogFragment.Listen } } + private fun restorePendingModelConfirmation() { + val descriptor = pendingModelPackageCatalogId?.let { catalogId -> + SpeechModelCatalog.models.singleOrNull { it.catalogId == catalogId } + } ?: return + showModelPackageConfirmation(descriptor) + } + + private fun showModelPackageConfirmation(descriptor: SpeechModelDescriptor) { + val uri = pendingModelPackageUri ?: return + val size = SpeechModelRepository.formatBytes(descriptor.approximateDownloadBytes) + val maturity = descriptor.maturity.name.lowercase().replaceFirstChar(Char::uppercase) + val dialog = MaterialAlertDialogBuilder(this) + .setTitle("Install ${descriptor.displayName}?") + .setMessage("$maturity model · ${descriptor.languages.summary} · approximately $size") + .setNegativeButton(android.R.string.cancel) { _, _ -> + pendingModelPackageUri = null + pendingModelPackageCatalogId = null + SpeechModelImportPermission.releaseOwnedIfUnused(this) + } + .setPositiveButton("Install", null) + .create() + dialog.setOnShowListener { + val confirm = dialog.getButton(AlertDialog.BUTTON_POSITIVE) + confirm.isEnabled = !transcriptionActive() + confirm.setOnClickListener { + if (transcriptionActive()) return@setOnClickListener + pendingModelPackageUri = null + pendingModelPackageCatalogId = null + modelPresentationKnown = true + modelReady = false + modelSetupState = ModelSetupSnapshotState.INSTALLING + setModelUi("Installing ${descriptor.displayName}", canDownload = false) + updateControls() + SpeechModelImportWorker.enqueue(this, uri, descriptor) + dialog.dismiss() + } + } + dialog.show() + } + private fun scanFolder( origin: FolderScanOrigin = FolderScanOrigin.USER, existingSyncGeneration: Long? = null, @@ -935,7 +991,7 @@ class MainActivity : AppCompatActivity(), StartupProcessingDialogFragment.Listen val bytes = info.progress.getLong(SpeechModelInstallationWork.KEY_BYTES_DOWNLOADED, 0) val total = info.progress.getLong( SpeechModelInstallationWork.KEY_TOTAL_BYTES, - EmbeddedSpeechModel.manifest.totalSizeBytes, + speechModel.approximateDownloadBytes, ) modelMessage = info.progress.getString(SpeechModelInstallationWork.KEY_MESSAGE) @@ -1307,6 +1363,12 @@ class MainActivity : AppCompatActivity(), StartupProcessingDialogFragment.Listen progressPercent = modelDownloadProgress, downloadAvailable = modelDownloadAvailable, canCancel = modelInstallCanCancel, + selectedModel = SpeechModelPackageIdentity( + schemaVersion = SpeechModelCatalog.PACKAGE_SCHEMA_VERSION, + catalogId = selectedDownloadModel.catalogId, + modelVersion = selectedDownloadModel.manifest.version, + ), + downloadChoices = SpeechModelCatalog.networkDownloadChoices(SpeechModelPlatform.ANDROID), ) val outputSnapshot = OutputSetupSnapshot( state = outputState, @@ -1375,14 +1437,39 @@ class MainActivity : AppCompatActivity(), StartupProcessingDialogFragment.Listen ) renderingFilter = false importAudio.isEnabled = state.importEnabled + updateTaskListAnimation(state.transcriptionActive) taskAdapter.submitList(TaskListDisplayItems.from(state.taskList, state.onboardingHint)) invalidateOptionsMenu() } + private fun updateTaskListAnimation(transcriptionActive: Boolean) { + val suppress = AndroidTaskListAnimationPolicy.suppressStructuralAnimations(transcriptionActive) + if (suppress == taskListAnimatorSuppressed) return + taskListAnimatorSuppressed = suppress + if (suppress) { + taskList.itemAnimator?.endAnimations() + taskList.itemAnimator = null + } else { + taskList.itemAnimator = taskListItemAnimator + } + } + private fun handleTaskAction(request: AndroidTaskActionRequest) { taskActionRouter.route(request) } + private fun openDocumentation() { + runCatching { + startActivity(Intent(Intent.ACTION_VIEW, Uri.parse(VoiceInboxPublicLinks.DOCUMENTATION))) + }.onFailure { error -> + if (error is android.content.ActivityNotFoundException) { + Toast.makeText(this, R.string.settings_about_link_error, Toast.LENGTH_LONG).show() + } else { + throw error + } + } + } + private fun dismissOnboardingHint() { if (onboardingHintLifecycle != AndroidOnboardingHintLifecycle.ACTIVE) return onboardingHintLifecycle = AndroidOnboardingHintLifecycle.DISMISSED @@ -1394,7 +1481,8 @@ class MainActivity : AppCompatActivity(), StartupProcessingDialogFragment.Listen when (kind) { TaskActionKind.DOWNLOAD_MODEL, TaskActionKind.RETRY_MODEL_DOWNLOAD, - -> SpeechModelDownloadWorker.enqueue(this) + -> SpeechModelDownloadWorker.enqueue(this, selectedDownloadModel) + TaskActionKind.SELECT_DOWNLOAD_MODEL -> chooseDownloadModel() TaskActionKind.IMPORT_MODEL -> modelFolderPicker.launch(null) TaskActionKind.CANCEL_MODEL_DOWNLOAD -> SpeechModelDownloadWorker.cancel(this) TaskActionKind.CREATE_OUTPUT -> launchOutputCreatorIfEnabled() @@ -1413,6 +1501,23 @@ class MainActivity : AppCompatActivity(), StartupProcessingDialogFragment.Listen } } + private fun chooseDownloadModel() { + if (modelSetupState == ModelSetupSnapshotState.INSTALLING) return + val choices = SpeechModelCatalog.modelsFor(SpeechModelPlatform.ANDROID) + .filter { it.distribution.networkDownloadAvailable } + MaterialAlertDialogBuilder(this) + .setTitle("Download speech model") + .setSingleChoiceItems( + choices.map { it.displayName }.toTypedArray(), + choices.indexOfFirst { it.catalogId == selectedDownloadModel.catalogId }, + ) { dialog, which -> + selectedDownloadModel = choices[which] + publishTaskState() + dialog.dismiss() + } + .show() + } + private fun transcribeEntry(entry: AudioCatalogEntry) { val output = outputUri ?: return if (entry.state != AudioFileState.PENDING || !outputAccessReady) return @@ -1586,6 +1691,9 @@ class MainActivity : AppCompatActivity(), StartupProcessingDialogFragment.Listen private companion object { const val STATE_STARTUP_PROCESSING_STAGE = "startup-processing-stage" + const val STATE_PENDING_MODEL_URI = "pending-model-uri" + const val STATE_PENDING_MODEL_CATALOG_ID = "pending-model-catalog-id" + const val EXTRA_OPEN_MODEL_FOLDER_PICKER = "open-model-folder-picker" const val REFRESH_SHOW_DELAY_MS = 180L const val REFRESH_MIN_VISIBLE_MS = 450L const val REFRESH_ROTATION_MS = 800L @@ -1600,7 +1708,7 @@ class MainActivity : AppCompatActivity(), StartupProcessingDialogFragment.Listen private val handledModelInstallSuccessIds = mutableSetOf() private val handledTranscriptionFailureIds = mutableSetOf() - fun getSharedModelReadiness(repository: SpeechModelRepository): SpeechModelReadinessManager = + fun getSharedModelReadiness(repository: () -> SpeechModelRepository): SpeechModelReadinessManager = sharedModelReadiness ?: synchronized(this) { sharedModelReadiness ?: SpeechModelReadinessManager( repository = repository, diff --git a/app/src/main/java/me/maxistar/voiceinbox/NativeTranscriptionBridge.kt b/app/src/main/java/me/maxistar/voiceinbox/NativeTranscriptionBridge.kt index 94a920e..1a2d510 100644 --- a/app/src/main/java/me/maxistar/voiceinbox/NativeTranscriptionBridge.kt +++ b/app/src/main/java/me/maxistar/voiceinbox/NativeTranscriptionBridge.kt @@ -20,7 +20,19 @@ object NativeTranscriptionBridge { System.loadLibrary("notes_recognition") } - external fun initialize(modelDirectory: String): Boolean + external fun initialize( + backend: String, + installationIdentity: String, + modelDirectory: String, + primaryFile: String, + ): Boolean + + fun initialize(model: InstalledSpeechModelState.Ready): Boolean = initialize( + backend = model.descriptor.backend.name, + installationIdentity = "${model.descriptor.catalogId}:${model.descriptor.manifest.version}", + modelDirectory = model.directory.absolutePath, + primaryFile = model.descriptor.manifest.files.first().name, + ) external fun reset() diff --git a/app/src/main/java/me/maxistar/voiceinbox/ScheduledTranscriptionWorker.kt b/app/src/main/java/me/maxistar/voiceinbox/ScheduledTranscriptionWorker.kt index 837580a..0f5c7ca 100644 --- a/app/src/main/java/me/maxistar/voiceinbox/ScheduledTranscriptionWorker.kt +++ b/app/src/main/java/me/maxistar/voiceinbox/ScheduledTranscriptionWorker.kt @@ -51,7 +51,7 @@ class ScheduledTranscriptionWorker( documentAccess.requireAppendable(output) folderScanner.requireReadable(folder) - if (SpeechModelRepository( + if (SpeechModelRepository.forActive( applicationContext.noBackupFilesDir.resolve("models"), ).inspectLightweight() !is InstalledSpeechModelState.Ready) return diff --git a/app/src/main/java/me/maxistar/voiceinbox/SettingsActivity.kt b/app/src/main/java/me/maxistar/voiceinbox/SettingsActivity.kt index f88b0e7..a5cd686 100644 --- a/app/src/main/java/me/maxistar/voiceinbox/SettingsActivity.kt +++ b/app/src/main/java/me/maxistar/voiceinbox/SettingsActivity.kt @@ -23,6 +23,7 @@ import java.util.concurrent.Executors internal object VoiceInboxPublicLinks { const val WEBSITE = "https://voiceinbox.simpleditor.org/" + const val DOCUMENTATION = "https://voiceinbox.simpleditor.org/docs/" const val LEGAL = "https://voiceinbox.simpleditor.org/legal/" } @@ -34,6 +35,7 @@ class SettingsActivity : AppCompatActivity() { private lateinit var folderScanner: AudioFolderScanner private lateinit var folderDetail: TextView private lateinit var outputDetail: TextView + private lateinit var modelDetail: TextView private lateinit var scheduledSwitch: SwitchCompat private lateinit var scheduledTime: TextView private lateinit var scheduledTimeDetail: TextView @@ -88,6 +90,7 @@ class SettingsActivity : AppCompatActivity() { settings = settingsStore.load() folderDetail = findViewById(R.id.settingsFolderDetail) outputDetail = findViewById(R.id.settingsOutputDetail) + modelDetail = findViewById(R.id.settingsModelDetail) scheduledTime = findViewById(R.id.scheduledTime) scheduledTimeDetail = findViewById(R.id.scheduledTimeDetail) scheduledSwitch = findViewById(R.id.scheduledSwitch) @@ -101,6 +104,12 @@ class SettingsActivity : AppCompatActivity() { findViewById(R.id.settingsOutputRow).setOnClickListener { showOutputDocumentOptions() } + findViewById(R.id.settingsModelRow).setOnClickListener { + startActivity( + Intent(this, MainActivity::class.java) + .putExtra("open-model-folder-picker", true), + ) + } findViewById(R.id.settingsWebsiteRow).setOnClickListener { openExternalUrl(VoiceInboxPublicLinks.WEBSITE) } @@ -234,6 +243,16 @@ class SettingsActivity : AppCompatActivity() { private fun renderStorage() { renderOutput() renderFolder() + renderModel() + } + + private fun renderModel() { + modelDetail.text = when (val state = SpeechModelRepository.forActive( + noBackupFilesDir.resolve("models"), + ).inspectLightweight()) { + is InstalledSpeechModelState.Ready -> "Active: ${state.descriptor.displayName}. Select a supported local model package to replace it." + else -> "Select a supported local model package." + } } private fun renderOutput() { diff --git a/app/src/main/java/me/maxistar/voiceinbox/SingleFileTranscriber.kt b/app/src/main/java/me/maxistar/voiceinbox/SingleFileTranscriber.kt index 517289f..8ea0a09 100644 --- a/app/src/main/java/me/maxistar/voiceinbox/SingleFileTranscriber.kt +++ b/app/src/main/java/me/maxistar/voiceinbox/SingleFileTranscriber.kt @@ -85,8 +85,8 @@ private class AndroidPlatformAudioDecoder( } private object AndroidPlatformNativeTranscriber : PlatformNativeTranscriber { - override fun initialize(modelDirectory: String): Boolean = - NativeTranscriptionBridge.initialize(modelDirectory) + // The worker prepares the descriptor-aware native engine before it claims audio work. + override fun initialize(modelDirectory: String): Boolean = true override fun transcribeChunk(samples: FloatArray): String? = NativeTranscriptionBridge.transcribeChunk(samples)?.text diff --git a/app/src/main/java/me/maxistar/voiceinbox/SpeechModelDirectoryReader.kt b/app/src/main/java/me/maxistar/voiceinbox/SpeechModelDirectoryReader.kt index 7a38d3d..9abb3c4 100644 --- a/app/src/main/java/me/maxistar/voiceinbox/SpeechModelDirectoryReader.kt +++ b/app/src/main/java/me/maxistar/voiceinbox/SpeechModelDirectoryReader.kt @@ -3,8 +3,10 @@ package me.maxistar.voiceinbox import android.content.ContentResolver import android.net.Uri import android.provider.DocumentsContract -import me.maxistar.voiceinbox.core.SpeechModelManifest +import me.maxistar.voiceinbox.core.* +import java.io.ByteArrayOutputStream import java.io.IOException +import java.io.InputStream data class SpeechModelSourceDocument( val name: String, @@ -12,13 +14,36 @@ data class SpeechModelSourceDocument( val mimeType: String?, ) +data class SpeechModelFolderPackage( + val descriptor: SpeechModelDescriptor, + val requiredDocuments: Map, +) + class SpeechModelDirectoryReader( private val resolver: ContentResolver, ) { - fun requiredDocuments( - treeUri: Uri, - manifest: SpeechModelManifest, - ): Map { + fun inspectPackage(treeUri: Uri): SpeechModelFolderPackage { + val documents = listDocuments(treeUri) + val manifests = documents.filter { + it.mimeType != DocumentsContract.Document.MIME_TYPE_DIR && + it.name == SpeechModelCatalog.PACKAGE_MANIFEST_FILENAME + } + if (manifests.size > 1) throw IOException("Multiple model package manifests were found") + val descriptor = if (manifests.isEmpty()) { + resolveLegacyParakeet(documents) + } else { + val manifestText = resolver.openInputStream(Uri.parse(manifests.single().uri))?.use { input -> + readBounded(input, MAX_PACKAGE_MANIFEST_BYTES).toString(Charsets.UTF_8) + } ?: throw IOException("The model package manifest cannot be read") + resolvePackageIdentity(manifestText) + } + return SpeechModelFolderPackage(descriptor, matchRequiredDocuments(documents, descriptor.manifest)) + } + + fun requiredDocuments(treeUri: Uri, manifest: SpeechModelManifest): Map = + matchRequiredDocuments(listDocuments(treeUri), manifest) + + private fun listDocuments(treeUri: Uri): List { val treeId = runCatching { DocumentsContract.getTreeDocumentId(treeUri) } .getOrElse { throw IOException("The selected model folder is not readable", it) } val childrenUri = DocumentsContract.buildChildDocumentsUriUsingTree(treeUri, treeId) @@ -29,11 +54,9 @@ class SpeechModelDirectoryReader( DocumentsContract.Document.COLUMN_DISPLAY_NAME, DocumentsContract.Document.COLUMN_MIME_TYPE, ), - null, - null, - null, + null, null, null, ) ?: throw IOException("The selected model folder cannot be enumerated") - val documents = cursor.use { + return cursor.use { val id = it.getColumnIndexOrThrow(DocumentsContract.Document.COLUMN_DOCUMENT_ID) val name = it.getColumnIndexOrThrow(DocumentsContract.Document.COLUMN_DISPLAY_NAME) val mime = it.getColumnIndex(DocumentsContract.Document.COLUMN_MIME_TYPE) @@ -42,25 +65,76 @@ class SpeechModelDirectoryReader( val displayName = it.getString(name) ?: continue add( SpeechModelSourceDocument( - name = displayName, - uri = DocumentsContract.buildDocumentUriUsingTree(treeUri, it.getString(id)).toString(), - mimeType = if (mime >= 0 && !it.isNull(mime)) it.getString(mime) else null, + displayName, + DocumentsContract.buildDocumentUriUsingTree(treeUri, it.getString(id)).toString(), + if (mime >= 0 && !it.isNull(mime)) it.getString(mime) else null, ), ) } } } - return matchRequiredDocuments(documents, manifest) } companion object { + const val MAX_PACKAGE_MANIFEST_BYTES = 16 * 1024 + + internal fun readBounded(input: InputStream, maximumBytes: Int): ByteArray { + require(maximumBytes >= 0) { "maximumBytes must not be negative" } + val output = ByteArrayOutputStream(minOf(maximumBytes, 4 * 1024)) + val buffer = ByteArray(minOf(4 * 1024, maximumBytes + 1).coerceAtLeast(1)) + var total = 0 + while (true) { + val read = input.read(buffer) + if (read < 0) break + if (read == 0) continue + total += read + if (total > maximumBytes) { + throw IOException("The model package manifest is too large") + } + output.write(buffer, 0, read) + } + return output.toByteArray() + } + + fun resolvePackageIdentity(json: String): SpeechModelDescriptor { + val allowed = setOf("schemaVersion", "catalogId", "modelVersion") + val trimmed = json.trim() + if (!trimmed.startsWith("{") || !trimmed.endsWith("}")) { + throw IOException("The model package manifest is malformed") + } + val keys = Regex("\"([^\"]+)\"\\s*:").findAll(trimmed).map { it.groupValues[1] }.toList() + if (keys.size != allowed.size || keys.toSet() != allowed) { + throw IOException("The model package manifest has an unsupported structure") + } + fun string(name: String): String? = + Regex("\"$name\"\\s*:\\s*\"([^\"]+)\"").find(trimmed)?.groupValues?.get(1) + val schemaVersion = Regex("\"schemaVersion\"\\s*:\\s*(\\d+)") + .find(trimmed)?.groupValues?.get(1)?.toIntOrNull() + val identity = SpeechModelPackageIdentity( + schemaVersion = schemaVersion ?: throw IOException("The model package manifest is malformed"), + catalogId = string("catalogId") ?: throw IOException("The model package manifest is malformed"), + modelVersion = string("modelVersion") ?: throw IOException("The model package manifest is malformed"), + ) + return SpeechModelCatalog.resolvePackage(identity, SpeechModelPlatform.ANDROID) + ?: throw IOException("This model package is unknown or unsupported on Android") + } + + fun resolveLegacyParakeet(documents: List): SpeechModelDescriptor { + val regularNames = documents + .filter { it.mimeType != DocumentsContract.Document.MIME_TYPE_DIR } + .map { it.name } + val expected = SpeechModelCatalog.defaultModel.manifest.files.map { it.name } + if (regularNames.size != expected.size || regularNames.toSet() != expected.toSet()) { + throw IOException("voice-inbox-model.json is missing from the selected folder") + } + return SpeechModelCatalog.defaultModel + } + fun matchRequiredDocuments( documents: List, manifest: SpeechModelManifest, ): Map { - val regularFiles = documents.filter { - it.mimeType != DocumentsContract.Document.MIME_TYPE_DIR - } + val regularFiles = documents.filter { it.mimeType != DocumentsContract.Document.MIME_TYPE_DIR } return manifest.files.associate { entry -> val matches = regularFiles.filter { it.name == entry.name } when (matches.size) { diff --git a/app/src/main/java/me/maxistar/voiceinbox/SpeechModelDownloadWorker.kt b/app/src/main/java/me/maxistar/voiceinbox/SpeechModelDownloadWorker.kt index 18cb1a5..71e173c 100644 --- a/app/src/main/java/me/maxistar/voiceinbox/SpeechModelDownloadWorker.kt +++ b/app/src/main/java/me/maxistar/voiceinbox/SpeechModelDownloadWorker.kt @@ -21,8 +21,19 @@ class SpeechModelDownloadWorker( appContext: Context, params: WorkerParameters, ) : CoroutineWorker(appContext, params) { + private val model: SpeechModelDescriptor by lazy { + SpeechModelCatalog.resolveNetworkDownload( + SpeechModelPackageIdentity( + schemaVersion = inputData.getInt(KEY_SCHEMA_VERSION, -1), + catalogId = inputData.getString(KEY_CATALOG_ID).orEmpty(), + modelVersion = inputData.getString(KEY_MODEL_VERSION).orEmpty(), + ), + SpeechModelPlatform.ANDROID, + ) ?: throw IllegalArgumentException("Selected speech model is unavailable for download") + } private val repository = SpeechModelRepository( root = applicationContext.noBackupFilesDir.resolve("models"), + descriptor = model, ) private val client = OkHttpClient.Builder() .connectTimeout(30, TimeUnit.SECONDS) @@ -33,6 +44,8 @@ class SpeechModelDownloadWorker( override suspend fun doWork(): Result { return try { installModel() + } catch (error: IllegalArgumentException) { + failure(error.message ?: "Selected speech model is unavailable for download") } catch (error: ForegroundPromotionException) { failure(error.userMessage) } @@ -165,12 +178,26 @@ class SpeechModelDownloadWorker( const val KEY_MESSAGE = SpeechModelInstallationWork.KEY_MESSAGE const val KEY_ERROR = SpeechModelInstallationWork.KEY_ERROR const val KEY_MODEL_PATH = SpeechModelInstallationWork.KEY_MODEL_PATH + const val KEY_SCHEMA_VERSION = "speech-model-schema-version" + const val KEY_CATALOG_ID = "speech-model-catalog-id" + const val KEY_MODEL_VERSION = "speech-model-version" private const val MAX_ATTEMPTS = 3 private const val RETRY_DELAY_MS = 2_000L private const val PROGRESS_STEP_BYTES = 2L * 1024L * 1024L - fun enqueue(context: Context) { - val request = OneTimeWorkRequestBuilder().build() + fun enqueue(context: Context, descriptor: SpeechModelDescriptor = SpeechModelCatalog.defaultModel) { + require(descriptor.distribution.networkDownloadAvailable) { + "Selected speech model is unavailable for download" + } + val request = OneTimeWorkRequestBuilder() + .setInputData( + workDataOf( + KEY_SCHEMA_VERSION to SpeechModelCatalog.PACKAGE_SCHEMA_VERSION, + KEY_CATALOG_ID to descriptor.catalogId, + KEY_MODEL_VERSION to descriptor.manifest.version, + ), + ) + .build() WorkManager.getInstance(context).enqueueUniqueWork( UNIQUE_WORK_NAME, ExistingWorkPolicy.KEEP, diff --git a/app/src/main/java/me/maxistar/voiceinbox/SpeechModelImportWorker.kt b/app/src/main/java/me/maxistar/voiceinbox/SpeechModelImportWorker.kt index 7ee62c6..fc07f0a 100644 --- a/app/src/main/java/me/maxistar/voiceinbox/SpeechModelImportWorker.kt +++ b/app/src/main/java/me/maxistar/voiceinbox/SpeechModelImportWorker.kt @@ -8,18 +8,26 @@ import androidx.work.OneTimeWorkRequestBuilder import androidx.work.WorkManager import androidx.work.WorkerParameters import androidx.work.workDataOf +import me.maxistar.voiceinbox.core.SpeechModelCatalog class SpeechModelImportWorker( appContext: Context, params: WorkerParameters, ) : CoroutineWorker(appContext, params) { - private val repository = SpeechModelRepository( - root = applicationContext.noBackupFilesDir.resolve("models"), - ) - override suspend fun doWork(): Result { val treeUri = inputData.getString(KEY_TREE_URI)?.let(Uri::parse) ?: return failure("No model folder was selected") + val catalogId = inputData.getString(KEY_CATALOG_ID) + ?: return failure("No speech model was selected") + val modelVersion = inputData.getString(KEY_MODEL_VERSION) + ?: return failure("No speech model version was selected") + val descriptor = SpeechModelCatalog.resolveInstallation(catalogId, modelVersion) + ?.takeIf { it.distribution.localImportAvailable } + ?: return failure("The selected speech model is no longer supported") + val repository = SpeechModelRepository( + root = applicationContext.noBackupFilesDir.resolve("models"), + descriptor = descriptor, + ) return try { SpeechModelInstallationWork.promote( worker = this, @@ -31,7 +39,7 @@ class SpeechModelImportWorker( val installed = SpeechModelLocalImporter( resolver = applicationContext.contentResolver, repository = repository, - ).import(treeUri.toString()) { progress -> publishProgress(progress) }.getOrElse { + ).import(treeUri.toString()) { progress -> publishProgress(progress, repository) }.getOrElse { return failure(it.message ?: "Could not import speech model") } SpeechModelPreparation.invalidate(NativeTranscriptionBridge::reset) @@ -45,7 +53,10 @@ class SpeechModelImportWorker( } } - private suspend fun publishProgress(progress: SpeechModelImportProgress) { + private suspend fun publishProgress( + progress: SpeechModelImportProgress, + repository: SpeechModelRepository, + ) { val total = repository.manifest.totalSizeBytes val percent = ((progress.bytesCopied.coerceIn(0, total) * 100) / total).toInt() setProgress( @@ -70,10 +81,22 @@ class SpeechModelImportWorker( companion object { const val KEY_TREE_URI = "tree-uri" + const val KEY_CATALOG_ID = "catalog-id" + const val KEY_MODEL_VERSION = "model-version" - fun enqueue(context: Context, treeUri: Uri) { + fun enqueue( + context: Context, + treeUri: Uri, + descriptor: me.maxistar.voiceinbox.core.SpeechModelDescriptor, + ) { val request = OneTimeWorkRequestBuilder() - .setInputData(workDataOf(KEY_TREE_URI to treeUri.toString())) + .setInputData( + workDataOf( + KEY_TREE_URI to treeUri.toString(), + KEY_CATALOG_ID to descriptor.catalogId, + KEY_MODEL_VERSION to descriptor.manifest.version, + ), + ) .build() WorkManager.getInstance(context).enqueueUniqueWork( SpeechModelInstallationWork.UNIQUE_WORK_NAME, diff --git a/app/src/main/java/me/maxistar/voiceinbox/SpeechModelLocalImporter.kt b/app/src/main/java/me/maxistar/voiceinbox/SpeechModelLocalImporter.kt index 117a148..58bfd92 100644 --- a/app/src/main/java/me/maxistar/voiceinbox/SpeechModelLocalImporter.kt +++ b/app/src/main/java/me/maxistar/voiceinbox/SpeechModelLocalImporter.kt @@ -23,7 +23,12 @@ class SpeechModelLocalImporter( ) : this( repository = repository, requiredDocuments = { treeUri, manifest -> - SpeechModelDirectoryReader(resolver).requiredDocuments(android.net.Uri.parse(treeUri), manifest) + val source = SpeechModelDirectoryReader(resolver).inspectPackage(android.net.Uri.parse(treeUri)) + check(source.descriptor.catalogId == repository.descriptor.catalogId && + source.descriptor.manifest.version == repository.descriptor.manifest.version) { + "The selected model package changed before it could be imported" + } + source.requiredDocuments }, openInputStream = { resolver.openInputStream(android.net.Uri.parse(it)) }, ) diff --git a/app/src/main/java/me/maxistar/voiceinbox/SpeechModelPreparation.kt b/app/src/main/java/me/maxistar/voiceinbox/SpeechModelPreparation.kt index 87e7a71..a3d4432 100644 --- a/app/src/main/java/me/maxistar/voiceinbox/SpeechModelPreparation.kt +++ b/app/src/main/java/me/maxistar/voiceinbox/SpeechModelPreparation.kt @@ -4,14 +4,14 @@ import java.io.File object SpeechModelPreparation { private val lock = Any() - private var preparedDirectory: String? = null + private var preparedInstallation: String? = null fun prepare( repository: SpeechModelRepository, - initializeModel: (File) -> Boolean, + initializeModel: (InstalledSpeechModelState.Ready) -> Boolean, ): Result = synchronized(lock) { - val expected = repository.installedDirectory.canonicalPath - if (preparedDirectory == expected) { + val expected = "${repository.descriptor.backend}:${repository.descriptor.catalogId}:${repository.manifest.version}:${repository.installedDirectory.canonicalPath}" + if (preparedInstallation == expected) { return@synchronized Result.success(repository.installedDirectory) } runCatching { @@ -23,15 +23,15 @@ object SpeechModelPreparation { is InstalledSpeechModelState.Ready -> error("unreachable") } } - check(initializeModel(installed.directory)) { "Speech model failed to load" } - preparedDirectory = installed.directory.canonicalPath + check(initializeModel(installed)) { "Speech model failed to load" } + preparedInstallation = expected installed.directory } } fun invalidate(resetNative: () -> Unit = {}) { synchronized(lock) { - preparedDirectory = null + preparedInstallation = null resetNative() } } diff --git a/app/src/main/java/me/maxistar/voiceinbox/SpeechModelReadinessManager.kt b/app/src/main/java/me/maxistar/voiceinbox/SpeechModelReadinessManager.kt index a4a886d..3caa057 100644 --- a/app/src/main/java/me/maxistar/voiceinbox/SpeechModelReadinessManager.kt +++ b/app/src/main/java/me/maxistar/voiceinbox/SpeechModelReadinessManager.kt @@ -12,9 +12,14 @@ sealed interface SpeechModelReadinessState { } class SpeechModelReadinessManager( - private val repository: SpeechModelRepository, + private val repository: () -> SpeechModelRepository, private val executor: Executor, ) { + constructor( + repository: SpeechModelRepository, + executor: Executor, + ) : this(repository = { repository }, executor = executor) + private val lock = Any() private var cachedState: SpeechModelReadinessState? = null private var checking = false @@ -51,6 +56,7 @@ class SpeechModelReadinessManager( private fun checkModel() { val state = runCatching { + val repository = repository() repository.cleanupStaleState() when (val installed = repository.inspectLightweight()) { is InstalledSpeechModelState.Ready -> SpeechModelReadinessState.Ready(installed.directory) diff --git a/app/src/main/java/me/maxistar/voiceinbox/SpeechModelRepository.kt b/app/src/main/java/me/maxistar/voiceinbox/SpeechModelRepository.kt index 294b21a..f2d5373 100644 --- a/app/src/main/java/me/maxistar/voiceinbox/SpeechModelRepository.kt +++ b/app/src/main/java/me/maxistar/voiceinbox/SpeechModelRepository.kt @@ -9,9 +9,40 @@ import java.nio.file.Files import java.nio.file.StandardCopyOption import java.util.UUID +data class ActiveSpeechModelIdentity( + val catalogId: String, + val modelVersion: String, + val backend: SpeechModelBackend, +) { + fun serialize(): String = + """{"schemaVersion":1,"catalogId":"$catalogId","modelVersion":"$modelVersion","backend":"${backend.name}"}""" + + companion object { + fun parse(text: String): ActiveSpeechModelIdentity? { + fun string(name: String): String? = + Regex("\\\"$name\\\"\\s*:\\s*\\\"([^\\\"]+)\\\"").find(text)?.groupValues?.get(1) + if (Regex("\\\"schemaVersion\\\"\\s*:\\s*(\\d+)").find(text)?.groupValues?.get(1) != "1") return null + val backend = string("backend")?.let { runCatching { SpeechModelBackend.valueOf(it) }.getOrNull() } + ?: return null + return ActiveSpeechModelIdentity( + catalogId = string("catalogId") ?: return null, + modelVersion = string("modelVersion") ?: return null, + backend = backend, + ) + } + } +} + +private data class ActivationTransaction( + val previous: ActiveSpeechModelIdentity?, + val candidateCatalogId: String, + val candidateVersion: String, +) + sealed interface InstalledSpeechModelState { data class Ready( val directory: File, + val descriptor: SpeechModelDescriptor, val verification: Verification = Verification.VERIFIED, ) : InstalledSpeechModelState { enum class Verification { @@ -25,12 +56,20 @@ sealed interface InstalledSpeechModelState { class SpeechModelRepository( private val root: File, - val manifest: SpeechModelManifest = EmbeddedSpeechModel.manifest, + val descriptor: SpeechModelDescriptor, private val usableSpace: (File) -> Long = { it.usableSpace }, private val moveDirectory: (File, File) -> Boolean = { source, destination -> source.renameTo(destination) }, ) { + constructor( + root: File, + manifest: SpeechModelManifest, + usableSpace: (File) -> Long = { it.usableSpace }, + moveDirectory: (File, File) -> Boolean = { source, destination -> source.renameTo(destination) }, + ) : this(root, descriptorForManifest(manifest), usableSpace, moveDirectory) + + val manifest: SpeechModelManifest = descriptor.manifest private val stagingRoot = File(root, "staging") private val installedRoot = File(root, "installed") private val activeVersionFile = File(root, "active-model") @@ -54,7 +93,10 @@ class SpeechModelRepository( return if (missing == null) { InstalledSpeechModelState.Ready( directory = installedDirectory, - verification = if (activeVersionFile.takeIf(File::isFile)?.readText()?.trim() == manifest.version) { + descriptor = descriptor, + verification = if (readActiveIdentity()?.let { + it.catalogId == descriptor.catalogId && it.modelVersion == manifest.version + } == true) { InstalledSpeechModelState.Ready.Verification.VERIFIED } else { InstalledSpeechModelState.Ready.Verification.LEGACY_UNVERIFIED @@ -67,14 +109,14 @@ class SpeechModelRepository( fun inspect(): InstalledSpeechModelState { recoverInterruptedActivation() - val activeVersion = activeVersionFile.takeIf(File::isFile)?.readText()?.trim() - if (activeVersion == manifest.version) { + val activeIdentity = readActiveIdentity() + if (activeIdentity?.catalogId == descriptor.catalogId && activeIdentity.modelVersion == manifest.version) { return recordValidation(validateDirectory(installedDirectory)) } return when (val installed = validateDirectory(installedDirectory)) { is InstalledSpeechModelState.Ready -> { - writeActiveVersion(manifest.version) + writeActiveIdentity(descriptor) invalidModelFile.delete() installed } @@ -156,9 +198,14 @@ class SpeechModelRepository( } installedRoot.mkdirs() + val previousIdentity = readActiveIdentity() + val previousDescriptor = previousIdentity?.let { + SpeechModelCatalog.resolveInstallation(it.catalogId, it.modelVersion) + } + val previousDirectory = previousDescriptor?.let { File(installedRoot, it.manifest.version) } backupDirectory.deleteRecursively() val replacing = installedDirectory.exists() - activationMarker.writeText(if (replacing) MARKER_REPLACEMENT else MARKER_FRESH) + writeActivationMarker(previousIdentity, descriptor) if (replacing) { check(moveDirectory(installedDirectory, backupDirectory)) { "Failed to back up installed model" @@ -168,7 +215,7 @@ class SpeechModelRepository( check(moveDirectory(stagingDirectory, installedDirectory)) { "Failed to activate staged model" } - writeActiveVersion(manifest.version) + writeActiveIdentity(descriptor) invalidModelFile.delete() activationMarker.delete() backupDirectory.deleteRecursively() @@ -178,18 +225,19 @@ class SpeechModelRepository( check(moveDirectory(backupDirectory, installedDirectory)) { "Failed to restore previous speech model" } - writeActiveVersion(manifest.version) + previousIdentity?.let(::writeActiveIdentity) invalidModelFile.delete() } else { - activeVersionFile.delete() + if (previousIdentity == null) activeVersionFile.delete() } activationMarker.delete() throw error } installedRoot.listFiles() - ?.filter { it.name != manifest.version } + ?.filter { it != previousDirectory && it.name != manifest.version } ?.forEach(File::deleteRecursively) + previousDirectory?.takeIf { it != installedDirectory }?.deleteRecursively() installedDirectory } @@ -216,24 +264,50 @@ class SpeechModelRepository( if (installedDirectory.exists()) { backupDirectory.deleteRecursively() } else if (moveDirectory(backupDirectory, installedDirectory)) { - writeActiveVersion(manifest.version) + writeActiveIdentity(descriptor) } } return } - val replacement = activationMarker.readText().trim() == MARKER_REPLACEMENT - if (replacement && backupDirectory.exists()) { + if (activationMarker.readText().trim() == "replacement") { installedDirectory.deleteRecursively() - check(moveDirectory(backupDirectory, installedDirectory)) { + if (backupDirectory.exists()) { + check(moveDirectory(backupDirectory, installedDirectory)) { + "Failed to recover previous speech model" + } + writeActiveIdentity(descriptor) + invalidModelFile.delete() + } + activationMarker.delete() + return + } + + val marker = readActivationMarker() ?: run { + activationMarker.delete() + return + } + val candidateDescriptor = SpeechModelCatalog.resolveInstallation( + marker.candidateCatalogId, + marker.candidateVersion, + ) + val candidateDirectory = File(installedRoot, marker.candidateVersion) + val candidateBackup = File(installedRoot, "${marker.candidateVersion}.backup") + val active = readActiveIdentity() + if (active?.catalogId == marker.candidateCatalogId && active.modelVersion == marker.candidateVersion) { + activationMarker.delete() + candidateBackup.deleteRecursively() + installedRoot.listFiles()?.filter { it.name != marker.candidateVersion }?.forEach(File::deleteRecursively) + return + } + candidateDirectory.deleteRecursively() + if (candidateBackup.exists()) { + check(moveDirectory(candidateBackup, candidateDirectory)) { "Failed to recover previous speech model" } - writeActiveVersion(manifest.version) - invalidModelFile.delete() - } else if (!replacement) { - installedDirectory.deleteRecursively() - activeVersionFile.delete() } + if (marker.previous == null) activeVersionFile.delete() else writeActiveIdentity(marker.previous) + if (candidateDescriptor != null) invalidModelFile.delete() activationMarker.delete() } @@ -251,6 +325,7 @@ class SpeechModelRepository( } return InstalledSpeechModelState.Ready( directory = directory, + descriptor = descriptor, verification = InstalledSpeechModelState.Ready.Verification.VERIFIED, ) } @@ -258,10 +333,39 @@ class SpeechModelRepository( private fun isValidFile(file: File, entry: SpeechModelFile): Boolean = verifyFile(file, entry).isSuccess - private fun writeActiveVersion(version: String) { + private fun readActiveIdentity(): ActiveSpeechModelIdentity? = readActiveIdentity(activeVersionFile) + + private fun writeActivationMarker( + previous: ActiveSpeechModelIdentity?, + candidate: SpeechModelDescriptor, + ) { + val previousText = previous?.serialize()?.replace("\n", "") ?: "null" + activationMarker.writeText( + """{"schemaVersion":1,"previous":$previousText,"candidateCatalogId":"${candidate.catalogId}","candidateVersion":"${candidate.manifest.version}"}""", + ) + } + + private fun readActivationMarker(): ActivationTransaction? { + val text = activationMarker.takeIf(File::isFile)?.readText().orEmpty() + fun string(name: String): String? = + Regex("\\\"$name\\\"\\s*:\\s*\\\"([^\\\"]+)\\\"").find(text)?.groupValues?.get(1) + val previousObject = Regex("\\\"previous\\\"\\s*:\\s*(\\{.*?})\\s*,\\s*\\\"candidateCatalogId", RegexOption.DOT_MATCHES_ALL) + .find(text)?.groupValues?.get(1) + return ActivationTransaction( + previous = previousObject?.let(ActiveSpeechModelIdentity::parse), + candidateCatalogId = string("candidateCatalogId") ?: return null, + candidateVersion = string("candidateVersion") ?: return null, + ) + } + + private fun writeActiveIdentity(descriptor: SpeechModelDescriptor) = writeActiveIdentity( + ActiveSpeechModelIdentity(descriptor.catalogId, descriptor.manifest.version, descriptor.backend), + ) + + private fun writeActiveIdentity(identity: ActiveSpeechModelIdentity) { root.mkdirs() val temporary = File(root, "active-model.${UUID.randomUUID()}.tmp") - temporary.writeText(version) + temporary.writeText(identity.serialize()) runCatching { Files.move( temporary.toPath(), @@ -300,8 +404,38 @@ class SpeechModelRepository( } companion object { - private const val MARKER_REPLACEMENT = "replacement" - private const val MARKER_FRESH = "fresh" + fun forActive(root: File): SpeechModelRepository { + val descriptor = readActiveIdentity(File(root, "active-model"))?.let { identity -> + SpeechModelCatalog.resolveInstallation(identity.catalogId, identity.modelVersion) + ?.takeIf { it.backend == identity.backend } + } ?: SpeechModelCatalog.defaultModel + return SpeechModelRepository(root, descriptor) + } + + private fun descriptorForManifest(manifest: SpeechModelManifest): SpeechModelDescriptor = + SpeechModelCatalog.models.firstOrNull { it.manifest == manifest } ?: SpeechModelDescriptor( + catalogId = manifest.modelId, + displayName = manifest.modelId, + backend = SpeechModelBackend.PARAKEET_TDT_ONNX, + manifest = manifest, + distribution = SpeechModelDistribution(false, true), + languages = SpeechModelLanguageCoverage("Test model", emptyList()), + maturity = SpeechModelMaturity.EXPERIMENTAL, + attribution = SpeechModelAttribution("", "", "", "", "", ""), + supportedPlatforms = setOf(SpeechModelPlatform.ANDROID), + ) + + private fun readActiveIdentity(file: File): ActiveSpeechModelIdentity? { + val text = file.takeIf(File::isFile)?.readText()?.trim().orEmpty() + if (text.isEmpty()) return null + if (!text.startsWith("{")) { + val legacy = SpeechModelCatalog.models.singleOrNull { it.manifest.version == text } + ?: return null + return ActiveSpeechModelIdentity(legacy.catalogId, legacy.manifest.version, legacy.backend) + } + return ActiveSpeechModelIdentity.parse(text) + } + fun sha256(file: File): String { val digest = MessageDigest.getInstance("SHA-256") FileInputStream(file).use { input -> diff --git a/app/src/main/java/me/maxistar/voiceinbox/TranscriptionWorker.kt b/app/src/main/java/me/maxistar/voiceinbox/TranscriptionWorker.kt index 2ece957..dee7158 100644 --- a/app/src/main/java/me/maxistar/voiceinbox/TranscriptionWorker.kt +++ b/app/src/main/java/me/maxistar/voiceinbox/TranscriptionWorker.kt @@ -38,13 +38,12 @@ class TranscriptionWorker( foregroundInfo = foreground("Preparing transcription", 0, true), source = SpeechModelInstallationWork.Source.TRANSCRIPTION, ) - val modelRepository = SpeechModelRepository( + val modelRepository = SpeechModelRepository.forActive( applicationContext.noBackupFilesDir.resolve("models"), ) publish("Preparing speech model", null, null, 0, 0, null, null) - SpeechModelPreparation.prepare(modelRepository) { directory -> - NativeTranscriptionBridge.initialize(directory.absolutePath) - }.getOrElse { return@withContext failure(it.message ?: "Speech model preparation failed") } + SpeechModelPreparation.prepare(modelRepository, NativeTranscriptionBridge::initialize) + .getOrElse { return@withContext failure(it.message ?: "Speech model preparation failed") } val batch = BatchTranscriptionUseCase( catalog = catalog, diff --git a/app/src/main/res/layout/activity_settings.xml b/app/src/main/res/layout/activity_settings.xml index aae5355..0aa3378 100644 --- a/app/src/main/res/layout/activity_settings.xml +++ b/app/src/main/res/layout/activity_settings.xml @@ -80,6 +80,31 @@ android:text="@string/output_not_selected" /> + + + + + + + + + diff --git a/app/src/main/res/values/strings.xml b/app/src/main/res/values/strings.xml index dfc7c77..76435d7 100644 --- a/app/src/main/res/values/strings.xml +++ b/app/src/main/res/values/strings.xml @@ -8,6 +8,7 @@ Audio folder refresh is idle Refreshing audio folder Settings + Documentation Settings Storage Audio folder diff --git a/app/src/test/java/me/maxistar/voiceinbox/AndroidTaskActionRouterTest.kt b/app/src/test/java/me/maxistar/voiceinbox/AndroidTaskActionRouterTest.kt index 433ffc6..7180f94 100644 --- a/app/src/test/java/me/maxistar/voiceinbox/AndroidTaskActionRouterTest.kt +++ b/app/src/test/java/me/maxistar/voiceinbox/AndroidTaskActionRouterTest.kt @@ -25,8 +25,8 @@ class AndroidTaskActionRouterTest { val calls = mutableListOf>() val router = AndroidTaskActionRouter({ state }) { kind, entry -> calls += kind to entry } - assertTrue(router.route(request("setup:model", null, TaskActionKind.IMPORT_MODEL))) - assertEquals(TaskActionKind.IMPORT_MODEL, calls.last().first) + assertTrue(router.route(request("setup:model", null, TaskActionKind.DOWNLOAD_MODEL))) + assertEquals(TaskActionKind.DOWNLOAD_MODEL, calls.last().first) assertNull(calls.last().second) state = state(entries = listOf(entry(9))) @@ -188,7 +188,10 @@ class AndroidTaskActionRouterTest { ) = AndroidTaskListSnapshotMapper.state( AndroidMainScreenInput( filter = filter, - model = ModelSetupSnapshot(if (modelReady) ModelSetupSnapshotState.READY else ModelSetupSnapshotState.REQUIRED), + model = ModelSetupSnapshot( + if (modelReady) ModelSetupSnapshotState.READY else ModelSetupSnapshotState.REQUIRED, + downloadAvailable = !modelReady, + ), output = OutputSetupSnapshot(OutputSetupSnapshotState.READY), folder = FolderSetupSnapshot(FolderSetupSnapshotState.READY), entries = entries, diff --git a/app/src/test/java/me/maxistar/voiceinbox/SpeechModelDirectoryReaderTest.kt b/app/src/test/java/me/maxistar/voiceinbox/SpeechModelDirectoryReaderTest.kt index e86b277..9ef0857 100644 --- a/app/src/test/java/me/maxistar/voiceinbox/SpeechModelDirectoryReaderTest.kt +++ b/app/src/test/java/me/maxistar/voiceinbox/SpeechModelDirectoryReaderTest.kt @@ -3,9 +3,11 @@ package me.maxistar.voiceinbox import android.provider.DocumentsContract import me.maxistar.voiceinbox.core.SpeechModelFile import me.maxistar.voiceinbox.core.SpeechModelManifest +import me.maxistar.voiceinbox.core.SpeechModelCatalog import org.junit.Assert.assertEquals import org.junit.Assert.assertTrue import org.junit.Test +import java.io.ByteArrayInputStream class SpeechModelDirectoryReaderTest { @Test @@ -57,6 +59,64 @@ class SpeechModelDirectoryReaderTest { assertTrue(error?.message?.contains("model.bin") == true) } + @Test + fun packageIdentitySelectsExactCatalogModel() { + val descriptor = SpeechModelDirectoryReader.resolvePackageIdentity( + """{"schemaVersion":1,"catalogId":"whisper-tiny-multilingual","modelVersion":"whisper-tiny-ggml-f16-r1"}""", + ) + + assertEquals(SpeechModelCatalog.whisperTinyMultilingual, descriptor) + } + + @Test + fun malformedUnknownAndUntrustedPackageMetadataAreRejected() { + listOf( + "not-json", + """{"schemaVersion":2,"catalogId":"whisper-tiny-multilingual","modelVersion":"whisper-tiny-ggml-f16-r1"}""", + """{"schemaVersion":1,"catalogId":"unknown","modelVersion":"unknown"}""", + """{"schemaVersion":1,"catalogId":"whisper-tiny-multilingual","modelVersion":"whisper-tiny-ggml-f16-r1","files":[]}""", + ).forEach { json -> + assertTrue(runCatching { SpeechModelDirectoryReader.resolvePackageIdentity(json) }.isFailure) + } + } + + @Test + fun legacyLayoutIsRestrictedToExactParakeetFiles() { + val files = SpeechModelCatalog.defaultModel.manifest.files.map { document(it.name) } + assertEquals( + SpeechModelCatalog.defaultModel, + SpeechModelDirectoryReader.resolveLegacyParakeet(files), + ) + assertTrue( + runCatching { + SpeechModelDirectoryReader.resolveLegacyParakeet(files + document("extra.txt")) + }.isFailure, + ) + assertTrue( + runCatching { + SpeechModelDirectoryReader.resolveLegacyParakeet( + SpeechModelCatalog.whisperTinyMultilingual.manifest.files.map { document(it.name) }, + ) + }.isFailure, + ) + } + + @Test + fun boundedManifestReadingWorksWithoutApi33InputStreamMethods() { + val payload = "model package".toByteArray() + + assertTrue( + payload.contentEquals( + SpeechModelDirectoryReader.readBounded(ByteArrayInputStream(payload), payload.size), + ), + ) + assertTrue( + runCatching { + SpeechModelDirectoryReader.readBounded(ByteArrayInputStream(payload), payload.size - 1) + }.exceptionOrNull()?.message?.contains("too large") == true, + ) + } + private fun document(name: String, mime: String = "application/octet-stream") = SpeechModelSourceDocument(name, "content://test/$name/${System.nanoTime()}", mime) diff --git a/app/src/test/java/me/maxistar/voiceinbox/SpeechModelReadinessManagerTest.kt b/app/src/test/java/me/maxistar/voiceinbox/SpeechModelReadinessManagerTest.kt index 8669454..98990d4 100644 --- a/app/src/test/java/me/maxistar/voiceinbox/SpeechModelReadinessManagerTest.kt +++ b/app/src/test/java/me/maxistar/voiceinbox/SpeechModelReadinessManagerTest.kt @@ -60,6 +60,34 @@ class SpeechModelReadinessManagerTest { assertTrue(cachedStates.single() is SpeechModelReadinessState.Ready) } + @Test + fun invalidationInspectsTheNewRepositoryAfterActiveModelChanges() { + val executor = QueueingExecutor() + val firstRepository = readyRepository() + val secondRepository = SpeechModelRepository( + root = File(temporaryFolder.root, "missing-model"), + manifest = testManifest, + usableSpace = { Long.MAX_VALUE }, + ) + var activeRepository = firstRepository + val manager = SpeechModelReadinessManager( + repository = { activeRepository }, + executor = executor, + ) + val states = mutableListOf() + + manager.refresh { states += it } + executor.runNext() + assertTrue(states.last() is SpeechModelReadinessState.Ready) + + activeRepository = secondRepository + manager.invalidate() + manager.refresh { states += it } + executor.runNext() + + assertEquals(SpeechModelReadinessState.Missing, states.last()) + } + private fun readyRepository(): SpeechModelRepository { val repository = SpeechModelRepository( root = File(temporaryFolder.root, "models"), diff --git a/app/src/test/java/me/maxistar/voiceinbox/SpeechModelRepositoryTest.kt b/app/src/test/java/me/maxistar/voiceinbox/SpeechModelRepositoryTest.kt index 6f2b40b..bc867cd 100644 --- a/app/src/test/java/me/maxistar/voiceinbox/SpeechModelRepositoryTest.kt +++ b/app/src/test/java/me/maxistar/voiceinbox/SpeechModelRepositoryTest.kt @@ -70,6 +70,27 @@ class SpeechModelRepositoryTest { assertTrue(File(temporaryFolder.root, "models/active-model").isFile) } + @Test + fun catalogDescriptorRecognizesExistingProductionLayoutWithoutMigration() { + val descriptor = SpeechModelCatalog.defaultModel + val root = File(temporaryFolder.root, "models") + val installed = File(root, "installed/${descriptor.manifest.version}").apply { mkdirs() } + descriptor.manifest.files.forEach { installed.resolve(it.name).createNewFile() } + File(root, "active-model").apply { + parentFile?.mkdirs() + writeText(descriptor.manifest.version) + } + + val state = SpeechModelRepository(root, descriptor.manifest).inspectLightweight() + + assertTrue(state is InstalledSpeechModelState.Ready) + assertEquals( + InstalledSpeechModelState.Ready.Verification.VERIFIED, + (state as InstalledSpeechModelState.Ready).verification, + ) + assertEquals(installed.canonicalFile, state.directory.canonicalFile) + } + @Test fun corruptAndIncompleteModelsAreRejected() { val repository = repository() @@ -219,6 +240,63 @@ class SpeechModelRepositoryTest { assertTrue(repository.inspectLightweight() is InstalledSpeechModelState.Ready) } + @Test + fun activeReceiptRoundTripsBackendAndIdentity() { + val identity = ActiveSpeechModelIdentity( + "whisper-tiny-multilingual", + "whisper-tiny-ggml-f16-r1", + SpeechModelBackend.WHISPER_CPP, + ) + + assertEquals(identity, ActiveSpeechModelIdentity.parse(identity.serialize())) + assertEquals(null, ActiveSpeechModelIdentity.parse("not-json")) + } + + @Test + fun crossModelActivationCommitsTypedReceiptAndRemovesPreviousPayload() { + val root = File(temporaryFolder.root, "switch-models") + val first = SpeechModelRepository(root, descriptor("first", "first-version"), { Long.MAX_VALUE }) + first.prepareFreshImport().getOrThrow() + writeValidStaging(first) + val oldDirectory = first.activate().getOrThrow() + + val second = SpeechModelRepository(root, descriptor("second", "second-version"), { Long.MAX_VALUE }) + second.prepareFreshImport().getOrThrow() + writeValidStaging(second) + second.activate().getOrThrow() + + assertFalse(oldDirectory.exists()) + assertTrue(second.inspectLightweight() is InstalledSpeechModelState.Ready) + val receipt = root.resolve("active-model").readText() + assertTrue(receipt.contains("\"catalogId\":\"second\"")) + assertTrue(receipt.contains("\"backend\":\"PARAKEET_TDT_ONNX\"")) + } + + @Test + fun failedCrossModelActivationPreservesPreviousPayloadAndReceipt() { + val root = File(temporaryFolder.root, "failed-switch-models") + val first = SpeechModelRepository(root, descriptor("first", "first-version"), { Long.MAX_VALUE }) + first.prepareFreshImport().getOrThrow() + writeValidStaging(first) + val oldDirectory = first.activate().getOrThrow() + val oldReceipt = root.resolve("active-model").readText() + val second = SpeechModelRepository( + root, + descriptor("second", "second-version"), + { Long.MAX_VALUE }, + { source, destination -> + if (source.path.contains("${File.separator}staging${File.separator}")) false + else source.renameTo(destination) + }, + ) + second.prepareFreshImport().getOrThrow() + writeValidStaging(second) + + assertTrue(second.activate().isFailure) + assertTrue(oldDirectory.isDirectory) + assertEquals(oldReceipt, root.resolve("active-model").readText()) + } + private fun repository() = SpeechModelRepository( root = File(temporaryFolder.root, "models"), manifest = testManifest, @@ -232,6 +310,18 @@ class SpeechModelRepositoryTest { } } + private fun descriptor(catalogId: String, version: String) = SpeechModelDescriptor( + catalogId = catalogId, + displayName = catalogId, + backend = SpeechModelBackend.PARAKEET_TDT_ONNX, + manifest = testManifest.copy(version = version), + distribution = SpeechModelDistribution(false, true), + languages = SpeechModelLanguageCoverage("Test", emptyList()), + maturity = SpeechModelMaturity.EXPERIMENTAL, + attribution = SpeechModelAttribution("", "", "", "", "", ""), + supportedPlatforms = setOf(SpeechModelPlatform.ANDROID), + ) + companion object { private val testFiles = linkedMapOf( "model.bin" to "model".toByteArray(), diff --git a/app/src/test/java/me/maxistar/voiceinbox/TaskListDisplayItemsTest.kt b/app/src/test/java/me/maxistar/voiceinbox/TaskListDisplayItemsTest.kt index 623cd5c..1567da9 100644 --- a/app/src/test/java/me/maxistar/voiceinbox/TaskListDisplayItemsTest.kt +++ b/app/src/test/java/me/maxistar/voiceinbox/TaskListDisplayItemsTest.kt @@ -21,6 +21,12 @@ import org.junit.Assert.assertTrue import org.junit.Test class TaskListDisplayItemsTest { + @Test + fun activeBatchSuppressesStructuralRecyclerViewAnimationsOnlyWhileActive() { + assertTrue(AndroidTaskListAnimationPolicy.suppressStructuralAnimations(true)) + assertFalse(AndroidTaskListAnimationPolicy.suppressStructuralAnimations(false)) + } + @Test fun setupBatchAndAudioItemsHaveStableKindsAndPlacement() { val items = items( diff --git a/app/src/test/java/me/maxistar/voiceinbox/VoiceInboxPublicLinksTest.kt b/app/src/test/java/me/maxistar/voiceinbox/VoiceInboxPublicLinksTest.kt index f03799a..4ff7feb 100644 --- a/app/src/test/java/me/maxistar/voiceinbox/VoiceInboxPublicLinksTest.kt +++ b/app/src/test/java/me/maxistar/voiceinbox/VoiceInboxPublicLinksTest.kt @@ -5,8 +5,9 @@ import org.junit.Test class VoiceInboxPublicLinksTest { @Test - fun settingsLinksUseCanonicalVoiceInboxOrigin() { + fun publicLinksUseCanonicalVoiceInboxOrigin() { assertEquals("https://voiceinbox.simpleditor.org/", VoiceInboxPublicLinks.WEBSITE) + assertEquals("https://voiceinbox.simpleditor.org/docs/", VoiceInboxPublicLinks.DOCUMENTATION) assertEquals("https://voiceinbox.simpleditor.org/legal/", VoiceInboxPublicLinks.LEGAL) } } diff --git a/iosApp/VoiceInbox.xcodeproj/project.pbxproj b/iosApp/VoiceInbox.xcodeproj/project.pbxproj index 528dbce..f8f2280 100644 --- a/iosApp/VoiceInbox.xcodeproj/project.pbxproj +++ b/iosApp/VoiceInbox.xcodeproj/project.pbxproj @@ -29,6 +29,7 @@ C30000000000000000000001 /* IosAudioMetadataFormatter.swift in Sources */ = {isa = PBXBuildFile; fileRef = C30000000000000000000002 /* IosAudioMetadataFormatter.swift */; }; C30000000000000000000003 /* IosTranscriptReview.swift in Sources */ = {isa = PBXBuildFile; fileRef = C30000000000000000000004 /* IosTranscriptReview.swift */; }; C30000000000000000000005 /* IosTranscriptReviewTests.swift in Sources */ = {isa = PBXBuildFile; fileRef = C30000000000000000000006 /* IosTranscriptReviewTests.swift */; }; + D40000000000000000000001 /* WhisperMobileSpikeTests.swift in Sources */ = {isa = PBXBuildFile; fileRef = D40000000000000000000002 /* WhisperMobileSpikeTests.swift */; }; /* End PBXBuildFile section */ /* Begin PBXContainerItemProxy section */ @@ -90,6 +91,7 @@ C30000000000000000000002 /* IosAudioMetadataFormatter.swift */ = {isa = PBXFileReference; lastKnownFileType = sourcecode.swift; path = IosAudioMetadataFormatter.swift; sourceTree = ""; }; C30000000000000000000004 /* IosTranscriptReview.swift */ = {isa = PBXFileReference; lastKnownFileType = sourcecode.swift; path = IosTranscriptReview.swift; sourceTree = ""; }; C30000000000000000000006 /* IosTranscriptReviewTests.swift */ = {isa = PBXFileReference; lastKnownFileType = sourcecode.swift; path = IosTranscriptReviewTests.swift; sourceTree = ""; }; + D40000000000000000000002 /* WhisperMobileSpikeTests.swift */ = {isa = PBXFileReference; lastKnownFileType = sourcecode.swift; path = WhisperMobileSpikeTests.swift; sourceTree = ""; }; /* End PBXFileReference section */ /* Begin PBXFrameworksBuildPhase section */ @@ -178,6 +180,7 @@ A10000000000000000000002 /* DeferredSpeechModelLoadingTests.swift */, B20000000000000000000004 /* IosInlineOnboardingTests.swift */, C30000000000000000000006 /* IosTranscriptReviewTests.swift */, + D40000000000000000000002 /* WhisperMobileSpikeTests.swift */, ); path = VoiceInboxTests; sourceTree = ""; @@ -388,6 +391,7 @@ A10000000000000000000001 /* DeferredSpeechModelLoadingTests.swift in Sources */, B20000000000000000000003 /* IosInlineOnboardingTests.swift in Sources */, C30000000000000000000005 /* IosTranscriptReviewTests.swift in Sources */, + D40000000000000000000001 /* WhisperMobileSpikeTests.swift in Sources */, ); runOnlyForDeploymentPostprocessing = 0; }; @@ -558,6 +562,8 @@ "-lc++", "-weak_framework", CoreML, + "-framework", + Accelerate, ); PRODUCT_BUNDLE_IDENTIFIER = me.maxistar.voiceinbox.ios; PRODUCT_NAME = "$(TARGET_NAME)"; @@ -607,6 +613,8 @@ "-lc++", "-weak_framework", CoreML, + "-framework", + Accelerate, ); PRODUCT_BUNDLE_IDENTIFIER = me.maxistar.voiceinbox.ios; PRODUCT_NAME = "$(TARGET_NAME)"; diff --git a/iosApp/VoiceInbox.xcodeproj/xcshareddata/xcschemes/VoiceInbox.xcscheme b/iosApp/VoiceInbox.xcodeproj/xcshareddata/xcschemes/VoiceInbox.xcscheme index ea1695e..b02006b 100644 --- a/iosApp/VoiceInbox.xcodeproj/xcshareddata/xcschemes/VoiceInbox.xcscheme +++ b/iosApp/VoiceInbox.xcodeproj/xcshareddata/xcschemes/VoiceInbox.xcscheme @@ -68,6 +68,38 @@ debugServiceExtension = "internal" allowLocationSimulation = "YES" queueDebuggingEnableBacktraceRecording = "Yes"> + + + + + + + + + + + + + + Void + let onCancel: () -> Void + + func makeUIViewController(context: Context) -> UIDocumentPickerViewController { + context.coordinator.makeDocumentPicker() + } + + func updateUIViewController(_ uiViewController: UIDocumentPickerViewController, context: Context) { + } + + func makeCoordinator() -> Coordinator { + Coordinator(onPick: onPick, onCancel: onCancel) + } + + final class Coordinator: NSObject, UIDocumentPickerDelegate { + private let onPick: (URL) -> Void + private let onCancel: () -> Void + + init(onPick: @escaping (URL) -> Void, onCancel: @escaping () -> Void) { + self.onPick = onPick + self.onCancel = onCancel + } + + func makeDocumentPicker() -> UIDocumentPickerViewController { + let temporaryDirectory = FileManager.default.temporaryDirectory + .appendingPathComponent("VoiceInboxOutput-\(UUID().uuidString)", isDirectory: true) + let sourceURL = temporaryDirectory.appendingPathComponent("Voice Inbox Transcripts.md") + do { + try FileManager.default.createDirectory( + at: temporaryDirectory, + withIntermediateDirectories: true + ) + try Data().write(to: sourceURL, options: .atomic) + self.temporaryDirectory = temporaryDirectory + let controller = UIDocumentPickerViewController(forExporting: [sourceURL], asCopy: true) + controller.delegate = self + return controller + } catch { + onCancel() + let controller = UIDocumentPickerViewController( + forOpeningContentTypes: [.plainText], + asCopy: false + ) + controller.delegate = self + return controller + } + } + + func documentPicker(_ controller: UIDocumentPickerViewController, didPickDocumentsAt urls: [URL]) { + guard let url = urls.first else { + finishCancelled() + return + } + cleanUpTemporarySource() + onPick(url) + } + + func documentPickerWasCancelled(_ controller: UIDocumentPickerViewController) { + finishCancelled() + } + + private var temporaryDirectory: URL? + + private func finishCancelled() { + cleanUpTemporarySource() + onCancel() + } + + private func cleanUpTemporarySource() { + guard let temporaryDirectory else { return } + try? FileManager.default.removeItem(at: temporaryDirectory) + self.temporaryDirectory = nil + } + } +} + struct IosSelectedOutputDocument { let id: String let url: URL diff --git a/iosApp/VoiceInbox/IosMainScreenShellState.swift b/iosApp/VoiceInbox/IosMainScreenShellState.swift index 041bc81..96f0874 100644 --- a/iosApp/VoiceInbox/IosMainScreenShellState.swift +++ b/iosApp/VoiceInbox/IosMainScreenShellState.swift @@ -53,6 +53,7 @@ enum IosTaskActionRoute: Equatable { case modelDownload case modelImport case modelCancel + case outputCreation case outputSelection case folderSelection case folderRefresh @@ -70,6 +71,7 @@ enum IosTaskActionRouter { case .downloadModel, .retryModelDownload: .modelDownload case .importModel: .modelImport case .cancelModelDownload: .modelCancel + case .createOutput: .outputCreation case .selectOutput: .outputSelection case .selectFolder: .folderSelection case .refreshFolder: .folderRefresh @@ -205,7 +207,9 @@ final class IosMainScreenShellState { installationPhase: installing ? installationPhase : nil, progressPercent: progress.map { KotlinInt(int: Int32($0)) }, downloadAvailable: downloadAvailable, - canCancel: canCancel + canCancel: canCancel, + selectedModel: nil, + downloadChoices: [] ) } diff --git a/iosApp/VoiceInbox/IosSingleFileTranscriptionController.swift b/iosApp/VoiceInbox/IosSingleFileTranscriptionController.swift index 90d13f2..9855401 100644 --- a/iosApp/VoiceInbox/IosSingleFileTranscriptionController.swift +++ b/iosApp/VoiceInbox/IosSingleFileTranscriptionController.swift @@ -17,6 +17,14 @@ enum IosTranscriptionPreparationGate { @_silgen_name("voiceinbox_transcription_initialize") private func voiceinbox_transcription_initialize(_ modelDirectory: UnsafePointer) -> Bool +@_silgen_name("voiceinbox_transcription_initialize_configured") +private func voiceinbox_transcription_initialize_configured( + _ backend: UnsafePointer, + _ installationIdentity: UnsafePointer, + _ modelDirectory: UnsafePointer, + _ primaryFile: UnsafePointer +) -> Bool + @_silgen_name("voiceinbox_transcription_transcribe_chunk_json") private func voiceinbox_transcription_transcribe_chunk_json( _ samples: UnsafePointer?, @@ -560,6 +568,28 @@ final class IosNativeTranscriber: PlatformNativeTranscriber { modelDirectory.withCString { voiceinbox_transcription_initialize($0) } } + static func prepare( + backend: String, + installationIdentity: String, + modelDirectory: String, + primaryFile: String + ) -> Bool { + backend.withCString { backendPointer in + installationIdentity.withCString { identityPointer in + modelDirectory.withCString { directoryPointer in + primaryFile.withCString { primaryPointer in + voiceinbox_transcription_initialize_configured( + backendPointer, + identityPointer, + directoryPointer, + primaryPointer + ) + } + } + } + } + } + static func resetModel() { voiceinbox_transcription_reset() } diff --git a/iosApp/VoiceInbox/IosSpeechModelStore.swift b/iosApp/VoiceInbox/IosSpeechModelStore.swift index 29d68f4..cadedf3 100644 --- a/iosApp/VoiceInbox/IosSpeechModelStore.swift +++ b/iosApp/VoiceInbox/IosSpeechModelStore.swift @@ -10,10 +10,100 @@ enum IosSpeechModelInstallationState: Equatable { case invalid } +struct IosSpeechModelFileDescriptor: Equatable, Hashable { + let name: String + let sizeBytes: Int64 + let sha256: String + let downloadURL: String +} + +struct IosSpeechModelDescriptor: Equatable, Hashable, Identifiable { + let catalogId: String + let displayName: String + let modelVersion: String + let backend: String + let maturity: String + let languageSummary: String + let networkDownloadAvailable: Bool + let localImportAvailable: Bool + let safetyMarginBytes: Int64 + let files: [IosSpeechModelFileDescriptor] + + var id: String { "\(catalogId):\(modelVersion)" } + var primaryFile: String { files.first?.name ?? "" } + var totalSizeBytes: Int64 { files.reduce(0) { $0 + $1.sizeBytes } } + + init(_ descriptor: SpeechModelDescriptor) { + catalogId = descriptor.catalogId + displayName = descriptor.displayName + modelVersion = descriptor.manifest.version + backend = descriptor.backend.name + maturity = descriptor.maturity.name.capitalized + languageSummary = descriptor.languages.summary + networkDownloadAvailable = descriptor.distribution.networkDownloadAvailable + localImportAvailable = descriptor.distribution.localImportAvailable + safetyMarginBytes = descriptor.manifest.safetyMarginBytes + files = descriptor.manifest.files.map { + IosSpeechModelFileDescriptor( + name: $0.name, + sizeBytes: $0.sizeBytes, + sha256: $0.sha256, + downloadURL: descriptor.manifest.downloadUrl(file: $0) + ) + } + } + + static var supported: [IosSpeechModelDescriptor] { + SpeechModelCatalog.shared.modelsFor(platform: .ios).map(IosSpeechModelDescriptor.init) + } + + static var defaultModel: IosSpeechModelDescriptor { + IosSpeechModelDescriptor(SpeechModelCatalog.shared.defaultModel) + } + + static func resolve(catalogId: String, modelVersion: String) -> IosSpeechModelDescriptor? { + supported.first { $0.catalogId == catalogId && $0.modelVersion == modelVersion } + } +} + +struct IosSpeechModelInstallation: Codable, Equatable { + static let receiptSchemaVersion = 2 + + let receiptSchemaVersion: Int + let packageSchemaVersion: Int + let catalogId: String + let modelVersion: String + let backend: String + let installationGeneration: String + + var identity: String { + "\(catalogId):\(modelVersion):\(backend):\(installationGeneration)" + } +} + +struct IosSpeechModelCandidate: Identifiable, Equatable { + let descriptor: IosSpeechModelDescriptor + let sourceURL: URL + var id: String { descriptor.id } +} + struct IosSpeechModelStatus { let directory: URL let installationState: IosSpeechModelInstallationState let missingFiles: [String] + let activeInstallation: IosSpeechModelInstallation? + + init( + directory: URL, + installationState: IosSpeechModelInstallationState, + missingFiles: [String], + activeInstallation: IosSpeechModelInstallation? = nil + ) { + self.directory = directory + self.installationState = installationState + self.missingFiles = missingFiles + self.activeInstallation = activeInstallation + } var isReady: Bool { installationState == .installedVerified || installationState == .installedLegacy @@ -67,6 +157,10 @@ enum IosSpeechModelPaths { static var invalidFile: URL { applicationSupportDirectory.appendingPathComponent("SpeechModel.invalid") } + + static var backupReceiptFile: URL { + applicationSupportDirectory.appendingPathComponent("SpeechModel.receipt.previous") + } } struct IosSpeechModelDownloadProgress { @@ -85,11 +179,36 @@ final class IosSpeechModelStore: ObservableObject { @Published private(set) var status: IosSpeechModelStatus @Published private(set) var isInstalling = false @Published private(set) var downloadProgress: IosSpeechModelDownloadProgress? + @Published var selectedDownloadModel: IosSpeechModelDescriptor = .defaultModel @Published private(set) var runtimeState = SpeechModelRuntimeState.unloaded @Published var message: String? + @Published var pendingCandidate: IosSpeechModelCandidate? + @Published private(set) var installationError: String? + + var availableModels: [IosSpeechModelDescriptor] { IosSpeechModelDescriptor.supported } + var activeDescriptor: IosSpeechModelDescriptor? { + guard let active = status.activeInstallation else { + return status.isReady ? .defaultModel : nil + } + return .resolve(catalogId: active.catalogId, modelVersion: active.modelVersion) + } + var actionableErrorMessage: String? { + if let installationError { return installationError } + guard let message else { return nil } + let lower = message.lowercased() + return ["could not", "not enough", "not supported", "malformed", "missing", "invalid", "wait for"] + .contains(where: lower.contains) ? message : nil + } + + func clearActionableError() { + installationError = nil + if actionableErrorMessage != nil { message = nil } + } private var downloadTask: Task? private var preparationTask: Task? + private var pendingSecurityScopeActive = false + private var pendingSecurityScopeURL: URL? private let installationDirectory: URL private let inspectInstallation: @Sendable (URL) -> IosSpeechModelStatus private let validateInstallation: @Sendable (URL) -> [String] @@ -103,10 +222,29 @@ final class IosSpeechModelStore: ObservableObject { self.init( directory: IosSpeechModelPaths.modelDirectory, inspectInstallation: { Self.inspectLightweight(directory: $0) }, - validateInstallation: { Self.validateModelFiles(in: $0).missingFiles }, - prepareNative: { IosNativeTranscriber.prepare(modelDirectory: $0) }, + validateInstallation: { directory in + guard let descriptor = Self.inspectLightweight(directory: directory).activeInstallation.flatMap({ + IosSpeechModelDescriptor.resolve(catalogId: $0.catalogId, modelVersion: $0.modelVersion) + }) ?? (Self.inspectLightweight(directory: directory).isReady ? .defaultModel : nil) else { + return ["active model identity"] + } + return Self.validateModelFiles(in: directory, descriptor: descriptor).missingFiles + }, + prepareNative: { directory in + let status = Self.inspectLightweight(directory: URL(fileURLWithPath: directory)) + guard let descriptor = status.activeInstallation.flatMap({ + IosSpeechModelDescriptor.resolve(catalogId: $0.catalogId, modelVersion: $0.modelVersion) + }) ?? (status.isReady ? .defaultModel : nil) else { return false } + let identity = status.activeInstallation?.identity ?? "legacy:\(descriptor.id)" + return IosNativeTranscriber.prepare( + backend: descriptor.backend, + installationIdentity: identity, + modelDirectory: directory, + primaryFile: descriptor.primaryFile + ) + }, nativeError: { IosNativeTranscriber.consumeLastError() }, - recordVerified: { Self.recordVerifiedInstallation() }, + recordVerified: { Self.recordVerifiedInstallationIfNeeded() }, recordInvalid: { Self.recordInvalidInstallation($0) }, resetNative: { IosNativeTranscriber.resetModel() } ) @@ -153,38 +291,111 @@ final class IosSpeechModelStore: ObservableObject { status = inspectInstallation(installationDirectory) } - func installModel(from sourceURL: URL) { + func inspectModelPackage(from sourceURL: URL) { guard !isBusy else { return } + releasePendingSecurityScope() + let accessed = sourceURL.startAccessingSecurityScopedResource() + do { + let descriptor = try Self.resolvePackage(in: sourceURL) + pendingCandidate = IosSpeechModelCandidate(descriptor: descriptor, sourceURL: sourceURL) + pendingSecurityScopeActive = accessed + pendingSecurityScopeURL = accessed ? sourceURL : nil + installationError = nil + message = nil + } catch { + if accessed { sourceURL.stopAccessingSecurityScopedResource() } + pendingCandidate = nil + installationError = error.localizedDescription + message = nil + } + } + + func cancelPendingInstallation() { + releasePendingSecurityScope() + pendingCandidate = nil + } + + func confirmPendingInstallation( + candidate confirmedCandidate: IosSpeechModelCandidate? = nil, + replacementAllowed: Bool = true + ) { + guard replacementAllowed else { + message = "Wait for the current transcription to finish before replacing the speech model." + return + } + guard let candidate = confirmedCandidate ?? pendingCandidate, !isBusy else { return } + let sourceAccessAlreadyActive = pendingSecurityScopeActive + pendingSecurityScopeActive = false + pendingSecurityScopeURL = nil + pendingCandidate = nil + isInstalling = true + installationError = nil downloadProgress = nil - message = "Installing speech model..." + message = "Installing \(candidate.descriptor.displayName)..." Task { let result = await Task.detached(priority: .userInitiated) { - Self.installModelFiles(from: sourceURL) + Self.installModelFiles( + from: candidate.sourceURL, + descriptor: candidate.descriptor, + sourceAccessAlreadyActive: sourceAccessAlreadyActive + ) }.value isInstalling = false status = Self.inspectLightweight(directory: IosSpeechModelPaths.modelDirectory) - runtimeState = .unloaded + if result.committed { + invalidateRuntimeAfterReplacement() + installationError = nil + } else { + installationError = result.message + } message = result.message } } - func downloadModel() { + private func releasePendingSecurityScope() { + if pendingSecurityScopeActive, let url = pendingSecurityScopeURL { + url.stopAccessingSecurityScopedResource() + } + pendingSecurityScopeActive = false + pendingSecurityScopeURL = nil + } + + // Compatibility adapter used by existing tests and callers. Production UI + // uses inspect + explicit confirmation. + func installModel(from sourceURL: URL) { + inspectModelPackage(from: sourceURL) + } + + func downloadModel(_ requested: IosSpeechModelDescriptor? = nil) { guard !isBusy else { return } + let descriptor = requested ?? selectedDownloadModel + guard IosSpeechModelDescriptor.resolve( + catalogId: descriptor.catalogId, + modelVersion: descriptor.modelVersion + )?.networkDownloadAvailable == true else { + message = "Selected speech model is not available for download." + return + } + guard descriptor.networkDownloadAvailable else { + message = "Network download is not available for this model." + return + } + isInstalling = true message = "Preparing speech model download..." downloadProgress = IosSpeechModelDownloadProgress( message: "Preparing speech model download...", bytesDownloaded: 0, - totalBytes: Self.manifestTotalSizeBytes() + totalBytes: descriptor.totalSizeBytes ) downloadTask = Task { - let result = await Self.downloadAndInstallModel { [weak self] progress in + let result = await Self.downloadAndInstallModel(descriptor: descriptor) { [weak self] progress in Task { @MainActor in self?.downloadProgress = progress self?.message = progress.message @@ -198,7 +409,9 @@ final class IosSpeechModelStore: ObservableObject { isInstalling = false downloadTask = nil status = Self.inspectLightweight(directory: IosSpeechModelPaths.modelDirectory) - runtimeState = .unloaded + if result.committed { + invalidateRuntimeAfterReplacement() + } downloadProgress = nil message = result.message } @@ -280,6 +493,7 @@ final class IosSpeechModelStore: ObservableObject { invalidFile: URL = IosSpeechModelPaths.invalidFile, requiredFileNames: [String]? = nil ) -> IosSpeechModelStatus { + recoverInterruptedActivationIfNeeded() if let reason = try? String(contentsOf: invalidFile, encoding: .utf8), !reason.trimmingCharacters(in: .whitespacesAndNewlines).isEmpty { return IosSpeechModelStatus( @@ -297,7 +511,27 @@ final class IosSpeechModelStore: ObservableObject { missingFiles: ["model directory"] ) } - let missing = (requiredFileNames ?? manifestFiles().map(\.name)).compactMap { fileName in + let receiptText = try? String(contentsOf: receiptFile, encoding: .utf8) + let receipt = receiptText.flatMap(decodeReceipt) + let descriptor: IosSpeechModelDescriptor? = receipt.flatMap { + guard $0.receiptSchemaVersion == IosSpeechModelInstallation.receiptSchemaVersion, + $0.packageSchemaVersion == 1, + let resolved = IosSpeechModelDescriptor.resolve( + catalogId: $0.catalogId, + modelVersion: $0.modelVersion + ), resolved.backend == $0.backend else { return nil } + return resolved + } + if receipt != nil && descriptor == nil { + return IosSpeechModelStatus( + directory: directory, + installationState: .invalid, + missingFiles: ["unsupported active model receipt"] + ) + } + let legacyDescriptor = IosSpeechModelDescriptor.defaultModel + let expectedNames = requiredFileNames ?? (descriptor ?? legacyDescriptor).files.map(\.name) + let missing = expectedNames.compactMap { fileName in fileManager.isReadableFile(atPath: directory.appendingPathComponent(fileName).path) ? nil : fileName @@ -309,36 +543,139 @@ final class IosSpeechModelStore: ObservableObject { missingFiles: missing ) } - let receipt = try? String(contentsOf: receiptFile, encoding: .utf8) - let state: IosSpeechModelInstallationState = receipt?.trimmingCharacters(in: .whitespacesAndNewlines) - == EmbeddedSpeechModel.shared.manifest.version - ? .installedVerified - : .installedLegacy + let legacyReceiptMatches = receiptText?.trimmingCharacters(in: .whitespacesAndNewlines) + == SpeechModelCatalog.shared.defaultModel.manifest.version + let legacyInstallation = IosSpeechModelInstallation( + receiptSchemaVersion: IosSpeechModelInstallation.receiptSchemaVersion, + packageSchemaVersion: 1, + catalogId: legacyDescriptor.catalogId, + modelVersion: legacyDescriptor.modelVersion, + backend: legacyDescriptor.backend, + installationGeneration: "legacy" + ) + let state: IosSpeechModelInstallationState = (receipt != nil || legacyReceiptMatches) + ? .installedVerified : .installedLegacy return IosSpeechModelStatus( directory: directory, installationState: state, - missingFiles: [] + missingFiles: [], + activeInstallation: receipt ?? legacyInstallation ) } nonisolated private static func recordVerifiedInstallation() { - try? FileManager.default.createDirectory( + let descriptor = IosSpeechModelDescriptor.defaultModel + recordVerifiedInstallation(descriptor: descriptor, generation: "legacy-migrated-\(UUID().uuidString)") + } + + nonisolated private static func recordVerifiedInstallationIfNeeded() { + if let text = try? String(contentsOf: IosSpeechModelPaths.receiptFile, encoding: .utf8), + decodeReceipt(text) != nil { + try? removeItemIfExists(IosSpeechModelPaths.invalidFile) + return + } + recordVerifiedInstallation() + } + + nonisolated private static func recordVerifiedInstallation( + descriptor: IosSpeechModelDescriptor, + generation: String = UUID().uuidString + ) { + try? writeVerifiedInstallation(descriptor: descriptor, generation: generation) + } + + nonisolated private static func writeVerifiedInstallation( + descriptor: IosSpeechModelDescriptor, + generation: String = UUID().uuidString + ) throws { + try FileManager.default.createDirectory( at: IosSpeechModelPaths.applicationSupportDirectory, withIntermediateDirectories: true ) - try? EmbeddedSpeechModel.shared.manifest.version.write( - to: IosSpeechModelPaths.receiptFile, - atomically: true, - encoding: .utf8 + let receipt = IosSpeechModelInstallation( + receiptSchemaVersion: IosSpeechModelInstallation.receiptSchemaVersion, + packageSchemaVersion: 1, + catalogId: descriptor.catalogId, + modelVersion: descriptor.modelVersion, + backend: descriptor.backend, + installationGeneration: generation + ) + let data = try JSONEncoder().encode(receipt) + try data.write(to: IosSpeechModelPaths.receiptFile, options: Data.WritingOptions.atomic) + try removeItemIfExists(IosSpeechModelPaths.invalidFile) + } + + nonisolated private static func decodeReceipt(_ text: String) -> IosSpeechModelInstallation? { + guard let data = text.data(using: .utf8) else { return nil } + return try? JSONDecoder().decode(IosSpeechModelInstallation.self, from: data) + } + + nonisolated static func resolvePackage(in directory: URL) throws -> IosSpeechModelDescriptor { + let manifestURL = directory.appendingPathComponent("voice-inbox-model.json") + guard let data = try? Data(contentsOf: manifestURL) else { + let expected = Set(IosSpeechModelDescriptor.defaultModel.files.map(\.name)) + let contents = try? FileManager.default.contentsOfDirectory( + at: directory, + includingPropertiesForKeys: [.isRegularFileKey], + options: [.skipsHiddenFiles] + ) + let regularNames = Set((contents ?? []).compactMap { url -> String? in + guard (try? url.resourceValues(forKeys: [.isRegularFileKey]).isRegularFile) == true else { + return nil + } + return url.lastPathComponent + }) + guard regularNames == expected else { + throw ModelPackageError.missingManifest + } + return .defaultModel + } + guard let object = try? JSONSerialization.jsonObject(with: data) as? [String: Any] else { + throw ModelPackageError.malformedManifest + } + let allowed = Set(["schemaVersion", "catalogId", "modelVersion"]) + guard Set(object.keys).isSubset(of: allowed), + object.keys.count == allowed.count, + let schema = object["schemaVersion"] as? Int, + let catalogId = object["catalogId"] as? String, + let modelVersion = object["modelVersion"] as? String else { + throw ModelPackageError.malformedManifest + } + guard schema == 1, + let descriptor = IosSpeechModelDescriptor.resolve( + catalogId: catalogId, + modelVersion: modelVersion + ), descriptor.localImportAvailable else { + throw ModelPackageError.unsupportedIdentity + } + return descriptor + } + + nonisolated private static func recoverInterruptedActivationIfNeeded() { + let fileManager = FileManager.default + guard !fileManager.fileExists(atPath: IosSpeechModelPaths.modelDirectory.path), + fileManager.fileExists(atPath: IosSpeechModelPaths.backupDirectory.path) else { return } + try? fileManager.moveItem( + at: IosSpeechModelPaths.backupDirectory, + to: IosSpeechModelPaths.modelDirectory ) - try? removeItemIfExists(IosSpeechModelPaths.invalidFile) + if fileManager.fileExists(atPath: IosSpeechModelPaths.backupReceiptFile.path) { + try? removeItemIfExists(IosSpeechModelPaths.receiptFile) + try? fileManager.moveItem( + at: IosSpeechModelPaths.backupReceiptFile, + to: IosSpeechModelPaths.receiptFile + ) + } } nonisolated private static func recordInvalidInstallation(_ reason: String) { try? reason.write(to: IosSpeechModelPaths.invalidFile, atomically: true, encoding: .utf8) } - nonisolated private static func validateModelFiles(in directory: URL) -> ModelValidationResult { + nonisolated private static func validateModelFiles( + in directory: URL, + descriptor: IosSpeechModelDescriptor + ) -> ModelValidationResult { let fileManager = FileManager.default var isDirectory: ObjCBool = false guard fileManager.fileExists(atPath: directory.path, isDirectory: &isDirectory), isDirectory.boolValue else { @@ -346,7 +683,7 @@ final class IosSpeechModelStore: ObservableObject { } var validationIssues = [String]() - for entry in manifestFiles() { + for entry in descriptor.files { let fileURL = directory.appendingPathComponent(entry.name) if !fileManager.isReadableFile(atPath: fileURL.path) { validationIssues.append(entry.name) @@ -360,18 +697,29 @@ final class IosSpeechModelStore: ObservableObject { return ModelValidationResult(missingFiles: validationIssues) } - nonisolated private static func installModelFiles(from sourceURL: URL) -> InstallResult { + nonisolated private static func installModelFiles( + from sourceURL: URL, + descriptor: IosSpeechModelDescriptor, + sourceAccessAlreadyActive: Bool = false + ) -> InstallResult { let fileManager = FileManager.default - let accessed = sourceURL.startAccessingSecurityScopedResource() + let accessed = sourceAccessAlreadyActive || sourceURL.startAccessingSecurityScopedResource() defer { if accessed { sourceURL.stopAccessingSecurityScopedResource() } } - let sourceValidation = validateModelFiles(in: sourceURL) + let sourceValidation = validateModelFiles(in: sourceURL, descriptor: descriptor) guard sourceValidation.missingFiles.isEmpty else { - return InstallResult(message: "Selected folder is not a valid speech model. Missing: \(sourceValidation.missingFiles.joined(separator: ", "))") + return InstallResult(message: "Selected folder is not a valid \(descriptor.displayName) package: \(sourceValidation.missingFiles.joined(separator: ", "))") + } + + let requiredBytes = descriptor.totalSizeBytes + descriptor.safetyMarginBytes + if let values = try? IosSpeechModelPaths.applicationSupportDirectory.resourceValues(forKeys: [.volumeAvailableCapacityForImportantUsageKey]), + let available = values.volumeAvailableCapacityForImportantUsage, + available < requiredBytes { + return InstallResult(message: "Not enough free storage to install \(descriptor.displayName).") } do { @@ -387,9 +735,9 @@ final class IosSpeechModelStore: ObservableObject { withIntermediateDirectories: true ) - try copyRequiredFiles(from: sourceURL, to: IosSpeechModelPaths.installDirectory) + try copyRequiredFiles(from: sourceURL, to: IosSpeechModelPaths.installDirectory, descriptor: descriptor) - let installedValidation = validateModelFiles(in: IosSpeechModelPaths.installDirectory) + let installedValidation = validateModelFiles(in: IosSpeechModelPaths.installDirectory, descriptor: descriptor) guard installedValidation.missingFiles.isEmpty else { try removeItemIfExists(IosSpeechModelPaths.installDirectory) return InstallResult(message: "Copied model is incomplete. Missing: \(installedValidation.missingFiles.joined(separator: ", "))") @@ -398,17 +746,29 @@ final class IosSpeechModelStore: ObservableObject { if fileManager.fileExists(atPath: IosSpeechModelPaths.modelDirectory.path) { try fileManager.moveItem(at: IosSpeechModelPaths.modelDirectory, to: IosSpeechModelPaths.backupDirectory) } + if fileManager.fileExists(atPath: IosSpeechModelPaths.receiptFile.path) { + try removeItemIfExists(IosSpeechModelPaths.backupReceiptFile) + try fileManager.copyItem( + at: IosSpeechModelPaths.receiptFile, + to: IosSpeechModelPaths.backupReceiptFile + ) + } do { try fileManager.moveItem(at: IosSpeechModelPaths.installDirectory, to: IosSpeechModelPaths.modelDirectory) + try writeVerifiedInstallation(descriptor: descriptor) try removeItemIfExists(IosSpeechModelPaths.backupDirectory) - recordVerifiedInstallation() - IosNativeTranscriber.resetModel() - return InstallResult(message: "Speech model installed.") + try removeItemIfExists(IosSpeechModelPaths.backupReceiptFile) + return InstallResult(message: "\(descriptor.displayName) installed.", committed: true) } catch { + try? removeItemIfExists(IosSpeechModelPaths.modelDirectory) if fileManager.fileExists(atPath: IosSpeechModelPaths.backupDirectory.path) { try? fileManager.moveItem(at: IosSpeechModelPaths.backupDirectory, to: IosSpeechModelPaths.modelDirectory) } + if fileManager.fileExists(atPath: IosSpeechModelPaths.backupReceiptFile.path) { + try? removeItemIfExists(IosSpeechModelPaths.receiptFile) + try? fileManager.moveItem(at: IosSpeechModelPaths.backupReceiptFile, to: IosSpeechModelPaths.receiptFile) + } throw error } } catch { @@ -417,8 +777,12 @@ final class IosSpeechModelStore: ObservableObject { } } - nonisolated private static func copyRequiredFiles(from sourceURL: URL, to destinationURL: URL) throws { - for entry in manifestFiles() { + nonisolated private static func copyRequiredFiles( + from sourceURL: URL, + to destinationURL: URL, + descriptor: IosSpeechModelDescriptor + ) throws { + for entry in descriptor.files { try copyFile(named: entry.name, from: sourceURL, to: destinationURL) } } @@ -431,11 +795,12 @@ final class IosSpeechModelStore: ObservableObject { } nonisolated private static func downloadAndInstallModel( + descriptor: IosSpeechModelDescriptor, progress: @escaping @Sendable (IosSpeechModelDownloadProgress) -> Void ) async -> InstallResult { let fileManager = FileManager.default - let files = manifestFiles() - let totalBytes = manifestTotalSizeBytes() + let files = descriptor.files + let totalBytes = descriptor.totalSizeBytes do { try fileManager.createDirectory( @@ -486,7 +851,7 @@ final class IosSpeechModelStore: ObservableObject { totalBytes: totalBytes )) - let stagedValidation = validateModelFiles(in: IosSpeechModelPaths.stagingDirectory) + let stagedValidation = validateModelFiles(in: IosSpeechModelPaths.stagingDirectory, descriptor: descriptor) guard stagedValidation.missingFiles.isEmpty else { return InstallResult( message: "Downloaded model is incomplete. Missing: \(stagedValidation.missingFiles.joined(separator: ", "))" @@ -494,14 +859,16 @@ final class IosSpeechModelStore: ObservableObject { } try activateStagedModel() - recordVerifiedInstallation() - IosNativeTranscriber.resetModel() - return InstallResult(message: "Speech model downloaded and installed.") + try writeVerifiedInstallation(descriptor: descriptor) + try removeItemIfExists(IosSpeechModelPaths.backupDirectory) + try removeItemIfExists(IosSpeechModelPaths.backupReceiptFile) + return InstallResult(message: "\(descriptor.displayName) downloaded and installed.", committed: true) } catch is CancellationError { try? removeItemIfExists(IosSpeechModelPaths.stagingDirectory) return InstallResult(message: "Speech model download cancelled.") } catch { try? removeItemIfExists(IosSpeechModelPaths.stagingDirectory) + restorePreviousInstallationIfNeeded() return InstallResult(message: "Could not download speech model: \(error.localizedDescription)") } } @@ -513,13 +880,13 @@ final class IosSpeechModelStore: ObservableObject { } nonisolated private static func downloadFile( - _ entry: IosSpeechModelManifestFile, + _ entry: IosSpeechModelFileDescriptor, completedBytes: Int64, totalBytes: Int64, progress: @escaping @Sendable (IosSpeechModelDownloadProgress) -> Void ) async throws { - guard let url = URL(string: entry.downloadUrl) else { - throw ModelDownloadError.invalidUrl(entry.downloadUrl) + guard let url = URL(string: entry.downloadURL) else { + throw ModelDownloadError.invalidUrl(entry.downloadURL) } let (bytes, response) = try await URLSession.shared.bytes(from: url) @@ -573,23 +940,49 @@ final class IosSpeechModelStore: ObservableObject { to: IosSpeechModelPaths.installDirectory ) try removeItemIfExists(IosSpeechModelPaths.backupDirectory) + try removeItemIfExists(IosSpeechModelPaths.backupReceiptFile) if fileManager.fileExists(atPath: IosSpeechModelPaths.modelDirectory.path) { try fileManager.moveItem(at: IosSpeechModelPaths.modelDirectory, to: IosSpeechModelPaths.backupDirectory) } + if fileManager.fileExists(atPath: IosSpeechModelPaths.receiptFile.path) { + try fileManager.copyItem(at: IosSpeechModelPaths.receiptFile, to: IosSpeechModelPaths.backupReceiptFile) + } do { try fileManager.moveItem(at: IosSpeechModelPaths.installDirectory, to: IosSpeechModelPaths.modelDirectory) - try removeItemIfExists(IosSpeechModelPaths.backupDirectory) } catch { + try? removeItemIfExists(IosSpeechModelPaths.modelDirectory) if fileManager.fileExists(atPath: IosSpeechModelPaths.backupDirectory.path) { try? fileManager.moveItem(at: IosSpeechModelPaths.backupDirectory, to: IosSpeechModelPaths.modelDirectory) } + if fileManager.fileExists(atPath: IosSpeechModelPaths.backupReceiptFile.path) { + try? removeItemIfExists(IosSpeechModelPaths.receiptFile) + try? fileManager.moveItem(at: IosSpeechModelPaths.backupReceiptFile, to: IosSpeechModelPaths.receiptFile) + } throw error } } - nonisolated private static func cleanupPartialFile(for entry: IosSpeechModelManifestFile) throws { + nonisolated private static func restorePreviousInstallationIfNeeded() { + let fileManager = FileManager.default + if fileManager.fileExists(atPath: IosSpeechModelPaths.backupDirectory.path) { + try? removeItemIfExists(IosSpeechModelPaths.modelDirectory) + try? fileManager.moveItem( + at: IosSpeechModelPaths.backupDirectory, + to: IosSpeechModelPaths.modelDirectory + ) + } + if fileManager.fileExists(atPath: IosSpeechModelPaths.backupReceiptFile.path) { + try? removeItemIfExists(IosSpeechModelPaths.receiptFile) + try? fileManager.moveItem( + at: IosSpeechModelPaths.backupReceiptFile, + to: IosSpeechModelPaths.receiptFile + ) + } + } + + nonisolated private static func cleanupPartialFile(for entry: IosSpeechModelFileDescriptor) throws { try removeItemIfExists(temporaryFile(for: entry)) let destination = IosSpeechModelPaths.stagingDirectory.appendingPathComponent(entry.name) if !isValidFile(destination, entry: entry) { @@ -597,11 +990,11 @@ final class IosSpeechModelStore: ObservableObject { } } - nonisolated private static func temporaryFile(for entry: IosSpeechModelManifestFile) -> URL { + nonisolated private static func temporaryFile(for entry: IosSpeechModelFileDescriptor) -> URL { IosSpeechModelPaths.stagingDirectory.appendingPathComponent("\(entry.name).part") } - nonisolated private static func isValidFile(_ url: URL, entry: IosSpeechModelManifestFile) -> Bool { + nonisolated private static func isValidFile(_ url: URL, entry: IosSpeechModelFileDescriptor) -> Bool { guard FileManager.default.isReadableFile(atPath: url.path) else { return false } @@ -640,28 +1033,18 @@ final class IosSpeechModelStore: ObservableObject { return digest.map { String(format: "%02x", $0) }.joined() } - nonisolated private static func manifestTotalSizeBytes() -> Int64 { - manifestFiles().reduce(0) { $0 + $1.sizeBytes } - } - - nonisolated private static func manifestFiles() -> [IosSpeechModelManifestFile] { - let manifest = EmbeddedSpeechModel.shared.manifest - return manifest.files.map { file in - return IosSpeechModelManifestFile( - name: file.name, - sizeBytes: file.sizeBytes, - sha256: file.sha256, - downloadUrl: manifest.downloadUrl(file: file) - ) - } - } - private struct ModelValidationResult { let missingFiles: [String] } private struct InstallResult { let message: String + let committed: Bool + + init(message: String, committed: Bool = false) { + self.message = message + self.committed = committed + } func cleanup() { try? IosSpeechModelStore.removeItemIfExists(IosSpeechModelPaths.stagingDirectory) @@ -674,13 +1057,6 @@ final class IosSpeechModelStore: ObservableObject { let message: String? } - private struct IosSpeechModelManifestFile { - let name: String - let sizeBytes: Int64 - let sha256: String - let downloadUrl: String - } - private enum ModelDownloadError: LocalizedError { case invalidUrl(String) case httpStatus(Int) @@ -694,6 +1070,23 @@ final class IosSpeechModelStore: ObservableObject { } } } + + private enum ModelPackageError: LocalizedError { + case missingManifest + case malformedManifest + case unsupportedIdentity + + var errorDescription: String? { + switch self { + case .missingManifest: + return "voice-inbox-model.json is missing from the selected folder." + case .malformedManifest: + return "voice-inbox-model.json is malformed or contains unsupported fields." + case .unsupportedIdentity: + return "This speech model package is not supported on iOS." + } + } + } } private extension Comparable { diff --git a/iosApp/VoiceInbox/NativeBridge/VoiceInboxNativeTranscription.h b/iosApp/VoiceInbox/NativeBridge/VoiceInboxNativeTranscription.h index 10c637a..8cc1313 100644 --- a/iosApp/VoiceInbox/NativeBridge/VoiceInboxNativeTranscription.h +++ b/iosApp/VoiceInbox/NativeBridge/VoiceInboxNativeTranscription.h @@ -5,6 +5,12 @@ #include bool voiceinbox_transcription_initialize(const char *model_directory); +bool voiceinbox_transcription_initialize_configured( + const char *backend, + const char *installation_identity, + const char *model_directory, + const char *primary_file +); char *voiceinbox_transcription_transcribe_chunk_json(const float *samples, size_t sample_count); char *voiceinbox_transcription_last_error(void); void voiceinbox_transcription_string_free(char *value); diff --git a/iosApp/VoiceInbox/SettingsView.swift b/iosApp/VoiceInbox/SettingsView.swift index c348ee5..248a4f8 100644 --- a/iosApp/VoiceInbox/SettingsView.swift +++ b/iosApp/VoiceInbox/SettingsView.swift @@ -58,11 +58,14 @@ struct SettingsView: View { @ObservedObject var importStore: IosAudioImportStore @ObservedObject var outputStore: IosOutputDocumentStore @ObservedObject var startupPolicyStore: IosStartupProcessingPolicyStore + @ObservedObject var speechModelStore: IosSpeechModelStore let selectInboxFolder: () -> Void let selectOutputFile: () -> Void + let installModelPackage: () -> Void - private let websiteURL = URL(string: "https://projects.maxistar.me/Voice-Inbox/")! - private let legalURL = URL(string: "https://projects.maxistar.me/Voice-Inbox/legal/")! + private let websiteURL = URL(string: "https://voiceinbox.simpleditor.org/")! + private let documentationURL = URL(string: "https://voiceinbox.simpleditor.org/docs/")! + private let legalURL = URL(string: "https://voiceinbox.simpleditor.org/legal/")! var body: some View { Form { @@ -122,9 +125,29 @@ struct SettingsView: View { .foregroundStyle(.secondary) } + Section("Speech Model") { + if let active = speechModelStore.activeDescriptor { + VStack(alignment: .leading, spacing: 4) { + Text(active.displayName).font(.headline) + Text("\(active.languageSummary) · \(active.maturity)") + .font(.footnote) + .foregroundStyle(.secondary) + } + } else { + Text("No speech model installed") + } + Button { + installModelPackage() + } label: { + Label("Install model package from folder", systemImage: "folder.badge.plus") + } + .disabled(speechModelStore.isBusy) + } + Section("About") { LabeledContent("Version", value: appVersion) Link("Website", destination: websiteURL) + Link("Documentation", destination: documentationURL) Link("Legal information", destination: legalURL) } } diff --git a/iosApp/VoiceInboxTests/DeferredSpeechModelLoadingTests.swift b/iosApp/VoiceInboxTests/DeferredSpeechModelLoadingTests.swift index d3a95cf..a8cbb70 100644 --- a/iosApp/VoiceInboxTests/DeferredSpeechModelLoadingTests.swift +++ b/iosApp/VoiceInboxTests/DeferredSpeechModelLoadingTests.swift @@ -1,9 +1,151 @@ import Foundation import Shared +import UniformTypeIdentifiers +import UIKit import XCTest @testable import VoiceInbox final class DeferredSpeechModelLoadingTests: XCTestCase { + func testIosCatalogResolvesStrictParakeetAndWhisperPackageIdentities() throws { + let root = FileManager.default.temporaryDirectory + .appendingPathComponent(UUID().uuidString, isDirectory: true) + defer { try? FileManager.default.removeItem(at: root) } + try FileManager.default.createDirectory(at: root, withIntermediateDirectories: true) + + func write(_ json: String) throws { + try Data(json.utf8).write(to: root.appendingPathComponent("voice-inbox-model.json")) + } + + try write(#"{"schemaVersion":1,"catalogId":"whisper-tiny-multilingual","modelVersion":"whisper-tiny-ggml-f16-r1"}"#) + let whisper = try IosSpeechModelStore.resolvePackage(in: root) + XCTAssertEqual(whisper.backend, "WHISPER_CPP") + XCTAssertTrue(whisper.networkDownloadAvailable) + XCTAssertTrue(whisper.localImportAvailable) + + try write(#"{"schemaVersion":1,"catalogId":"parakeet-tdt-0.6b-v3-int8","modelVersion":"parakeet-tdt-0.6b-v3-int8-r1"}"#) + let parakeet = try IosSpeechModelStore.resolvePackage(in: root) + XCTAssertEqual(parakeet.backend, "PARAKEET_TDT_ONNX") + XCTAssertTrue(parakeet.networkDownloadAvailable) + XCTAssertEqual( + Set(IosSpeechModelDescriptor.supported.filter(\.networkDownloadAvailable).map(\.catalogId)), + Set(["parakeet-tdt-0.6b-v3-int8", "whisper-tiny-multilingual"]) + ) + + for invalid in [ + #"{"schemaVersion":2,"catalogId":"whisper-tiny-multilingual","modelVersion":"whisper-tiny-ggml-f16-r1"}"#, + #"{"schemaVersion":1,"catalogId":"unknown","modelVersion":"unknown"}"#, + #"{"schemaVersion":1,"catalogId":"whisper-tiny-multilingual","modelVersion":"whisper-tiny-ggml-f16-r1","backend":"PARAKEET_TDT_ONNX"}"#, + ] { + try write(invalid) + XCTAssertThrowsError(try IosSpeechModelStore.resolvePackage(in: root)) + } + } + + func testIosPackageResolutionAcceptsAndroidCompatibleLegacyParakeetFolder() throws { + let root = FileManager.default.temporaryDirectory + .appendingPathComponent(UUID().uuidString, isDirectory: true) + defer { try? FileManager.default.removeItem(at: root) } + try FileManager.default.createDirectory(at: root, withIntermediateDirectories: true) + + for file in IosSpeechModelDescriptor.defaultModel.files { + FileManager.default.createFile( + atPath: root.appendingPathComponent(file.name).path, + contents: Data() + ) + } + + let descriptor = try IosSpeechModelStore.resolvePackage(in: root) + XCTAssertEqual(descriptor.catalogId, IosSpeechModelDescriptor.defaultModel.catalogId) + + FileManager.default.createFile( + atPath: root.appendingPathComponent("unexpected.txt").path, + contents: Data() + ) + XCTAssertThrowsError(try IosSpeechModelStore.resolvePackage(in: root)) + } + + @MainActor + func testConfirmedCandidateSurvivesAlertBindingDismissal() async { + let root = FileManager.default.temporaryDirectory + let store = IosSpeechModelStore( + directory: root, + inspectInstallation: { directory in + IosSpeechModelStatus(directory: directory, installationState: .missing, missingFiles: []) + }, + validateInstallation: { _ in [] }, + prepareNative: { _ in true }, + nativeError: { nil }, + recordVerified: {}, + recordInvalid: { _ in }, + resetNative: {} + ) + let candidate = IosSpeechModelCandidate( + descriptor: .defaultModel, + sourceURL: root + ) + + // SwiftUI clears an alert(item:) binding as it dismisses the alert. + store.pendingCandidate = nil + store.confirmPendingInstallation(candidate: candidate) + + XCTAssertTrue(store.isInstalling) + XCTAssertEqual(store.message, "Installing \(candidate.descriptor.displayName)...") + + while store.isInstalling { + await Task.yield() + } + } + + func testLightweightRestoreUsesVersionedWhisperReceiptWithoutHashingPayload() throws { + let root = FileManager.default.temporaryDirectory + .appendingPathComponent(UUID().uuidString, isDirectory: true) + let model = root.appendingPathComponent("SpeechModel", isDirectory: true) + let receipt = root.appendingPathComponent("SpeechModel.receipt") + let invalid = root.appendingPathComponent("SpeechModel.invalid") + defer { try? FileManager.default.removeItem(at: root) } + try FileManager.default.createDirectory(at: model, withIntermediateDirectories: true) + try Data("not model weights".utf8).write(to: model.appendingPathComponent("ggml-tiny.bin")) + let active = IosSpeechModelInstallation( + receiptSchemaVersion: 2, + packageSchemaVersion: 1, + catalogId: "whisper-tiny-multilingual", + modelVersion: "whisper-tiny-ggml-f16-r1", + backend: "WHISPER_CPP", + installationGeneration: "test-generation" + ) + try JSONEncoder().encode(active).write(to: receipt) + + let status = IosSpeechModelStore.inspectLightweight( + directory: model, + receiptFile: receipt, + invalidFile: invalid, + requiredFileNames: ["ggml-tiny.bin"] + ) + XCTAssertEqual(status.installationState, .installedVerified) + XCTAssertEqual(status.activeInstallation, active) + XCTAssertTrue(status.isReady) + } + + func testCatalogPreservesProductionParakeetManifest() { + let descriptor = SpeechModelCatalog.shared.defaultModel + let manifest = descriptor.manifest + + XCTAssertEqual(SpeechModelCatalog.shared.models.count, 2) + XCTAssertEqual(SpeechModelCatalog.shared.modelsFor(platform: .ios).count, 2) + XCTAssertEqual(descriptor.catalogId, "parakeet-tdt-0.6b-v3-int8") + XCTAssertEqual(manifest.modelId, "istupakov/parakeet-tdt-0.6b-v3-onnx") + XCTAssertEqual(manifest.version, "parakeet-tdt-0.6b-v3-int8-r1") + XCTAssertEqual(manifest.repositoryRevision, "8f23f0c03c8761650bdb5b40aaf3e40d2c15f1ce") + XCTAssertEqual(manifest.totalSizeBytes, 670_619_803) + XCTAssertEqual(Set(manifest.files.map(\.name)), Set([ + "encoder-model.int8.onnx", + "decoder_joint-model.int8.onnx", + "nemo128.onnx", + "vocab.txt", + "config.json", + ])) + } + func testLightweightInspectionDistinguishesMissingLegacyVerifiedAndKnownInvalidWithoutReadingPayloads() throws { let root = FileManager.default.temporaryDirectory .appendingPathComponent(UUID().uuidString, isDirectory: true) @@ -32,7 +174,11 @@ final class DeferredSpeechModelLoadingTests: XCTestCase { XCTAssertEqual(legacy.installationState, .installedLegacy) XCTAssertTrue(legacy.isReady) - try EmbeddedSpeechModel.shared.manifest.version.write(to: receipt, atomically: true, encoding: .utf8) + try SpeechModelCatalog.shared.defaultModel.manifest.version.write( + to: receipt, + atomically: true, + encoding: .utf8 + ) let verified = IosSpeechModelStore.inspectLightweight( directory: model, receiptFile: receipt, @@ -80,7 +226,7 @@ final class DeferredSpeechModelLoadingTests: XCTestCase { }, nativeError: { nil }, recordVerified: { - try? EmbeddedSpeechModel.shared.manifest.version.write( + try? SpeechModelCatalog.shared.defaultModel.manifest.version.write( to: receipt, atomically: true, encoding: .utf8 @@ -228,6 +374,7 @@ final class DeferredSpeechModelLoadingTests: XCTestCase { func testTypedActionsHaveExplicitIosRoutes() { XCTAssertEqual(IosTaskActionRouter.route(.downloadModel), .modelDownload) + XCTAssertEqual(IosTaskActionRouter.route(.createOutput), .outputCreation) XCTAssertEqual(IosTaskActionRouter.route(.selectOutput), .outputSelection) XCTAssertEqual(IosTaskActionRouter.route(.selectFolder), .folderSelection) XCTAssertEqual(IosTaskActionRouter.route(.transcribe), .transcribe) @@ -235,6 +382,28 @@ final class DeferredSpeechModelLoadingTests: XCTestCase { XCTAssertEqual(IosTaskActionRouter.route(.showText), .showText) } + func testOutputDocumentCreatorCoordinatorPreservesCancellationAndRoutesCreatedDocument() { + let controller = UIDocumentPickerViewController(forOpeningContentTypes: [.plainText]) + var pickedURL: URL? + var cancellationCount = 0 + let coordinator = IosOutputDocumentCreator.Coordinator( + onPick: { pickedURL = $0 }, + onCancel: { cancellationCount += 1 } + ) + + coordinator.documentPicker(controller, didPickDocumentsAt: []) + XCTAssertNil(pickedURL) + XCTAssertEqual(cancellationCount, 1) + + let createdURL = URL(fileURLWithPath: "/tmp/Voice Inbox Transcripts.md") + coordinator.documentPicker(controller, didPickDocumentsAt: [createdURL]) + XCTAssertEqual(pickedURL, createdURL) + XCTAssertEqual(cancellationCount, 1) + + coordinator.documentPickerWasCancelled(controller) + XCTAssertEqual(cancellationCount, 2) + } + func testRoutineImportAndScanSummariesDoNotRequestAnAlert() { XCTAssertNil(IosAudioImportSummary(imported: 2, skipped: 0, failed: 0).alertMessage) XCTAssertNil(IosAudioImportSummary(imported: 0, skipped: 4, failed: 0).alertMessage) diff --git a/iosApp/VoiceInboxTests/IosInlineOnboardingTests.swift b/iosApp/VoiceInboxTests/IosInlineOnboardingTests.swift index acf1342..0f46a05 100644 --- a/iosApp/VoiceInboxTests/IosInlineOnboardingTests.swift +++ b/iosApp/VoiceInboxTests/IosInlineOnboardingTests.swift @@ -246,7 +246,9 @@ final class IosInlineOnboardingTests: XCTestCase { installationPhase: nil, progressPercent: nil, downloadAvailable: downloadAvailable, - canCancel: false + canCancel: false, + selectedModel: nil, + downloadChoices: [] ) } diff --git a/iosApp/VoiceInboxTests/WhisperMobileSpikeTests.swift b/iosApp/VoiceInboxTests/WhisperMobileSpikeTests.swift new file mode 100644 index 0000000..0c4d6e8 --- /dev/null +++ b/iosApp/VoiceInboxTests/WhisperMobileSpikeTests.swift @@ -0,0 +1,260 @@ +import Darwin +import Foundation +import XCTest + +/// Opt-in developer harness. Native spike symbols are resolved dynamically so +/// ordinary app and test builds do not require the experimental backend. +final class WhisperMobileSpikeTests: XCTestCase { + private typealias InitializeFunction = @convention(c) ( + UnsafePointer? + ) -> UnsafeMutablePointer? + private typealias TranscribeFunction = @convention(c) ( + UnsafePointer?, Int, UnsafePointer? + ) -> UnsafeMutablePointer? + private typealias ResetFunction = @convention(c) () -> Void + private typealias FreeFunction = @convention(c) (UnsafeMutablePointer?) -> Void + + func testPhysicalDeviceSmokeTranscription() throws { + #if targetEnvironment(simulator) + throw XCTSkip("The simulator is build evidence only; run this test on a physical iPhone") + #else + let environment = ProcessInfo.processInfo.environment + let resultURL = FileManager.default.temporaryDirectory + .appendingPathComponent("ios-whisper-spike-smoke.json") + var result = baseResult(environment: environment) + let availableMemoryBefore = os_proc_available_memory() + + guard let modelPathValue = environment["WHISPER_MODEL_PATH"], !modelPathValue.isEmpty else { + try fail( + "WHISPER_MODEL_PATH was not provided; use an absolute path or a path relative to the test runner container", + result: &result, + resultURL: resultURL + ) + } + guard let pcmPathValue = environment["WHISPER_PCM_PATH"], !pcmPathValue.isEmpty else { + try fail( + "WHISPER_PCM_PATH was not provided; physical evidence requires representative float32 mono 16 kHz PCM", + result: &result, + resultURL: resultURL + ) + } + let modelPath = provisionedPath(modelPathValue) + let pcmPath = provisionedPath(pcmPathValue) + + guard + let initialize: InitializeFunction = resolve("voiceinbox_whisper_spike_initialize_json"), + let transcribe: TranscribeFunction = resolve("voiceinbox_whisper_spike_transcribe_json"), + let reset: ResetFunction = resolve("voiceinbox_whisper_spike_reset"), + let free: FreeFunction = resolve("voiceinbox_transcription_string_free") + else { + try fail( + "Whisper spike symbols are missing; build with VOICEINBOX_WHISPER_MOBILE_SPIKE=1 and force-load the spike archive members", + result: &result, + resultURL: resultURL + ) + } + defer { reset() } + + let initialization = modelPath.withCString { path in + takeJSON(initialize(path), free: free) + } + result["load_ms"] = initialization["load_ms"] ?? NSNull() + guard initialization["status"] as? String == "ok" else { + try fail( + initialization["error"] as? String ?? "Whisper initialization failed without a diagnostic", + result: &result, + resultURL: resultURL, + diagnostics: ["initialization": initialization] + ) + } + + let samples: [Float] + do { + samples = try loadSamples(path: pcmPath) + } catch { + try fail( + "Could not load PCM input: \(error.localizedDescription)", + result: &result, + resultURL: resultURL, + diagnostics: ["initialization": initialization] + ) + } + let language = environment["WHISPER_LANGUAGE"] ?? "" + let chunks = chunk(samples: samples) + var chunkResults = [[String: Any]]() + for chunk in chunks { + let chunkResult = chunk.withUnsafeBufferPointer { samples in + language.withCString { language in + takeJSON(transcribe(samples.baseAddress, samples.count, language), free: free) + } + } + guard chunkResult["status"] as? String == "ok" else { + try fail( + chunkResult["error"] as? String ?? "Whisper inference failed without a diagnostic", + result: &result, + resultURL: resultURL, + diagnostics: [ + "initialization": initialization, + "completed_chunks": chunkResults, + "failed_chunk": chunkResult, + ] + ) + } + chunkResults.append(chunkResult) + } + + let inferenceTimes = chunkResults.compactMap { $0["inference_ms"] as? Double } + let inferenceMilliseconds = inferenceTimes.reduce(0, +) + let audioDurationSeconds = Double(samples.count) / 16_000.0 + let transcript = chunkResults + .compactMap { $0["transcript"] as? String } + .joined(separator: " ") + .trimmingCharacters(in: .whitespacesAndNewlines) + + result["status"] = "ok" + result["load_ms"] = initialization["load_ms"] ?? NSNull() + result["chunk_inference_ms"] = inferenceTimes + result["inference_ms"] = inferenceMilliseconds + result["audio_duration_seconds"] = audioDurationSeconds + result["real_time_factor"] = inferenceMilliseconds / 1_000.0 / audioDurationSeconds + result["detected_language"] = chunkResults.compactMap { $0["detected_language"] as? String }.last ?? NSNull() + result["transcript"] = transcript + result["error"] = NSNull() + result["diagnostics"] = [ + "initialization": initialization, + "sample_rate_hz": 16_000, + "sample_count": samples.count, + "chunk_count": chunks.count, + "available_memory_before_bytes": availableMemoryBefore, + "available_memory_after_bytes": os_proc_available_memory(), + "memory_note": "Peak process memory is unavailable in this XCTest harness; available process memory is recorded before and after the run.", + "test_result": "Physical XCTest completed without host-process termination.", + ] + try write(result, to: resultURL) + print("Whisper spike result: \(resultURL.path)") + XCTAssertFalse(transcript.isEmpty, "Physical speech input produced an empty transcript; result: \(resultURL.path)") + #endif + } + + func testChunkingUsesThirtySecondsWithOneSecondOverlap() { + let samples = Array(repeating: Float.zero, count: 31 * 16_000) + let chunks = chunk(samples: samples) + XCTAssertEqual(chunks.map(\.count), [30 * 16_000, 2 * 16_000]) + } + + func testRelativeProvisioningPathUsesTestRunnerHome() { + XCTAssertEqual( + provisionedPath("Documents/whisper/ggml-tiny.bin"), + URL(fileURLWithPath: NSHomeDirectory()) + .appendingPathComponent("Documents/whisper/ggml-tiny.bin").path + ) + } + + private func baseResult(environment: [String: String]) -> [String: Any] { + [ + "schema_version": 1, + "backend": "whisper.cpp", + "model_id": "openai/whisper-tiny", + "model_revision": "5359861c739e955e79d9a303bcbc70fb988958b1", + "platform": "ios", + "device": environment["WHISPER_DEVICE_NAME"] ?? "physical iPhone", + "os_version": ProcessInfo.processInfo.operatingSystemVersionString, + "build_revision": environment["WHISPER_BUILD_REVISION"] ?? "workspace", + "build_profile": "release-native/xctest", + "corpus_item_id": environment["WHISPER_CORPUS_ITEM_ID"] ?? "physical-ios-smoke", + "status": "error", + "load_ms": NSNull(), + "chunk_inference_ms": NSNull(), + "inference_ms": NSNull(), + "audio_duration_seconds": 0.0, + "real_time_factor": NSNull(), + "peak_memory_bytes": NSNull(), + "detected_language": NSNull(), + "transcript": NSNull(), + "quality": [ + "wer": NSNull(), + "cer": NSNull(), + "punctuation": NSNull(), + "capitalization": NSNull(), + "silence_hallucination": NSNull(), + "overlap_boundary": NSNull(), + ], + "diagnostics": NSNull(), + "error": "Physical XCTest did not complete", + ] + } + + private func fail( + _ message: String, + result: inout [String: Any], + resultURL: URL, + diagnostics: [String: Any] = [:], + file: StaticString = #filePath, + line: UInt = #line + ) throws -> Never { + result["status"] = "error" + result["error"] = message + result["diagnostics"] = diagnostics + try write(result, to: resultURL) + XCTFail("\(message). Structured result: \(resultURL.path)", file: file, line: line) + throw NSError(domain: "WhisperMobileSpike", code: 1, userInfo: [NSLocalizedDescriptionKey: message]) + } + + private func write(_ result: [String: Any], to url: URL) throws { + let data = try JSONSerialization.data(withJSONObject: result, options: [.prettyPrinted, .sortedKeys]) + try data.write(to: url, options: .atomic) + } + + private func provisionedPath(_ value: String) -> String { + if value.hasPrefix("/") { return value } + return URL(fileURLWithPath: NSHomeDirectory()).appendingPathComponent(value).path + } + + private func resolve(_ symbol: String) -> T? { + guard let handle = dlopen(nil, RTLD_NOW), let pointer = dlsym(handle, symbol) else { + return nil + } + return unsafeBitCast(pointer, to: T.self) + } + + private func takeJSON( + _ pointer: UnsafeMutablePointer?, + free: FreeFunction + ) -> [String: Any] { + guard let pointer else { + return ["status": "error", "error": "native function returned null"] + } + defer { free(pointer) } + let data = Data(String(cString: pointer).utf8) + return (try? JSONSerialization.jsonObject(with: data)) as? [String: Any] + ?? ["status": "error", "error": "native function returned invalid JSON"] + } + + private func loadSamples(path: String) throws -> [Float] { + let data = try Data(contentsOf: URL(fileURLWithPath: path)) + guard !data.isEmpty, data.count.isMultiple(of: MemoryLayout.size) else { + throw NSError( + domain: "WhisperMobileSpike", + code: 2, + userInfo: [NSLocalizedDescriptionKey: "PCM must be non-empty little-endian float32"] + ) + } + return data.withUnsafeBytes { bytes in + Array(bytes.bindMemory(to: Float.self)) + } + } + + private func chunk(samples: [Float]) -> [[Float]] { + let chunkSamples = 16_000 * 30 + let overlapSamples = 16_000 + var result = [[Float]]() + var start = 0 + while start < samples.count { + let end = min(start + chunkSamples, samples.count) + result.append(Array(samples[start../dev/null 2>&1 && [[ -z "${CMAKE:-}" ]]; then + ANDROID_SDK_ROOT_CANDIDATE="${ANDROID_SDK_ROOT:-${ANDROID_HOME:-${HOME}/Library/Android/sdk}}" + CMAKE_CANDIDATE="$(find "${ANDROID_SDK_ROOT_CANDIDATE}/cmake" -maxdepth 3 -type f -name cmake 2>/dev/null | sort -r | head -n 1 || true)" + if [[ -n "${CMAKE_CANDIDATE}" ]]; then + export CMAKE="${CMAKE_CANDIDATE}" + export CMAKE_GENERATOR="Ninja" + export CMAKE_MAKE_PROGRAM="$(dirname "${CMAKE_CANDIDATE}")/ninja" + echo " CMake=${CMAKE}" + fi +fi + +if [[ "${VOICEINBOX_WHISPER_MOBILE_SPIKE:-0}" == "1" ]]; then + CARGO_FEATURES+=(whisper-mobile-spike) + echo " Whisper mobile spike enabled (test-only build)" + # whisper.cpp is built through CMake, which initializes Apple's deployment + # target from MACOSX_DEPLOYMENT_TARGET even for an iOS toolchain. + export MACOSX_DEPLOYMENT_TARGET="${IPHONEOS_DEPLOYMENT_TARGET:-16.0}" + echo " Native deployment target=${MACOSX_DEPLOYMENT_TARGET}" +fi + +if [[ ${#CARGO_FEATURES[@]} -gt 0 ]]; then + FEATURES_CSV="$(IFS=,; echo "${CARGO_FEATURES[*]}")" + CARGO_FEATURE_ARGS=(--no-default-features --features "${FEATURES_CSV}") +fi + if ! command -v cargo >/dev/null 2>&1; then echo "Cargo was not found. Install Rust or make sure ${HOME}/.cargo/bin is available to Xcode." >&2 exit 1 @@ -81,7 +113,12 @@ cd "${ROOT_DIR}" build_target() { local target="$1" - local cargo_env=() + local cargo_env=( + env + "IPHONEOS_DEPLOYMENT_TARGET=${IPHONEOS_DEPLOYMENT_TARGET:-16.0}" + "MACOSX_DEPLOYMENT_TARGET=${IPHONEOS_DEPLOYMENT_TARGET:-16.0}" + "CMAKE_OSX_DEPLOYMENT_TARGET=${IPHONEOS_DEPLOYMENT_TARGET:-16.0}" + ) echo "Building Rust target ${target}" if [[ -n "${IOS_ONNX_RUNTIME_XCFRAMEWORK}" ]]; then @@ -115,7 +152,7 @@ build_target() { local link_dir="${OUT_DIR}/ort-link/${target}" mkdir -p "${link_dir}" lipo "${framework_binary}" -thin "${arch}" -output "${link_dir}/libonnxruntime.a" - cargo_env=(env ORT_LIB_LOCATION="${link_dir}" ORT_SKIP_DOWNLOAD=1 ORT_PREFER_DYNAMIC_LINK=0) + cargo_env+=(ORT_LIB_LOCATION="${link_dir}" ORT_SKIP_DOWNLOAD=1 ORT_PREFER_DYNAMIC_LINK=0) fi if [[ ${#cargo_env[@]} -gt 0 ]]; then diff --git a/scripts/validate-whisper-tiny.sh b/scripts/validate-whisper-tiny.sh new file mode 100755 index 0000000..e17f7d5 --- /dev/null +++ b/scripts/validate-whisper-tiny.sh @@ -0,0 +1,42 @@ +#!/usr/bin/env bash +set -euo pipefail + +EXPECTED_NAME="ggml-tiny.bin" +EXPECTED_SIZE="77691713" +EXPECTED_SHA256="be07e048e1e599ad46341c8d2a135645097a538221678b7acdd1b1919c6e1b21" +MODEL_PATH="${1:-}" + +if [[ -z "${MODEL_PATH}" ]]; then + echo "Usage: $0 /absolute/path/to/${EXPECTED_NAME}" >&2 + exit 2 +fi +if [[ ! -f "${MODEL_PATH}" ]]; then + echo "Model file does not exist: ${MODEL_PATH}" >&2 + exit 1 +fi +if [[ "$(basename "${MODEL_PATH}")" != "${EXPECTED_NAME}" ]]; then + echo "Expected filename ${EXPECTED_NAME}, got $(basename "${MODEL_PATH}")" >&2 + exit 1 +fi + +ACTUAL_SIZE="$(wc -c < "${MODEL_PATH}" | tr -d ' ')" +if [[ "${ACTUAL_SIZE}" != "${EXPECTED_SIZE}" ]]; then + echo "Size mismatch: expected ${EXPECTED_SIZE}, got ${ACTUAL_SIZE}" >&2 + exit 1 +fi + +if command -v shasum >/dev/null 2>&1; then + ACTUAL_SHA256="$(shasum -a 256 "${MODEL_PATH}" | awk '{print $1}')" +elif command -v sha256sum >/dev/null 2>&1; then + ACTUAL_SHA256="$(sha256sum "${MODEL_PATH}" | awk '{print $1}')" +else + echo "Neither shasum nor sha256sum is available" >&2 + exit 2 +fi + +if [[ "${ACTUAL_SHA256}" != "${EXPECTED_SHA256}" ]]; then + echo "SHA-256 mismatch: expected ${EXPECTED_SHA256}, got ${ACTUAL_SHA256}" >&2 + exit 1 +fi + +echo "Validated ${MODEL_PATH} (${ACTUAL_SIZE} bytes, SHA-256 ${ACTUAL_SHA256})" diff --git a/scripts/verify-whisper-spike-isolation.sh b/scripts/verify-whisper-spike-isolation.sh new file mode 100755 index 0000000..27b9ab6 --- /dev/null +++ b/scripts/verify-whisper-spike-isolation.sh @@ -0,0 +1,38 @@ +#!/usr/bin/env bash +set -euo pipefail + +ROOT_DIR="$(cd "$(dirname "$0")/.." && pwd)" +ARTIFACT="${1:-}" + +cd "${ROOT_DIR}" + +if [[ -n "${ARTIFACT}" ]]; then + if [[ ! -e "${ARTIFACT}" ]]; then + echo "Artifact does not exist: ${ARTIFACT}" >&2 + exit 2 + fi + SYMBOL_ARTIFACT="${ARTIFACT}" + if [[ -d "${ARTIFACT}" && -f "${ARTIFACT}/VoiceInbox" ]]; then + SYMBOL_ARTIFACT="${ARTIFACT}/VoiceInbox" + fi + if command -v nm >/dev/null 2>&1 && nm -g "${SYMBOL_ARTIFACT}" 2>/dev/null | grep -q 'voiceinbox_whisper_spike'; then + echo "Ordinary artifact unexpectedly exports Whisper evaluation symbols: ${ARTIFACT}" >&2 + exit 1 + fi + case "${ARTIFACT}" in + *.app) + if find "${ARTIFACT}" -type f \( -name 'ggml-tiny.bin' -o -name '*.pcm' -o -name 'encoder-model*.onnx' \) | grep -q .; then + echo "Application bundle unexpectedly contains model weights or test PCM: ${ARTIFACT}" >&2 + exit 1 + fi + ;; + *.ipa|*.apk|*.aab|*.zip) + if unzip -l "${ARTIFACT}" 2>/dev/null | grep -q 'ggml-tiny\.bin'; then + echo "Default package unexpectedly contains ggml-tiny.bin: ${ARTIFACT}" >&2 + exit 1 + fi + ;; + esac +fi + +echo "Build contains no Whisper evaluation exports or model weights" diff --git a/scripts/whisper-android-cmake/android.toolchain.cmake b/scripts/whisper-android-cmake/android.toolchain.cmake new file mode 100644 index 0000000..f1e335a --- /dev/null +++ b/scripts/whisper-android-cmake/android.toolchain.cmake @@ -0,0 +1,11 @@ +# Compatibility wrapper for whisper-rs-sys/cmake-rs Android cross-compilation. +set(ANDROID_USE_LEGACY_TOOLCHAIN_FILE FALSE CACHE BOOL "Use CMake's Android toolchain support" FORCE) +set(CMAKE_ANDROID_ARCH_ABI "arm64-v8a" CACHE STRING "Voice Inbox Whisper ABI" FORCE) +set(ANDROID_ABI "arm64-v8a" CACHE STRING "Voice Inbox Whisper ABI compatibility" FORCE) +set(ANDROID_PLATFORM "android-24" CACHE STRING "Voice Inbox minimum API" FORCE) + +if(NOT DEFINED ENV{ANDROID_NDK_HOME}) + message(FATAL_ERROR "ANDROID_NDK_HOME is required") +endif() + +include("$ENV{ANDROID_NDK_HOME}/build/cmake/android.toolchain.cmake") diff --git a/shared/src/commonMain/kotlin/me/maxistar/voiceinbox/core/BatchTranscriptionUseCase.kt b/shared/src/commonMain/kotlin/me/maxistar/voiceinbox/core/BatchTranscriptionUseCase.kt index eecd85c..dd94432 100644 --- a/shared/src/commonMain/kotlin/me/maxistar/voiceinbox/core/BatchTranscriptionUseCase.kt +++ b/shared/src/commonMain/kotlin/me/maxistar/voiceinbox/core/BatchTranscriptionUseCase.kt @@ -52,6 +52,47 @@ interface BatchClock { fun currentTimeMillis(): Long } +/** + * Limits decoder-level progress to a cadence presentation layers can draw reliably. + * Lifecycle changes bypass the interval so ownership and terminal outcomes remain prompt. + */ +class BatchProgressCoalescer( + private val clock: BatchClock, + private val minimumIntervalMillis: Long = CONTINUOUS_PROGRESS_INTERVAL_MILLIS, +) { + private var lastPublished: BatchTranscriptionProgress? = null + private var lastContinuousPublicationMillis: Long? = null + + fun shouldPublish(progress: BatchTranscriptionProgress): Boolean { + val previous = lastPublished + val lifecycleChanged = previous == null || + previous.phase != progress.phase || + previous.activeEntryId != progress.activeEntryId || + previous.filename != progress.filename || + previous.completed != progress.completed || + previous.total != progress.total || + previous.failed != progress.failed + if (lifecycleChanged) { + lastPublished = progress + lastContinuousPublicationMillis = clock.currentTimeMillis() + return true + } + + val now = clock.currentTimeMillis() + val lastContinuous = lastContinuousPublicationMillis + if (lastContinuous == null || now - lastContinuous >= minimumIntervalMillis) { + lastPublished = progress + lastContinuousPublicationMillis = now + return true + } + return false + } + + companion object { + const val CONTINUOUS_PROGRESS_INTERVAL_MILLIS = 250L + } +} + class BatchTranscriptionUseCase( private val catalog: AudioCatalogQueuePort, private val transcriber: BatchEntryTranscriber, @@ -61,6 +102,10 @@ class BatchTranscriptionUseCase( input: BatchTranscriptionInput, onProgress: (BatchTranscriptionProgress) -> Unit, ): BatchTranscriptionResult { + val progressCoalescer = BatchProgressCoalescer(clock) + fun publish(progress: BatchTranscriptionProgress) { + if (progressCoalescer.shouldPublish(progress)) onProgress(progress) + } catalog.recoverInterrupted() val total = if (input.retryEntryId == null) { catalog.pendingCount(input.sourceScope) @@ -81,7 +126,7 @@ class BatchTranscriptionUseCase( outputId = input.outputId, runId = input.runId, ) { progress -> - onProgress( + publish( BatchTranscriptionProgress( phase = progress.phase, activeEntryId = entry.id, @@ -117,7 +162,7 @@ class BatchTranscriptionUseCase( } completed += 1 currentEntry = null - onProgress(summaryProgress(completed, total, failed)) + publish(summaryProgress(completed, total, failed)) if (input.retryEntryId != null) break } } catch (cancelled: CancellationException) { diff --git a/shared/src/commonMain/kotlin/me/maxistar/voiceinbox/core/SpeechModelManifest.kt b/shared/src/commonMain/kotlin/me/maxistar/voiceinbox/core/SpeechModelManifest.kt index 8ec710c..61c5336 100644 --- a/shared/src/commonMain/kotlin/me/maxistar/voiceinbox/core/SpeechModelManifest.kt +++ b/shared/src/commonMain/kotlin/me/maxistar/voiceinbox/core/SpeechModelManifest.kt @@ -21,8 +21,67 @@ data class SpeechModelManifest( } } -object EmbeddedSpeechModel { - val manifest = SpeechModelManifest( +enum class SpeechModelBackend { + PARAKEET_TDT_ONNX, + WHISPER_CPP, +} + +enum class SpeechModelMaturity { + STABLE, + EXPERIMENTAL, +} + +enum class SpeechModelPlatform { + ANDROID, + IOS, +} + +data class SpeechModelDistribution( + val networkDownloadAvailable: Boolean, + val localImportAvailable: Boolean, +) + +data class SpeechModelLanguageCoverage( + val summary: String, + val languageTags: List, +) + +data class SpeechModelAttribution( + val sourceName: String, + val sourceUrl: String, + val upstreamName: String, + val upstreamUrl: String, + val licenseName: String, + val licenseUrl: String, +) + +data class SpeechModelDescriptor( + val catalogId: String, + val displayName: String, + val backend: SpeechModelBackend, + val manifest: SpeechModelManifest, + val distribution: SpeechModelDistribution, + val languages: SpeechModelLanguageCoverage, + val maturity: SpeechModelMaturity, + val attribution: SpeechModelAttribution, + val supportedPlatforms: Set, +) { + val approximateDownloadBytes: Long = manifest.totalSizeBytes + val requiredStorageBytes: Long = manifest.requiredFreeBytes +} + +data class SpeechModelPackageIdentity( + val schemaVersion: Int, + val catalogId: String, + val modelVersion: String, +) + +object SpeechModelCatalog { + val parakeetTdt06bV3Int8 = SpeechModelDescriptor( + catalogId = "parakeet-tdt-0.6b-v3-int8", + displayName = "Parakeet TDT 0.6B v3 INT8", + backend = SpeechModelBackend.PARAKEET_TDT_ONNX, + manifest = SpeechModelManifest( modelId = "istupakov/parakeet-tdt-0.6b-v3-onnx", version = "parakeet-tdt-0.6b-v3-int8-r1", repositoryRevision = "8f23f0c03c8761650bdb5b40aaf3e40d2c15f1ce", @@ -54,5 +113,119 @@ object EmbeddedSpeechModel { ), ), safetyMarginBytes = 64L * 1024L * 1024L, + ), + distribution = SpeechModelDistribution( + networkDownloadAvailable = true, + localImportAvailable = true, + ), + languages = SpeechModelLanguageCoverage( + summary = "25 European languages", + languageTags = listOf( + "bg", "cs", "da", "de", "el", "en", "es", "et", "fi", "fr", "hr", "hu", + "it", "lt", "lv", "mt", "nl", "pl", "pt", "ro", "ru", "sk", "sl", "sv", "uk", + ), + ), + maturity = SpeechModelMaturity.STABLE, + attribution = SpeechModelAttribution( + sourceName = "istupakov/parakeet-tdt-0.6b-v3-onnx", + sourceUrl = "https://huggingface.co/istupakov/parakeet-tdt-0.6b-v3-onnx", + upstreamName = "NVIDIA Parakeet TDT 0.6B v3", + upstreamUrl = "https://huggingface.co/nvidia/parakeet-tdt-0.6b-v3", + licenseName = "CC BY 4.0", + licenseUrl = "https://creativecommons.org/licenses/by/4.0/", + ), + supportedPlatforms = setOf(SpeechModelPlatform.ANDROID, SpeechModelPlatform.IOS), + ) + + val whisperTinyMultilingual = SpeechModelDescriptor( + catalogId = "whisper-tiny-multilingual", + displayName = "Whisper Tiny Multilingual", + backend = SpeechModelBackend.WHISPER_CPP, + manifest = SpeechModelManifest( + modelId = "ggerganov/whisper.cpp", + version = "whisper-tiny-ggml-f16-r1", + repositoryRevision = "5359861c739e955e79d9a303bcbc70fb988958b1", + files = listOf( + SpeechModelFile( + name = "ggml-tiny.bin", + sizeBytes = 77_691_713, + sha256 = "be07e048e1e599ad46341c8d2a135645097a538221678b7acdd1b1919c6e1b21", + ), + ), + safetyMarginBytes = 64L * 1024L * 1024L, + ), + distribution = SpeechModelDistribution( + networkDownloadAvailable = true, + localImportAvailable = true, + ), + languages = SpeechModelLanguageCoverage( + summary = "Multilingual", + languageTags = emptyList(), + ), + maturity = SpeechModelMaturity.EXPERIMENTAL, + attribution = SpeechModelAttribution( + sourceName = "whisper.cpp Whisper Tiny model", + sourceUrl = "https://huggingface.co/ggerganov/whisper.cpp", + upstreamName = "OpenAI Whisper Tiny", + upstreamUrl = "https://github.com/openai/whisper", + licenseName = "MIT", + licenseUrl = "https://github.com/openai/whisper/blob/main/LICENSE", + ), + supportedPlatforms = setOf(SpeechModelPlatform.ANDROID, SpeechModelPlatform.IOS), ) + + val models: List = listOf( + parakeetTdt06bV3Int8, + whisperTinyMultilingual, + ) + val defaultModel: SpeechModelDescriptor = parakeetTdt06bV3Int8 + + fun resolvePackage( + identity: SpeechModelPackageIdentity, + platform: SpeechModelPlatform, + ): SpeechModelDescriptor? { + if (identity.schemaVersion != PACKAGE_SCHEMA_VERSION) return null + return models.singleOrNull { + it.catalogId == identity.catalogId && + it.manifest.version == identity.modelVersion && + platform in it.supportedPlatforms && + it.distribution.localImportAvailable + } + } + + fun resolveInstallation(catalogId: String, modelVersion: String): SpeechModelDescriptor? = + models.singleOrNull { it.catalogId == catalogId && it.manifest.version == modelVersion } + + fun modelsFor(platform: SpeechModelPlatform): List = + models.filter { platform in it.supportedPlatforms } + + fun networkDownloadChoices(platform: SpeechModelPlatform): List = + modelsFor(platform) + .filter { it.distribution.networkDownloadAvailable } + .map { descriptor -> + SpeechModelDownloadChoice( + identity = SpeechModelPackageIdentity( + schemaVersion = PACKAGE_SCHEMA_VERSION, + catalogId = descriptor.catalogId, + modelVersion = descriptor.manifest.version, + ), + displayName = descriptor.displayName, + languageSummary = descriptor.languages.summary, + maturity = descriptor.maturity.name.lowercase().replaceFirstChar { it.uppercase() }, + downloadBytes = descriptor.approximateDownloadBytes, + requiredStorageBytes = descriptor.requiredStorageBytes, + ) + } + + fun resolveNetworkDownload( + identity: SpeechModelPackageIdentity, + platform: SpeechModelPlatform, + ): SpeechModelDescriptor? = modelsFor(platform).singleOrNull { + it.catalogId == identity.catalogId && + it.manifest.version == identity.modelVersion && + it.distribution.networkDownloadAvailable + } + + const val PACKAGE_SCHEMA_VERSION = 1 + const val PACKAGE_MANIFEST_FILENAME = "voice-inbox-model.json" } diff --git a/shared/src/commonMain/kotlin/me/maxistar/voiceinbox/core/TaskListPresentationController.kt b/shared/src/commonMain/kotlin/me/maxistar/voiceinbox/core/TaskListPresentationController.kt index cce19f5..bfd6049 100644 --- a/shared/src/commonMain/kotlin/me/maxistar/voiceinbox/core/TaskListPresentationController.kt +++ b/shared/src/commonMain/kotlin/me/maxistar/voiceinbox/core/TaskListPresentationController.kt @@ -33,6 +33,7 @@ enum class TaskRetention { enum class TaskActionKind { DOWNLOAD_MODEL, + SELECT_DOWNLOAD_MODEL, IMPORT_MODEL, CANCEL_MODEL_DOWNLOAD, RETRY_MODEL_DOWNLOAD, @@ -116,6 +117,17 @@ data class ModelSetupSnapshot( val progressPercent: Int? = null, val downloadAvailable: Boolean = false, val canCancel: Boolean = false, + val selectedModel: SpeechModelPackageIdentity? = null, + val downloadChoices: List = emptyList(), +) + +data class SpeechModelDownloadChoice( + val identity: SpeechModelPackageIdentity, + val displayName: String, + val languageSummary: String, + val maturity: String, + val downloadBytes: Long, + val requiredStorageBytes: Long, ) enum class OutputSetupSnapshotState { @@ -256,14 +268,20 @@ object TaskListPresentationController { ) active -> emptyList() error -> listOf( + TaskActionPresentation( + TaskActionKind.SELECT_DOWNLOAD_MODEL, + selectedModelLabel(snapshot), + snapshot.downloadChoices.isNotEmpty(), + ), TaskActionPresentation(TaskActionKind.RETRY_MODEL_DOWNLOAD, "Retry Download", snapshot.downloadAvailable), - TaskActionPresentation(TaskActionKind.IMPORT_MODEL, "Install Manually"), ) else -> buildList { + if (snapshot.downloadChoices.size > 1) { + add(TaskActionPresentation(TaskActionKind.SELECT_DOWNLOAD_MODEL, selectedModelLabel(snapshot))) + } if (snapshot.downloadAvailable) { - add(TaskActionPresentation(TaskActionKind.DOWNLOAD_MODEL, "Download Model")) + add(TaskActionPresentation(TaskActionKind.DOWNLOAD_MODEL, "Download")) } - add(TaskActionPresentation(TaskActionKind.IMPORT_MODEL, "Install Manually")) } } return SetupTaskPresentation( @@ -275,7 +293,7 @@ object TaskListPresentationController { else -> SetupTaskState.REQUIRED }, title = "Install Speech Model", - detail = snapshot.detail.takeUnless { active }, + detail = selectedModelDetail(snapshot) ?: snapshot.detail.takeUnless { active }, badge = if (active) "Installing" else if (error) "Needs attention" else "Required", progress = if (active) { TaskProgressPresentation( @@ -290,6 +308,25 @@ object TaskListPresentationController { ) } + private fun selectedModelDetail(snapshot: ModelSetupSnapshot): String? { + val selected = snapshot.selectedModel ?: return null + val choice = snapshot.downloadChoices.firstOrNull { it.identity == selected } ?: return null + val download = formatBytes(choice.downloadBytes) + val storage = formatBytes(choice.requiredStorageBytes) + return "${choice.displayName} · ${choice.languageSummary} · ${choice.maturity} · $download download · $storage free" + } + + private fun selectedModelLabel(snapshot: ModelSetupSnapshot): String = + snapshot.selectedModel + ?.let { selected -> snapshot.downloadChoices.firstOrNull { it.identity == selected } } + ?.displayName + ?: "Choose Model" + + private fun formatBytes(bytes: Long): String = when { + bytes >= 1024L * 1024L * 1024L -> "${bytes / (1024L * 1024L * 1024L)} GB" + else -> "${bytes / (1024L * 1024L)} MB" + } + private fun outputTask(snapshot: OutputSetupSnapshot): SetupTaskPresentation? { if (snapshot.state == OutputSetupSnapshotState.READY) return null val error = snapshot.state == OutputSetupSnapshotState.INVALID @@ -431,12 +468,13 @@ object TaskListPresentationController { } private fun audioComparator(filter: TaskListFilter): Comparator = when (filter) { - TaskListFilter.PROCESSED -> compareByDescending { - it.terminalAtMillis ?: it.importedAtMillis - }.thenByDescending { it.importedAtMillis }.thenByDescending { it.entryId } - TaskListFilter.NEW, + TaskListFilter.PROCESSED, TaskListFilter.ALL, - -> compareByDescending { it.importedAtMillis }.thenByDescending { it.entryId } + -> compareByDescending { it.terminalAtMillis ?: it.importedAtMillis } + .thenByDescending { it.importedAtMillis } + .thenByDescending { it.entryId } + TaskListFilter.NEW -> compareByDescending { it.importedAtMillis } + .thenByDescending { it.entryId } } private fun emptyMessage(filter: TaskListFilter): String = when (filter) { diff --git a/shared/src/commonTest/kotlin/me/maxistar/voiceinbox/core/BatchTranscriptionUseCaseTest.kt b/shared/src/commonTest/kotlin/me/maxistar/voiceinbox/core/BatchTranscriptionUseCaseTest.kt index 4958817..450eed0 100644 --- a/shared/src/commonTest/kotlin/me/maxistar/voiceinbox/core/BatchTranscriptionUseCaseTest.kt +++ b/shared/src/commonTest/kotlin/me/maxistar/voiceinbox/core/BatchTranscriptionUseCaseTest.kt @@ -7,6 +7,38 @@ import kotlin.test.assertFailsWith import kotlin.test.assertNull class BatchTranscriptionUseCaseTest { + @Test + fun progressCoalescerBoundsContinuousUpdatesAndPublishesLifecycleChangesImmediately() { + val clock = MutableClock() + val coalescer = BatchProgressCoalescer(clock) + val base = BatchTranscriptionProgress( + phase = "Transcribing", + activeEntryId = 1, + filename = "one.wav", + completed = 0, + total = 2, + failed = 0, + processedUs = 1, + durationUs = 100, + progress = 1, + ) + + assertEquals(true, coalescer.shouldPublish(base)) + clock.now = 100 + assertEquals(false, coalescer.shouldPublish(base.copy(processedUs = 10, progress = 10))) + clock.now = 250 + assertEquals(true, coalescer.shouldPublish(base.copy(processedUs = 25, progress = 25))) + clock.now = 251 + assertEquals( + true, + coalescer.shouldPublish(base.copy(phase = "Appending", processedUs = 100, progress = 100)), + ) + assertEquals( + true, + coalescer.shouldPublish(base.copy(completed = 1, activeEntryId = null, filename = null)), + ) + } + @Test fun transcribeAllProcessesPendingEntriesAndReportsResult() { val catalog = FakeCatalog( @@ -172,6 +204,10 @@ class BatchTranscriptionUseCaseTest { retryEntryId = retryEntryId, ) + private class MutableClock(var now: Long = 0) : BatchClock { + override fun currentTimeMillis(): Long = now + } + private class FakeCatalog( vararg entries: AudioCatalogEntry, ) : AudioCatalogQueuePort { diff --git a/shared/src/commonTest/kotlin/me/maxistar/voiceinbox/core/SpeechModelManifestTest.kt b/shared/src/commonTest/kotlin/me/maxistar/voiceinbox/core/SpeechModelManifestTest.kt index d8cbcd7..bcd74e6 100644 --- a/shared/src/commonTest/kotlin/me/maxistar/voiceinbox/core/SpeechModelManifestTest.kt +++ b/shared/src/commonTest/kotlin/me/maxistar/voiceinbox/core/SpeechModelManifestTest.kt @@ -7,8 +7,25 @@ import kotlin.test.Test class SpeechModelManifestTest { @Test - fun embeddedManifestPinsAllRuntimeFiles() { - val manifest = EmbeddedSpeechModel.manifest + fun productionCatalogKeepsStableParakeetDefaultAndAddsDownloadableMobileWhisper() { + val descriptor = SpeechModelCatalog.defaultModel + val manifest = descriptor.manifest + + assertEquals( + listOf(descriptor, SpeechModelCatalog.whisperTinyMultilingual), + SpeechModelCatalog.models, + ) + assertEquals("parakeet-tdt-0.6b-v3-int8", descriptor.catalogId) + assertEquals(SpeechModelBackend.PARAKEET_TDT_ONNX, descriptor.backend) + assertEquals(SpeechModelMaturity.STABLE, descriptor.maturity) + assertEquals( + setOf(SpeechModelPlatform.ANDROID, SpeechModelPlatform.IOS), + descriptor.supportedPlatforms, + ) + assertTrue(descriptor.distribution.networkDownloadAvailable) + assertTrue(descriptor.distribution.localImportAvailable) + assertTrue(descriptor.languages.languageTags.containsAll(listOf("en", "de", "ru", "uk"))) + assertEquals("CC BY 4.0", descriptor.attribution.licenseName) assertEquals(40, manifest.repositoryRevision.length) assertEquals( @@ -24,5 +41,52 @@ class SpeechModelManifestTest { assertTrue(manifest.files.all { it.sizeBytes > 0 }) assertTrue(manifest.files.all { it.sha256.matches(Regex("[0-9a-f]{64}")) }) assertEquals(670_619_803, manifest.totalSizeBytes) + assertEquals(manifest.totalSizeBytes, descriptor.approximateDownloadBytes) + assertEquals(manifest.requiredFreeBytes, descriptor.requiredStorageBytes) + assertEquals( + "https://huggingface.co/istupakov/parakeet-tdt-0.6b-v3-onnx/resolve/" + + "8f23f0c03c8761650bdb5b40aaf3e40d2c15f1ce/encoder-model.int8.onnx?download=true", + manifest.downloadUrl(manifest.files.first()), + ) + + val whisper = SpeechModelCatalog.whisperTinyMultilingual + assertEquals(SpeechModelBackend.WHISPER_CPP, whisper.backend) + assertEquals(SpeechModelMaturity.EXPERIMENTAL, whisper.maturity) + assertEquals( + setOf(SpeechModelPlatform.ANDROID, SpeechModelPlatform.IOS), + whisper.supportedPlatforms, + ) + assertTrue(whisper.distribution.localImportAvailable) + assertTrue(whisper.distribution.networkDownloadAvailable) + assertEquals("ggml-tiny.bin", whisper.manifest.files.single().name) + assertEquals(77_691_713, whisper.manifest.totalSizeBytes) + assertEquals(listOf(descriptor, whisper), SpeechModelCatalog.modelsFor(SpeechModelPlatform.IOS)) + assertEquals( + listOf(descriptor, whisper), + SpeechModelCatalog.networkDownloadChoices(SpeechModelPlatform.ANDROID).map { choice -> + SpeechModelCatalog.resolveNetworkDownload(choice.identity, SpeechModelPlatform.ANDROID) + }, + ) + } + + @Test + fun packageResolutionRequiresExactTrustedIdentityAndPlatform() { + val identity = SpeechModelPackageIdentity( + schemaVersion = 1, + catalogId = "whisper-tiny-multilingual", + modelVersion = "whisper-tiny-ggml-f16-r1", + ) + + assertEquals( + SpeechModelCatalog.whisperTinyMultilingual, + SpeechModelCatalog.resolvePackage(identity, SpeechModelPlatform.ANDROID), + ) + assertEquals( + SpeechModelCatalog.whisperTinyMultilingual, + SpeechModelCatalog.resolvePackage(identity, SpeechModelPlatform.IOS), + ) + assertEquals(null, SpeechModelCatalog.resolvePackage(identity.copy(schemaVersion = 2), SpeechModelPlatform.ANDROID)) + assertEquals(null, SpeechModelCatalog.resolvePackage(identity.copy(modelVersion = "latest"), SpeechModelPlatform.ANDROID)) + assertEquals(null, SpeechModelCatalog.resolvePackage(identity.copy(catalogId = "unknown"), SpeechModelPlatform.ANDROID)) } } diff --git a/shared/src/commonTest/kotlin/me/maxistar/voiceinbox/core/TaskListPresentationControllerTest.kt b/shared/src/commonTest/kotlin/me/maxistar/voiceinbox/core/TaskListPresentationControllerTest.kt index 82c7c3b..5f535bd 100644 --- a/shared/src/commonTest/kotlin/me/maxistar/voiceinbox/core/TaskListPresentationControllerTest.kt +++ b/shared/src/commonTest/kotlin/me/maxistar/voiceinbox/core/TaskListPresentationControllerTest.kt @@ -8,6 +8,35 @@ import kotlin.test.assertNull import kotlin.test.assertTrue class TaskListPresentationControllerTest { + @Test + fun modelTaskShowsSelectedDownloadDescriptorAndLocksChoiceDuringInstall() { + val whisper = SpeechModelCatalog.networkDownloadChoices(SpeechModelPlatform.ANDROID) + .single { it.identity.catalogId == "whisper-tiny-multilingual" } + val required = assertIs( + state( + model = ModelSetupSnapshot( + state = ModelSetupSnapshotState.REQUIRED, + downloadAvailable = true, + selectedModel = whisper.identity, + downloadChoices = SpeechModelCatalog.networkDownloadChoices(SpeechModelPlatform.ANDROID), + ), + ).tasks.single(), + ) + assertTrue(required.detail?.contains("Whisper Tiny Multilingual") == true) + assertTrue(required.actions.any { it.kind == TaskActionKind.SELECT_DOWNLOAD_MODEL }) + + val installing = assertIs( + state( + model = ModelSetupSnapshot( + state = ModelSetupSnapshotState.INSTALLING, + selectedModel = whisper.identity, + downloadChoices = SpeechModelCatalog.networkDownloadChoices(SpeechModelPlatform.ANDROID), + ), + ).tasks.single(), + ) + assertFalse(installing.actions.any { it.kind == TaskActionKind.SELECT_DOWNLOAD_MODEL }) + } + @Test fun setupTasksAreSynthesizedInKindOrderAndCompletedTasksDisappear() { val state = state( @@ -111,7 +140,7 @@ class TaskListPresentationControllerTest { assertEquals("Verification failed", task.errorMessage) assertFalse(task.actions.single { it.kind == TaskActionKind.RETRY_MODEL_DOWNLOAD }.enabled) - assertTrue(task.actions.single { it.kind == TaskActionKind.IMPORT_MODEL }.enabled) + assertFalse(task.actions.any { it.kind == TaskActionKind.IMPORT_MODEL }) } @Test @@ -215,14 +244,32 @@ class TaskListPresentationControllerTest { } @Test - fun orderingUsesImportTimeExceptProcessedUsesTerminalTime() { + fun terminalOnlyAllUsesTheSameTerminalTimeOrderAsProcessed() { val audio = listOf( audio(1, AudioFileState.PROCESSED, importedAt = 300, terminalAt = 100), audio(2, AudioFileState.FAILED, importedAt = 100, terminalAt = 400), ) assertEquals(listOf("audio:2", "audio:1"), state(TaskListFilter.PROCESSED, audio = audio).tasks.map { it.stableId }) - assertEquals(listOf("audio:1", "audio:2"), state(TaskListFilter.ALL, audio = audio).tasks.map { it.stableId }) + assertEquals(listOf("audio:2", "audio:1"), state(TaskListFilter.ALL, audio = audio).tasks.map { it.stableId }) + } + + @Test + fun newKeepsImportTimeOrderWhenAllUsesLifecycleTime() { + val audio = listOf( + audio(1, AudioFileState.PENDING, importedAt = 300), + audio(2, AudioFileState.PENDING, importedAt = 100), + audio(3, AudioFileState.PROCESSED, importedAt = 50, terminalAt = 400), + ) + + assertEquals( + listOf("audio:1", "audio:2"), + state(TaskListFilter.NEW, audio = audio).tasks.map { it.stableId }, + ) + assertEquals( + listOf("audio:3", "audio:1", "audio:2"), + state(TaskListFilter.ALL, audio = audio).tasks.map { it.stableId }, + ) } @Test diff --git a/src/engine.rs b/src/engine.rs index 4a8e044..b002c63 100644 --- a/src/engine.rs +++ b/src/engine.rs @@ -1,22 +1,67 @@ use once_cell::sync::Lazy; use std::path::PathBuf; use std::sync::{Arc, Condvar, Mutex}; -use transcribe_rs::engines::parakeet::ParakeetEngine; -use transcribe_rs::TranscriptionEngine; +use transcribe_rs::engines::parakeet::{ + ParakeetEngine, ParakeetInferenceParams, ParakeetModelParams, TimestampGranularity, +}; +#[cfg(any( + all(target_os = "android", feature = "android-dual-backend"), + all(target_os = "ios", feature = "ios-dual-backend") +))] +use transcribe_rs::engines::whisper::WhisperEngine; +use transcribe_rs::{TranscriptionEngine, TranscriptionResult}; #[cfg(target_os = "android")] use jni::objects::{GlobalRef, JObject}; #[cfg(target_os = "android")] use jni::JNIEnv; -static GLOBAL_ENGINE: Lazy>>>> = +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct ModelConfiguration { + pub backend: String, + pub identity: String, + pub directory: PathBuf, + pub primary_file: String, +} + +pub enum ActiveSpeechEngine { + Parakeet(ParakeetEngine), + #[cfg(any( + all(target_os = "android", feature = "android-dual-backend"), + all(target_os = "ios", feature = "ios-dual-backend") + ))] + Whisper(WhisperEngine), +} + +impl ActiveSpeechEngine { + pub fn transcribe_samples(&mut self, samples: Vec) -> Result { + match self { + Self::Parakeet(engine) => engine + .transcribe_samples( + samples, + Some(ParakeetInferenceParams { + timestamp_granularity: TimestampGranularity::Word, + }), + ) + .map_err(|error| error.to_string()), + #[cfg(any( + all(target_os = "android", feature = "android-dual-backend"), + all(target_os = "ios", feature = "ios-dual-backend") + ))] + Self::Whisper(engine) => engine + .transcribe_samples(samples, None) + .map_err(|error| error.to_string()), + } + } +} + +static GLOBAL_ENGINE: Lazy>>>> = Lazy::new(|| Mutex::new(None)); -static MODEL_DIRECTORY: Lazy>> = Lazy::new(|| Mutex::new(None)); +static CONFIGURATION: Lazy>> = Lazy::new(|| Mutex::new(None)); static LOAD_STATE: Lazy<(Mutex, Condvar)> = Lazy::new(|| (Mutex::new(LoadState::Idle), Condvar::new())); #[derive(Debug, Clone, PartialEq)] -#[allow(dead_code)] enum LoadState { Idle, Loading, @@ -24,22 +69,32 @@ enum LoadState { Failed(String), } -pub fn get_engine() -> Option>> { +pub fn get_engine() -> Option>> { GLOBAL_ENGINE.lock().unwrap().clone() } -pub fn is_engine_loaded() -> bool { +fn is_engine_loaded() -> bool { GLOBAL_ENGINE.lock().unwrap().is_some() } -pub fn configure_model_directory(path: PathBuf) { - let path = std::fs::canonicalize(&path).unwrap_or(path); - let unchanged = MODEL_DIRECTORY.lock().unwrap().as_ref() == Some(&path); - if unchanged { +pub fn configure_model(mut configuration: ModelConfiguration) { + configuration.directory = + std::fs::canonicalize(&configuration.directory).unwrap_or(configuration.directory); + if CONFIGURATION.lock().unwrap().as_ref() == Some(&configuration) { return; } invalidate_loaded_model(); - *MODEL_DIRECTORY.lock().unwrap() = Some(path); + *CONFIGURATION.lock().unwrap() = Some(configuration); +} + +#[cfg(target_os = "ios")] +pub fn configure_model_directory(path: PathBuf) { + configure_model(ModelConfiguration { + backend: "PARAKEET_TDT_ONNX".to_string(), + identity: path.to_string_lossy().into_owned(), + directory: path, + primary_file: "encoder-model.int8.onnx".to_string(), + }); } pub fn invalidate_loaded_model() { @@ -57,65 +112,62 @@ fn ensure_loaded_with_status(mut status: impl FnMut(&str)) -> Result<(), String> status("Ready"); return Ok(()); } - let (lock, cvar) = &*LOAD_STATE; let mut state = lock.lock().unwrap(); - if is_engine_loaded() { - status("Ready"); - return Ok(()); - } - if *state == LoadState::Loading { status("Waiting for model..."); while *state == LoadState::Loading { state = cvar.wait(state).unwrap(); } return match &*state { - LoadState::Done if is_engine_loaded() => { - status("Ready"); - Ok(()) - } - LoadState::Failed(message) => { - status(&format!("Error: {message}")); - Err(message.clone()) - } + LoadState::Done if is_engine_loaded() => Ok(()), + LoadState::Failed(message) => Err(message.clone()), _ => Err("Model loading was interrupted".to_string()), }; } - *state = LoadState::Loading; drop(state); status("Loading model..."); - let result = load_configured_engine(); let mut state = lock.lock().unwrap(); - match &result { - Ok(()) => { - *state = LoadState::Done; - status("Ready"); - } - Err(message) => { - *state = LoadState::Failed(message.clone()); - status(&format!("Error: {message}")); - } - } + *state = match &result { + Ok(()) => LoadState::Done, + Err(message) => LoadState::Failed(message.clone()), + }; + status(if result.is_ok() { "Ready" } else { "Model loading failed" }); cvar.notify_all(); result } fn load_configured_engine() -> Result<(), String> { - let path = MODEL_DIRECTORY - .lock() - .unwrap() - .clone() - .ok_or_else(|| "Model directory was not configured".to_string())?; - let mut engine = ParakeetEngine::new(); - engine - .load_model_with_params( - &path, - transcribe_rs::engines::parakeet::ParakeetModelParams::int8(), - ) - .map_err(|error| format!("Model error: {error}"))?; + let configuration = CONFIGURATION + .lock().unwrap().clone().ok_or_else(|| "Model was not configured".to_string())?; + let engine = match configuration.backend.as_str() { + "PARAKEET_TDT_ONNX" => { + let mut engine = ParakeetEngine::new(); + engine.load_model_with_params(&configuration.directory, ParakeetModelParams::int8()) + .map_err(|error| format!("Model error: {error}"))?; + ActiveSpeechEngine::Parakeet(engine) + } + "WHISPER_CPP" => { + #[cfg(any( + all(target_os = "android", feature = "android-dual-backend"), + all(target_os = "ios", feature = "ios-dual-backend") + ))] + { + let mut engine = WhisperEngine::new(); + engine.load_model(&configuration.directory.join(&configuration.primary_file)) + .map_err(|error| format!("Model error: {error}"))?; + ActiveSpeechEngine::Whisper(engine) + } + #[cfg(not(any( + all(target_os = "android", feature = "android-dual-backend"), + all(target_os = "ios", feature = "ios-dual-backend") + )))] + return Err("Whisper backend is not available on this platform".to_string()); + } + backend => return Err(format!("Unsupported speech backend: {backend}")), + }; *GLOBAL_ENGINE.lock().unwrap() = Some(Arc::new(Mutex::new(engine))); Ok(()) } @@ -123,12 +175,7 @@ fn load_configured_engine() -> Result<(), String> { #[cfg(target_os = "android")] fn notify_status(env: &mut JNIEnv, obj: &JObject, msg: &str) { if let Ok(jmsg) = env.new_string(msg) { - let _ = env.call_method( - obj, - "onStatusUpdate", - "(Ljava/lang/String;)V", - &[(&jmsg).into()], - ); + let _ = env.call_method(obj, "onStatusUpdate", "(Ljava/lang/String;)V", &[(&jmsg).into()]); } } @@ -138,10 +185,7 @@ pub fn ensure_loaded(env: &mut JNIEnv, context: &JObject) -> Result<(), String> } #[cfg(target_os = "android")] -pub fn ensure_loaded_from_thread( - jvm: &Arc, - target_ref: &GlobalRef, -) -> Result<(), String> { +pub fn ensure_loaded_from_thread(jvm: &Arc, target_ref: &GlobalRef) -> Result<(), String> { ensure_loaded_with_status(|message| { if let Ok(mut env) = jvm.attach_current_thread() { notify_status(&mut env, target_ref.as_obj(), message); @@ -153,33 +197,22 @@ pub fn ensure_loaded_from_thread( mod tests { use super::*; - static TEST_LOCK: Lazy> = Lazy::new(|| Mutex::new(())); - #[test] - fn configuring_same_directory_preserves_load_state() { - let _guard = TEST_LOCK.lock().unwrap(); + fn same_complete_configuration_preserves_state_and_key_changes_reset_it() { invalidate_loaded_model(); - configure_model_directory(PathBuf::from("test-model")); + let first = ModelConfiguration { + backend: "PARAKEET_TDT_ONNX".into(), identity: "one".into(), + directory: PathBuf::from("model"), primary_file: "encoder.onnx".into(), + }; + configure_model(first.clone()); *LOAD_STATE.0.lock().unwrap() = LoadState::Done; - - configure_model_directory(PathBuf::from("test-model")); - + configure_model(first); assert_eq!(*LOAD_STATE.0.lock().unwrap(), LoadState::Done); - invalidate_loaded_model(); - } - - #[test] - fn changing_directory_and_explicit_invalidation_reset_state() { - let _guard = TEST_LOCK.lock().unwrap(); - invalidate_loaded_model(); - configure_model_directory(PathBuf::from("first-model")); - *LOAD_STATE.0.lock().unwrap() = LoadState::Done; - - configure_model_directory(PathBuf::from("second-model")); + configure_model(ModelConfiguration { + backend: "WHISPER_CPP".into(), identity: "two".into(), + directory: PathBuf::from("model"), primary_file: "ggml.bin".into(), + }); assert_eq!(*LOAD_STATE.0.lock().unwrap(), LoadState::Idle); - - *LOAD_STATE.0.lock().unwrap() = LoadState::Failed("failed".to_string()); invalidate_loaded_model(); - assert_eq!(*LOAD_STATE.0.lock().unwrap(), LoadState::Idle); } } diff --git a/src/lib.rs b/src/lib.rs index 9725736..734d5c1 100644 --- a/src/lib.rs +++ b/src/lib.rs @@ -1,6 +1,9 @@ #[cfg(any(not(target_os = "ios"), all(target_os = "ios", feature = "ios-onnx")))] pub mod engine; +#[cfg(feature = "whisper-mobile-spike")] +mod whisper_mobile_spike; + #[cfg(target_os = "ios")] use once_cell::sync::Lazy; #[cfg(target_os = "ios")] @@ -12,6 +15,26 @@ use std::path::Path; #[cfg(target_os = "ios")] use std::sync::Mutex; +#[cfg(feature = "whisper-mobile-spike")] +fn run_whisper_spike_operation( + operation: &str, + action: impl FnOnce() -> String + std::panic::UnwindSafe, +) -> String { + std::panic::catch_unwind(action).unwrap_or_else(|_| { + serde_json::json!({ + "schema_version": 1, + "backend": "whisper.cpp", + "operation": operation, + "status": "error", + "error": "Whisper spike operation panicked", + }) + .to_string() + }) +} + +#[cfg(all(target_os = "ios", feature = "whisper-mobile-spike"))] +use std::path::PathBuf; + #[cfg(target_os = "android")] use jni::objects::{JClass, JFloatArray, JString}; #[cfg(target_os = "android")] @@ -20,8 +43,6 @@ use jni::sys::{jboolean, jstring, JNI_FALSE, JNI_TRUE}; use jni::JNIEnv; #[cfg(target_os = "android")] use std::path::PathBuf; -#[cfg(any(target_os = "android", all(target_os = "ios", feature = "ios-onnx")))] -use transcribe_rs::TranscriptionEngine; #[cfg(any(not(target_os = "ios"), all(target_os = "ios", feature = "ios-onnx")))] fn serialize_chunk_result(result: transcribe_rs::TranscriptionResult) -> String { @@ -49,22 +70,30 @@ fn serialize_chunk_result(result: transcribe_rs::TranscriptionResult) -> String pub unsafe extern "system" fn Java_me_maxistar_voiceinbox_NativeTranscriptionBridge_initialize( mut env: JNIEnv, _class: JClass, + backend: JString, + installation_identity: JString, model_directory: JString, + primary_file: JString, ) -> jboolean { android_logger::init_once( android_logger::Config::default().with_max_level(log::LevelFilter::Info), ); let _ = ort::init().commit(); - let model_directory: String = match env.get_string(&model_directory) { - Ok(path) => path.into(), - Err(error) => { - log::error!("Failed to read model directory from JNI: {error}"); - return JNI_FALSE; - } + let read = |env: &mut JNIEnv, value: &JString| -> Result { + env.get_string(value).map(Into::into).map_err(|error| error.to_string()) + }; + let configuration = match ( + read(&mut env, &backend), read(&mut env, &installation_identity), + read(&mut env, &model_directory), read(&mut env, &primary_file), + ) { + (Ok(backend), Ok(identity), Ok(directory), Ok(primary_file)) => engine::ModelConfiguration { + backend, identity, directory: PathBuf::from(directory), primary_file, + }, + _ => return JNI_FALSE, }; - engine::configure_model_directory(PathBuf::from(model_directory)); + engine::configure_model(configuration); match engine::ensure_loaded_without_callback() { Ok(()) => JNI_TRUE, Err(error) => { @@ -105,18 +134,8 @@ pub unsafe extern "system" fn Java_me_maxistar_voiceinbox_NativeTranscriptionBri let result = engine::get_engine() .ok_or_else(|| "Model is not loaded".to_string()) .and_then(|engine| { - engine - .lock() - .unwrap() - .transcribe_samples( - buffer, - Some(transcribe_rs::engines::parakeet::ParakeetInferenceParams { - timestamp_granularity: - transcribe_rs::engines::parakeet::TimestampGranularity::Word, - }), - ) + engine.lock().unwrap().transcribe_samples(buffer) .map(serialize_chunk_result) - .map_err(|error| error.to_string()) }); match result.and_then(|text| env.new_string(text).map_err(|error| error.to_string())) { @@ -165,6 +184,18 @@ fn read_model_directory(model_directory: *const c_char) -> Result Result { + if value.is_null() { + return Err(format!("{label} was not provided")); + } + match unsafe { CStr::from_ptr(value) }.to_str() { + Ok(text) if !text.is_empty() => Ok(text.to_string()), + Ok(_) => Err(format!("{label} was empty")), + Err(_) => Err(format!("{label} was not valid UTF-8")), + } +} + #[cfg(all(target_os = "ios", feature = "ios-onnx"))] fn validate_ios_model_directory(model_directory: &str) -> Result<(), String> { let path = Path::new(model_directory); @@ -238,6 +269,59 @@ pub unsafe extern "C" fn voiceinbox_transcription_initialize( true } +#[cfg(all(target_os = "ios", feature = "ios-onnx"))] +#[no_mangle] +pub unsafe extern "C" fn voiceinbox_transcription_initialize_configured( + backend: *const c_char, + installation_identity: *const c_char, + model_directory: *const c_char, + primary_file: *const c_char, +) -> bool { + let configuration = match ( + read_required_c_string(backend, "Speech backend"), + read_required_c_string(installation_identity, "Installation identity"), + read_required_c_string(model_directory, "Model directory"), + read_required_c_string(primary_file, "Primary model file"), + ) { + (Ok(backend), Ok(identity), Ok(directory), Ok(primary_file)) => engine::ModelConfiguration { + backend, + identity, + directory: directory.into(), + primary_file, + }, + values => { + let error = [values.0.err(), values.1.err(), values.2.err(), values.3.err()] + .into_iter().flatten().next() + .unwrap_or_else(|| "Invalid speech model configuration".to_string()); + set_ios_error(error); + return false; + } + }; + + let validation = match configuration.backend.as_str() { + "PARAKEET_TDT_ONNX" => validate_ios_model_directory( + configuration.directory.to_string_lossy().as_ref(), + ), + "WHISPER_CPP" => { + let model = configuration.directory.join(&configuration.primary_file); + if model.is_file() { Ok(()) } else { Err(format!("Whisper model file does not exist: {}", model.display())) } + } + backend => Err(format!("Unsupported speech backend: {backend}")), + }; + let result = validation.and_then(|_| { + if configuration.backend == "PARAKEET_TDT_ONNX" { + initialize_ios_onnx_runtime()?; + } + engine::configure_model(configuration); + engine::ensure_loaded_without_callback() + }); + if let Err(error) = result { + set_ios_error(error); + return false; + } + true +} + #[cfg(all(target_os = "ios", feature = "ios-onnx"))] #[no_mangle] pub extern "C" fn voiceinbox_transcription_reset() { @@ -267,6 +351,18 @@ pub unsafe extern "C" fn voiceinbox_transcription_initialize( false } +#[cfg(all(target_os = "ios", not(feature = "ios-onnx")))] +#[no_mangle] +pub unsafe extern "C" fn voiceinbox_transcription_initialize_configured( + _backend: *const c_char, + _installation_identity: *const c_char, + _model_directory: *const c_char, + _primary_file: *const c_char, +) -> bool { + set_ios_error("iOS speech runtimes are not linked in this build"); + false +} + #[cfg(all(target_os = "ios", feature = "ios-onnx"))] #[no_mangle] pub unsafe extern "C" fn voiceinbox_transcription_transcribe_chunk_json( @@ -282,18 +378,8 @@ pub unsafe extern "C" fn voiceinbox_transcription_transcribe_chunk_json( let result = engine::get_engine() .ok_or_else(|| "Model is not loaded".to_string()) .and_then(|engine| { - engine - .lock() - .unwrap() - .transcribe_samples( - buffer, - Some(transcribe_rs::engines::parakeet::ParakeetInferenceParams { - timestamp_granularity: - transcribe_rs::engines::parakeet::TimestampGranularity::Word, - }), - ) + engine.lock().unwrap().transcribe_samples(buffer) .map(serialize_chunk_result) - .map_err(|error| error.to_string()) }); match result { @@ -334,6 +420,80 @@ pub unsafe extern "C" fn voiceinbox_transcription_string_free(value: *mut c_char } } +#[cfg(all(target_os = "ios", feature = "whisper-mobile-spike"))] +#[no_mangle] +pub unsafe extern "C" fn voiceinbox_whisper_spike_initialize_json( + model_path: *const c_char, +) -> *mut c_char { + let path = match read_model_directory(model_path) { + Ok(path) => PathBuf::from(path), + Err(error) => { + return into_c_string( + serde_json::json!({ + "schema_version": 1, + "backend": "whisper.cpp", + "operation": "initialize", + "status": "error", + "error": error, + }) + .to_string(), + ) + } + }; + into_c_string(run_whisper_spike_operation("initialize", || { + whisper_mobile_spike::initialize(&path) + })) +} + +#[cfg(all(target_os = "ios", feature = "whisper-mobile-spike"))] +#[no_mangle] +pub unsafe extern "C" fn voiceinbox_whisper_spike_transcribe_json( + samples: *const c_float, + sample_count: usize, + language: *const c_char, +) -> *mut c_char { + if samples.is_null() || sample_count == 0 { + return into_c_string( + serde_json::json!({ + "schema_version": 1, + "backend": "whisper.cpp", + "operation": "transcribe", + "status": "error", + "error": "No PCM samples were provided", + }) + .to_string(), + ); + } + let buffer = std::slice::from_raw_parts(samples, sample_count).to_vec(); + let language = if language.is_null() { + None + } else { + CStr::from_ptr(language) + .to_str() + .ok() + .map(str::to_owned) + .filter(|value| !value.is_empty()) + }; + into_c_string(run_whisper_spike_operation("transcribe", || { + whisper_mobile_spike::transcribe(buffer, language) + })) +} + +#[cfg(all(target_os = "ios", feature = "whisper-mobile-spike"))] +#[no_mangle] +pub extern "C" fn voiceinbox_whisper_spike_diagnostics_json() -> *mut c_char { + into_c_string(run_whisper_spike_operation( + "diagnostics", + whisper_mobile_spike::diagnostics, + )) +} + +#[cfg(all(target_os = "ios", feature = "whisper-mobile-spike"))] +#[no_mangle] +pub extern "C" fn voiceinbox_whisper_spike_reset() { + let _ = std::panic::catch_unwind(whisper_mobile_spike::reset); +} + #[cfg(test)] mod tests { use super::serialize_chunk_result; diff --git a/src/whisper_mobile_spike.rs b/src/whisper_mobile_spike.rs new file mode 100644 index 0000000..0e7113a --- /dev/null +++ b/src/whisper_mobile_spike.rs @@ -0,0 +1,249 @@ +//! Disabled-by-default Whisper Tiny physical-device evaluation backend. + +use once_cell::sync::Lazy; +use serde_json::{json, Value}; +use sha2::{Digest, Sha256}; +use std::fs::File; +use std::io::Read; +use std::path::{Path, PathBuf}; +use std::sync::Mutex; +use std::time::Instant; +use transcribe_rs::engines::whisper::{WhisperEngine, WhisperInferenceParams}; +use transcribe_rs::TranscriptionEngine; + +pub const MODEL_ID: &str = "openai/whisper-tiny"; +pub const MODEL_REPOSITORY: &str = "ggerganov/whisper.cpp"; +pub const MODEL_REVISION: &str = "5359861c739e955e79d9a303bcbc70fb988958b1"; +pub const MODEL_FILENAME: &str = "ggml-tiny.bin"; +pub const MODEL_SIZE_BYTES: u64 = 77_691_713; +pub const MODEL_SHA256: &str = "be07e048e1e599ad46341c8d2a135645097a538221678b7acdd1b1919c6e1b21"; +pub const SAMPLE_RATE_HZ: usize = 16_000; +pub const MAX_SAMPLE_COUNT: usize = SAMPLE_RATE_HZ * 30; + +struct SpikeEngine { + engine: WhisperEngine, + model_path: PathBuf, + load_ms: f64, +} + +static SPIKE_ENGINE: Lazy>> = Lazy::new(|| Mutex::new(None)); + +fn error(operation: &str, message: impl Into) -> String { + json!({ + "schema_version": 1, + "backend": "whisper.cpp", + "model_id": MODEL_ID, + "model_revision": MODEL_REVISION, + "operation": operation, + "status": "error", + "error": message.into(), + }) + .to_string() +} + +fn sha256(path: &Path) -> Result { + let mut file = File::open(path).map_err(|cause| format!("cannot open model: {cause}"))?; + let mut digest = Sha256::new(); + let mut buffer = [0_u8; 64 * 1024]; + loop { + let read = file + .read(&mut buffer) + .map_err(|cause| format!("cannot read model: {cause}"))?; + if read == 0 { + break; + } + digest.update(&buffer[..read]); + } + Ok(format!("{:x}", digest.finalize())) +} + +pub fn validate_model(path: &Path) -> Result<(), String> { + if !path.is_file() { + return Err(format!("model file does not exist: {}", path.display())); + } + if path.file_name().and_then(|name| name.to_str()) != Some(MODEL_FILENAME) { + return Err(format!("expected model filename {MODEL_FILENAME}")); + } + let size = path + .metadata() + .map_err(|cause| format!("cannot inspect model: {cause}"))? + .len(); + if size != MODEL_SIZE_BYTES { + return Err(format!( + "model size mismatch: expected {MODEL_SIZE_BYTES}, got {size}" + )); + } + let actual_sha256 = sha256(path)?; + if actual_sha256 != MODEL_SHA256 { + return Err(format!( + "model SHA-256 mismatch: expected {MODEL_SHA256}, got {actual_sha256}" + )); + } + Ok(()) +} + +pub fn initialize(model_path: &Path) -> String { + if let Err(message) = validate_model(model_path) { + return error("initialize", message); + } + + let started = Instant::now(); + let mut engine = WhisperEngine::new(); + if let Err(cause) = engine.load_model(model_path) { + return error("initialize", format!("Whisper model load failed: {cause}")); + } + let load_ms = started.elapsed().as_secs_f64() * 1_000.0; + let canonical_path = model_path + .canonicalize() + .unwrap_or_else(|_| model_path.to_path_buf()); + *SPIKE_ENGINE.lock().unwrap() = Some(SpikeEngine { + engine, + model_path: canonical_path.clone(), + load_ms, + }); + + json!({ + "schema_version": 1, + "backend": "whisper.cpp", + "model_id": MODEL_ID, + "model_repository": MODEL_REPOSITORY, + "model_revision": MODEL_REVISION, + "model_filename": MODEL_FILENAME, + "model_sha256": MODEL_SHA256, + "model_size_bytes": MODEL_SIZE_BYTES, + "model_path": canonical_path, + "operation": "initialize", + "status": "ok", + "load_ms": load_ms, + "peak_memory_bytes": Value::Null, + "error": Value::Null, + }) + .to_string() +} + +pub fn reset() { + *SPIKE_ENGINE.lock().unwrap() = None; +} + +pub fn diagnostics() -> String { + let state = SPIKE_ENGINE.lock().unwrap(); + let (loaded, path, load_ms) = state + .as_ref() + .map(|state| (true, json!(state.model_path), json!(state.load_ms))) + .unwrap_or((false, Value::Null, Value::Null)); + json!({ + "schema_version": 1, + "backend": "whisper.cpp", + "model_id": MODEL_ID, + "model_revision": MODEL_REVISION, + "operation": "diagnostics", + "status": "ok", + "loaded": loaded, + "model_path": path, + "load_ms": load_ms, + "sample_rate_hz": SAMPLE_RATE_HZ, + "maximum_chunk_seconds": 30, + "error": Value::Null, + }) + .to_string() +} + +pub fn transcribe(samples: Vec, language: Option) -> String { + if samples.is_empty() { + return error("transcribe", "no PCM samples were provided"); + } + if samples.len() > MAX_SAMPLE_COUNT { + return error( + "transcribe", + format!( + "chunk exceeds 30 seconds: {} samples at {SAMPLE_RATE_HZ} Hz", + samples.len() + ), + ); + } + + let sample_count = samples.len(); + let audio_duration_seconds = sample_count as f64 / SAMPLE_RATE_HZ as f64; + let mut state = SPIKE_ENGINE.lock().unwrap(); + let Some(state) = state.as_mut() else { + return error("transcribe", "Whisper spike model is not loaded"); + }; + let started = Instant::now(); + let inference = state.engine.transcribe_samples( + samples, + Some(WhisperInferenceParams { + language: language.filter(|value| !value.is_empty()), + ..WhisperInferenceParams::default() + }), + ); + let inference_ms = started.elapsed().as_secs_f64() * 1_000.0; + + match inference { + Ok(result) => { + let segments = result + .segments + .unwrap_or_default() + .into_iter() + .map(|segment| { + json!({ + "text": segment.text, + "start_seconds": segment.start, + "end_seconds": segment.end, + }) + }) + .collect::>(); + json!({ + "schema_version": 1, + "backend": "whisper.cpp", + "model_id": MODEL_ID, + "model_revision": MODEL_REVISION, + "operation": "transcribe", + "status": "ok", + "sample_rate_hz": SAMPLE_RATE_HZ, + "sample_count": sample_count, + "audio_duration_seconds": audio_duration_seconds, + "inference_ms": inference_ms, + "real_time_factor": inference_ms / 1_000.0 / audio_duration_seconds, + "detected_language": Value::Null, + "transcript": result.text, + "segments": segments, + "peak_memory_bytes": Value::Null, + "error": Value::Null, + }) + .to_string() + } + Err(cause) => error("transcribe", format!("Whisper inference failed: {cause}")), + } +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn rejects_missing_model_without_downloading() { + let result = initialize(Path::new("missing/ggml-tiny.bin")); + let json: Value = serde_json::from_str(&result).unwrap(); + assert_eq!(json["status"], "error"); + assert!(json["error"].as_str().unwrap().contains("does not exist")); + } + + #[test] + fn rejects_empty_pcm_with_structured_diagnostic() { + let result = transcribe(Vec::new(), None); + let json: Value = serde_json::from_str(&result).unwrap(); + assert_eq!(json["status"], "error"); + assert_eq!(json["operation"], "transcribe"); + } + + #[test] + fn rejects_chunks_longer_than_thirty_seconds() { + let result = transcribe(vec![0.0; MAX_SAMPLE_COUNT + 1], None); + let json: Value = serde_json::from_str(&result).unwrap(); + assert_eq!(json["status"], "error"); + assert!(json["error"] + .as_str() + .unwrap() + .contains("exceeds 30 seconds")); + } +} diff --git a/transcribe-rs/Cargo.toml b/transcribe-rs/Cargo.toml index 69ec3f2..b70b430 100755 --- a/transcribe-rs/Cargo.toml +++ b/transcribe-rs/Cargo.toml @@ -6,11 +6,16 @@ description = "A simple library to help you transcribe audio" license = "MIT" repository = "https://github.com/cjpais/transcribe-rs" +[features] +default = ["parakeet", "whisper"] +parakeet = ["dep:ort"] +whisper = ["dep:whisper-rs"] + [dependencies] hound = "=3.5.1" log = "=0.4.28" ndarray = "=0.16.1" -ort = { version = "=2.0.0-rc.10" } +ort = { version = "=2.0.0-rc.10", optional = true } env_logger = "=0.10.0" regex = "=1.11.2" thiserror = "=2.0.16" @@ -21,11 +26,13 @@ async-trait = { version = "=0.1.89" } derive_builder = { version = "=0.20.2" } [target.'cfg(target_os = "macos")'.dependencies] -whisper-rs = { version = "=0.13.2", features = ["metal"] } +whisper-rs = { version = "=0.16.0", features = ["metal"], optional = true } [target.'cfg(target_os = "windows")'.dependencies] -whisper-rs = { version = "=0.13.2", features = ["vulkan"] } +whisper-rs = { version = "=0.16.0", features = ["vulkan"], optional = true } [target.'cfg(target_os = "linux")'.dependencies] -whisper-rs = { version = "=0.13.2", features = ["vulkan"] } +whisper-rs = { version = "=0.16.0", features = ["vulkan"], optional = true } +[target.'cfg(any(target_os = "android", target_os = "ios"))'.dependencies] +whisper-rs = { version = "=0.16.0", optional = true } diff --git a/transcribe-rs/src/engines/mod.rs b/transcribe-rs/src/engines/mod.rs index bc1d94a..4751722 100755 --- a/transcribe-rs/src/engines/mod.rs +++ b/transcribe-rs/src/engines/mod.rs @@ -42,6 +42,7 @@ //! # Ok::<(), Box>(()) //! ``` +#[cfg(feature = "parakeet")] pub mod parakeet; -#[cfg(not(any(target_os = "android", target_os = "ios")))] +#[cfg(feature = "whisper")] pub mod whisper; diff --git a/transcribe-rs/src/engines/whisper.rs b/transcribe-rs/src/engines/whisper.rs index c5be777..df49340 100755 --- a/transcribe-rs/src/engines/whisper.rs +++ b/transcribe-rs/src/engines/whisper.rs @@ -244,7 +244,7 @@ impl TranscriptionEngine for WhisperEngine { full_params.set_print_realtime(whisper_params.print_realtime); full_params.set_print_timestamps(whisper_params.print_timestamps); full_params.set_suppress_blank(whisper_params.suppress_blank); - full_params.set_suppress_non_speech_tokens(whisper_params.suppress_non_speech_tokens); + full_params.set_suppress_nst(whisper_params.suppress_non_speech_tokens); full_params.set_no_speech_thold(whisper_params.no_speech_thold); if let Some(ref prompt) = whisper_params.initial_prompt { @@ -253,17 +253,13 @@ impl TranscriptionEngine for WhisperEngine { state.full(full_params, &samples)?; - let num_segments = state - .full_n_segments() - .expect("failed to get number of segments"); - - let mut segments = Vec::new(); - let mut full_text = String::new(); - - for i in 0..num_segments { - let text = state.full_get_segment_text(i)?; - let start = state.full_get_segment_t0(i)? as f32 / 100.0; - let end = state.full_get_segment_t1(i)? as f32 / 100.0; + let mut segments = Vec::new(); + let mut full_text = String::new(); + + for segment in state.as_iter() { + let text = segment.to_str_lossy()?.into_owned(); + let start = segment.start_timestamp() as f32 / 100.0; + let end = segment.end_timestamp() as f32 / 100.0; segments.push(TranscriptionSegment { start, diff --git a/website/src/pages/docs.astro b/website/src/pages/docs.astro new file mode 100644 index 0000000..b711056 --- /dev/null +++ b/website/src/pages/docs.astro @@ -0,0 +1,68 @@ +--- +import Layout from '../layouts/Layout.astro'; + +const base = import.meta.env.BASE_URL.replace(/\/$/, ''); +--- + +
+

Voice Inbox documentation

+

From audio files to a local text inbox.

+

+ Voice Inbox is designed for accumulated recordings, not live dictation. Install a speech model, + choose where transcripts are written, add audio files, then transcribe when you are ready. +

+ +
+ +
+
+
+

1. Set up Voice Inbox

+

Install or download a compatible model, create or select an output document, and import audio.

+ Read the setup guide → +
+
+

2. Get answers

+

Find help with device requirements, audio files, model installation, transcripts, and failures.

+ Read the FAQ → +
+ + +
+

Legal and acknowledgements

+

Review app information, source licensing, third-party notices, and model attribution.

+ Read legal information → +
+
+
+ +
+

Platform notes

+

+ Android has the broader workflow today, including individual imports and sharing, optional folder + processing, and scheduled transcription. iOS is an active MVP with local import, model installation, + preview, transcription, and text output; its automation features are still evolving. +

+
+ + +