Standalone Shonar Desktop: vendor portable sources + local engine; decouple from ~/Projects/Shonar
- shared/ = portable Android-origin sources vendored from deferred/desktop-server (app/build.gradle.kts srcDir repointed; PlaybackController.kt excluded as Android-only) - backend/ = bundled-lite engine (SQLite + inline queue); .venv symlinked from the old checkout, PYTHONPATH pins THIS backend's code over any editable install - repoRoot() resolves this project dir (env SHONAR_REPO still wins); desktop-dev.sh watches shared/ + backend/ - Verified: :app:compileKotlin + :app:test green (23 tests); engine boots on :8010, self-migrates, /healthz ok
This commit is contained in:
commit
76c867fca4
136 changed files with 21099 additions and 0 deletions
9
.gitignore
vendored
Normal file
9
.gitignore
vendored
Normal file
|
|
@ -0,0 +1,9 @@
|
||||||
|
.gradle/
|
||||||
|
build/
|
||||||
|
app/build/
|
||||||
|
.kotlin/
|
||||||
|
backend/.venv
|
||||||
|
backend/__pycache__/
|
||||||
|
**/__pycache__/
|
||||||
|
*.pyc
|
||||||
|
.dev-restart-flag
|
||||||
99
app/build.gradle.kts
Normal file
99
app/build.gradle.kts
Normal file
|
|
@ -0,0 +1,99 @@
|
||||||
|
plugins {
|
||||||
|
kotlin("jvm") version "2.0.21"
|
||||||
|
id("org.jetbrains.kotlin.plugin.compose") version "2.0.21"
|
||||||
|
id("org.jetbrains.compose") version "1.7.3"
|
||||||
|
}
|
||||||
|
|
||||||
|
group = "com.shonar"
|
||||||
|
version = "0.1.0"
|
||||||
|
|
||||||
|
repositories {
|
||||||
|
mavenCentral()
|
||||||
|
google()
|
||||||
|
}
|
||||||
|
|
||||||
|
kotlin {
|
||||||
|
jvmToolchain(17)
|
||||||
|
}
|
||||||
|
|
||||||
|
sourceSets {
|
||||||
|
main {
|
||||||
|
kotlin {
|
||||||
|
// Portable shared sources vendored into this project (formerly
|
||||||
|
// compiled in place from the sibling Android repo). Android-only
|
||||||
|
// files excluded below.
|
||||||
|
srcDir("../shared")
|
||||||
|
exclude(
|
||||||
|
"com/shonar/ShonarApplication.kt",
|
||||||
|
"com/shonar/MainActivity.kt",
|
||||||
|
// settings store impls are Android (DataStore/Keystore);
|
||||||
|
// desktop provides its own (desktopSettingsStore.kt).
|
||||||
|
"com/shonar/settings/SettingsStore.kt",
|
||||||
|
// recording: only AiContent.kt is portable.
|
||||||
|
"com/shonar/recording/MigrationRunner.kt",
|
||||||
|
"com/shonar/recording/RecordingDao.kt",
|
||||||
|
"com/shonar/recording/RecordingEntity.kt",
|
||||||
|
"com/shonar/recording/RecordingNotificationHelper.kt",
|
||||||
|
"com/shonar/recording/PlaybackController.kt",
|
||||||
|
"com/shonar/recording/RecordingPermissionHelper.kt",
|
||||||
|
"com/shonar/recording/RecordingRepository.kt",
|
||||||
|
"com/shonar/recording/RecordingService.kt",
|
||||||
|
"com/shonar/recording/RecordingStateStore.kt",
|
||||||
|
"com/shonar/recording/ShonarDatabase.kt",
|
||||||
|
"com/shonar/recording/SyncDrain.kt",
|
||||||
|
"com/shonar/recording/SyncSlots.kt",
|
||||||
|
"com/shonar/recording/SyncStateConverter.kt",
|
||||||
|
"com/shonar/recording/SyncWorker.kt",
|
||||||
|
// widget + rename-after-save UI: Android-only (AppWidget,
|
||||||
|
// ShonarApplication references).
|
||||||
|
"com/shonar/widget/**",
|
||||||
|
"com/shonar/ui/rename/**",
|
||||||
|
// ui: only ui/folder/FolderList.kt is portable.
|
||||||
|
"com/shonar/ui/detail/**",
|
||||||
|
"com/shonar/ui/home/**",
|
||||||
|
"com/shonar/ui/provider/**",
|
||||||
|
"com/shonar/ui/settings/**",
|
||||||
|
"com/shonar/ui/theme/**",
|
||||||
|
"com/shonar/ui/folder/FolderBrowserScreen.kt",
|
||||||
|
)
|
||||||
|
srcDir("src/main/kotlin")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
test {
|
||||||
|
kotlin {
|
||||||
|
srcDir("src/test/kotlin")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
dependencies {
|
||||||
|
implementation(compose.desktop.currentOs)
|
||||||
|
implementation(compose.material3)
|
||||||
|
implementation(compose.materialIconsExtended)
|
||||||
|
implementation("org.jetbrains.kotlinx:kotlinx-coroutines-swing:1.9.0")
|
||||||
|
implementation("org.jetbrains.kotlinx:kotlinx-serialization-json:1.7.3")
|
||||||
|
implementation("com.squareup.okhttp3:okhttp:4.12.0")
|
||||||
|
// Real org.json for JVM (Android stubs throw "not mocked").
|
||||||
|
implementation("org.json:json:20240303")
|
||||||
|
|
||||||
|
testImplementation("junit:junit:4.13.2")
|
||||||
|
testImplementation("org.jetbrains.kotlinx:kotlinx-coroutines-test:1.9.0")
|
||||||
|
testImplementation("com.squareup.okhttp3:mockwebserver:4.12.0")
|
||||||
|
}
|
||||||
|
|
||||||
|
compose.desktop {
|
||||||
|
application {
|
||||||
|
mainClass = "com.shonar.desktop.MainKt"
|
||||||
|
nativeDistributions {
|
||||||
|
targetFormats(org.jetbrains.compose.desktop.application.dsl.TargetFormat.Dmg)
|
||||||
|
targetFormats(org.jetbrains.compose.desktop.application.dsl.TargetFormat.Msi)
|
||||||
|
targetFormats(org.jetbrains.compose.desktop.application.dsl.TargetFormat.Deb)
|
||||||
|
packageName = "shonar-desktop"
|
||||||
|
packageVersion = "0.1.0"
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
tasks.withType<Test> {
|
||||||
|
useJUnit()
|
||||||
|
}
|
||||||
1057
app/src/main/kotlin/com/shonar/desktop/DesktopState.kt
Normal file
1057
app/src/main/kotlin/com/shonar/desktop/DesktopState.kt
Normal file
File diff suppressed because it is too large
Load diff
96
app/src/main/kotlin/com/shonar/desktop/LibraryQueue.kt
Normal file
96
app/src/main/kotlin/com/shonar/desktop/LibraryQueue.kt
Normal file
|
|
@ -0,0 +1,96 @@
|
||||||
|
package com.shonar.desktop
|
||||||
|
|
||||||
|
import java.io.File
|
||||||
|
|
||||||
|
/**
|
||||||
|
* Pure library-queue logic (unit-tested; no coroutines, no Android).
|
||||||
|
*
|
||||||
|
* [plan] decides which files on disk are transcription candidates, and
|
||||||
|
* the queue pump consumes them in a stable order. All policy lives here
|
||||||
|
* so the state class stays plumbing.
|
||||||
|
*/
|
||||||
|
|
||||||
|
enum class FileStatus {
|
||||||
|
/** No report yet — a candidate. */
|
||||||
|
NEW,
|
||||||
|
|
||||||
|
/** Queued for auto-transcription. */
|
||||||
|
QUEUED,
|
||||||
|
|
||||||
|
/** Upload/transcribe/summarize running now. */
|
||||||
|
RUNNING,
|
||||||
|
|
||||||
|
/** Report written next to the audio. */
|
||||||
|
DONE,
|
||||||
|
|
||||||
|
/** Last attempt failed (error message kept beside it). */
|
||||||
|
FAILED,
|
||||||
|
;
|
||||||
|
}
|
||||||
|
|
||||||
|
/**
|
||||||
|
* Diff between two scans of the library folder. [added] preserves sorted
|
||||||
|
* file-name order so the queue is deterministic across rescans.
|
||||||
|
*/
|
||||||
|
class LibraryDiff(
|
||||||
|
val added: List<File>,
|
||||||
|
/** Names that disappeared since the previous scan. */
|
||||||
|
val removed: Set<String>,
|
||||||
|
)
|
||||||
|
|
||||||
|
object LibraryQueue {
|
||||||
|
|
||||||
|
/**
|
||||||
|
* Validate renaming [oldFile] to base name [newBase] (no extension;
|
||||||
|
* a typed extension is stripped). Returns the target File or a
|
||||||
|
* human-readable failure. The extension always stays the original's.
|
||||||
|
*/
|
||||||
|
fun renameTarget(oldFile: File, newBase: String): Result<File> {
|
||||||
|
val base = newBase.trim()
|
||||||
|
.let { s ->
|
||||||
|
val ext = oldFile.extension
|
||||||
|
if (ext.isNotEmpty() && s.length > ext.length + 1 &&
|
||||||
|
s.endsWith(".$ext", ignoreCase = true)
|
||||||
|
) s.dropLast(ext.length + 1) else s
|
||||||
|
}
|
||||||
|
.trim()
|
||||||
|
fun fail(msg: String): Result<File> =
|
||||||
|
Result.failure(IllegalArgumentException(msg))
|
||||||
|
if (base.isEmpty()) return fail("Name is empty.")
|
||||||
|
if (base.any { it == '/' || it == '\\' }) return fail("Name can't contain a slash.")
|
||||||
|
if (base == "." || base == "..") return fail("Not a valid name.")
|
||||||
|
if (base.startsWith(".")) return fail("Name can't start with a dot.")
|
||||||
|
if (base.length > 120) return fail("Name is too long.")
|
||||||
|
val target = File(oldFile.parentFile, "$base.${oldFile.extension}")
|
||||||
|
if (target == oldFile) return fail("That's already the name.")
|
||||||
|
if (target.exists()) return fail("A file with that name already exists.")
|
||||||
|
return Result.success(target)
|
||||||
|
}
|
||||||
|
|
||||||
|
/** Audio candidates: regular files, known extension, not hidden. */
|
||||||
|
fun isAudioCandidate(f: File, audioExts: Set<String>): Boolean =
|
||||||
|
f.isFile && !f.name.startsWith(".") &&
|
||||||
|
f.extension.lowercase() in audioExts
|
||||||
|
|
||||||
|
fun diff(before: Set<String>, after: List<File>): LibraryDiff {
|
||||||
|
val afterNames = after.map { it.name }.toSet()
|
||||||
|
val added = after.filter { it.name !in before }
|
||||||
|
val removed = before - afterNames
|
||||||
|
return LibraryDiff(added, removed)
|
||||||
|
}
|
||||||
|
|
||||||
|
/**
|
||||||
|
* Files to (re)queue for auto-transcription: candidates without a
|
||||||
|
* report, minus files already in-flight or previously failed since
|
||||||
|
* the last scan (a failure must not hot-loop; a rescan after the
|
||||||
|
* user edits/renames clears it naturally when the report appears).
|
||||||
|
*/
|
||||||
|
fun autoQueueCandidates(
|
||||||
|
entries: List<File>,
|
||||||
|
reportExists: (File) -> Boolean,
|
||||||
|
inFlightOrDone: Set<String>,
|
||||||
|
failed: Set<String>,
|
||||||
|
): List<File> = entries.filter {
|
||||||
|
!reportExists(it) && it.name !in inFlightOrDone && it.name !in failed
|
||||||
|
}
|
||||||
|
}
|
||||||
87
app/src/main/kotlin/com/shonar/desktop/Main.kt
Normal file
87
app/src/main/kotlin/com/shonar/desktop/Main.kt
Normal file
|
|
@ -0,0 +1,87 @@
|
||||||
|
package com.shonar.desktop
|
||||||
|
|
||||||
|
import androidx.compose.foundation.layout.Box
|
||||||
|
import androidx.compose.foundation.layout.fillMaxSize
|
||||||
|
import androidx.compose.foundation.layout.padding
|
||||||
|
import androidx.compose.material3.Surface
|
||||||
|
import androidx.compose.runtime.CompositionLocalProvider
|
||||||
|
import androidx.compose.runtime.collectAsState
|
||||||
|
import androidx.compose.runtime.getValue
|
||||||
|
import androidx.compose.runtime.mutableStateOf
|
||||||
|
import androidx.compose.runtime.remember
|
||||||
|
import androidx.compose.runtime.setValue
|
||||||
|
import androidx.compose.ui.Modifier
|
||||||
|
import androidx.compose.ui.input.key.Key
|
||||||
|
import androidx.compose.ui.input.key.KeyEventType
|
||||||
|
import androidx.compose.ui.input.key.isCtrlPressed
|
||||||
|
import androidx.compose.ui.input.key.key
|
||||||
|
import androidx.compose.ui.input.key.onPreviewKeyEvent
|
||||||
|
import androidx.compose.ui.input.key.type
|
||||||
|
import androidx.compose.ui.platform.LocalDensity
|
||||||
|
import androidx.compose.ui.unit.Density
|
||||||
|
import androidx.compose.ui.unit.dp
|
||||||
|
import androidx.compose.ui.window.Window
|
||||||
|
import androidx.compose.ui.window.application
|
||||||
|
import androidx.compose.ui.window.rememberWindowState
|
||||||
|
|
||||||
|
/** Scales every dp/sp in the app uniformly (fonts, paddings, controls). */
|
||||||
|
private data class ScaledDensity(val base: Density, val zoom: Float) : Density {
|
||||||
|
override val density: Float get() = base.density * zoom
|
||||||
|
override val fontScale: Float get() = base.fontScale * zoom
|
||||||
|
}
|
||||||
|
|
||||||
|
fun main() = application {
|
||||||
|
val state = remember { DesktopState() }
|
||||||
|
val screen by state.screen.collectAsState()
|
||||||
|
// Ctrl+= / Ctrl+- zoom the whole app; Ctrl+0 resets (after a zoom).
|
||||||
|
var zoom by remember { mutableStateOf(1f) }
|
||||||
|
Window(
|
||||||
|
onCloseRequest = ::exitApplication,
|
||||||
|
title = "SHONAR Desktop" + run {
|
||||||
|
// Build stamp so "which window is this?" is answerable at a glance.
|
||||||
|
val jar = java.io.File(
|
||||||
|
System.getProperty("java.class.path", "").split(":")
|
||||||
|
.firstOrNull { it.endsWith(".jar") } ?: "")
|
||||||
|
if (jar.exists()) {
|
||||||
|
val fmt = java.time.format.DateTimeFormatter.ofPattern("HH:mm")
|
||||||
|
.withZone(java.time.ZoneId.systemDefault())
|
||||||
|
" · build " + fmt.format(java.time.Instant.ofEpochMilli(jar.lastModified()))
|
||||||
|
} else ""
|
||||||
|
},
|
||||||
|
state = rememberWindowState(width = 1100.dp, height = 800.dp),
|
||||||
|
) {
|
||||||
|
val baseDensity = LocalDensity.current
|
||||||
|
ShonarTheme {
|
||||||
|
Surface(
|
||||||
|
Modifier.fillMaxSize().onPreviewKeyEvent { e ->
|
||||||
|
if (e.type != KeyEventType.KeyDown || !e.isCtrlPressed) return@onPreviewKeyEvent false
|
||||||
|
when (e.key) {
|
||||||
|
Key.Equals, Key.Plus, Key.NumPadAdd -> {
|
||||||
|
zoom = (zoom * 1.1f).coerceAtMost(3f); true
|
||||||
|
}
|
||||||
|
Key.Minus, Key.NumPadSubtract -> {
|
||||||
|
zoom = (zoom / 1.1f).coerceAtLeast(0.4f); true
|
||||||
|
}
|
||||||
|
Key.Zero, Key.NumPad0 -> {
|
||||||
|
if (zoom != 1f) { zoom = 1f; true } else false
|
||||||
|
}
|
||||||
|
else -> false
|
||||||
|
}
|
||||||
|
},
|
||||||
|
) {
|
||||||
|
CompositionLocalProvider(
|
||||||
|
LocalDensity provides ScaledDensity(baseDensity, zoom),
|
||||||
|
) {
|
||||||
|
Box(Modifier.fillMaxSize().padding(20.dp)) {
|
||||||
|
when (screen) {
|
||||||
|
DesktopState.Screen.ENGINE -> EngineScreen(state)
|
||||||
|
DesktopState.Screen.LIBRARY -> LibraryScreen(state)
|
||||||
|
DesktopState.Screen.DETAIL -> DetailScreen(state)
|
||||||
|
DesktopState.Screen.SETTINGS -> SettingsScreen(state)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
51
app/src/main/kotlin/com/shonar/desktop/RemoteMapping.kt
Normal file
51
app/src/main/kotlin/com/shonar/desktop/RemoteMapping.kt
Normal file
|
|
@ -0,0 +1,51 @@
|
||||||
|
package com.shonar.desktop
|
||||||
|
|
||||||
|
import java.io.File
|
||||||
|
import org.json.JSONObject
|
||||||
|
|
||||||
|
/**
|
||||||
|
* Sidecar mapping a library file to its server recording id. Written next
|
||||||
|
* to the audio as `<name>.shonar.json` (excluded from the library listing
|
||||||
|
* by the audio-extension filter; travels with the file through renames
|
||||||
|
* and Syncthing).
|
||||||
|
*
|
||||||
|
* Why persisted: without it, "Re-transcribe" after a restart uploads a
|
||||||
|
* brand-new recording (duplicate on the server). With it, the desktop
|
||||||
|
* re-runs the pipeline in place via POST /reprocess.
|
||||||
|
*
|
||||||
|
* Pure logic — unit-tested.
|
||||||
|
*/
|
||||||
|
data class RemoteMapping(
|
||||||
|
val recordingId: String,
|
||||||
|
/** Model override the last transcribe used, if any (informational). */
|
||||||
|
val model: String? = null,
|
||||||
|
val title: String? = null,
|
||||||
|
)
|
||||||
|
|
||||||
|
fun mappingFile(audio: File): File =
|
||||||
|
File(audio.parentFile, "${audio.nameWithoutExtension}.shonar.json")
|
||||||
|
|
||||||
|
fun loadMapping(audio: File): RemoteMapping? = runCatching {
|
||||||
|
val o = JSONObject(mappingFile(audio).readText())
|
||||||
|
val id = o.optString("recording_id")
|
||||||
|
if (id.isBlank()) null
|
||||||
|
else RemoteMapping(
|
||||||
|
recordingId = id,
|
||||||
|
model = o.optString("model").takeIf { it.isNotBlank() },
|
||||||
|
title = o.optString("title").takeIf { it.isNotBlank() },
|
||||||
|
)
|
||||||
|
}.getOrNull()
|
||||||
|
|
||||||
|
fun saveMapping(audio: File, mapping: RemoteMapping) {
|
||||||
|
val o = JSONObject()
|
||||||
|
o.put("recording_id", mapping.recordingId)
|
||||||
|
mapping.model?.let { o.put("model", it) }
|
||||||
|
mapping.title?.let { o.put("title", it) }
|
||||||
|
mappingFile(audio).writeText(o.toString(2) + "\n")
|
||||||
|
}
|
||||||
|
|
||||||
|
/** Carry the mapping across a file rename (returns true if one moved). */
|
||||||
|
fun moveMapping(old: File, new: File): Boolean {
|
||||||
|
val src = mappingFile(old)
|
||||||
|
return src.exists() && src.renameTo(mappingFile(new))
|
||||||
|
}
|
||||||
48
app/src/main/kotlin/com/shonar/desktop/Reports.kt
Normal file
48
app/src/main/kotlin/com/shonar/desktop/Reports.kt
Normal file
|
|
@ -0,0 +1,48 @@
|
||||||
|
package com.shonar.desktop
|
||||||
|
|
||||||
|
import com.shonar.recording.SUMMARY_LIST_KEYS
|
||||||
|
import com.shonar.recording.SummaryData
|
||||||
|
import com.shonar.recording.TranscriptData
|
||||||
|
|
||||||
|
/**
|
||||||
|
* Plaud-style Markdown report written next to the audio file.
|
||||||
|
* Pure logic — unit-tested.
|
||||||
|
*/
|
||||||
|
fun renderReport(title: String, transcript: TranscriptData?, summary: SummaryData?): String {
|
||||||
|
val out = mutableListOf("# $title", "")
|
||||||
|
if (summary != null) {
|
||||||
|
if (summary.short.isNotBlank()) {
|
||||||
|
out += listOf("## Summary", "", summary.short, "")
|
||||||
|
}
|
||||||
|
for (key in SUMMARY_LIST_KEYS) {
|
||||||
|
val items = summary.list(key)
|
||||||
|
if (items.isNotEmpty()) {
|
||||||
|
out += listOf("## " + key.replace('_', ' ').replaceFirstChar { it.uppercase() }, "")
|
||||||
|
out += items.map { "- $it" } + listOf("")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if (summary.detailed.isNotBlank()) {
|
||||||
|
out += listOf("## Details", "", summary.detailed, "")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if (transcript != null) {
|
||||||
|
out += listOf("## Transcript", "")
|
||||||
|
val segs = transcript.segments
|
||||||
|
if (segs.isNotEmpty()) {
|
||||||
|
segs.forEach { s -> out += "[${fmtTs(s.startSec)}] ${s.text}" }
|
||||||
|
out += ""
|
||||||
|
} else {
|
||||||
|
out += listOf(transcript.text, "")
|
||||||
|
}
|
||||||
|
val prov = transcript.provider.ifBlank { "?" }
|
||||||
|
out += "*Transcribed with $prov, v${transcript.version}.*"
|
||||||
|
} else {
|
||||||
|
out += listOf("## Transcript", "", "_No transcript available._", "")
|
||||||
|
}
|
||||||
|
return out.joinToString("\n")
|
||||||
|
}
|
||||||
|
|
||||||
|
fun fmtTs(sec: Double): String {
|
||||||
|
val total = sec.toLong().coerceAtLeast(0)
|
||||||
|
return "%02d:%02d:%02d".format(total / 3600, (total % 3600) / 60, total % 60)
|
||||||
|
}
|
||||||
525
app/src/main/kotlin/com/shonar/desktop/Screens.kt
Normal file
525
app/src/main/kotlin/com/shonar/desktop/Screens.kt
Normal file
|
|
@ -0,0 +1,525 @@
|
||||||
|
package com.shonar.desktop
|
||||||
|
|
||||||
|
import androidx.compose.foundation.clickable
|
||||||
|
import androidx.compose.foundation.layout.Arrangement
|
||||||
|
import androidx.compose.foundation.layout.Column
|
||||||
|
import androidx.compose.foundation.layout.Row
|
||||||
|
import androidx.compose.foundation.layout.Spacer
|
||||||
|
import androidx.compose.foundation.layout.fillMaxSize
|
||||||
|
import androidx.compose.foundation.layout.fillMaxWidth
|
||||||
|
import androidx.compose.foundation.layout.height
|
||||||
|
import androidx.compose.foundation.layout.padding
|
||||||
|
import androidx.compose.foundation.layout.width
|
||||||
|
import androidx.compose.foundation.lazy.LazyColumn
|
||||||
|
import androidx.compose.foundation.lazy.items
|
||||||
|
import androidx.compose.material3.AlertDialog
|
||||||
|
import androidx.compose.material3.Button
|
||||||
|
import androidx.compose.material3.Card
|
||||||
|
import androidx.compose.material3.Checkbox
|
||||||
|
import androidx.compose.material3.CircularProgressIndicator
|
||||||
|
import androidx.compose.material3.LinearProgressIndicator
|
||||||
|
import androidx.compose.material3.MaterialTheme
|
||||||
|
import androidx.compose.material3.OutlinedButton
|
||||||
|
import androidx.compose.material3.OutlinedTextField
|
||||||
|
import androidx.compose.material3.Text
|
||||||
|
import androidx.compose.material3.TextButton
|
||||||
|
import androidx.compose.runtime.collectAsState
|
||||||
|
import androidx.compose.runtime.getValue
|
||||||
|
import androidx.compose.runtime.mutableStateOf
|
||||||
|
import androidx.compose.runtime.remember
|
||||||
|
import androidx.compose.runtime.setValue
|
||||||
|
import androidx.compose.ui.Alignment
|
||||||
|
import androidx.compose.ui.Modifier
|
||||||
|
import androidx.compose.ui.unit.dp
|
||||||
|
import androidx.compose.ui.text.style.TextOverflow
|
||||||
|
import java.io.File
|
||||||
|
import javax.swing.JFileChooser
|
||||||
|
|
||||||
|
// ---- engine -----------------------------------------------------------------
|
||||||
|
|
||||||
|
@androidx.compose.runtime.Composable
|
||||||
|
fun EngineScreen(state: DesktopState) {
|
||||||
|
val engine by state.engine.collectAsState()
|
||||||
|
val error by state.engineError.collectAsState()
|
||||||
|
|
||||||
|
Column(Modifier.fillMaxSize(), verticalArrangement = Arrangement.spacedBy(12.dp)) {
|
||||||
|
Text("Starting the transcription engine…",
|
||||||
|
style = MaterialTheme.typography.headlineSmall)
|
||||||
|
Text(
|
||||||
|
"SHONAR transcribes on this computer with local models — nothing " +
|
||||||
|
"is uploaded anywhere. The engine runs in the background while " +
|
||||||
|
"this app is open.",
|
||||||
|
style = MaterialTheme.typography.bodyMedium,
|
||||||
|
)
|
||||||
|
when (engine) {
|
||||||
|
DesktopState.EngineState.STARTING, DesktopState.EngineState.UNKNOWN -> {
|
||||||
|
Row(verticalAlignment = Alignment.CenterVertically) {
|
||||||
|
CircularProgressIndicator(Modifier.width(20.dp).height(20.dp))
|
||||||
|
Spacer(Modifier.width(10.dp))
|
||||||
|
Text("Starting…")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
else -> {
|
||||||
|
error?.let { Text(it, color = MaterialTheme.colorScheme.error) }
|
||||||
|
Row(horizontalArrangement = Arrangement.spacedBy(8.dp)) {
|
||||||
|
Button({ state.startEngine() }) { Text("Start engine") }
|
||||||
|
OutlinedButton({ state.ensureReady() }) { Text("Retry") }
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// ---- library ---------------------------------------------------------------
|
||||||
|
|
||||||
|
private fun chooseFolder(): File? {
|
||||||
|
val chooser = JFileChooser().apply {
|
||||||
|
dialogTitle = "Choose your recordings folder"
|
||||||
|
fileSelectionMode = JFileChooser.DIRECTORIES_ONLY
|
||||||
|
isAcceptAllFileFilterUsed = false
|
||||||
|
// Without an explicit size the dialog comes up tiny on HiDPI
|
||||||
|
// compositors (and there is no resize handle), so pin a usable one.
|
||||||
|
preferredSize = java.awt.Dimension(1040, 700)
|
||||||
|
}
|
||||||
|
return if (chooser.showOpenDialog(null) == JFileChooser.APPROVE_OPTION) chooser.selectedFile else null
|
||||||
|
}
|
||||||
|
|
||||||
|
@androidx.compose.runtime.Composable
|
||||||
|
fun LibraryScreen(state: DesktopState) {
|
||||||
|
val folder by state.folder.collectAsState()
|
||||||
|
val entries by state.entries.collectAsState()
|
||||||
|
val query by state.query.collectAsState()
|
||||||
|
val hits by state.searchResults.collectAsState()
|
||||||
|
val connected by state.connected.collectAsState()
|
||||||
|
var renaming by remember { mutableStateOf<File?>(null) }
|
||||||
|
|
||||||
|
Column(Modifier.fillMaxSize(), verticalArrangement = Arrangement.spacedBy(12.dp)) {
|
||||||
|
Row(verticalAlignment = Alignment.CenterVertically) {
|
||||||
|
Text("Recordings", style = MaterialTheme.typography.headlineSmall,
|
||||||
|
modifier = Modifier.weight(1f))
|
||||||
|
TextButton({ state.go(DesktopState.Screen.SETTINGS) }) { Text("Settings") }
|
||||||
|
}
|
||||||
|
if (!connected) {
|
||||||
|
Card(Modifier.fillMaxWidth()) {
|
||||||
|
Row(Modifier.fillMaxWidth().padding(12.dp),
|
||||||
|
verticalAlignment = Alignment.CenterVertically) {
|
||||||
|
Text("Transcription engine isn't running.",
|
||||||
|
modifier = Modifier.weight(1f))
|
||||||
|
TextButton({ state.go(DesktopState.Screen.ENGINE) }) {
|
||||||
|
Text("Start it")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
Row(verticalAlignment = Alignment.CenterVertically) {
|
||||||
|
Text(folder?.absolutePath ?: "No folder chosen",
|
||||||
|
style = MaterialTheme.typography.bodySmall,
|
||||||
|
modifier = Modifier.weight(1f))
|
||||||
|
Spacer(Modifier.width(8.dp))
|
||||||
|
OutlinedButton({ chooseFolder()?.let { state.pickFolder(it) } }) {
|
||||||
|
Text("Choose folder…")
|
||||||
|
}
|
||||||
|
Spacer(Modifier.width(8.dp))
|
||||||
|
Row(verticalAlignment = Alignment.CenterVertically) {
|
||||||
|
val auto = state.autoTranscribe.collectAsState().value
|
||||||
|
Checkbox(auto, { state.setAutoTranscribe(it) })
|
||||||
|
Text("Auto-transcribe new files",
|
||||||
|
style = MaterialTheme.typography.bodySmall)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
OutlinedTextField(query, { state.setQuery(it) },
|
||||||
|
label = { Text("Search saved transcripts") },
|
||||||
|
singleLine = true, modifier = Modifier.fillMaxWidth())
|
||||||
|
if (query.isNotBlank()) {
|
||||||
|
if (hits.isEmpty()) {
|
||||||
|
Text("No matches.", style = MaterialTheme.typography.bodySmall)
|
||||||
|
} else {
|
||||||
|
LazyColumn(Modifier.fillMaxWidth().weight(1f),
|
||||||
|
verticalArrangement = Arrangement.spacedBy(4.dp)) {
|
||||||
|
items(hits, key = { "${it.audioPath ?: "text"}:${it.report}:${it.line}" }) { h ->
|
||||||
|
val openable = h.audioPath?.let { File(it).exists() } == true
|
||||||
|
Card(
|
||||||
|
Modifier.fillMaxWidth().then(
|
||||||
|
if (openable) Modifier.clickable {
|
||||||
|
state.openDetail(File(h.audioPath))
|
||||||
|
} else Modifier
|
||||||
|
),
|
||||||
|
) {
|
||||||
|
Column(Modifier.padding(10.dp)) {
|
||||||
|
Text(h.report, style = MaterialTheme.typography.titleSmall)
|
||||||
|
Text(
|
||||||
|
if (h.line == 0) h.snippet else "line ${h.line}: …${h.snippet}…",
|
||||||
|
style = MaterialTheme.typography.bodySmall,
|
||||||
|
)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
} else if (folder == null) {
|
||||||
|
Text("Pick the folder with your .m4a recordings to begin.",
|
||||||
|
color = MaterialTheme.colorScheme.onSurfaceVariant)
|
||||||
|
} else if (entries.isEmpty()) {
|
||||||
|
Text("No audio files in this folder.",
|
||||||
|
color = MaterialTheme.colorScheme.onSurfaceVariant)
|
||||||
|
} else {
|
||||||
|
LazyColumn(Modifier.fillMaxWidth().weight(1f),
|
||||||
|
verticalArrangement = Arrangement.spacedBy(6.dp)) {
|
||||||
|
items(entries, key = { it.file.absolutePath }) { e ->
|
||||||
|
Card(Modifier.fillMaxWidth().clickable { state.openDetail(e.file) }) {
|
||||||
|
Row(Modifier.fillMaxWidth().padding(12.dp),
|
||||||
|
verticalAlignment = Alignment.CenterVertically) {
|
||||||
|
Column(Modifier.weight(1f)) {
|
||||||
|
Text(e.file.nameWithoutExtension,
|
||||||
|
style = MaterialTheme.typography.titleSmall,
|
||||||
|
maxLines = 1,
|
||||||
|
overflow = TextOverflow.Ellipsis)
|
||||||
|
Text(
|
||||||
|
"%.1f MB • %s".format(
|
||||||
|
e.file.length() / 1e6,
|
||||||
|
when (e.status) {
|
||||||
|
FileStatus.DONE -> "transcript saved"
|
||||||
|
FileStatus.QUEUED -> "queued for transcription"
|
||||||
|
// statusNote carries the live
|
||||||
|
// "transcribing… 42%" label.
|
||||||
|
FileStatus.RUNNING ->
|
||||||
|
e.statusNote ?: "transcribing…"
|
||||||
|
FileStatus.FAILED ->
|
||||||
|
"failed" + (e.statusNote?.let { " ($it)" } ?: "")
|
||||||
|
FileStatus.NEW -> "not transcribed"
|
||||||
|
},
|
||||||
|
),
|
||||||
|
style = MaterialTheme.typography.bodySmall,
|
||||||
|
color = when (e.status) {
|
||||||
|
FileStatus.FAILED -> MaterialTheme.colorScheme.error
|
||||||
|
FileStatus.DONE -> MaterialTheme.colorScheme.primary
|
||||||
|
else -> MaterialTheme.colorScheme.onSurfaceVariant
|
||||||
|
},
|
||||||
|
)
|
||||||
|
if (e.status == FileStatus.RUNNING) {
|
||||||
|
Spacer(Modifier.height(6.dp))
|
||||||
|
val frac = e.progress
|
||||||
|
if (frac != null) {
|
||||||
|
LinearProgressIndicator(
|
||||||
|
progress = { frac },
|
||||||
|
modifier = Modifier.fillMaxWidth().height(4.dp),
|
||||||
|
)
|
||||||
|
} else {
|
||||||
|
LinearProgressIndicator(
|
||||||
|
modifier = Modifier.fillMaxWidth().height(4.dp),
|
||||||
|
)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if (e.status == FileStatus.NEW || e.status == FileStatus.FAILED) {
|
||||||
|
TextButton({ state.pumpFile(e.file) }) { Text("Transcribe") }
|
||||||
|
}
|
||||||
|
if (e.status != FileStatus.RUNNING &&
|
||||||
|
e.status != FileStatus.QUEUED) {
|
||||||
|
TextButton({ renaming = e.file }) { Text("Rename") }
|
||||||
|
}
|
||||||
|
TextButton({ state.openDetail(e.file) }) { Text("Open") }
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
renaming?.let { target ->
|
||||||
|
var text by remember(target) { mutableStateOf(target.nameWithoutExtension) }
|
||||||
|
var error by remember(target) { mutableStateOf<String?>(null) }
|
||||||
|
AlertDialog(
|
||||||
|
onDismissRequest = { renaming = null },
|
||||||
|
title = { Text("Rename recording") },
|
||||||
|
text = {
|
||||||
|
Column {
|
||||||
|
OutlinedTextField(
|
||||||
|
value = text,
|
||||||
|
onValueChange = { text = it; error = null },
|
||||||
|
label = { Text("New name") },
|
||||||
|
singleLine = true,
|
||||||
|
)
|
||||||
|
error?.let {
|
||||||
|
Text(it, color = MaterialTheme.colorScheme.error,
|
||||||
|
style = MaterialTheme.typography.bodySmall)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
},
|
||||||
|
confirmButton = {
|
||||||
|
TextButton({
|
||||||
|
error = state.renameFile(target, text)
|
||||||
|
if (error == null) renaming = null
|
||||||
|
}) { Text("Rename") }
|
||||||
|
},
|
||||||
|
dismissButton = {
|
||||||
|
TextButton({ renaming = null }) { Text("Cancel") }
|
||||||
|
},
|
||||||
|
)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// ---- detail -----------------------------------------------------------------
|
||||||
|
|
||||||
|
@androidx.compose.runtime.Composable
|
||||||
|
fun DetailScreen(state: DesktopState) {
|
||||||
|
val detail = state.detail.collectAsState().value ?: run {
|
||||||
|
state.go(DesktopState.Screen.LIBRARY)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
val models = state.models.collectAsState().value
|
||||||
|
var showModels by remember { mutableStateOf(false) }
|
||||||
|
|
||||||
|
Column(Modifier.fillMaxSize(), verticalArrangement = Arrangement.spacedBy(12.dp)) {
|
||||||
|
Row(verticalAlignment = Alignment.CenterVertically) {
|
||||||
|
TextButton({ state.go(DesktopState.Screen.LIBRARY) }) { Text("← Library") }
|
||||||
|
Text(detail.file.nameWithoutExtension,
|
||||||
|
style = MaterialTheme.typography.headlineSmall,
|
||||||
|
modifier = Modifier.weight(1f))
|
||||||
|
}
|
||||||
|
// Model override: "Use default" or a specific size.
|
||||||
|
Row(verticalAlignment = Alignment.CenterVertically) {
|
||||||
|
Text("Model: ", style = MaterialTheme.typography.bodyMedium)
|
||||||
|
val current = detail.overrideModel
|
||||||
|
?: models?.defaultModel
|
||||||
|
?: "base"
|
||||||
|
TextButton({ showModels = !showModels }) {
|
||||||
|
Text(if (detail.overrideModel == null) "Use default ($current)" else current)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if (showModels && models != null) {
|
||||||
|
Column(verticalArrangement = Arrangement.spacedBy(4.dp)) {
|
||||||
|
ModelRow(state, detail, null, models.defaultModel,
|
||||||
|
"Use default (${models.defaultModel})", "")
|
||||||
|
models.models.forEach { m ->
|
||||||
|
ModelRow(state, detail, m.name, models.defaultModel,
|
||||||
|
m.displayName, m.description +
|
||||||
|
(if (!m.available) " — not downloaded" else ""))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
Spacer(Modifier.height(4.dp))
|
||||||
|
}
|
||||||
|
detail.uploadProgress?.let {
|
||||||
|
LinearProgressIndicator(progress = { it }, modifier = Modifier.fillMaxWidth())
|
||||||
|
Text("Uploading… ${(it * 100).toInt()}%")
|
||||||
|
}
|
||||||
|
detail.busy?.takeIf { detail.uploadProgress == null }?.let {
|
||||||
|
Row(verticalAlignment = Alignment.CenterVertically) {
|
||||||
|
CircularProgressIndicator(Modifier.width(20.dp).height(20.dp))
|
||||||
|
Spacer(Modifier.width(10.dp))
|
||||||
|
Text(it)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
val tJob = detail.jobs.firstOrNull { it.jobType == "transcribe" }
|
||||||
|
tJob?.let {
|
||||||
|
val running = it.status == "running"
|
||||||
|
val label = when {
|
||||||
|
it.status == "succeeded" -> "Transcribed"
|
||||||
|
it.status == "failed" -> "Failed"
|
||||||
|
it.stage == "loading-model" -> "Loading model…"
|
||||||
|
running && it.progress != null -> "Transcribing… ${it.progress}%"
|
||||||
|
it.stage == "transcribing" -> "Transcribing…"
|
||||||
|
running -> "Working… (attempt ${it.attempt}/${it.maxAttempts})"
|
||||||
|
else -> "Queued"
|
||||||
|
}
|
||||||
|
Text(label, style = MaterialTheme.typography.bodyMedium,
|
||||||
|
color = if (it.status == "failed") MaterialTheme.colorScheme.error
|
||||||
|
else MaterialTheme.colorScheme.primary)
|
||||||
|
if (running) {
|
||||||
|
Spacer(Modifier.height(4.dp))
|
||||||
|
if (it.progress != null) {
|
||||||
|
LinearProgressIndicator(
|
||||||
|
progress = { it.progress / 100f },
|
||||||
|
modifier = Modifier.fillMaxWidth().height(4.dp),
|
||||||
|
)
|
||||||
|
} else {
|
||||||
|
LinearProgressIndicator(
|
||||||
|
modifier = Modifier.fillMaxWidth().height(4.dp),
|
||||||
|
)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
val sJob = detail.jobs.firstOrNull { it.jobType == "summarize" }
|
||||||
|
sJob?.takeIf { it.status == "running" || it.status == "queued" }?.let {
|
||||||
|
val label = if (it.status == "running") "Summarizing…" else "Summary queued"
|
||||||
|
Text(label, style = MaterialTheme.typography.bodyMedium,
|
||||||
|
color = MaterialTheme.colorScheme.primary)
|
||||||
|
Spacer(Modifier.height(4.dp))
|
||||||
|
LinearProgressIndicator(
|
||||||
|
modifier = Modifier.fillMaxWidth().height(4.dp),
|
||||||
|
)
|
||||||
|
}
|
||||||
|
detail.error?.let {
|
||||||
|
Text(it, color = MaterialTheme.colorScheme.error)
|
||||||
|
if (tJob?.status == "failed") {
|
||||||
|
Button({ state.retry() }) { Text("Retry transcription") }
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if (detail.busy == null && detail.transcript == null &&
|
||||||
|
detail.reportText == null && detail.error == null) {
|
||||||
|
Button({ state.transcribe() }) { Text("Transcribe") }
|
||||||
|
}
|
||||||
|
detail.transcript?.let { t ->
|
||||||
|
Text("Transcript", style = MaterialTheme.typography.titleMedium)
|
||||||
|
LazyColumn(Modifier.fillMaxWidth().weight(1f),
|
||||||
|
verticalArrangement = Arrangement.spacedBy(4.dp)) {
|
||||||
|
if (t.segments.isNotEmpty()) {
|
||||||
|
items(t.segments.size) { i ->
|
||||||
|
val s = t.segments[i]
|
||||||
|
Text("[${fmtTs(s.startSec)}] ${s.text}",
|
||||||
|
style = MaterialTheme.typography.bodyMedium)
|
||||||
|
}
|
||||||
|
} else {
|
||||||
|
item { Text(t.text) }
|
||||||
|
}
|
||||||
|
}
|
||||||
|
Text("Saved to ${detail.file.nameWithoutExtension}.transcript.md",
|
||||||
|
style = MaterialTheme.typography.bodySmall,
|
||||||
|
color = MaterialTheme.colorScheme.onSurfaceVariant)
|
||||||
|
}
|
||||||
|
detail.summary?.let { s ->
|
||||||
|
if (s.short.isNotBlank()) {
|
||||||
|
Text("Summary", style = MaterialTheme.typography.titleMedium)
|
||||||
|
Text(s.short)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
// No live transcript (file opened from the library after the fact):
|
||||||
|
// show the saved report so summary + transcript are actually visible.
|
||||||
|
val savedReport = detail.reportText
|
||||||
|
if (detail.transcript == null && savedReport != null) {
|
||||||
|
Text("Saved report", style = MaterialTheme.typography.titleMedium)
|
||||||
|
val reportLines = savedReport.lines()
|
||||||
|
LazyColumn(Modifier.fillMaxWidth().weight(1f),
|
||||||
|
verticalArrangement = Arrangement.spacedBy(2.dp)) {
|
||||||
|
items(reportLines.size) { i ->
|
||||||
|
Text(reportLines[i],
|
||||||
|
style = MaterialTheme.typography.bodyMedium)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
Text("Saved to ${detail.file.nameWithoutExtension}.transcript.md",
|
||||||
|
style = MaterialTheme.typography.bodySmall,
|
||||||
|
color = MaterialTheme.colorScheme.onSurfaceVariant)
|
||||||
|
}
|
||||||
|
Spacer(Modifier.height(4.dp))
|
||||||
|
if (detail.transcript != null || detail.reportText != null) {
|
||||||
|
// Honest labels: "Re-…" only when that output actually exists.
|
||||||
|
val report = detail.reportText
|
||||||
|
val hasTranscript = detail.transcript != null ||
|
||||||
|
(report != null && !report.contains("_No transcript available._"))
|
||||||
|
val hasSummary = detail.summary != null ||
|
||||||
|
(report?.contains("## Summary") == true)
|
||||||
|
val remoteKnown = state.remoteIdFor(detail.file) != null
|
||||||
|
Row(horizontalArrangement = Arrangement.spacedBy(8.dp)) {
|
||||||
|
TextButton({ state.transcribe() }) {
|
||||||
|
Text(if (hasTranscript) "Re-transcribe" else "Transcribe")
|
||||||
|
}
|
||||||
|
TextButton({ state.summarize() },
|
||||||
|
enabled = remoteKnown && hasTranscript) {
|
||||||
|
Text(if (hasSummary) "Re-summarize" else "Summarize")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if (!remoteKnown || !hasTranscript) {
|
||||||
|
Text(if (!hasTranscript)
|
||||||
|
"Summarize unlocks once this recording has a transcript."
|
||||||
|
else
|
||||||
|
"Summarize unlocks once this file is uploaded (press Transcribe).",
|
||||||
|
style = MaterialTheme.typography.bodySmall,
|
||||||
|
color = MaterialTheme.colorScheme.onSurfaceVariant)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
@androidx.compose.runtime.Composable
|
||||||
|
private fun ModelRow(
|
||||||
|
state: DesktopState,
|
||||||
|
detail: DesktopState.DetailUi,
|
||||||
|
value: String?,
|
||||||
|
default: String,
|
||||||
|
title: String,
|
||||||
|
subtitle: String,
|
||||||
|
) {
|
||||||
|
val selected = detail.overrideModel == value ||
|
||||||
|
(value == null && detail.overrideModel == null)
|
||||||
|
Card(
|
||||||
|
Modifier.fillMaxWidth().clickable { state.setOverride(value) },
|
||||||
|
) {
|
||||||
|
Row(Modifier.fillMaxWidth().padding(10.dp),
|
||||||
|
verticalAlignment = Alignment.CenterVertically) {
|
||||||
|
Text(if (selected) "◉ " else "○ ")
|
||||||
|
Column {
|
||||||
|
Text(title, style = MaterialTheme.typography.bodyMedium)
|
||||||
|
if (subtitle.isNotBlank()) {
|
||||||
|
Text(subtitle, style = MaterialTheme.typography.bodySmall,
|
||||||
|
color = MaterialTheme.colorScheme.onSurfaceVariant)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// ---- settings -----------------------------------------------------------------
|
||||||
|
|
||||||
|
@androidx.compose.runtime.Composable
|
||||||
|
fun SettingsScreen(state: DesktopState) {
|
||||||
|
val models = state.models.collectAsState().value
|
||||||
|
val modelsError = state.modelsError.collectAsState().value
|
||||||
|
val url by state.serverUrl.collectAsState()
|
||||||
|
|
||||||
|
Column(Modifier.fillMaxSize(), verticalArrangement = Arrangement.spacedBy(12.dp)) {
|
||||||
|
Row(verticalAlignment = Alignment.CenterVertically) {
|
||||||
|
TextButton({ state.go(DesktopState.Screen.LIBRARY) }) { Text("← Library") }
|
||||||
|
Text("Settings", style = MaterialTheme.typography.headlineSmall)
|
||||||
|
}
|
||||||
|
Text("Local engine: $url", style = MaterialTheme.typography.bodySmall)
|
||||||
|
Spacer(Modifier.height(4.dp))
|
||||||
|
Text("Default transcription model", style = MaterialTheme.typography.titleMedium)
|
||||||
|
Text("Applies to future transcriptions. Each recording keeps the model it used.",
|
||||||
|
style = MaterialTheme.typography.bodySmall,
|
||||||
|
color = MaterialTheme.colorScheme.onSurfaceVariant)
|
||||||
|
if (models == null) {
|
||||||
|
Row(verticalAlignment = Alignment.CenterVertically) {
|
||||||
|
CircularProgressIndicator(Modifier.width(20.dp).height(20.dp))
|
||||||
|
Spacer(Modifier.width(10.dp))
|
||||||
|
Text("Loading models…")
|
||||||
|
}
|
||||||
|
} else {
|
||||||
|
if (!models.fasterWhisperInstalled) {
|
||||||
|
Text("faster-whisper is not installed on the server " +
|
||||||
|
"(pip install shonar-backend[faster-whisper]).",
|
||||||
|
color = MaterialTheme.colorScheme.error)
|
||||||
|
}
|
||||||
|
models.models.forEach { m ->
|
||||||
|
Card(Modifier.fillMaxWidth()) {
|
||||||
|
Column(Modifier.padding(10.dp),
|
||||||
|
verticalArrangement = Arrangement.spacedBy(4.dp)) {
|
||||||
|
Row(verticalAlignment = Alignment.CenterVertically) {
|
||||||
|
Text(
|
||||||
|
(if (m.isDefault) "◉ " else "○ ") + m.displayName,
|
||||||
|
style = MaterialTheme.typography.bodyMedium,
|
||||||
|
modifier = Modifier.weight(1f),
|
||||||
|
)
|
||||||
|
if (!m.isDefault) {
|
||||||
|
TextButton({ state.setDefaultModel(m.name) }) {
|
||||||
|
Text("Set default")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if (!m.downloaded) {
|
||||||
|
TextButton({ state.downloadModel(m.name) }) {
|
||||||
|
Text("Download")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
Text(m.description, style = MaterialTheme.typography.bodySmall,
|
||||||
|
color = MaterialTheme.colorScheme.onSurfaceVariant)
|
||||||
|
Text("${m.params} • ${m.approxMemory} • ${m.relativeSpeed}" +
|
||||||
|
(if (m.available) " • ready" else " • not downloaded"),
|
||||||
|
style = MaterialTheme.typography.bodySmall,
|
||||||
|
color = MaterialTheme.colorScheme.onSurfaceVariant)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
modelsError?.let { Text(it, color = MaterialTheme.colorScheme.error) }
|
||||||
|
}
|
||||||
|
}
|
||||||
69
app/src/main/kotlin/com/shonar/desktop/Theme.kt
Normal file
69
app/src/main/kotlin/com/shonar/desktop/Theme.kt
Normal file
|
|
@ -0,0 +1,69 @@
|
||||||
|
package com.shonar.desktop
|
||||||
|
|
||||||
|
import androidx.compose.material3.MaterialTheme
|
||||||
|
import androidx.compose.material3.darkColorScheme
|
||||||
|
import androidx.compose.runtime.Composable
|
||||||
|
import androidx.compose.ui.graphics.Color
|
||||||
|
|
||||||
|
/**
|
||||||
|
* SHONAR Desktop theme — "warm graphite": near-black charcoal surfaces
|
||||||
|
* with an olive-gold accent, inspired by the Skills Hub hero design.
|
||||||
|
* Desktop-only for now (Android keeps its own ui/theme/Theme.kt).
|
||||||
|
*
|
||||||
|
* Status semantics in the library list map to scheme roles:
|
||||||
|
* done -> primary (gold), failed -> error (red), queued/running ->
|
||||||
|
* onSurfaceVariant with a primary progress bar.
|
||||||
|
*/
|
||||||
|
private val ShonarGold = Color(0xFFE0B23A) // text-safe gold (brighter for AA on cards)
|
||||||
|
private val ShonarGoldBright = Color(0xFFF0C419)
|
||||||
|
private val GoldInk = Color(0xFF1F1800) // text on gold fills (AA)
|
||||||
|
private val GraphiteBg = Color(0xFF1A1A1A)
|
||||||
|
private val GraphiteSurface = Color(0xFF222222)
|
||||||
|
private val GraphiteHigh = Color(0xFF2B2B2B)
|
||||||
|
private val InkPrimary = Color(0xFFF4F4F4)
|
||||||
|
private val InkSecondary = Color(0xFFB3B3B3)
|
||||||
|
private val InkMuted = Color(0xFF8A8A8A)
|
||||||
|
private val OutlineWarm = Color(0xFF3A3A3A)
|
||||||
|
private val DangerWarm = Color(0xFFE5544F)
|
||||||
|
|
||||||
|
private val ShonarDarkScheme = darkColorScheme(
|
||||||
|
primary = ShonarGold,
|
||||||
|
onPrimary = GoldInk,
|
||||||
|
primaryContainer = Color(0xFF3A3216),
|
||||||
|
onPrimaryContainer = ShonarGoldBright,
|
||||||
|
|
||||||
|
secondary = Color(0xFFB99B5E),
|
||||||
|
onSecondary = GoldInk,
|
||||||
|
secondaryContainer = Color(0xFF332C1A),
|
||||||
|
onSecondaryContainer = Color(0xFFE8D9AE),
|
||||||
|
|
||||||
|
tertiary = Color(0xFF8FA6B2), // cool counterpoint (info-ish)
|
||||||
|
onTertiary = Color(0xFF0E1A20),
|
||||||
|
tertiaryContainer = Color(0xFF2A3B45),
|
||||||
|
onTertiaryContainer = Color(0xFFCDE2EC),
|
||||||
|
|
||||||
|
background = GraphiteBg,
|
||||||
|
onBackground = InkPrimary,
|
||||||
|
|
||||||
|
surface = GraphiteSurface,
|
||||||
|
onSurface = InkPrimary,
|
||||||
|
surfaceVariant = Color(0xFF2E2A21),
|
||||||
|
onSurfaceVariant = InkSecondary,
|
||||||
|
surfaceContainer = GraphiteHigh, // cards lift off the bg
|
||||||
|
surfaceContainerHigh = Color(0xFF333333),
|
||||||
|
surfaceContainerHighest = Color(0xFF3A3A3A),
|
||||||
|
|
||||||
|
outline = OutlineWarm,
|
||||||
|
outlineVariant = Color(0xFF2E2E2E),
|
||||||
|
|
||||||
|
error = DangerWarm,
|
||||||
|
onError = Color(0xFF2B0A08),
|
||||||
|
errorContainer = Color(0xFF47201E),
|
||||||
|
onErrorContainer = Color(0xFFFFB4AF),
|
||||||
|
)
|
||||||
|
|
||||||
|
/** Wrap the app content in the SHONAR dark-gold theme. */
|
||||||
|
@Composable
|
||||||
|
fun ShonarTheme(content: @Composable () -> Unit) {
|
||||||
|
MaterialTheme(colorScheme = ShonarDarkScheme, content = content)
|
||||||
|
}
|
||||||
101
app/src/main/kotlin/com/shonar/settings/DesktopSettingsStore.kt
Normal file
101
app/src/main/kotlin/com/shonar/settings/DesktopSettingsStore.kt
Normal file
|
|
@ -0,0 +1,101 @@
|
||||||
|
package com.shonar.settings
|
||||||
|
|
||||||
|
import java.io.File
|
||||||
|
import kotlinx.coroutines.flow.MutableStateFlow
|
||||||
|
import kotlinx.coroutines.flow.StateFlow
|
||||||
|
import kotlinx.coroutines.flow.asStateFlow
|
||||||
|
import kotlinx.coroutines.sync.Mutex
|
||||||
|
import kotlinx.coroutines.sync.withLock
|
||||||
|
import org.json.JSONObject
|
||||||
|
|
||||||
|
/**
|
||||||
|
* Desktop copy of the settings-store contract. The Android original
|
||||||
|
* (DataStore/Keystore implementations) cannot compile on JVM desktop, so
|
||||||
|
* this file provides the same fully-qualified names backed by plain files:
|
||||||
|
* [SettingsStore], [InMemorySettingsStore], and [FileSettingsStore].
|
||||||
|
*
|
||||||
|
* Secrets note: desktop v1 keeps server tokens in a 0600 JSON file next to
|
||||||
|
* ordinary settings. Good enough for a local single-user app; a keyring
|
||||||
|
* backend can replace [FileSettingsStore] later without touching callers.
|
||||||
|
*/
|
||||||
|
interface SettingsStore {
|
||||||
|
suspend fun getString(key: String): String?
|
||||||
|
suspend fun putString(key: String, value: String)
|
||||||
|
suspend fun remove(key: String)
|
||||||
|
suspend fun keys(): Set<String>
|
||||||
|
|
||||||
|
/** Emits whenever any value changes (for reactive UI). */
|
||||||
|
val changes: StateFlow<Long>
|
||||||
|
}
|
||||||
|
|
||||||
|
/** In-memory store: used by unit tests. */
|
||||||
|
class InMemorySettingsStore : SettingsStore {
|
||||||
|
private val map = mutableMapOf<String, String>()
|
||||||
|
private val changeFlow = MutableStateFlow(0L)
|
||||||
|
override val changes: StateFlow<Long> = changeFlow.asStateFlow()
|
||||||
|
|
||||||
|
override suspend fun getString(key: String): String? = map[key]
|
||||||
|
override suspend fun putString(key: String, value: String) {
|
||||||
|
map[key] = value
|
||||||
|
changeFlow.value += 1
|
||||||
|
}
|
||||||
|
|
||||||
|
override suspend fun remove(key: String) {
|
||||||
|
map.remove(key)
|
||||||
|
changeFlow.value += 1
|
||||||
|
}
|
||||||
|
|
||||||
|
override suspend fun keys(): Set<String> = map.keys.toSet()
|
||||||
|
}
|
||||||
|
|
||||||
|
/** JSON file store (0600). Keys map 1:1 to the JSON object members. */
|
||||||
|
class FileSettingsStore(private val file: File) : SettingsStore {
|
||||||
|
private val mutex = Mutex()
|
||||||
|
private val changeFlow = MutableStateFlow(0L)
|
||||||
|
override val changes: StateFlow<Long> = changeFlow.asStateFlow()
|
||||||
|
|
||||||
|
private fun readAll(): MutableMap<String, String> {
|
||||||
|
if (!file.isFile) return mutableMapOf()
|
||||||
|
return runCatching {
|
||||||
|
val obj = JSONObject(file.readText())
|
||||||
|
buildMap {
|
||||||
|
for (key in obj.keys()) put(key, obj.optString(key, ""))
|
||||||
|
}.toMutableMap()
|
||||||
|
}.getOrDefault(mutableMapOf())
|
||||||
|
}
|
||||||
|
|
||||||
|
private fun writeAll(map: Map<String, String>) {
|
||||||
|
file.parentFile?.mkdirs()
|
||||||
|
file.writeText(JSONObject(map).toString())
|
||||||
|
runCatching {
|
||||||
|
file.setReadable(false, false)
|
||||||
|
file.setReadable(true, true)
|
||||||
|
file.setWritable(false, false)
|
||||||
|
file.setWritable(true, true)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
override suspend fun getString(key: String): String? = mutex.withLock {
|
||||||
|
readAll()[key]
|
||||||
|
}
|
||||||
|
|
||||||
|
override suspend fun putString(key: String, value: String) = mutex.withLock {
|
||||||
|
val map = readAll()
|
||||||
|
map[key] = value
|
||||||
|
writeAll(map)
|
||||||
|
changeFlow.value += 1
|
||||||
|
}
|
||||||
|
|
||||||
|
override suspend fun remove(key: String) {
|
||||||
|
mutex.withLock {
|
||||||
|
val map = readAll()
|
||||||
|
map.remove(key)
|
||||||
|
writeAll(map)
|
||||||
|
changeFlow.value += 1
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
override suspend fun keys(): Set<String> = mutex.withLock {
|
||||||
|
readAll().keys.toSet()
|
||||||
|
}
|
||||||
|
}
|
||||||
48
app/src/test/kotlin/com/shonar/desktop/JobProgressTest.kt
Normal file
48
app/src/test/kotlin/com/shonar/desktop/JobProgressTest.kt
Normal file
|
|
@ -0,0 +1,48 @@
|
||||||
|
package com.shonar.desktop
|
||||||
|
|
||||||
|
import com.shonar.recording.JobInfo
|
||||||
|
import org.junit.Assert.assertEquals
|
||||||
|
import org.junit.Assert.assertNull
|
||||||
|
import org.junit.Test
|
||||||
|
|
||||||
|
class JobProgressTest {
|
||||||
|
|
||||||
|
private fun job(type: String, status: String, stage: String? = null, progress: Int? = null) =
|
||||||
|
JobInfo(jobType = type, status = status, stage = stage, progress = progress)
|
||||||
|
|
||||||
|
@Test fun `idle when nothing running`() {
|
||||||
|
assertNull(DesktopState.jobProgress(emptyList()))
|
||||||
|
assertNull(DesktopState.jobProgress(listOf(job("transcribe", "succeeded"))))
|
||||||
|
}
|
||||||
|
|
||||||
|
@Test fun `loading model has no percentage`() {
|
||||||
|
val p = DesktopState.jobProgress(
|
||||||
|
listOf(job("transcribe", "running", stage = "loading-model")),
|
||||||
|
)
|
||||||
|
assertEquals("loading model…", p?.label)
|
||||||
|
assertNull(p?.fraction)
|
||||||
|
}
|
||||||
|
|
||||||
|
@Test fun `transcribing shows backend percentage`() {
|
||||||
|
val p = DesktopState.jobProgress(
|
||||||
|
listOf(job("transcribe", "running", stage = "transcribing", progress = 42)),
|
||||||
|
)
|
||||||
|
assertEquals("transcribing… 42%", p?.label)
|
||||||
|
assertEquals(0.42f, p!!.fraction!!, 0.0001f)
|
||||||
|
}
|
||||||
|
|
||||||
|
@Test fun `transcribe failure falls back to indeterminate`() {
|
||||||
|
val p = DesktopState.jobProgress(
|
||||||
|
listOf(job("transcribe", "running", stage = "transcribing")),
|
||||||
|
)
|
||||||
|
assertEquals("transcribing…", p?.label)
|
||||||
|
assertNull(p?.fraction)
|
||||||
|
}
|
||||||
|
|
||||||
|
@Test fun `summarizing stage`() {
|
||||||
|
val p = DesktopState.jobProgress(
|
||||||
|
listOf(job("transcribe", "succeeded"), job("summarize", "running")),
|
||||||
|
)
|
||||||
|
assertEquals("summarizing…", p?.label)
|
||||||
|
}
|
||||||
|
}
|
||||||
51
app/src/test/kotlin/com/shonar/desktop/LibraryQueueTest.kt
Normal file
51
app/src/test/kotlin/com/shonar/desktop/LibraryQueueTest.kt
Normal file
|
|
@ -0,0 +1,51 @@
|
||||||
|
package com.shonar.desktop
|
||||||
|
|
||||||
|
import java.io.File
|
||||||
|
import kotlin.io.path.createTempDirectory
|
||||||
|
import org.junit.Assert.assertEquals
|
||||||
|
import org.junit.Assert.assertTrue
|
||||||
|
import org.junit.Test
|
||||||
|
|
||||||
|
class LibraryQueueTest {
|
||||||
|
|
||||||
|
private val exts = setOf("m4a", "wav", "mp3")
|
||||||
|
|
||||||
|
@Test fun candidateOnlyForKnownAudioExt() {
|
||||||
|
val dir = createTempDirectory("lq").toFile()
|
||||||
|
val audio = File(dir, "a.m4a").apply { writeText("x") }
|
||||||
|
val note = File(dir, "notes.txt").apply { writeText("x") }
|
||||||
|
val hidden = File(dir, ".hidden.m4a").apply { writeText("x") }
|
||||||
|
assertTrue(LibraryQueue.isAudioCandidate(audio, exts))
|
||||||
|
assertTrue(!LibraryQueue.isAudioCandidate(note, exts))
|
||||||
|
assertTrue(!LibraryQueue.isAudioCandidate(hidden, exts))
|
||||||
|
assertTrue(!LibraryQueue.isAudioCandidate(File(dir, "sub"), exts))
|
||||||
|
}
|
||||||
|
|
||||||
|
@Test fun diffDetectsAddedAndRemoved() {
|
||||||
|
val dir = createTempDirectory("lq").toFile()
|
||||||
|
val a = File(dir, "a.m4a").apply { writeText("x") }
|
||||||
|
val b = File(dir, "b.m4a").apply { writeText("x") }
|
||||||
|
val d = LibraryQueue.diff(setOf("gone.m4a"), listOf(a, b))
|
||||||
|
assertEquals(listOf("a.m4a", "b.m4a"), d.added.map { it.name })
|
||||||
|
assertEquals(setOf("gone.m4a"), d.removed)
|
||||||
|
}
|
||||||
|
|
||||||
|
@Test fun autoQueueSkipsReportsInFlightAndFailed() {
|
||||||
|
val dir = createTempDirectory("lq").toFile()
|
||||||
|
val fresh = File(dir, "fresh.m4a").apply { writeText("x") }
|
||||||
|
val done = File(dir, "done.m4a").apply { writeText("x") }
|
||||||
|
File(dir, "done.transcript.md").writeText("# done")
|
||||||
|
val running = File(dir, "running.m4a").apply { writeText("x") }
|
||||||
|
val broken = File(dir, "broken.m4a").apply { writeText("x") }
|
||||||
|
|
||||||
|
val picked = LibraryQueue.autoQueueCandidates(
|
||||||
|
entries = listOf(fresh, done, running, broken),
|
||||||
|
reportExists = { f ->
|
||||||
|
File(f.parentFile, "${f.nameWithoutExtension}.transcript.md").exists()
|
||||||
|
},
|
||||||
|
inFlightOrDone = setOf("running.m4a"),
|
||||||
|
failed = setOf("broken.m4a"),
|
||||||
|
)
|
||||||
|
assertEquals(listOf("fresh.m4a"), picked.map { it.name })
|
||||||
|
}
|
||||||
|
}
|
||||||
59
app/src/test/kotlin/com/shonar/desktop/RemoteMappingTest.kt
Normal file
59
app/src/test/kotlin/com/shonar/desktop/RemoteMappingTest.kt
Normal file
|
|
@ -0,0 +1,59 @@
|
||||||
|
package com.shonar.desktop
|
||||||
|
|
||||||
|
import java.io.File
|
||||||
|
import kotlin.io.path.createTempDirectory
|
||||||
|
import org.junit.Assert.assertEquals
|
||||||
|
import org.junit.Assert.assertNull
|
||||||
|
import org.junit.Assert.assertTrue
|
||||||
|
import org.junit.Test
|
||||||
|
|
||||||
|
class RemoteMappingTest {
|
||||||
|
private fun tmpAudio(name: String = "rec.m4a"): File {
|
||||||
|
val dir = createTempDirectory().toFile()
|
||||||
|
val f = File(dir, name)
|
||||||
|
f.writeText("audio")
|
||||||
|
return f
|
||||||
|
}
|
||||||
|
|
||||||
|
@Test fun `round trip save and load`() {
|
||||||
|
val f = tmpAudio()
|
||||||
|
saveMapping(f, RemoteMapping("uuid-1", model = "small", title = "rec"))
|
||||||
|
val m = loadMapping(f)
|
||||||
|
assertEquals("uuid-1", m?.recordingId)
|
||||||
|
assertEquals("small", m?.model)
|
||||||
|
assertEquals("rec", m?.title)
|
||||||
|
}
|
||||||
|
|
||||||
|
@Test fun `load returns null when absent`() {
|
||||||
|
assertNull(loadMapping(tmpAudio()))
|
||||||
|
}
|
||||||
|
|
||||||
|
@Test fun `load returns null on corrupt json`() {
|
||||||
|
val f = tmpAudio()
|
||||||
|
mappingFile(f).writeText("{not json")
|
||||||
|
assertNull(loadMapping(f))
|
||||||
|
}
|
||||||
|
|
||||||
|
@Test fun `load returns null on blank id`() {
|
||||||
|
val f = tmpAudio()
|
||||||
|
mappingFile(f).writeText("""{"recording_id": ""}""")
|
||||||
|
assertNull(loadMapping(f))
|
||||||
|
}
|
||||||
|
|
||||||
|
@Test fun `move follows the audio rename`() {
|
||||||
|
val f = tmpAudio("old.m4a")
|
||||||
|
saveMapping(f, RemoteMapping("uuid-2"))
|
||||||
|
val renamed = File(f.parentFile, "new.m4a")
|
||||||
|
f.renameTo(renamed)
|
||||||
|
assertTrue(moveMapping(f, renamed))
|
||||||
|
assertEquals("uuid-2", loadMapping(renamed)?.recordingId)
|
||||||
|
assertNull(loadMapping(f))
|
||||||
|
}
|
||||||
|
|
||||||
|
@Test fun `move is a no-op without an existing mapping`() {
|
||||||
|
val f = tmpAudio("a.m4a")
|
||||||
|
val g = File(f.parentFile, "b.m4a")
|
||||||
|
assertTrue(!moveMapping(f, g))
|
||||||
|
assertNull(loadMapping(g))
|
||||||
|
}
|
||||||
|
}
|
||||||
43
app/src/test/kotlin/com/shonar/desktop/RenameTest.kt
Normal file
43
app/src/test/kotlin/com/shonar/desktop/RenameTest.kt
Normal file
|
|
@ -0,0 +1,43 @@
|
||||||
|
package com.shonar.desktop
|
||||||
|
|
||||||
|
import java.io.File
|
||||||
|
import kotlin.io.path.createTempDirectory
|
||||||
|
import org.junit.Assert.assertEquals
|
||||||
|
import org.junit.Assert.assertTrue
|
||||||
|
import org.junit.Test
|
||||||
|
|
||||||
|
class RenameTest {
|
||||||
|
|
||||||
|
private fun tmp(): File = createTempDirectory("rn").toFile()
|
||||||
|
|
||||||
|
@Test fun `plain rename keeps extension`() {
|
||||||
|
val dir = tmp()
|
||||||
|
val f = File(dir, "Recording 20260803.m4a").apply { writeText("x") }
|
||||||
|
val target = LibraryQueue.renameTarget(f, "standup intro").getOrThrow()
|
||||||
|
assertEquals(File(dir, "standup intro.m4a"), target)
|
||||||
|
}
|
||||||
|
|
||||||
|
@Test fun `typed extension is stripped`() {
|
||||||
|
val dir = tmp()
|
||||||
|
val f = File(dir, "a.m4a").apply { writeText("x") }
|
||||||
|
assertEquals(File(dir, "b.m4a"), LibraryQueue.renameTarget(f, "b.M4A").getOrThrow())
|
||||||
|
assertEquals(File(dir, "b.m4a"), LibraryQueue.renameTarget(f, "b.m4a").getOrThrow())
|
||||||
|
}
|
||||||
|
|
||||||
|
@Test fun `rejects empty path and dot names`() {
|
||||||
|
val dir = tmp()
|
||||||
|
val f = File(dir, "a.m4a").apply { writeText("x") }
|
||||||
|
for (bad in listOf("", " ", "x/y", "x\\y", ".", "..", ".hidden")) {
|
||||||
|
val r = LibraryQueue.renameTarget(f, bad)
|
||||||
|
assertTrue("should reject '$bad'", r.isFailure)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
@Test fun `rejects collisions and no-op renames`() {
|
||||||
|
val dir = tmp()
|
||||||
|
val f = File(dir, "a.m4a").apply { writeText("x") }
|
||||||
|
File(dir, "taken.m4a").writeText("y")
|
||||||
|
assertTrue(LibraryQueue.renameTarget(f, "taken").isFailure)
|
||||||
|
assertTrue(LibraryQueue.renameTarget(f, "a").isFailure)
|
||||||
|
}
|
||||||
|
}
|
||||||
60
app/src/test/kotlin/com/shonar/desktop/ReportsTest.kt
Normal file
60
app/src/test/kotlin/com/shonar/desktop/ReportsTest.kt
Normal file
|
|
@ -0,0 +1,60 @@
|
||||||
|
package com.shonar.desktop
|
||||||
|
|
||||||
|
import com.shonar.recording.SummaryData
|
||||||
|
import com.shonar.recording.TranscriptData
|
||||||
|
import com.shonar.recording.TranscriptSegment
|
||||||
|
import java.io.File
|
||||||
|
import java.nio.file.Files
|
||||||
|
import org.junit.Assert.assertEquals
|
||||||
|
import org.junit.Assert.assertTrue
|
||||||
|
import org.junit.Test
|
||||||
|
|
||||||
|
class ReportsTest {
|
||||||
|
|
||||||
|
@Test fun renderTranscriptWithTimestamps() {
|
||||||
|
val t = TranscriptData(
|
||||||
|
version = 1, provider = "faster_whisper", model = "base",
|
||||||
|
text = "hello world",
|
||||||
|
segments = listOf(
|
||||||
|
TranscriptSegment(0.0, 1.5, "hello"),
|
||||||
|
TranscriptSegment(65.0, 67.0, "world"),
|
||||||
|
),
|
||||||
|
)
|
||||||
|
val md = renderReport("Evening notes", t, null)
|
||||||
|
assertTrue(md.startsWith("# Evening notes"))
|
||||||
|
assertTrue("[00:00:00] hello" in md)
|
||||||
|
assertTrue("[00:01:05] world" in md)
|
||||||
|
assertTrue("faster_whisper, v1" in md)
|
||||||
|
}
|
||||||
|
|
||||||
|
@Test fun renderEmptyTranscript() {
|
||||||
|
val md = renderReport("x", null, null)
|
||||||
|
assertTrue("_No transcript available._" in md)
|
||||||
|
}
|
||||||
|
|
||||||
|
@Test fun fmtTsShape() {
|
||||||
|
assertEquals("00:00:00", fmtTs(0.0))
|
||||||
|
assertEquals("00:01:05", fmtTs(65.0))
|
||||||
|
assertEquals("01:02:03", fmtTs(3723.0))
|
||||||
|
}
|
||||||
|
|
||||||
|
@Test fun reportFileName() {
|
||||||
|
val dir = Files.createTempDirectory("rep").toFile()
|
||||||
|
try {
|
||||||
|
val audio = File(dir, "chat.m4a")
|
||||||
|
assertEquals("chat.transcript.md", DesktopState.reportFile(audio).name)
|
||||||
|
} finally {
|
||||||
|
dir.deleteRecursively()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
@Test fun sharedParsersWorkOnDesktop() {
|
||||||
|
val raw = """{"version":2,"provider":"user","text":"hi","segments":[],
|
||||||
|
"edited_by_user":true}"""
|
||||||
|
val t = com.shonar.recording.parseTranscript(raw)!!
|
||||||
|
assertEquals(2, t.version)
|
||||||
|
assertTrue(t.editedByUser)
|
||||||
|
val summary = SummaryData(version = 1, content = mapOf("short" to "s"))
|
||||||
|
assertEquals("s", summary.short)
|
||||||
|
}
|
||||||
|
}
|
||||||
16
backend/README.md
Normal file
16
backend/README.md
Normal file
|
|
@ -0,0 +1,16 @@
|
||||||
|
# S.H.O.N.A.R. backend
|
||||||
|
|
||||||
|
FastAPI + PostgreSQL backend for the S.H.O.N.A.R. Android app.
|
||||||
|
|
||||||
|
See the repository root [README](../README.md) and `docs/` for full documentation.
|
||||||
|
|
||||||
|
## Quick start (development)
|
||||||
|
|
||||||
|
```bash
|
||||||
|
cd backend
|
||||||
|
uv venv .venv && uv pip install -e ".[dev]"
|
||||||
|
cp ../.env.example .env # then edit SHONAR_SECRET_KEY etc.
|
||||||
|
uvicorn shonar.main:app --reload --port 8000
|
||||||
|
```
|
||||||
|
|
||||||
|
Tests: `pytest` · Lint: `ruff check .` · Migrations: `alembic upgrade head`
|
||||||
37
backend/alembic.ini
Normal file
37
backend/alembic.ini
Normal file
|
|
@ -0,0 +1,37 @@
|
||||||
|
[alembic]
|
||||||
|
script_location = migrations
|
||||||
|
prepend_sys_path = .
|
||||||
|
# URL is injected from settings in env.py; this is a placeholder.
|
||||||
|
sqlalchemy.url =
|
||||||
|
|
||||||
|
[loggers]
|
||||||
|
keys = root,sqlalchemy,alembic
|
||||||
|
|
||||||
|
[handlers]
|
||||||
|
keys = console
|
||||||
|
|
||||||
|
[formatters]
|
||||||
|
keys = generic
|
||||||
|
|
||||||
|
[logger_root]
|
||||||
|
level = WARN
|
||||||
|
handlers = console
|
||||||
|
|
||||||
|
[logger_sqlalchemy]
|
||||||
|
level = WARN
|
||||||
|
handlers =
|
||||||
|
qualname = sqlalchemy.engine
|
||||||
|
|
||||||
|
[logger_alembic]
|
||||||
|
level = INFO
|
||||||
|
handlers =
|
||||||
|
qualname = alembic
|
||||||
|
|
||||||
|
[handler_console]
|
||||||
|
class = StreamHandler
|
||||||
|
args = (sys.stderr,)
|
||||||
|
level = NOTSET
|
||||||
|
formatter = generic
|
||||||
|
|
||||||
|
[formatter_generic]
|
||||||
|
format = %(levelname)-5.5s [%(name)s] %(message)s
|
||||||
62
backend/migrations/env.py
Normal file
62
backend/migrations/env.py
Normal file
|
|
@ -0,0 +1,62 @@
|
||||||
|
"""Alembic environment (async)."""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import asyncio
|
||||||
|
from logging.config import fileConfig
|
||||||
|
|
||||||
|
from alembic import context
|
||||||
|
from sqlalchemy import pool
|
||||||
|
from sqlalchemy.engine import Connection
|
||||||
|
from sqlalchemy.ext.asyncio import async_engine_from_config
|
||||||
|
|
||||||
|
from shonar.core.config import get_settings
|
||||||
|
from shonar.db import models # noqa: F401 (register tables)
|
||||||
|
from shonar.db.base import Base
|
||||||
|
|
||||||
|
config = context.config
|
||||||
|
if config.config_file_name is not None:
|
||||||
|
fileConfig(config.config_file_name)
|
||||||
|
|
||||||
|
settings = get_settings()
|
||||||
|
config.set_main_option("sqlalchemy.url", settings.database_url)
|
||||||
|
|
||||||
|
target_metadata = Base.metadata
|
||||||
|
|
||||||
|
|
||||||
|
def run_migrations_offline() -> None:
|
||||||
|
context.configure(
|
||||||
|
url=settings.database_url,
|
||||||
|
target_metadata=target_metadata,
|
||||||
|
literal_binds=True,
|
||||||
|
dialect_opts={"paramstyle": "named"},
|
||||||
|
)
|
||||||
|
with context.begin_transaction():
|
||||||
|
context.run_migrations()
|
||||||
|
|
||||||
|
|
||||||
|
def do_run_migrations(connection: Connection) -> None:
|
||||||
|
context.configure(connection=connection, target_metadata=target_metadata)
|
||||||
|
with context.begin_transaction():
|
||||||
|
context.run_migrations()
|
||||||
|
|
||||||
|
|
||||||
|
async def run_async_migrations() -> None:
|
||||||
|
connectable = async_engine_from_config(
|
||||||
|
config.get_section(config.config_ini_section, {}),
|
||||||
|
prefix="sqlalchemy.",
|
||||||
|
poolclass=pool.NullPool,
|
||||||
|
)
|
||||||
|
async with connectable.connect() as connection:
|
||||||
|
await connection.run_sync(do_run_migrations)
|
||||||
|
await connectable.dispose()
|
||||||
|
|
||||||
|
|
||||||
|
def run_migrations_online() -> None:
|
||||||
|
asyncio.run(run_async_migrations())
|
||||||
|
|
||||||
|
|
||||||
|
if context.is_offline_mode():
|
||||||
|
run_migrations_offline()
|
||||||
|
else:
|
||||||
|
run_migrations_online()
|
||||||
26
backend/migrations/script.py.mako
Normal file
26
backend/migrations/script.py.mako
Normal file
|
|
@ -0,0 +1,26 @@
|
||||||
|
"""${message}
|
||||||
|
|
||||||
|
Revision ID: ${up_revision}
|
||||||
|
Revises: ${down_revision | comma,n}
|
||||||
|
Create Date: ${create_date}
|
||||||
|
|
||||||
|
"""
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from collections.abc import Sequence
|
||||||
|
|
||||||
|
import sqlalchemy as sa
|
||||||
|
from alembic import op
|
||||||
|
${imports if imports else ""}
|
||||||
|
revision: str = ${repr(up_revision)}
|
||||||
|
down_revision: str | None = ${repr(down_revision)}
|
||||||
|
branch_labels: str | Sequence[str] | None = ${repr(branch_labels)}
|
||||||
|
depends_on: str | Sequence[str] | None = ${repr(depends_on)}
|
||||||
|
|
||||||
|
|
||||||
|
def upgrade() -> None:
|
||||||
|
${upgrades if upgrades else "pass"}
|
||||||
|
|
||||||
|
|
||||||
|
def downgrade() -> None:
|
||||||
|
${downgrades if downgrades else "pass"}
|
||||||
281
backend/migrations/versions/8d51af959ae4_initial_schema.py
Normal file
281
backend/migrations/versions/8d51af959ae4_initial_schema.py
Normal file
|
|
@ -0,0 +1,281 @@
|
||||||
|
"""initial schema
|
||||||
|
|
||||||
|
Revision ID: 8d51af959ae4
|
||||||
|
Revises:
|
||||||
|
Create Date: 2026-09-08 13:07:10.455837
|
||||||
|
|
||||||
|
"""
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from collections.abc import Sequence
|
||||||
|
|
||||||
|
import sqlalchemy as sa
|
||||||
|
from alembic import op
|
||||||
|
from sqlalchemy.dialects import postgresql
|
||||||
|
|
||||||
|
revision: str = '8d51af959ae4'
|
||||||
|
down_revision: str | None = None
|
||||||
|
branch_labels: str | Sequence[str] | None = None
|
||||||
|
depends_on: str | Sequence[str] | None = None
|
||||||
|
|
||||||
|
|
||||||
|
def _json_type() -> sa.types.TypeEngine:
|
||||||
|
"""JSONB on PostgreSQL, plain JSON on SQLite (desktop bundled engine)."""
|
||||||
|
if op.get_context().dialect.name == "postgresql":
|
||||||
|
return postgresql.JSONB(astext_type=sa.Text())
|
||||||
|
return sa.JSON()
|
||||||
|
|
||||||
|
|
||||||
|
def upgrade() -> None:
|
||||||
|
# ### commands auto generated by Alembic - please adjust! ###
|
||||||
|
op.create_table('users',
|
||||||
|
sa.Column('email', sa.String(length=320), nullable=False),
|
||||||
|
sa.Column('password_hash', sa.String(length=255), nullable=False),
|
||||||
|
sa.Column('display_name', sa.String(length=120), nullable=True),
|
||||||
|
sa.Column('is_active', sa.Boolean(), nullable=False),
|
||||||
|
sa.Column('location_storage_enabled', sa.Boolean(), nullable=False),
|
||||||
|
sa.Column('deleted_at', sa.DateTime(timezone=True), nullable=True),
|
||||||
|
sa.Column('id', sa.Uuid(), nullable=False),
|
||||||
|
sa.Column('created_at', sa.DateTime(timezone=True), nullable=False),
|
||||||
|
sa.Column('updated_at', sa.DateTime(timezone=True), nullable=False),
|
||||||
|
sa.PrimaryKeyConstraint('id')
|
||||||
|
)
|
||||||
|
op.create_index(op.f('ix_users_email'), 'users', ['email'], unique=True)
|
||||||
|
op.create_table('devices',
|
||||||
|
sa.Column('user_id', sa.Uuid(), nullable=False),
|
||||||
|
sa.Column('name', sa.String(length=120), nullable=False),
|
||||||
|
sa.Column('platform', sa.String(length=40), nullable=False),
|
||||||
|
sa.Column('last_seen_at', sa.DateTime(timezone=True), nullable=False),
|
||||||
|
sa.Column('revoked_at', sa.DateTime(timezone=True), nullable=True),
|
||||||
|
sa.Column('id', sa.Uuid(), nullable=False),
|
||||||
|
sa.Column('created_at', sa.DateTime(timezone=True), nullable=False),
|
||||||
|
sa.Column('updated_at', sa.DateTime(timezone=True), nullable=False),
|
||||||
|
sa.ForeignKeyConstraint(['user_id'], ['users.id'], ondelete='CASCADE'),
|
||||||
|
sa.PrimaryKeyConstraint('id')
|
||||||
|
)
|
||||||
|
op.create_index(op.f('ix_devices_user_id'), 'devices', ['user_id'], unique=False)
|
||||||
|
op.create_table('tags',
|
||||||
|
sa.Column('user_id', sa.Uuid(), nullable=False),
|
||||||
|
sa.Column('name', sa.String(length=80), nullable=False),
|
||||||
|
sa.Column('id', sa.Uuid(), nullable=False),
|
||||||
|
sa.Column('created_at', sa.DateTime(timezone=True), nullable=False),
|
||||||
|
sa.Column('updated_at', sa.DateTime(timezone=True), nullable=False),
|
||||||
|
sa.ForeignKeyConstraint(['user_id'], ['users.id'], ondelete='CASCADE'),
|
||||||
|
sa.PrimaryKeyConstraint('id'),
|
||||||
|
sa.UniqueConstraint('user_id', 'name')
|
||||||
|
)
|
||||||
|
op.create_index(op.f('ix_tags_user_id'), 'tags', ['user_id'], unique=False)
|
||||||
|
op.create_table('recordings',
|
||||||
|
sa.Column('user_id', sa.Uuid(), nullable=False),
|
||||||
|
sa.Column('device_id', sa.Uuid(), nullable=True),
|
||||||
|
sa.Column('client_recording_id', sa.String(length=64), nullable=True),
|
||||||
|
sa.Column('title', sa.String(length=300), nullable=False),
|
||||||
|
sa.Column('recorded_at', sa.DateTime(timezone=True), nullable=False),
|
||||||
|
sa.Column('duration_seconds', sa.Float(), nullable=False),
|
||||||
|
sa.Column('notes', sa.Text(), nullable=True),
|
||||||
|
sa.Column('latitude', sa.Float(), nullable=True),
|
||||||
|
sa.Column('longitude', sa.Float(), nullable=True),
|
||||||
|
sa.Column('location_accuracy_m', sa.Float(), nullable=True),
|
||||||
|
sa.Column('processing_status', sa.Enum('pending_upload', 'uploaded', 'queued', 'processing', 'completed', 'failed', 'ai_disabled', name='processing_status'), nullable=False),
|
||||||
|
sa.Column('processing_error', sa.Text(), nullable=True),
|
||||||
|
sa.Column('deleted_at', sa.DateTime(timezone=True), nullable=True),
|
||||||
|
sa.Column('id', sa.Uuid(), nullable=False),
|
||||||
|
sa.Column('created_at', sa.DateTime(timezone=True), nullable=False),
|
||||||
|
sa.Column('updated_at', sa.DateTime(timezone=True), nullable=False),
|
||||||
|
sa.ForeignKeyConstraint(['device_id'], ['devices.id'], ondelete='SET NULL'),
|
||||||
|
sa.ForeignKeyConstraint(['user_id'], ['users.id'], ondelete='CASCADE'),
|
||||||
|
sa.PrimaryKeyConstraint('id')
|
||||||
|
)
|
||||||
|
op.create_index(op.f('ix_recordings_client_recording_id'), 'recordings', ['client_recording_id'], unique=False)
|
||||||
|
op.create_index(op.f('ix_recordings_processing_status'), 'recordings', ['processing_status'], unique=False)
|
||||||
|
op.create_index('ix_recordings_user_client_id', 'recordings', ['user_id', 'client_recording_id'], unique=True, postgresql_where='client_recording_id IS NOT NULL', sqlite_where=sa.text('client_recording_id IS NOT NULL'))
|
||||||
|
op.create_index(op.f('ix_recordings_user_id'), 'recordings', ['user_id'], unique=False)
|
||||||
|
op.create_index('ix_recordings_user_recorded', 'recordings', ['user_id', 'recorded_at'], unique=False)
|
||||||
|
op.create_table('refresh_tokens',
|
||||||
|
sa.Column('user_id', sa.Uuid(), nullable=False),
|
||||||
|
sa.Column('token_hash', sa.String(length=128), nullable=False),
|
||||||
|
sa.Column('family', sa.Uuid(), nullable=False),
|
||||||
|
sa.Column('device_id', sa.Uuid(), nullable=True),
|
||||||
|
sa.Column('expires_at', sa.DateTime(timezone=True), nullable=False),
|
||||||
|
sa.Column('revoked_at', sa.DateTime(timezone=True), nullable=True),
|
||||||
|
sa.Column('replaced_by', sa.Uuid(), nullable=True),
|
||||||
|
sa.Column('id', sa.Uuid(), nullable=False),
|
||||||
|
sa.Column('created_at', sa.DateTime(timezone=True), nullable=False),
|
||||||
|
sa.Column('updated_at', sa.DateTime(timezone=True), nullable=False),
|
||||||
|
sa.ForeignKeyConstraint(['device_id'], ['devices.id'], ondelete='SET NULL'),
|
||||||
|
sa.ForeignKeyConstraint(['replaced_by'], ['refresh_tokens.id'], ondelete='SET NULL'),
|
||||||
|
sa.ForeignKeyConstraint(['user_id'], ['users.id'], ondelete='CASCADE'),
|
||||||
|
sa.PrimaryKeyConstraint('id'),
|
||||||
|
sa.UniqueConstraint('token_hash')
|
||||||
|
)
|
||||||
|
op.create_index(op.f('ix_refresh_tokens_family'), 'refresh_tokens', ['family'], unique=False)
|
||||||
|
op.create_index(op.f('ix_refresh_tokens_user_id'), 'refresh_tokens', ['user_id'], unique=False)
|
||||||
|
op.create_table('assets',
|
||||||
|
sa.Column('recording_id', sa.Uuid(), nullable=True),
|
||||||
|
sa.Column('user_id', sa.Uuid(), nullable=False),
|
||||||
|
sa.Column('kind', sa.Enum('original', 'normalized', 'export', name='asset_kind'), nullable=False),
|
||||||
|
sa.Column('storage_key', sa.String(length=500), nullable=False),
|
||||||
|
sa.Column('mime_type', sa.String(length=100), nullable=False),
|
||||||
|
sa.Column('size_bytes', sa.BigInteger(), nullable=False),
|
||||||
|
sa.Column('checksum_sha256', sa.String(length=64), nullable=False),
|
||||||
|
sa.Column('id', sa.Uuid(), nullable=False),
|
||||||
|
sa.Column('created_at', sa.DateTime(timezone=True), nullable=False),
|
||||||
|
sa.Column('updated_at', sa.DateTime(timezone=True), nullable=False),
|
||||||
|
sa.ForeignKeyConstraint(['recording_id'], ['recordings.id'], ondelete='CASCADE'),
|
||||||
|
sa.ForeignKeyConstraint(['user_id'], ['users.id'], ondelete='CASCADE'),
|
||||||
|
sa.PrimaryKeyConstraint('id')
|
||||||
|
)
|
||||||
|
op.create_index(op.f('ix_assets_recording_id'), 'assets', ['recording_id'], unique=False)
|
||||||
|
op.create_index(op.f('ix_assets_user_id'), 'assets', ['user_id'], unique=False)
|
||||||
|
op.create_index('uq_assets_one_original', 'assets', ['recording_id'], unique=True, postgresql_where="kind = 'original'", sqlite_where=sa.text("kind = 'original'"))
|
||||||
|
op.create_table('processing_jobs',
|
||||||
|
sa.Column('recording_id', sa.Uuid(), nullable=False),
|
||||||
|
sa.Column('job_type', sa.Enum('normalize_audio', 'transcribe', 'summarize', name='job_type'), nullable=False),
|
||||||
|
sa.Column('status', sa.Enum('queued', 'running', 'succeeded', 'failed', 'skipped', name='job_status'), nullable=False),
|
||||||
|
sa.Column('attempt', sa.Integer(), nullable=False),
|
||||||
|
sa.Column('max_attempts', sa.Integer(), nullable=False),
|
||||||
|
sa.Column('error', sa.Text(), nullable=True),
|
||||||
|
sa.Column('started_at', sa.DateTime(timezone=True), nullable=True),
|
||||||
|
sa.Column('finished_at', sa.DateTime(timezone=True), nullable=True),
|
||||||
|
sa.Column('task_handle', sa.String(length=120), nullable=True),
|
||||||
|
sa.Column('id', sa.Uuid(), nullable=False),
|
||||||
|
sa.Column('created_at', sa.DateTime(timezone=True), nullable=False),
|
||||||
|
sa.Column('updated_at', sa.DateTime(timezone=True), nullable=False),
|
||||||
|
sa.ForeignKeyConstraint(['recording_id'], ['recordings.id'], ondelete='CASCADE'),
|
||||||
|
sa.PrimaryKeyConstraint('id')
|
||||||
|
)
|
||||||
|
op.create_index(op.f('ix_processing_jobs_recording_id'), 'processing_jobs', ['recording_id'], unique=False)
|
||||||
|
op.create_index('ix_processing_jobs_recording_type', 'processing_jobs', ['recording_id', 'job_type'], unique=False)
|
||||||
|
op.create_index(op.f('ix_processing_jobs_status'), 'processing_jobs', ['status'], unique=False)
|
||||||
|
op.create_table('recording_tags',
|
||||||
|
sa.Column('recording_id', sa.Uuid(), nullable=False),
|
||||||
|
sa.Column('tag_id', sa.Uuid(), nullable=False),
|
||||||
|
sa.ForeignKeyConstraint(['recording_id'], ['recordings.id'], ondelete='CASCADE'),
|
||||||
|
sa.ForeignKeyConstraint(['tag_id'], ['tags.id'], ondelete='CASCADE'),
|
||||||
|
sa.PrimaryKeyConstraint('recording_id', 'tag_id')
|
||||||
|
)
|
||||||
|
op.create_table('summaries',
|
||||||
|
sa.Column('recording_id', sa.Uuid(), nullable=False),
|
||||||
|
sa.Column('version', sa.Integer(), nullable=False),
|
||||||
|
sa.Column('superseded_at', sa.DateTime(timezone=True), nullable=True),
|
||||||
|
sa.Column('provider', sa.String(length=80), nullable=False),
|
||||||
|
sa.Column('model', sa.String(length=120), nullable=True),
|
||||||
|
sa.Column('content', _json_type(), nullable=False),
|
||||||
|
sa.Column('edited_by_user', sa.Boolean(), nullable=False),
|
||||||
|
sa.Column('id', sa.Uuid(), nullable=False),
|
||||||
|
sa.Column('created_at', sa.DateTime(timezone=True), nullable=False),
|
||||||
|
sa.Column('updated_at', sa.DateTime(timezone=True), nullable=False),
|
||||||
|
sa.ForeignKeyConstraint(['recording_id'], ['recordings.id'], ondelete='CASCADE'),
|
||||||
|
sa.PrimaryKeyConstraint('id')
|
||||||
|
)
|
||||||
|
op.create_index(op.f('ix_summaries_recording_id'), 'summaries', ['recording_id'], unique=False)
|
||||||
|
op.create_table('transcripts',
|
||||||
|
sa.Column('recording_id', sa.Uuid(), nullable=False),
|
||||||
|
sa.Column('version', sa.Integer(), nullable=False),
|
||||||
|
sa.Column('superseded_at', sa.DateTime(timezone=True), nullable=True),
|
||||||
|
sa.Column('language', sa.String(length=16), nullable=True),
|
||||||
|
sa.Column('provider', sa.String(length=80), nullable=False),
|
||||||
|
sa.Column('model', sa.String(length=120), nullable=True),
|
||||||
|
sa.Column('text', sa.Text(), nullable=False),
|
||||||
|
sa.Column('segments', _json_type(), nullable=True),
|
||||||
|
sa.Column('edited_by_user', sa.Boolean(), nullable=False),
|
||||||
|
sa.Column('id', sa.Uuid(), nullable=False),
|
||||||
|
sa.Column('created_at', sa.DateTime(timezone=True), nullable=False),
|
||||||
|
sa.Column('updated_at', sa.DateTime(timezone=True), nullable=False),
|
||||||
|
sa.ForeignKeyConstraint(['recording_id'], ['recordings.id'], ondelete='CASCADE'),
|
||||||
|
sa.PrimaryKeyConstraint('id')
|
||||||
|
)
|
||||||
|
op.create_index(op.f('ix_transcripts_recording_id'), 'transcripts', ['recording_id'], unique=False)
|
||||||
|
op.create_index('ix_transcripts_recording_version', 'transcripts', ['recording_id', 'version'], unique=False)
|
||||||
|
op.create_table('export_jobs',
|
||||||
|
sa.Column('user_id', sa.Uuid(), nullable=False),
|
||||||
|
sa.Column('recording_id', sa.Uuid(), nullable=False),
|
||||||
|
sa.Column('export_type', sa.String(length=40), nullable=False),
|
||||||
|
sa.Column('status', sa.Enum('queued', 'running', 'succeeded', 'failed', 'skipped', name='export_job_status'), nullable=False),
|
||||||
|
sa.Column('asset_id', sa.Uuid(), nullable=True),
|
||||||
|
sa.Column('error', sa.Text(), nullable=True),
|
||||||
|
sa.Column('id', sa.Uuid(), nullable=False),
|
||||||
|
sa.Column('created_at', sa.DateTime(timezone=True), nullable=False),
|
||||||
|
sa.Column('updated_at', sa.DateTime(timezone=True), nullable=False),
|
||||||
|
sa.ForeignKeyConstraint(['asset_id'], ['assets.id'], ondelete='SET NULL'),
|
||||||
|
sa.ForeignKeyConstraint(['recording_id'], ['recordings.id'], ondelete='CASCADE'),
|
||||||
|
sa.ForeignKeyConstraint(['user_id'], ['users.id'], ondelete='CASCADE'),
|
||||||
|
sa.PrimaryKeyConstraint('id')
|
||||||
|
)
|
||||||
|
op.create_index(op.f('ix_export_jobs_recording_id'), 'export_jobs', ['recording_id'], unique=False)
|
||||||
|
op.create_index(op.f('ix_export_jobs_user_id'), 'export_jobs', ['user_id'], unique=False)
|
||||||
|
op.create_table('upload_sessions',
|
||||||
|
sa.Column('user_id', sa.Uuid(), nullable=False),
|
||||||
|
sa.Column('recording_id', sa.Uuid(), nullable=True),
|
||||||
|
sa.Column('client_recording_id', sa.String(length=64), nullable=True),
|
||||||
|
sa.Column('title', sa.String(length=300), nullable=True),
|
||||||
|
sa.Column('declared_mime_type', sa.String(length=100), nullable=False),
|
||||||
|
sa.Column('declared_size_bytes', sa.BigInteger(), nullable=False),
|
||||||
|
sa.Column('chunk_size_bytes', sa.BigInteger(), nullable=False),
|
||||||
|
sa.Column('status', sa.Enum('open', 'finalizing', 'completed', 'aborted', 'expired', name='upload_session_status'), nullable=False),
|
||||||
|
sa.Column('expires_at', sa.DateTime(timezone=True), nullable=False),
|
||||||
|
sa.Column('completed_asset_id', sa.Uuid(), nullable=True),
|
||||||
|
sa.Column('id', sa.Uuid(), nullable=False),
|
||||||
|
sa.Column('created_at', sa.DateTime(timezone=True), nullable=False),
|
||||||
|
sa.Column('updated_at', sa.DateTime(timezone=True), nullable=False),
|
||||||
|
sa.ForeignKeyConstraint(['completed_asset_id'], ['assets.id'], ondelete='SET NULL'),
|
||||||
|
sa.ForeignKeyConstraint(['recording_id'], ['recordings.id'], ondelete='CASCADE'),
|
||||||
|
sa.ForeignKeyConstraint(['user_id'], ['users.id'], ondelete='CASCADE'),
|
||||||
|
sa.PrimaryKeyConstraint('id')
|
||||||
|
)
|
||||||
|
op.create_index(op.f('ix_upload_sessions_user_id'), 'upload_sessions', ['user_id'], unique=False)
|
||||||
|
op.create_table('upload_chunks',
|
||||||
|
sa.Column('id', sa.Integer(), nullable=False),
|
||||||
|
sa.Column('session_id', sa.Uuid(), nullable=False),
|
||||||
|
sa.Column('chunk_index', sa.Integer(), nullable=False),
|
||||||
|
sa.Column('size_bytes', sa.BigInteger(), nullable=False),
|
||||||
|
sa.Column('checksum_sha256', sa.String(length=64), nullable=False),
|
||||||
|
sa.Column('storage_key', sa.String(length=500), nullable=False),
|
||||||
|
sa.Column('created_at', sa.DateTime(timezone=True), nullable=False),
|
||||||
|
sa.ForeignKeyConstraint(['session_id'], ['upload_sessions.id'], ondelete='CASCADE'),
|
||||||
|
sa.PrimaryKeyConstraint('id'),
|
||||||
|
sa.UniqueConstraint('session_id', 'chunk_index')
|
||||||
|
)
|
||||||
|
op.create_index(op.f('ix_upload_chunks_session_id'), 'upload_chunks', ['session_id'], unique=False)
|
||||||
|
# ### end Alembic commands ###
|
||||||
|
|
||||||
|
|
||||||
|
def downgrade() -> None:
|
||||||
|
# ### commands auto generated by Alembic - please adjust! ###
|
||||||
|
op.drop_index(op.f('ix_upload_chunks_session_id'), table_name='upload_chunks')
|
||||||
|
op.drop_table('upload_chunks')
|
||||||
|
op.drop_index(op.f('ix_upload_sessions_user_id'), table_name='upload_sessions')
|
||||||
|
op.drop_table('upload_sessions')
|
||||||
|
op.drop_index(op.f('ix_export_jobs_user_id'), table_name='export_jobs')
|
||||||
|
op.drop_index(op.f('ix_export_jobs_recording_id'), table_name='export_jobs')
|
||||||
|
op.drop_table('export_jobs')
|
||||||
|
op.drop_index('ix_transcripts_recording_version', table_name='transcripts')
|
||||||
|
op.drop_index(op.f('ix_transcripts_recording_id'), table_name='transcripts')
|
||||||
|
op.drop_table('transcripts')
|
||||||
|
op.drop_index(op.f('ix_summaries_recording_id'), table_name='summaries')
|
||||||
|
op.drop_table('summaries')
|
||||||
|
op.drop_table('recording_tags')
|
||||||
|
op.drop_index(op.f('ix_processing_jobs_status'), table_name='processing_jobs')
|
||||||
|
op.drop_index('ix_processing_jobs_recording_type', table_name='processing_jobs')
|
||||||
|
op.drop_index(op.f('ix_processing_jobs_recording_id'), table_name='processing_jobs')
|
||||||
|
op.drop_table('processing_jobs')
|
||||||
|
op.drop_index('uq_assets_one_original', table_name='assets', postgresql_where="kind = 'original'")
|
||||||
|
op.drop_index(op.f('ix_assets_user_id'), table_name='assets')
|
||||||
|
op.drop_index(op.f('ix_assets_recording_id'), table_name='assets')
|
||||||
|
op.drop_table('assets')
|
||||||
|
op.drop_index(op.f('ix_refresh_tokens_user_id'), table_name='refresh_tokens')
|
||||||
|
op.drop_index(op.f('ix_refresh_tokens_family'), table_name='refresh_tokens')
|
||||||
|
op.drop_table('refresh_tokens')
|
||||||
|
op.drop_index('ix_recordings_user_recorded', table_name='recordings')
|
||||||
|
op.drop_index(op.f('ix_recordings_user_id'), table_name='recordings')
|
||||||
|
op.drop_index('ix_recordings_user_client_id', table_name='recordings', postgresql_where='client_recording_id IS NOT NULL')
|
||||||
|
op.drop_index(op.f('ix_recordings_processing_status'), table_name='recordings')
|
||||||
|
op.drop_index(op.f('ix_recordings_client_recording_id'), table_name='recordings')
|
||||||
|
op.drop_table('recordings')
|
||||||
|
op.drop_index(op.f('ix_tags_user_id'), table_name='tags')
|
||||||
|
op.drop_table('tags')
|
||||||
|
op.drop_index(op.f('ix_devices_user_id'), table_name='devices')
|
||||||
|
op.drop_table('devices')
|
||||||
|
op.drop_index(op.f('ix_users_email'), table_name='users')
|
||||||
|
op.drop_table('users')
|
||||||
|
# ### end Alembic commands ###
|
||||||
86
backend/migrations/versions/fts0000000001_fts_columns.py
Normal file
86
backend/migrations/versions/fts0000000001_fts_columns.py
Normal file
|
|
@ -0,0 +1,86 @@
|
||||||
|
"""full-text search columns (PostgreSQL tsvector)
|
||||||
|
|
||||||
|
Revision ID: fts0000000001
|
||||||
|
Revises: 0c938663a363
|
||||||
|
|
||||||
|
Generated (STORED) tsvector columns + GIN indexes for search across title,
|
||||||
|
transcript text, summary content, tags, and action items. The SearchBackend
|
||||||
|
protocol keeps this swappable for Meilisearch/OpenSearch later.
|
||||||
|
"""
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from collections.abc import Sequence
|
||||||
|
|
||||||
|
from alembic import op
|
||||||
|
|
||||||
|
revision: str = "fts0000000001"
|
||||||
|
down_revision: str | None = "8d51af959ae4"
|
||||||
|
branch_labels: str | Sequence[str] | None = None
|
||||||
|
depends_on: str | Sequence[str] | None = None
|
||||||
|
|
||||||
|
|
||||||
|
def upgrade() -> None:
|
||||||
|
# tsvector is PostgreSQL-only. On SQLite (desktop bundled engine) full-
|
||||||
|
# text search is handled outside the DB (the app scans report files),
|
||||||
|
# so this migration is a no-op.
|
||||||
|
if op.get_context().dialect.name != "postgresql":
|
||||||
|
return
|
||||||
|
# recordings: title + notes + tag names are indexed; transcript/summary
|
||||||
|
# contribute via their own tables (joined at query time).
|
||||||
|
op.execute(
|
||||||
|
"""
|
||||||
|
ALTER TABLE recordings
|
||||||
|
ADD COLUMN search_vector tsvector
|
||||||
|
GENERATED ALWAYS AS (
|
||||||
|
setweight(to_tsvector('simple', coalesce(title, '')), 'A') ||
|
||||||
|
setweight(to_tsvector('simple', coalesce(notes, '')), 'B')
|
||||||
|
) STORED;
|
||||||
|
"""
|
||||||
|
)
|
||||||
|
op.execute("CREATE INDEX ix_recordings_search ON recordings USING GIN (search_vector);")
|
||||||
|
|
||||||
|
op.execute(
|
||||||
|
"""
|
||||||
|
ALTER TABLE transcripts
|
||||||
|
ADD COLUMN search_vector tsvector
|
||||||
|
GENERATED ALWAYS AS (
|
||||||
|
setweight(to_tsvector('simple', coalesce(text, '')), 'B')
|
||||||
|
) STORED;
|
||||||
|
"""
|
||||||
|
)
|
||||||
|
op.execute("CREATE INDEX ix_transcripts_search ON transcripts USING GIN (search_vector);")
|
||||||
|
|
||||||
|
# summary content JSON -> extractive text for FTS. Generated columns
|
||||||
|
# forbid subqueries/set-returning functions, so we index the JSON body
|
||||||
|
# with punctuation stripped (covers key_points/decisions/action_items/
|
||||||
|
# questions) plus weighted short/detailed fields.
|
||||||
|
op.execute(
|
||||||
|
"""
|
||||||
|
ALTER TABLE summaries
|
||||||
|
ADD COLUMN search_vector tsvector
|
||||||
|
GENERATED ALWAYS AS (
|
||||||
|
setweight(to_tsvector('simple', coalesce(content->>'short', '')), 'A') ||
|
||||||
|
setweight(to_tsvector('simple', coalesce(content->>'detailed', '')), 'B') ||
|
||||||
|
setweight(
|
||||||
|
to_tsvector('simple',
|
||||||
|
regexp_replace(coalesce(content::text, ''), '[\\[\\]{}"]', ' ', 'g')), 'C')
|
||||||
|
) STORED;
|
||||||
|
"""
|
||||||
|
)
|
||||||
|
op.execute("CREATE INDEX ix_summaries_search ON summaries USING GIN (search_vector);")
|
||||||
|
|
||||||
|
op.execute(
|
||||||
|
"""
|
||||||
|
ALTER TABLE tags ADD COLUMN IF NOT EXISTS search_vector tsvector
|
||||||
|
GENERATED ALWAYS AS (to_tsvector('simple', coalesce(name, ''))) STORED;
|
||||||
|
"""
|
||||||
|
)
|
||||||
|
op.execute("CREATE INDEX ix_tags_search ON tags USING GIN (search_vector);")
|
||||||
|
|
||||||
|
|
||||||
|
def downgrade() -> None:
|
||||||
|
if op.get_context().dialect.name != "postgresql":
|
||||||
|
return
|
||||||
|
for table in ("tags", "summaries", "transcripts", "recordings"):
|
||||||
|
op.execute(f"DROP INDEX IF EXISTS ix_{table}_search;")
|
||||||
|
op.execute(f"ALTER TABLE {table} DROP COLUMN IF EXISTS search_vector;")
|
||||||
69
backend/migrations/versions/m8models000001_model_support.py
Normal file
69
backend/migrations/versions/m8models000001_model_support.py
Normal file
|
|
@ -0,0 +1,69 @@
|
||||||
|
"""per-recording transcription model + job progress + app settings
|
||||||
|
|
||||||
|
Revision ID: m8models000001
|
||||||
|
Revises: fts0000000001
|
||||||
|
|
||||||
|
- recordings.transcription_model (nullable; NULL = server default)
|
||||||
|
- processing_jobs.stage / processing_jobs.progress (display-only)
|
||||||
|
- app_settings table (global default model, …)
|
||||||
|
"""
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from collections.abc import Sequence
|
||||||
|
|
||||||
|
import sqlalchemy as sa
|
||||||
|
from alembic import op
|
||||||
|
from sqlalchemy.dialects import postgresql
|
||||||
|
|
||||||
|
revision: str = 'm8models000001'
|
||||||
|
down_revision: str | None = 'fts0000000001'
|
||||||
|
branch_labels: str | Sequence[str] | None = None
|
||||||
|
depends_on: str | Sequence[str] | None = None
|
||||||
|
|
||||||
|
|
||||||
|
def _json_type() -> sa.types.TypeEngine:
|
||||||
|
"""JSONB on PostgreSQL, plain JSON on SQLite (desktop bundled engine)."""
|
||||||
|
if op.get_context().dialect.name == "postgresql":
|
||||||
|
return postgresql.JSONB(astext_type=sa.Text())
|
||||||
|
return sa.JSON()
|
||||||
|
|
||||||
|
|
||||||
|
def upgrade() -> None:
|
||||||
|
op.add_column(
|
||||||
|
"recordings",
|
||||||
|
sa.Column("transcription_model", sa.String(length=32), nullable=True),
|
||||||
|
)
|
||||||
|
op.add_column(
|
||||||
|
"upload_sessions",
|
||||||
|
sa.Column("transcription_model", sa.String(length=32), nullable=True),
|
||||||
|
)
|
||||||
|
op.add_column(
|
||||||
|
"processing_jobs",
|
||||||
|
sa.Column("stage", sa.String(length=32), nullable=True),
|
||||||
|
)
|
||||||
|
op.add_column(
|
||||||
|
"processing_jobs",
|
||||||
|
sa.Column("progress", sa.Integer(), nullable=True),
|
||||||
|
)
|
||||||
|
op.create_table(
|
||||||
|
"app_settings",
|
||||||
|
sa.Column("key", sa.String(length=120), nullable=False),
|
||||||
|
sa.Column(
|
||||||
|
"value",
|
||||||
|
_json_type(),
|
||||||
|
nullable=False,
|
||||||
|
),
|
||||||
|
sa.Column("id", sa.Uuid(), nullable=False),
|
||||||
|
sa.Column("created_at", sa.DateTime(timezone=True), nullable=False),
|
||||||
|
sa.Column("updated_at", sa.DateTime(timezone=True), nullable=False),
|
||||||
|
sa.PrimaryKeyConstraint("id"),
|
||||||
|
sa.UniqueConstraint("key"),
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def downgrade() -> None:
|
||||||
|
op.drop_table("app_settings")
|
||||||
|
op.drop_column("processing_jobs", "progress")
|
||||||
|
op.drop_column("processing_jobs", "stage")
|
||||||
|
op.drop_column("upload_sessions", "transcription_model")
|
||||||
|
op.drop_column("recordings", "transcription_model")
|
||||||
64
backend/pyproject.toml
Normal file
64
backend/pyproject.toml
Normal file
|
|
@ -0,0 +1,64 @@
|
||||||
|
[project]
|
||||||
|
name = "shonar-backend"
|
||||||
|
version = "0.1.0"
|
||||||
|
description = "S.H.O.N.A.R. — Self-hosted Oral Notes and Audio Recorder backend"
|
||||||
|
readme = "README.md"
|
||||||
|
requires-python = ">=3.11"
|
||||||
|
license = { text = "Apache-2.0" }
|
||||||
|
dependencies = [
|
||||||
|
"fastapi>=0.115",
|
||||||
|
"uvicorn[standard]>=0.30",
|
||||||
|
"sqlalchemy[asyncio]>=2.0",
|
||||||
|
"asyncpg>=0.29",
|
||||||
|
"aiosqlite>=0.20", # tests only in practice; harmless dep
|
||||||
|
"alembic>=1.13",
|
||||||
|
"pydantic>=2.8",
|
||||||
|
"pydantic-settings>=2.4",
|
||||||
|
"email-validator>=2.0",
|
||||||
|
"argon2-cffi>=23.1",
|
||||||
|
"pyjwt>=2.9",
|
||||||
|
"python-multipart>=0.0.9",
|
||||||
|
"slowapi>=0.1.9",
|
||||||
|
"redis>=5.0",
|
||||||
|
"arq>=0.26",
|
||||||
|
]
|
||||||
|
|
||||||
|
[project.optional-dependencies]
|
||||||
|
s3 = ["boto3>=1.34"]
|
||||||
|
faster-whisper = ["faster-whisper>=1.0"]
|
||||||
|
dev = [
|
||||||
|
"pytest>=8.0",
|
||||||
|
"pytest-asyncio>=0.23",
|
||||||
|
"httpx>=0.27",
|
||||||
|
"ruff>=0.5",
|
||||||
|
"mypy>=1.10",
|
||||||
|
]
|
||||||
|
|
||||||
|
[build-system]
|
||||||
|
requires = ["hatchling"]
|
||||||
|
build-backend = "hatchling.build"
|
||||||
|
|
||||||
|
[tool.hatch.build.targets.wheel]
|
||||||
|
packages = ["shonar"]
|
||||||
|
|
||||||
|
[tool.pytest.ini_options]
|
||||||
|
asyncio_mode = "auto"
|
||||||
|
asyncio_default_fixture_loop_scope = "session"
|
||||||
|
asyncio_default_test_loop_scope = "session"
|
||||||
|
testpaths = ["tests"]
|
||||||
|
|
||||||
|
[tool.ruff]
|
||||||
|
line-length = 100
|
||||||
|
target-version = "py311"
|
||||||
|
exclude = ["migrations"]
|
||||||
|
|
||||||
|
[tool.ruff.lint]
|
||||||
|
select = ["E", "F", "I", "UP", "B", "SIM"]
|
||||||
|
|
||||||
|
[tool.ruff.lint.per-file-ignores]
|
||||||
|
# FastAPI idiom: Query(...)/Depends() as parameter defaults.
|
||||||
|
"shonar/api/**" = ["B008"]
|
||||||
|
|
||||||
|
[tool.mypy]
|
||||||
|
python_version = "3.11"
|
||||||
|
ignore_missing_imports = true
|
||||||
3
backend/shonar/__init__.py
Normal file
3
backend/shonar/__init__.py
Normal file
|
|
@ -0,0 +1,3 @@
|
||||||
|
"""S.H.O.N.A.R. — Self-hosted Oral Notes and Audio Recorder."""
|
||||||
|
|
||||||
|
__version__ = "0.1.0"
|
||||||
0
backend/shonar/api/__init__.py
Normal file
0
backend/shonar/api/__init__.py
Normal file
49
backend/shonar/api/deps.py
Normal file
49
backend/shonar/api/deps.py
Normal file
|
|
@ -0,0 +1,49 @@
|
||||||
|
"""FastAPI dependencies: current user, session, rate limiting."""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import uuid
|
||||||
|
from typing import Annotated
|
||||||
|
|
||||||
|
from fastapi import Depends, HTTPException, Request, status
|
||||||
|
from fastapi.security import HTTPAuthorizationCredentials, HTTPBearer
|
||||||
|
from sqlalchemy.ext.asyncio import AsyncSession
|
||||||
|
|
||||||
|
from shonar.core.security import TokenError, decode_access_token
|
||||||
|
from shonar.db.models import User
|
||||||
|
from shonar.db.session import get_session
|
||||||
|
|
||||||
|
SessionDep = Annotated[AsyncSession, Depends(get_session)]
|
||||||
|
|
||||||
|
_bearer = HTTPBearer(auto_error=False)
|
||||||
|
|
||||||
|
|
||||||
|
async def get_current_user(
|
||||||
|
request: Request,
|
||||||
|
session: SessionDep,
|
||||||
|
creds: Annotated[HTTPAuthorizationCredentials | None, Depends(_bearer)] = None,
|
||||||
|
) -> User:
|
||||||
|
if creds is None or not creds.credentials:
|
||||||
|
raise HTTPException(
|
||||||
|
status_code=status.HTTP_401_UNAUTHORIZED,
|
||||||
|
detail="Not authenticated",
|
||||||
|
headers={"WWW-Authenticate": "Bearer"},
|
||||||
|
) from None
|
||||||
|
try:
|
||||||
|
payload = decode_access_token(creds.credentials)
|
||||||
|
except TokenError:
|
||||||
|
raise HTTPException(
|
||||||
|
status_code=status.HTTP_401_UNAUTHORIZED,
|
||||||
|
detail="Invalid or expired token",
|
||||||
|
headers={"WWW-Authenticate": "Bearer"},
|
||||||
|
) from None
|
||||||
|
user = await session.get(User, uuid.UUID(payload["sub"]))
|
||||||
|
if user is None or not user.is_active or user.deleted_at is not None:
|
||||||
|
raise HTTPException(
|
||||||
|
status_code=status.HTTP_401_UNAUTHORIZED, detail="Account unavailable"
|
||||||
|
) from None
|
||||||
|
request.state.user_id = user.id
|
||||||
|
return user
|
||||||
|
|
||||||
|
|
||||||
|
CurrentUser = Annotated[User, Depends(get_current_user)]
|
||||||
69
backend/shonar/api/schemas_common.py
Normal file
69
backend/shonar/api/schemas_common.py
Normal file
|
|
@ -0,0 +1,69 @@
|
||||||
|
"""Shared Pydantic schemas."""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import uuid
|
||||||
|
from datetime import datetime
|
||||||
|
|
||||||
|
from pydantic import BaseModel, ConfigDict, EmailStr, Field
|
||||||
|
|
||||||
|
|
||||||
|
class ORMModel(BaseModel):
|
||||||
|
model_config = ConfigDict(from_attributes=True)
|
||||||
|
|
||||||
|
|
||||||
|
# --- auth ---
|
||||||
|
|
||||||
|
|
||||||
|
class RegisterRequest(BaseModel):
|
||||||
|
email: EmailStr
|
||||||
|
password: str = Field(min_length=10, max_length=256)
|
||||||
|
display_name: str | None = Field(default=None, max_length=120)
|
||||||
|
|
||||||
|
|
||||||
|
class LoginRequest(BaseModel):
|
||||||
|
email: EmailStr
|
||||||
|
password: str
|
||||||
|
device_name: str | None = Field(default=None, max_length=120)
|
||||||
|
platform: str = Field(default="android", max_length=40)
|
||||||
|
|
||||||
|
|
||||||
|
class TokenPair(BaseModel):
|
||||||
|
access_token: str
|
||||||
|
token_type: str = "bearer"
|
||||||
|
expires_in: int
|
||||||
|
refresh_token: str
|
||||||
|
device_id: uuid.UUID | None = None
|
||||||
|
|
||||||
|
|
||||||
|
class RefreshRequest(BaseModel):
|
||||||
|
refresh_token: str
|
||||||
|
|
||||||
|
|
||||||
|
class LogoutRequest(BaseModel):
|
||||||
|
refresh_token: str
|
||||||
|
|
||||||
|
|
||||||
|
class UserOut(ORMModel):
|
||||||
|
id: uuid.UUID
|
||||||
|
email: EmailStr
|
||||||
|
display_name: str | None
|
||||||
|
location_storage_enabled: bool
|
||||||
|
created_at: datetime
|
||||||
|
|
||||||
|
|
||||||
|
class UserUpdate(BaseModel):
|
||||||
|
display_name: str | None = Field(default=None, max_length=120)
|
||||||
|
location_storage_enabled: bool | None = None
|
||||||
|
|
||||||
|
|
||||||
|
class DeviceOut(ORMModel):
|
||||||
|
id: uuid.UUID
|
||||||
|
name: str
|
||||||
|
platform: str
|
||||||
|
last_seen_at: datetime
|
||||||
|
revoked_at: datetime | None
|
||||||
|
|
||||||
|
|
||||||
|
class DeleteAccountRequest(BaseModel):
|
||||||
|
password: str
|
||||||
149
backend/shonar/api/schemas_recordings.py
Normal file
149
backend/shonar/api/schemas_recordings.py
Normal file
|
|
@ -0,0 +1,149 @@
|
||||||
|
"""Recording / upload schemas."""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import uuid
|
||||||
|
from datetime import datetime
|
||||||
|
|
||||||
|
from pydantic import BaseModel, ConfigDict, Field
|
||||||
|
|
||||||
|
|
||||||
|
class ORMModel(BaseModel):
|
||||||
|
model_config = ConfigDict(from_attributes=True)
|
||||||
|
|
||||||
|
|
||||||
|
class UploadSessionCreate(BaseModel):
|
||||||
|
declared_mime_type: str = Field(max_length=100)
|
||||||
|
declared_size_bytes: int = Field(gt=0)
|
||||||
|
title: str | None = Field(default=None, max_length=300)
|
||||||
|
client_recording_id: str | None = Field(default=None, max_length=64)
|
||||||
|
# Optional per-recording transcription model override ("Use default"
|
||||||
|
# when omitted). Validated against the model registry at finalize time.
|
||||||
|
transcription_model: str | None = Field(default=None, max_length=32)
|
||||||
|
|
||||||
|
|
||||||
|
class UploadSessionOut(ORMModel):
|
||||||
|
id: uuid.UUID
|
||||||
|
status: str
|
||||||
|
chunk_size_bytes: int
|
||||||
|
declared_mime_type: str
|
||||||
|
declared_size_bytes: int
|
||||||
|
expires_at: datetime
|
||||||
|
|
||||||
|
|
||||||
|
class UploadStatusOut(BaseModel):
|
||||||
|
id: uuid.UUID
|
||||||
|
status: str
|
||||||
|
received_chunk_indexes: list[int]
|
||||||
|
|
||||||
|
|
||||||
|
class RecordingFinalize(BaseModel):
|
||||||
|
recorded_at: datetime | None = None
|
||||||
|
duration_seconds: float = Field(default=0.0, ge=0)
|
||||||
|
# Location accepted ONLY if the user has location storage enabled
|
||||||
|
# (enforced in the route; silently dropped otherwise).
|
||||||
|
latitude: float | None = Field(default=None, ge=-90, le=90)
|
||||||
|
longitude: float | None = Field(default=None, ge=-180, le=180)
|
||||||
|
location_accuracy_m: float | None = Field(default=None, ge=0)
|
||||||
|
notes: str | None = Field(default=None, max_length=20000)
|
||||||
|
# Optional per-recording transcription model override. When both the
|
||||||
|
# session and the finalize body specify one, the finalize body wins.
|
||||||
|
transcription_model: str | None = Field(default=None, max_length=32)
|
||||||
|
|
||||||
|
|
||||||
|
class RecordingOut(ORMModel):
|
||||||
|
id: uuid.UUID
|
||||||
|
title: str
|
||||||
|
recorded_at: datetime
|
||||||
|
duration_seconds: float
|
||||||
|
notes: str | None
|
||||||
|
latitude: float | None
|
||||||
|
longitude: float | None
|
||||||
|
processing_status: str
|
||||||
|
processing_error: str | None
|
||||||
|
# Effective transcription model saved at finalize time (override or the
|
||||||
|
# global default then in force). None for pre-model rows.
|
||||||
|
transcription_model: str | None
|
||||||
|
has_audio: bool
|
||||||
|
mime_type: str | None
|
||||||
|
size_bytes: int | None
|
||||||
|
tags: list[str]
|
||||||
|
created_at: datetime
|
||||||
|
updated_at: datetime
|
||||||
|
|
||||||
|
|
||||||
|
class RecordingUpdate(BaseModel):
|
||||||
|
title: str | None = Field(default=None, min_length=1, max_length=300)
|
||||||
|
notes: str | None = Field(default=None, max_length=20000)
|
||||||
|
tags: list[str] | None = Field(default=None, max_length=20)
|
||||||
|
latitude: float | None = Field(default=None, ge=-90, le=90)
|
||||||
|
longitude: float | None = Field(default=None, ge=-180, le=180)
|
||||||
|
|
||||||
|
|
||||||
|
class RecordingListOut(BaseModel):
|
||||||
|
items: list[RecordingOut]
|
||||||
|
total: int
|
||||||
|
limit: int
|
||||||
|
offset: int
|
||||||
|
|
||||||
|
|
||||||
|
class SegmentOut(BaseModel):
|
||||||
|
start: float
|
||||||
|
end: float
|
||||||
|
text: str
|
||||||
|
speaker: str | None = None
|
||||||
|
|
||||||
|
|
||||||
|
class TranscriptOut(ORMModel):
|
||||||
|
version: int
|
||||||
|
language: str | None
|
||||||
|
provider: str
|
||||||
|
model: str | None
|
||||||
|
text: str
|
||||||
|
segments: list[SegmentOut]
|
||||||
|
edited_by_user: bool
|
||||||
|
created_at: datetime
|
||||||
|
updated_at: datetime
|
||||||
|
|
||||||
|
|
||||||
|
class SummaryOut(ORMModel):
|
||||||
|
version: int
|
||||||
|
provider: str
|
||||||
|
model: str | None
|
||||||
|
content: dict
|
||||||
|
edited_by_user: bool
|
||||||
|
created_at: datetime
|
||||||
|
updated_at: datetime
|
||||||
|
|
||||||
|
|
||||||
|
class ProcessingJobOut(ORMModel):
|
||||||
|
job_type: str
|
||||||
|
status: str
|
||||||
|
attempt: int
|
||||||
|
max_attempts: int
|
||||||
|
error: str | None
|
||||||
|
# Display-only phase within a running job ("loading-model",
|
||||||
|
# "transcribing") plus 0-100 progress when known.
|
||||||
|
stage: str | None
|
||||||
|
progress: int | None
|
||||||
|
started_at: datetime | None
|
||||||
|
finished_at: datetime | None
|
||||||
|
|
||||||
|
|
||||||
|
class TranscriptUpdate(BaseModel):
|
||||||
|
"""User edit: stored as a new version with edited_by_user=True.
|
||||||
|
|
||||||
|
The AI pipeline never overwrites a newest user-edited row, so edits
|
||||||
|
are verdicts. Segments are optional — when omitted the previous
|
||||||
|
segments are dropped (plain-text correction).
|
||||||
|
"""
|
||||||
|
|
||||||
|
text: str = Field(min_length=1, max_length=200000)
|
||||||
|
segments: list[SegmentOut] | None = Field(default=None, max_length=2000)
|
||||||
|
language: str | None = Field(default=None, max_length=16)
|
||||||
|
|
||||||
|
|
||||||
|
class SummaryUpdate(BaseModel):
|
||||||
|
"""User edit for the structured summary (same versioning rule)."""
|
||||||
|
|
||||||
|
content: dict = Field(max_length=100)
|
||||||
23
backend/shonar/api/schemas_search.py
Normal file
23
backend/shonar/api/schemas_search.py
Normal file
|
|
@ -0,0 +1,23 @@
|
||||||
|
"""Search + export schemas (M9)."""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import uuid
|
||||||
|
|
||||||
|
from pydantic import BaseModel
|
||||||
|
|
||||||
|
|
||||||
|
class SearchHitOut(BaseModel):
|
||||||
|
id: uuid.UUID
|
||||||
|
title: str
|
||||||
|
field: str # title | tag | notes | summary | transcript
|
||||||
|
snippet: str
|
||||||
|
|
||||||
|
|
||||||
|
class SearchOut(BaseModel):
|
||||||
|
query: str
|
||||||
|
scope: str
|
||||||
|
items: list[SearchHitOut]
|
||||||
|
total: int
|
||||||
|
limit: int
|
||||||
|
offset: int
|
||||||
25
backend/shonar/api/v1/__init__.py
Normal file
25
backend/shonar/api/v1/__init__.py
Normal file
|
|
@ -0,0 +1,25 @@
|
||||||
|
"""API v1 router aggregation."""
|
||||||
|
|
||||||
|
from fastapi import APIRouter
|
||||||
|
|
||||||
|
from shonar.api.v1 import (
|
||||||
|
auth,
|
||||||
|
health,
|
||||||
|
models,
|
||||||
|
provider_info,
|
||||||
|
recordings,
|
||||||
|
search_exports,
|
||||||
|
users,
|
||||||
|
)
|
||||||
|
|
||||||
|
api_router = APIRouter(prefix="/api/v1")
|
||||||
|
api_router.include_router(health.router)
|
||||||
|
api_router.include_router(auth.router)
|
||||||
|
api_router.include_router(users.router)
|
||||||
|
api_router.include_router(recordings.router)
|
||||||
|
api_router.include_router(models.router)
|
||||||
|
api_router.include_router(provider_info.router)
|
||||||
|
api_router.include_router(search_exports.router)
|
||||||
|
|
||||||
|
# Included as later milestones land:
|
||||||
|
# - tags (M9 follow-on)
|
||||||
91
backend/shonar/api/v1/auth.py
Normal file
91
backend/shonar/api/v1/auth.py
Normal file
|
|
@ -0,0 +1,91 @@
|
||||||
|
"""Auth endpoints (rate-limited)."""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from fastapi import APIRouter, HTTPException, Request
|
||||||
|
|
||||||
|
from shonar.api.deps import CurrentUser, SessionDep
|
||||||
|
from shonar.api.schemas_common import (
|
||||||
|
DeleteAccountRequest,
|
||||||
|
LoginRequest,
|
||||||
|
LogoutRequest,
|
||||||
|
RefreshRequest,
|
||||||
|
RegisterRequest,
|
||||||
|
TokenPair,
|
||||||
|
UserOut,
|
||||||
|
)
|
||||||
|
from shonar.core.config import get_settings
|
||||||
|
from shonar.core.ratelimit import auth_limit
|
||||||
|
from shonar.core.security import create_access_token
|
||||||
|
from shonar.services import auth as auth_service
|
||||||
|
from shonar.services.auth import AuthError
|
||||||
|
|
||||||
|
router = APIRouter(tags=["auth"])
|
||||||
|
|
||||||
|
|
||||||
|
@router.post("/auth/register", response_model=TokenPair, status_code=201)
|
||||||
|
@auth_limit
|
||||||
|
async def register(body: RegisterRequest, request: Request, session: SessionDep):
|
||||||
|
settings = get_settings()
|
||||||
|
if not settings.allow_registration:
|
||||||
|
raise HTTPException(403, "Registration is disabled on this server.") from None
|
||||||
|
try:
|
||||||
|
user = await auth_service.register_user(
|
||||||
|
session, body.email, body.password, body.display_name
|
||||||
|
)
|
||||||
|
except AuthError as e:
|
||||||
|
raise HTTPException(e.status_code, e.message) from None
|
||||||
|
access, ttl = create_access_token(user.id)
|
||||||
|
from shonar.services.auth import issue_refresh_token
|
||||||
|
|
||||||
|
refresh, _ = await issue_refresh_token(session, user.id, None, None)
|
||||||
|
return TokenPair(access_token=access, expires_in=ttl, refresh_token=refresh)
|
||||||
|
|
||||||
|
|
||||||
|
@router.post("/auth/login", response_model=TokenPair)
|
||||||
|
@auth_limit
|
||||||
|
async def login(body: LoginRequest, request: Request, session: SessionDep):
|
||||||
|
try:
|
||||||
|
user, refresh, device = await auth_service.login(
|
||||||
|
session, body.email, body.password, body.device_name, body.platform
|
||||||
|
)
|
||||||
|
except AuthError as e:
|
||||||
|
raise HTTPException(e.status_code, e.message) from None
|
||||||
|
access, ttl = create_access_token(user.id, device.id)
|
||||||
|
return TokenPair(
|
||||||
|
access_token=access, expires_in=ttl, refresh_token=refresh, device_id=device.id
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@router.post("/auth/refresh", response_model=TokenPair)
|
||||||
|
@auth_limit
|
||||||
|
async def refresh(body: RefreshRequest, request: Request, session: SessionDep):
|
||||||
|
try:
|
||||||
|
user, new_refresh, device_id = await auth_service.rotate_refresh_token(
|
||||||
|
session, body.refresh_token
|
||||||
|
)
|
||||||
|
except AuthError as e:
|
||||||
|
raise HTTPException(e.status_code, e.message) from None
|
||||||
|
access, ttl = create_access_token(user.id, device_id)
|
||||||
|
return TokenPair(
|
||||||
|
access_token=access, expires_in=ttl, refresh_token=new_refresh, device_id=device_id
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@router.post("/auth/logout", status_code=204)
|
||||||
|
async def logout(body: LogoutRequest, session: SessionDep):
|
||||||
|
await auth_service.logout(session, body.refresh_token)
|
||||||
|
|
||||||
|
|
||||||
|
@router.get("/auth/me", response_model=UserOut)
|
||||||
|
async def me(user: CurrentUser):
|
||||||
|
return user
|
||||||
|
|
||||||
|
|
||||||
|
@router.post("/auth/delete-account", status_code=202)
|
||||||
|
async def delete_account(body: DeleteAccountRequest, user: CurrentUser, session: SessionDep):
|
||||||
|
try:
|
||||||
|
await auth_service.delete_account(session, user, body.password)
|
||||||
|
except AuthError as e:
|
||||||
|
raise HTTPException(e.status_code, e.message) from None
|
||||||
|
return {"detail": "Account scheduled for deletion.", "grace_days": 30}
|
||||||
65
backend/shonar/api/v1/health.py
Normal file
65
backend/shonar/api/v1/health.py
Normal file
|
|
@ -0,0 +1,65 @@
|
||||||
|
"""Health and system endpoints."""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import time
|
||||||
|
|
||||||
|
from fastapi import APIRouter
|
||||||
|
from sqlalchemy import text
|
||||||
|
|
||||||
|
from shonar.api.deps import SessionDep
|
||||||
|
from shonar.core.config import get_settings
|
||||||
|
|
||||||
|
router = APIRouter(tags=["health"])
|
||||||
|
|
||||||
|
_STARTED = time.monotonic()
|
||||||
|
|
||||||
|
|
||||||
|
@router.get("/healthz")
|
||||||
|
async def healthz() -> dict:
|
||||||
|
"""Liveness: process up. No auth, no dependencies."""
|
||||||
|
return {"status": "ok", "uptime_seconds": round(time.monotonic() - _STARTED, 1)}
|
||||||
|
|
||||||
|
|
||||||
|
@router.get("/readyz")
|
||||||
|
async def readyz(session: SessionDep) -> dict:
|
||||||
|
"""Readiness: database reachable. Config warnings surfaced for admins
|
||||||
|
via /api/v1/system/status instead of failing readiness."""
|
||||||
|
try:
|
||||||
|
await session.execute(text("SELECT 1"))
|
||||||
|
db_ok = True
|
||||||
|
except Exception:
|
||||||
|
db_ok = False
|
||||||
|
status_code_body = {"status": "ok" if db_ok else "degraded", "database": db_ok}
|
||||||
|
return status_code_body
|
||||||
|
|
||||||
|
|
||||||
|
@router.get("/system/status")
|
||||||
|
async def system_status() -> dict:
|
||||||
|
"""Public-ish status: which AI features are enabled (never any secrets).
|
||||||
|
The app uses this to show honest AI-processing state to users."""
|
||||||
|
settings = get_settings()
|
||||||
|
return {
|
||||||
|
"app": settings.app_name,
|
||||||
|
"registration_enabled": settings.allow_registration,
|
||||||
|
"ai": {
|
||||||
|
"transcription_provider": settings.transcription_provider,
|
||||||
|
"transcription_enabled": settings.transcription_provider != "none",
|
||||||
|
"llm_provider": settings.llm_provider,
|
||||||
|
"llm_enabled": settings.llm_provider != "none",
|
||||||
|
# Explicit disclosure: are any external (non-local) calls made?
|
||||||
|
"external_ai_in_use": (
|
||||||
|
settings.transcription_provider == "whisper_http"
|
||||||
|
and "localhost" not in settings.transcription_base_url
|
||||||
|
and "127.0.0.1" not in settings.transcription_base_url
|
||||||
|
)
|
||||||
|
or (
|
||||||
|
settings.llm_provider == "openai_compat"
|
||||||
|
and "localhost" not in settings.llm_base_url
|
||||||
|
and "127.0.0.1" not in settings.llm_base_url
|
||||||
|
),
|
||||||
|
},
|
||||||
|
"storage_backend": settings.storage_backend,
|
||||||
|
"audio_conversion_enabled": settings.audio_conversion_enabled,
|
||||||
|
"config_warnings": settings.validate_production(),
|
||||||
|
}
|
||||||
132
backend/shonar/api/v1/models.py
Normal file
132
backend/shonar/api/v1/models.py
Normal file
|
|
@ -0,0 +1,132 @@
|
||||||
|
"""Transcription model registry + global default (Stage 1).
|
||||||
|
|
||||||
|
- GET /models — supported models with display metadata, download and
|
||||||
|
availability status, and which is the default.
|
||||||
|
- GET /models/default — the current global default model name.
|
||||||
|
- PUT /models/default — change the global default (affects future
|
||||||
|
recordings only; saved per-recording models are never rewritten).
|
||||||
|
- POST /models/{name}/download — fetch a model into the local cache
|
||||||
|
(needs internet once; runs synchronously and may take minutes for
|
||||||
|
large models).
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import asyncio
|
||||||
|
|
||||||
|
from fastapi import APIRouter, HTTPException
|
||||||
|
from pydantic import BaseModel, Field
|
||||||
|
|
||||||
|
from shonar.api.deps import CurrentUser, SessionDep
|
||||||
|
from shonar.services.ai import ProviderConfigError
|
||||||
|
from shonar.services.ai.model_registry import (
|
||||||
|
SUPPORTED_TRANSCRIPTION_MODELS,
|
||||||
|
get_global_default_model,
|
||||||
|
is_faster_whisper_installed,
|
||||||
|
is_model_downloaded,
|
||||||
|
set_global_default_model,
|
||||||
|
validate_model_name,
|
||||||
|
)
|
||||||
|
|
||||||
|
router = APIRouter(tags=["models"])
|
||||||
|
|
||||||
|
|
||||||
|
class TranscriptionModelOut(BaseModel):
|
||||||
|
name: str
|
||||||
|
display_name: str
|
||||||
|
description: str
|
||||||
|
params: str
|
||||||
|
approx_memory: str
|
||||||
|
relative_speed: str
|
||||||
|
is_default: bool
|
||||||
|
downloaded: bool
|
||||||
|
available: bool
|
||||||
|
|
||||||
|
|
||||||
|
class ModelsOut(BaseModel):
|
||||||
|
default_model: str
|
||||||
|
faster_whisper_installed: bool
|
||||||
|
models: list[TranscriptionModelOut]
|
||||||
|
|
||||||
|
|
||||||
|
class DefaultModelUpdate(BaseModel):
|
||||||
|
model: str = Field(min_length=1, max_length=32)
|
||||||
|
|
||||||
|
|
||||||
|
class DefaultModelOut(BaseModel):
|
||||||
|
default_model: str
|
||||||
|
|
||||||
|
|
||||||
|
async def _models_out(session: SessionDep) -> ModelsOut:
|
||||||
|
default = await get_global_default_model(session)
|
||||||
|
installed = is_faster_whisper_installed()
|
||||||
|
return ModelsOut(
|
||||||
|
default_model=default,
|
||||||
|
faster_whisper_installed=installed,
|
||||||
|
models=[
|
||||||
|
TranscriptionModelOut(
|
||||||
|
name=info.name,
|
||||||
|
display_name=info.display_name,
|
||||||
|
description=info.description,
|
||||||
|
params=info.params,
|
||||||
|
approx_memory=info.approx_memory,
|
||||||
|
relative_speed=info.relative_speed,
|
||||||
|
is_default=info.name == default,
|
||||||
|
downloaded=is_model_downloaded(info.name),
|
||||||
|
available=installed and is_model_downloaded(info.name),
|
||||||
|
)
|
||||||
|
for info in SUPPORTED_TRANSCRIPTION_MODELS.values()
|
||||||
|
],
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@router.get("/models", response_model=ModelsOut)
|
||||||
|
async def list_models(user: CurrentUser, session: SessionDep):
|
||||||
|
return await _models_out(session)
|
||||||
|
|
||||||
|
|
||||||
|
@router.get("/models/default", response_model=DefaultModelOut)
|
||||||
|
async def get_default_model(user: CurrentUser, session: SessionDep):
|
||||||
|
return DefaultModelOut(default_model=await get_global_default_model(session))
|
||||||
|
|
||||||
|
|
||||||
|
@router.put("/models/default", response_model=DefaultModelOut)
|
||||||
|
async def put_default_model(body: DefaultModelUpdate, user: CurrentUser, session: SessionDep):
|
||||||
|
try:
|
||||||
|
name = await set_global_default_model(session, body.model)
|
||||||
|
except ProviderConfigError as e:
|
||||||
|
raise HTTPException(422, str(e)) from None
|
||||||
|
return DefaultModelOut(default_model=name)
|
||||||
|
|
||||||
|
|
||||||
|
@router.post("/models/{name}/download", response_model=TranscriptionModelOut)
|
||||||
|
async def download_model(name: str, user: CurrentUser, session: SessionDep):
|
||||||
|
try:
|
||||||
|
clean = validate_model_name(name)
|
||||||
|
except ProviderConfigError as e:
|
||||||
|
raise HTTPException(422, str(e)) from None
|
||||||
|
if not is_faster_whisper_installed():
|
||||||
|
raise HTTPException(
|
||||||
|
501,
|
||||||
|
"faster-whisper is not installed on this server "
|
||||||
|
"(pip install shonar-backend[faster-whisper]).",
|
||||||
|
)
|
||||||
|
from shonar.services.ai.faster_whisper import _load_model
|
||||||
|
|
||||||
|
try:
|
||||||
|
await asyncio.to_thread(_load_model, clean)
|
||||||
|
except Exception as e: # noqa: BLE001 — surface download failures plainly
|
||||||
|
raise HTTPException(500, f"Model download failed: {e}") from None
|
||||||
|
default = await get_global_default_model(session)
|
||||||
|
info = SUPPORTED_TRANSCRIPTION_MODELS[clean]
|
||||||
|
return TranscriptionModelOut(
|
||||||
|
name=info.name,
|
||||||
|
display_name=info.display_name,
|
||||||
|
description=info.description,
|
||||||
|
params=info.params,
|
||||||
|
approx_memory=info.approx_memory,
|
||||||
|
relative_speed=info.relative_speed,
|
||||||
|
is_default=info.name == default,
|
||||||
|
downloaded=is_model_downloaded(clean),
|
||||||
|
available=is_model_downloaded(clean),
|
||||||
|
)
|
||||||
38
backend/shonar/api/v1/provider_info.py
Normal file
38
backend/shonar/api/v1/provider_info.py
Normal file
|
|
@ -0,0 +1,38 @@
|
||||||
|
"""Public provider handshake.
|
||||||
|
|
||||||
|
`GET /api/v1/provider-info` lets SHONAR clients (and platform probes on
|
||||||
|
Start9/Umbrel) identify this server as SHONAR-compatible and learn its
|
||||||
|
capabilities BEFORE any credentials are involved. Returns static deployment
|
||||||
|
metadata only — never user data, never configuration secrets.
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from fastapi import APIRouter
|
||||||
|
|
||||||
|
from shonar import __version__
|
||||||
|
from shonar.core.config import get_settings
|
||||||
|
|
||||||
|
router = APIRouter(tags=["provider"])
|
||||||
|
|
||||||
|
|
||||||
|
@router.get("/provider-info")
|
||||||
|
async def provider_info() -> dict:
|
||||||
|
"""Static, unauthenticated deployment identity for client probing."""
|
||||||
|
settings = get_settings()
|
||||||
|
# Capability flags reflect what this deployment can actually do right now.
|
||||||
|
transcription = settings.transcription_provider != "none"
|
||||||
|
llm = settings.llm_provider != "none"
|
||||||
|
return {
|
||||||
|
"kind": "shonar",
|
||||||
|
"version": __version__,
|
||||||
|
"api_version": "v1",
|
||||||
|
"capabilities": {
|
||||||
|
"chunked_upload": True,
|
||||||
|
"server_transcription": transcription,
|
||||||
|
"server_summary": llm,
|
||||||
|
"account_deletion": True,
|
||||||
|
},
|
||||||
|
# Storage backend family (never a path): "local" | "s3"
|
||||||
|
"storage_backend": settings.storage_backend,
|
||||||
|
}
|
||||||
528
backend/shonar/api/v1/recordings.py
Normal file
528
backend/shonar/api/v1/recordings.py
Normal file
|
|
@ -0,0 +1,528 @@
|
||||||
|
"""Upload sessions + recordings CRUD.
|
||||||
|
|
||||||
|
Every endpoint enforces ownership server-side; missing/foreign resources
|
||||||
|
return 404 (no existence leaks). Storage keys are never exposed.
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import contextlib
|
||||||
|
import uuid
|
||||||
|
from datetime import datetime
|
||||||
|
|
||||||
|
from fastapi import APIRouter, Header, HTTPException, Query, Request, Response
|
||||||
|
from sqlalchemy import func, select
|
||||||
|
|
||||||
|
from shonar.api.deps import CurrentUser, SessionDep
|
||||||
|
from shonar.api.schemas_recordings import (
|
||||||
|
ProcessingJobOut,
|
||||||
|
RecordingFinalize,
|
||||||
|
RecordingListOut,
|
||||||
|
RecordingOut,
|
||||||
|
RecordingUpdate,
|
||||||
|
SummaryOut,
|
||||||
|
SummaryUpdate,
|
||||||
|
TranscriptOut,
|
||||||
|
TranscriptUpdate,
|
||||||
|
UploadSessionCreate,
|
||||||
|
UploadSessionOut,
|
||||||
|
UploadStatusOut,
|
||||||
|
)
|
||||||
|
from shonar.db.models import (
|
||||||
|
Asset,
|
||||||
|
AssetKind,
|
||||||
|
ProcessingJob,
|
||||||
|
Recording,
|
||||||
|
RecordingTag,
|
||||||
|
Summary,
|
||||||
|
Tag,
|
||||||
|
Transcript,
|
||||||
|
utcnow,
|
||||||
|
)
|
||||||
|
from shonar.services import uploads as up
|
||||||
|
|
||||||
|
router = APIRouter(tags=["uploads", "recordings"])
|
||||||
|
|
||||||
|
|
||||||
|
async def _recording_out(session, rec: Recording) -> RecordingOut: # noqa: ANN001
|
||||||
|
original = await session.scalar(
|
||||||
|
select(Asset).where(Asset.recording_id == rec.id, Asset.kind == AssetKind.original)
|
||||||
|
)
|
||||||
|
tags = await session.scalars(
|
||||||
|
select(Tag.name)
|
||||||
|
.join(RecordingTag, RecordingTag.tag_id == Tag.id)
|
||||||
|
.where(RecordingTag.recording_id == rec.id)
|
||||||
|
.order_by(Tag.name)
|
||||||
|
)
|
||||||
|
return RecordingOut(
|
||||||
|
id=rec.id,
|
||||||
|
title=rec.title,
|
||||||
|
recorded_at=rec.recorded_at,
|
||||||
|
duration_seconds=rec.duration_seconds,
|
||||||
|
notes=rec.notes,
|
||||||
|
latitude=rec.latitude,
|
||||||
|
longitude=rec.longitude,
|
||||||
|
processing_status=rec.processing_status.value,
|
||||||
|
processing_error=rec.processing_error,
|
||||||
|
transcription_model=rec.transcription_model,
|
||||||
|
has_audio=original is not None,
|
||||||
|
mime_type=original.mime_type if original else None,
|
||||||
|
size_bytes=original.size_bytes if original else None,
|
||||||
|
tags=list(tags),
|
||||||
|
created_at=rec.created_at,
|
||||||
|
updated_at=rec.updated_at,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
# --- upload sessions ---------------------------------------------------------
|
||||||
|
|
||||||
|
|
||||||
|
@router.post("/uploads", response_model=UploadSessionOut, status_code=201)
|
||||||
|
async def create_upload(body: UploadSessionCreate, user: CurrentUser, session: SessionDep):
|
||||||
|
try:
|
||||||
|
us = await up.create_session(
|
||||||
|
session, user.id, body.declared_mime_type, body.declared_size_bytes,
|
||||||
|
body.title, body.client_recording_id, body.transcription_model,
|
||||||
|
)
|
||||||
|
except up.UploadError as e:
|
||||||
|
raise HTTPException(e.status_code, e.message) from None
|
||||||
|
return us
|
||||||
|
|
||||||
|
|
||||||
|
@router.get("/uploads/{session_id}", response_model=UploadStatusOut)
|
||||||
|
async def upload_status(session_id: uuid.UUID, user: CurrentUser, session: SessionDep):
|
||||||
|
try:
|
||||||
|
us = await up.get_owned_session(session, user.id, session_id)
|
||||||
|
indexes = await up.received_indexes(session, user.id, session_id)
|
||||||
|
except up.UploadError as e:
|
||||||
|
raise HTTPException(e.status_code, e.message) from None
|
||||||
|
return UploadStatusOut(id=us.id, status=us.status.value, received_chunk_indexes=indexes)
|
||||||
|
|
||||||
|
|
||||||
|
@router.put("/uploads/{session_id}/chunks/{chunk_index}", status_code=201)
|
||||||
|
async def put_chunk(
|
||||||
|
session_id: uuid.UUID,
|
||||||
|
chunk_index: int,
|
||||||
|
request: Request,
|
||||||
|
user: CurrentUser,
|
||||||
|
session: SessionDep,
|
||||||
|
x_chunk_sha256: str | None = Header(default=None),
|
||||||
|
):
|
||||||
|
data = await request.body()
|
||||||
|
try:
|
||||||
|
chunk = await up.put_chunk(session, user.id, session_id, chunk_index, data, x_chunk_sha256)
|
||||||
|
except up.UploadError as e:
|
||||||
|
raise HTTPException(e.status_code, e.message) from None
|
||||||
|
return {"chunk_index": chunk.chunk_index, "size_bytes": chunk.size_bytes}
|
||||||
|
|
||||||
|
|
||||||
|
@router.post("/uploads/{session_id}/finalize", response_model=RecordingOut, status_code=201)
|
||||||
|
async def finalize_upload(
|
||||||
|
session_id: uuid.UUID, body: RecordingFinalize, user: CurrentUser, session: SessionDep
|
||||||
|
):
|
||||||
|
# Location is stored ONLY with explicit per-user consent.
|
||||||
|
lat = body.latitude if user.location_storage_enabled else None
|
||||||
|
lon = body.longitude if user.location_storage_enabled else None
|
||||||
|
acc = body.location_accuracy_m if user.location_storage_enabled else None
|
||||||
|
try:
|
||||||
|
_us, rec, _asset = await up.finalize(
|
||||||
|
session, user.id, session_id,
|
||||||
|
recorded_at=body.recorded_at, duration_seconds=body.duration_seconds,
|
||||||
|
latitude=lat, longitude=lon, location_accuracy_m=acc, notes=body.notes,
|
||||||
|
transcription_model=body.transcription_model,
|
||||||
|
)
|
||||||
|
except up.UploadError as e:
|
||||||
|
raise HTTPException(e.status_code, e.message) from None
|
||||||
|
return await _recording_out(session, rec)
|
||||||
|
|
||||||
|
|
||||||
|
@router.delete("/uploads/{session_id}", status_code=204)
|
||||||
|
async def abort_upload(session_id: uuid.UUID, user: CurrentUser, session: SessionDep):
|
||||||
|
try:
|
||||||
|
await up.abort(session, user.id, session_id)
|
||||||
|
except up.UploadError as e:
|
||||||
|
raise HTTPException(e.status_code, e.message) from None
|
||||||
|
|
||||||
|
|
||||||
|
# --- recordings ----------------------------------------------------------------
|
||||||
|
|
||||||
|
|
||||||
|
@router.get("/recordings", response_model=RecordingListOut)
|
||||||
|
async def list_recordings(
|
||||||
|
user: CurrentUser,
|
||||||
|
session: SessionDep,
|
||||||
|
limit: int = Query(default=50, ge=1, le=200),
|
||||||
|
offset: int = Query(default=0, ge=0),
|
||||||
|
sort: str = Query(
|
||||||
|
default="recorded_at",
|
||||||
|
pattern="^(recorded_at|created_at|duration_seconds|title)$",
|
||||||
|
),
|
||||||
|
order: str = Query(default="desc", pattern="^(asc|desc)$"),
|
||||||
|
# M9 filters (AND-combined).
|
||||||
|
tag: str | None = Query(default=None, max_length=80),
|
||||||
|
status: str | None = Query(default=None, max_length=20),
|
||||||
|
from_date: datetime | None = Query(default=None, description="recorded_at >= (UTC)"),
|
||||||
|
to_date: datetime | None = Query(default=None, description="recorded_at <= (UTC)"),
|
||||||
|
):
|
||||||
|
where = [Recording.user_id == user.id, Recording.deleted_at.is_(None)]
|
||||||
|
if tag:
|
||||||
|
from shonar.db.models import RecordingTag as RT
|
||||||
|
from shonar.db.models import Tag as T
|
||||||
|
|
||||||
|
where.append(
|
||||||
|
Recording.id.in_(
|
||||||
|
select(RT.recording_id)
|
||||||
|
.join(T, T.id == RT.tag_id)
|
||||||
|
.where(T.user_id == user.id, T.name == tag.strip().lower())
|
||||||
|
)
|
||||||
|
)
|
||||||
|
if status:
|
||||||
|
from shonar.db.models import ProcessingStatus
|
||||||
|
|
||||||
|
try:
|
||||||
|
where.append(Recording.processing_status == ProcessingStatus(status))
|
||||||
|
except ValueError:
|
||||||
|
raise HTTPException(422, f"Unknown status: {status}") from None
|
||||||
|
if from_date is not None:
|
||||||
|
where.append(Recording.recorded_at >= from_date)
|
||||||
|
if to_date is not None:
|
||||||
|
where.append(Recording.recorded_at <= to_date)
|
||||||
|
total = await session.scalar(select(func.count(Recording.id)).where(*where))
|
||||||
|
col = getattr(Recording, sort)
|
||||||
|
col = col.desc() if order == "desc" else col.asc()
|
||||||
|
rows = await session.scalars(
|
||||||
|
select(Recording).where(*where).order_by(col).limit(limit).offset(offset)
|
||||||
|
)
|
||||||
|
items = [await _recording_out(session, r) for r in rows]
|
||||||
|
return RecordingListOut(items=items, total=total or 0, limit=limit, offset=offset)
|
||||||
|
|
||||||
|
|
||||||
|
@router.get("/recordings/{recording_id}", response_model=RecordingOut)
|
||||||
|
async def get_recording(recording_id: uuid.UUID, user: CurrentUser, session: SessionDep):
|
||||||
|
rec = await session.get(Recording, recording_id)
|
||||||
|
if rec is None or rec.user_id != user.id or rec.deleted_at is not None:
|
||||||
|
raise HTTPException(404, "Recording not found")
|
||||||
|
return await _recording_out(session, rec)
|
||||||
|
|
||||||
|
|
||||||
|
@router.patch("/recordings/{recording_id}", response_model=RecordingOut)
|
||||||
|
async def update_recording(
|
||||||
|
recording_id: uuid.UUID, body: RecordingUpdate, user: CurrentUser, session: SessionDep
|
||||||
|
):
|
||||||
|
rec = await session.get(Recording, recording_id)
|
||||||
|
if rec is None or rec.user_id != user.id or rec.deleted_at is not None:
|
||||||
|
raise HTTPException(404, "Recording not found")
|
||||||
|
if body.title is not None:
|
||||||
|
rec.title = body.title
|
||||||
|
if body.notes is not None:
|
||||||
|
rec.notes = body.notes
|
||||||
|
if body.latitude is not None and user.location_storage_enabled:
|
||||||
|
rec.latitude = body.latitude
|
||||||
|
if body.longitude is not None and user.location_storage_enabled:
|
||||||
|
rec.longitude = body.longitude
|
||||||
|
if body.tags is not None:
|
||||||
|
# Replace tag set. Tags are per-user, created on demand.
|
||||||
|
names = sorted(dict.fromkeys(t.strip().lower() for t in body.tags if t.strip()))[:20]
|
||||||
|
existing = list(
|
||||||
|
await session.scalars(select(Tag).where(Tag.user_id == user.id, Tag.name.in_(names)))
|
||||||
|
)
|
||||||
|
by_name = {t.name: t for t in existing}
|
||||||
|
for name in names:
|
||||||
|
if name not in by_name:
|
||||||
|
t = Tag(user_id=user.id, name=name)
|
||||||
|
session.add(t)
|
||||||
|
await session.flush()
|
||||||
|
by_name[name] = t
|
||||||
|
# Deterministic replace: drop all links for this recording, re-add.
|
||||||
|
from sqlalchemy import delete as sql_delete
|
||||||
|
|
||||||
|
await session.execute(
|
||||||
|
sql_delete(RecordingTag).where(RecordingTag.recording_id == rec.id)
|
||||||
|
)
|
||||||
|
for name in names:
|
||||||
|
session.add(RecordingTag(recording_id=rec.id, tag_id=by_name[name].id))
|
||||||
|
await session.flush()
|
||||||
|
return await _recording_out(session, rec)
|
||||||
|
|
||||||
|
|
||||||
|
@router.delete("/recordings/{recording_id}", status_code=204)
|
||||||
|
async def delete_recording(
|
||||||
|
recording_id: uuid.UUID,
|
||||||
|
user: CurrentUser,
|
||||||
|
session: SessionDep,
|
||||||
|
purge: bool = Query(default=False, description="true also deletes stored audio"),
|
||||||
|
):
|
||||||
|
"""Soft-delete by default; ?purge=true removes rows + stored files now."""
|
||||||
|
rec = await session.get(Recording, recording_id)
|
||||||
|
if rec is None or rec.user_id != user.id or rec.deleted_at is not None:
|
||||||
|
raise HTTPException(404, "Recording not found")
|
||||||
|
if not purge:
|
||||||
|
rec.deleted_at = utcnow()
|
||||||
|
await session.flush()
|
||||||
|
return Response(status_code=204)
|
||||||
|
|
||||||
|
from shonar.storage import get_storage
|
||||||
|
|
||||||
|
storage = get_storage()
|
||||||
|
assets = list(
|
||||||
|
await session.scalars(select(Asset).where(Asset.recording_id == rec.id))
|
||||||
|
)
|
||||||
|
await session.delete(rec) # cascades to assets/transcripts/summaries/jobs
|
||||||
|
await session.flush()
|
||||||
|
for a in assets:
|
||||||
|
with contextlib.suppress(Exception): # best effort
|
||||||
|
await storage.delete(a.storage_key)
|
||||||
|
return Response(status_code=204)
|
||||||
|
|
||||||
|
|
||||||
|
@router.get("/recordings/{recording_id}/audio")
|
||||||
|
async def download_audio(recording_id: uuid.UUID, user: CurrentUser, session: SessionDep):
|
||||||
|
rec = await session.get(Recording, recording_id)
|
||||||
|
if rec is None or rec.user_id != user.id or rec.deleted_at is not None:
|
||||||
|
raise HTTPException(404, "Recording not found")
|
||||||
|
original = await session.scalar(
|
||||||
|
select(Asset).where(Asset.recording_id == rec.id, Asset.kind == AssetKind.original)
|
||||||
|
)
|
||||||
|
if original is None:
|
||||||
|
raise HTTPException(404, "No audio stored for this recording")
|
||||||
|
from fastapi.responses import Response as RawResponse
|
||||||
|
|
||||||
|
from shonar.storage import get_storage
|
||||||
|
|
||||||
|
data = await get_storage().get(original.storage_key)
|
||||||
|
ext = original.storage_key[original.storage_key.rfind(".") :]
|
||||||
|
filename = f"{rec.recorded_at:%Y%m%d-%H%M%S}{ext}"
|
||||||
|
return RawResponse(
|
||||||
|
content=data,
|
||||||
|
media_type=original.mime_type,
|
||||||
|
headers={
|
||||||
|
"Content-Disposition": f'attachment; filename="{filename}"',
|
||||||
|
"Cache-Control": "private, no-store",
|
||||||
|
},
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
# --- AI outputs (M7): latest transcript / summary / job states --------------
|
||||||
|
|
||||||
|
|
||||||
|
async def _owned_recording(
|
||||||
|
session: SessionDep, user: CurrentUser, recording_id: uuid.UUID
|
||||||
|
) -> Recording:
|
||||||
|
rec = await session.get(Recording, recording_id)
|
||||||
|
if rec is None or rec.user_id != user.id or rec.deleted_at is not None:
|
||||||
|
raise HTTPException(404, "Recording not found")
|
||||||
|
return rec
|
||||||
|
|
||||||
|
|
||||||
|
@router.get("/recordings/{recording_id}/transcript", response_model=TranscriptOut)
|
||||||
|
async def get_transcript(
|
||||||
|
recording_id: uuid.UUID, user: CurrentUser, session: SessionDep
|
||||||
|
):
|
||||||
|
rec = await _owned_recording(session, user, recording_id)
|
||||||
|
row = await session.scalar(
|
||||||
|
select(Transcript)
|
||||||
|
.where(
|
||||||
|
Transcript.recording_id == rec.id,
|
||||||
|
Transcript.superseded_at.is_(None),
|
||||||
|
)
|
||||||
|
.order_by(Transcript.version.desc())
|
||||||
|
)
|
||||||
|
if row is None:
|
||||||
|
raise HTTPException(404, "No transcript yet")
|
||||||
|
return _transcript_out(row)
|
||||||
|
|
||||||
|
|
||||||
|
@router.get("/recordings/{recording_id}/summary", response_model=SummaryOut)
|
||||||
|
async def get_summary(recording_id: uuid.UUID, user: CurrentUser, session: SessionDep):
|
||||||
|
rec = await _owned_recording(session, user, recording_id)
|
||||||
|
row = await session.scalar(
|
||||||
|
select(Summary)
|
||||||
|
.where(
|
||||||
|
Summary.recording_id == rec.id,
|
||||||
|
Summary.superseded_at.is_(None),
|
||||||
|
)
|
||||||
|
.order_by(Summary.version.desc())
|
||||||
|
)
|
||||||
|
if row is None:
|
||||||
|
raise HTTPException(404, "No summary yet")
|
||||||
|
return row
|
||||||
|
|
||||||
|
|
||||||
|
@router.get("/recordings/{recording_id}/jobs", response_model=list[ProcessingJobOut])
|
||||||
|
async def list_jobs(recording_id: uuid.UUID, user: CurrentUser, session: SessionDep):
|
||||||
|
rec = await _owned_recording(session, user, recording_id)
|
||||||
|
rows = await session.scalars(
|
||||||
|
select(ProcessingJob)
|
||||||
|
.where(ProcessingJob.recording_id == rec.id)
|
||||||
|
.order_by(ProcessingJob.id)
|
||||||
|
)
|
||||||
|
return list(rows)
|
||||||
|
|
||||||
|
|
||||||
|
@router.post("/recordings/{recording_id}/reprocess", response_model=list[ProcessingJobOut])
|
||||||
|
async def reprocess_recording(
|
||||||
|
recording_id: uuid.UUID,
|
||||||
|
user: CurrentUser,
|
||||||
|
session: SessionDep,
|
||||||
|
job: str = Query(default="summarize", pattern="^(transcribe|summarize)$"),
|
||||||
|
model: str | None = Query(default=None, max_length=64),
|
||||||
|
):
|
||||||
|
"""Force one pipeline stage to run again (Summarize / Re-transcribe).
|
||||||
|
|
||||||
|
Unlike the enqueue-on-finalize path, this ignores prior success: a
|
||||||
|
summary the user wants regenerated (better model, new prompt) is a
|
||||||
|
deliberate request. Running jobs are left alone (409 instead of a
|
||||||
|
duplicate).
|
||||||
|
"""
|
||||||
|
from shonar.db.models import JobStatus, ProcessingStatus
|
||||||
|
from shonar.db.models import JobType as JT
|
||||||
|
from shonar.services import processing as proc
|
||||||
|
|
||||||
|
rec = await _owned_recording(session, user, recording_id)
|
||||||
|
job_type = JT(job)
|
||||||
|
if job_type is JT.summarize and await proc.latest_transcript_text(
|
||||||
|
session, rec.id
|
||||||
|
) is None:
|
||||||
|
raise HTTPException(409, "Transcribe first — there is nothing to summarize.")
|
||||||
|
existing = await session.scalar(
|
||||||
|
select(ProcessingJob)
|
||||||
|
.where(
|
||||||
|
ProcessingJob.recording_id == rec.id,
|
||||||
|
ProcessingJob.job_type == job_type,
|
||||||
|
)
|
||||||
|
.order_by(ProcessingJob.id.desc())
|
||||||
|
)
|
||||||
|
if existing is not None and existing.status in (
|
||||||
|
JobStatus.queued,
|
||||||
|
JobStatus.running,
|
||||||
|
):
|
||||||
|
raise HTTPException(409, "That stage is already running.")
|
||||||
|
if existing is None:
|
||||||
|
session.add(ProcessingJob(recording_id=rec.id, job_type=job_type))
|
||||||
|
else:
|
||||||
|
existing.status = JobStatus.queued
|
||||||
|
existing.attempt = 0
|
||||||
|
existing.error = None
|
||||||
|
existing.stage = None
|
||||||
|
existing.progress = None
|
||||||
|
existing.started_at = None
|
||||||
|
existing.finished_at = None
|
||||||
|
if job_type is JT.transcribe:
|
||||||
|
if model is not None:
|
||||||
|
# A re-transcribe may switch models; the saved per-recording
|
||||||
|
# override is what the worker reads, so persist it here.
|
||||||
|
from shonar.services.ai import ProviderConfigError
|
||||||
|
from shonar.services.ai.model_registry import validate_model_name
|
||||||
|
|
||||||
|
try:
|
||||||
|
rec.transcription_model = validate_model_name(model)
|
||||||
|
except ProviderConfigError as e:
|
||||||
|
raise HTTPException(422, str(e)) from e
|
||||||
|
rec.processing_status = ProcessingStatus.processing
|
||||||
|
rec.processing_error = None
|
||||||
|
await session.flush()
|
||||||
|
await proc.transport_enqueue(job_type, rec.id)
|
||||||
|
rows = await session.scalars(
|
||||||
|
select(ProcessingJob)
|
||||||
|
.where(ProcessingJob.recording_id == rec.id)
|
||||||
|
.order_by(ProcessingJob.id)
|
||||||
|
)
|
||||||
|
return list(rows)
|
||||||
|
|
||||||
|
|
||||||
|
# --- M8: user edits (new version, edited_by_user=True; pipeline won't clobber)
|
||||||
|
def _transcript_out(row: Transcript) -> TranscriptOut:
|
||||||
|
return TranscriptOut(
|
||||||
|
version=row.version,
|
||||||
|
language=row.language,
|
||||||
|
provider=row.provider,
|
||||||
|
model=row.model,
|
||||||
|
text=row.text,
|
||||||
|
segments=[
|
||||||
|
s for s in (row.segments or []) if isinstance(s, dict)
|
||||||
|
],
|
||||||
|
edited_by_user=row.edited_by_user,
|
||||||
|
created_at=row.created_at,
|
||||||
|
updated_at=row.updated_at,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@router.put("/recordings/{recording_id}/transcript", response_model=TranscriptOut)
|
||||||
|
async def update_transcript(
|
||||||
|
recording_id: uuid.UUID, body: TranscriptUpdate, user: CurrentUser, session: SessionDep
|
||||||
|
):
|
||||||
|
from sqlalchemy import func as sql_func
|
||||||
|
|
||||||
|
rec = await _owned_recording(session, user, recording_id)
|
||||||
|
existing = list(
|
||||||
|
await session.scalars(
|
||||||
|
select(Transcript)
|
||||||
|
.where(
|
||||||
|
Transcript.recording_id == rec.id,
|
||||||
|
Transcript.superseded_at.is_(None),
|
||||||
|
)
|
||||||
|
.order_by(Transcript.version.desc())
|
||||||
|
)
|
||||||
|
)
|
||||||
|
now = utcnow()
|
||||||
|
max_version = await session.scalar(
|
||||||
|
select(sql_func.max(Transcript.version)).where(Transcript.recording_id == rec.id)
|
||||||
|
)
|
||||||
|
for row in existing:
|
||||||
|
row.superseded_at = now
|
||||||
|
row = Transcript(
|
||||||
|
recording_id=rec.id,
|
||||||
|
version=(max_version or 0) + 1,
|
||||||
|
language=body.language,
|
||||||
|
provider="user",
|
||||||
|
model=None,
|
||||||
|
text=body.text,
|
||||||
|
segments=(
|
||||||
|
[
|
||||||
|
{"start": s.start, "end": s.end, "text": s.text, "speaker": s.speaker}
|
||||||
|
for s in (body.segments or [])
|
||||||
|
]
|
||||||
|
if body.segments is not None
|
||||||
|
else None
|
||||||
|
),
|
||||||
|
edited_by_user=True,
|
||||||
|
)
|
||||||
|
session.add(row)
|
||||||
|
await session.flush()
|
||||||
|
return _transcript_out(row)
|
||||||
|
|
||||||
|
|
||||||
|
@router.put("/recordings/{recording_id}/summary", response_model=SummaryOut)
|
||||||
|
async def update_summary(
|
||||||
|
recording_id: uuid.UUID, body: SummaryUpdate, user: CurrentUser, session: SessionDep
|
||||||
|
):
|
||||||
|
from sqlalchemy import func as sql_func
|
||||||
|
|
||||||
|
rec = await _owned_recording(session, user, recording_id)
|
||||||
|
existing = list(
|
||||||
|
await session.scalars(
|
||||||
|
select(Summary)
|
||||||
|
.where(
|
||||||
|
Summary.recording_id == rec.id,
|
||||||
|
Summary.superseded_at.is_(None),
|
||||||
|
)
|
||||||
|
.order_by(Summary.version.desc())
|
||||||
|
)
|
||||||
|
)
|
||||||
|
now = utcnow()
|
||||||
|
max_version = await session.scalar(
|
||||||
|
select(sql_func.max(Summary.version)).where(Summary.recording_id == rec.id)
|
||||||
|
)
|
||||||
|
for row in existing:
|
||||||
|
row.superseded_at = now
|
||||||
|
row = Summary(
|
||||||
|
recording_id=rec.id,
|
||||||
|
version=(max_version or 0) + 1,
|
||||||
|
provider="user",
|
||||||
|
model=None,
|
||||||
|
content=body.content,
|
||||||
|
edited_by_user=True,
|
||||||
|
)
|
||||||
|
session.add(row)
|
||||||
|
await session.flush()
|
||||||
|
return row
|
||||||
81
backend/shonar/api/v1/search_exports.py
Normal file
81
backend/shonar/api/v1/search_exports.py
Normal file
|
|
@ -0,0 +1,81 @@
|
||||||
|
"""Search + exports endpoints (M9).
|
||||||
|
|
||||||
|
Ownership is enforced everywhere: a search only ever sees the caller's own
|
||||||
|
non-deleted recordings, and an export 404s on anything the caller doesn't
|
||||||
|
own (no existence leak).
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import uuid
|
||||||
|
|
||||||
|
from fastapi import APIRouter, HTTPException, Query
|
||||||
|
from fastapi.responses import Response
|
||||||
|
|
||||||
|
from shonar.api.deps import CurrentUser, SessionDep
|
||||||
|
from shonar.api.schemas_search import SearchHitOut, SearchOut
|
||||||
|
from shonar.db.models import Recording
|
||||||
|
from shonar.services import exports as ex
|
||||||
|
from shonar.services import search as sr
|
||||||
|
|
||||||
|
router = APIRouter(tags=["search", "exports"])
|
||||||
|
|
||||||
|
|
||||||
|
# --- search -------------------------------------------------------------------
|
||||||
|
|
||||||
|
|
||||||
|
@router.get("/search", response_model=SearchOut)
|
||||||
|
async def search(
|
||||||
|
user: CurrentUser,
|
||||||
|
session: SessionDep,
|
||||||
|
q: str = Query(..., min_length=1, max_length=200, description="Search text"),
|
||||||
|
scope: str = Query(
|
||||||
|
default="all",
|
||||||
|
pattern="^(all|title|notes|transcript|summary|tag)$",
|
||||||
|
description="Restrict the search to one field (default: all)",
|
||||||
|
),
|
||||||
|
limit: int = Query(default=20, ge=1, le=100),
|
||||||
|
offset: int = Query(default=0, ge=0),
|
||||||
|
):
|
||||||
|
hits, total = await sr.search_recordings(
|
||||||
|
session, user.id, q, scope=scope, limit=limit, offset=offset
|
||||||
|
)
|
||||||
|
items = [
|
||||||
|
SearchHitOut(
|
||||||
|
id=h.recording.id,
|
||||||
|
title=h.recording.title,
|
||||||
|
field=h.field,
|
||||||
|
snippet=h.snippet,
|
||||||
|
)
|
||||||
|
for h in hits
|
||||||
|
]
|
||||||
|
return SearchOut(query=q, scope=scope, items=items, total=total, limit=limit, offset=offset)
|
||||||
|
|
||||||
|
|
||||||
|
# --- exports ------------------------------------------------------------------
|
||||||
|
|
||||||
|
|
||||||
|
@router.get("/recordings/{recording_id}/export")
|
||||||
|
async def export_recording(
|
||||||
|
recording_id: uuid.UUID,
|
||||||
|
user: CurrentUser,
|
||||||
|
session: SessionDep,
|
||||||
|
fmt: str = Query(default="zip", pattern="^(audio|txt|md|zip)$"),
|
||||||
|
):
|
||||||
|
"""Return one export artifact inline (audio / txt / md / zip bundle)."""
|
||||||
|
rec = await session.get(Recording, recording_id)
|
||||||
|
if rec is None or rec.user_id != user.id or rec.deleted_at is not None:
|
||||||
|
raise HTTPException(404, "Recording not found")
|
||||||
|
try:
|
||||||
|
result = await ex.build_export(session, rec, fmt)
|
||||||
|
except ex.ExportError as e:
|
||||||
|
raise HTTPException(e.status_code, e.message) from None
|
||||||
|
await ex.record_export(session, user.id, rec, fmt, result)
|
||||||
|
return Response(
|
||||||
|
content=result.data,
|
||||||
|
media_type=result.mime_type,
|
||||||
|
headers={
|
||||||
|
"Content-Disposition": f'attachment; filename="{result.filename}"',
|
||||||
|
"Cache-Control": "private, no-store",
|
||||||
|
},
|
||||||
|
)
|
||||||
54
backend/shonar/api/v1/users.py
Normal file
54
backend/shonar/api/v1/users.py
Normal file
|
|
@ -0,0 +1,54 @@
|
||||||
|
"""Users and devices."""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import uuid
|
||||||
|
|
||||||
|
from fastapi import APIRouter, HTTPException
|
||||||
|
from sqlalchemy import select, update
|
||||||
|
|
||||||
|
from shonar.api.deps import CurrentUser, SessionDep
|
||||||
|
from shonar.api.schemas_common import DeviceOut, UserOut, UserUpdate
|
||||||
|
from shonar.db.models import Device, RefreshToken, utcnow
|
||||||
|
|
||||||
|
router = APIRouter(tags=["users", "devices"])
|
||||||
|
|
||||||
|
|
||||||
|
@router.get("/users/me", response_model=UserOut)
|
||||||
|
async def get_me(user: CurrentUser):
|
||||||
|
return user
|
||||||
|
|
||||||
|
|
||||||
|
@router.patch("/users/me", response_model=UserOut)
|
||||||
|
async def update_me(body: UserUpdate, user: CurrentUser, session: SessionDep):
|
||||||
|
if body.display_name is not None:
|
||||||
|
user.display_name = body.display_name
|
||||||
|
if body.location_storage_enabled is not None:
|
||||||
|
user.location_storage_enabled = body.location_storage_enabled
|
||||||
|
await session.flush()
|
||||||
|
return user
|
||||||
|
|
||||||
|
|
||||||
|
@router.get("/devices", response_model=list[DeviceOut])
|
||||||
|
async def list_devices(user: CurrentUser, session: SessionDep):
|
||||||
|
rows = await session.scalars(
|
||||||
|
select(Device).where(Device.user_id == user.id).order_by(Device.last_seen_at.desc())
|
||||||
|
)
|
||||||
|
return list(rows)
|
||||||
|
|
||||||
|
|
||||||
|
@router.delete("/devices/{device_id}", status_code=204)
|
||||||
|
async def revoke_device(device_id: uuid.UUID, user: CurrentUser, session: SessionDep):
|
||||||
|
device = await session.get(Device, device_id)
|
||||||
|
# Ownership check — no cross-user access, and 404 (not 403) to avoid
|
||||||
|
# leaking existence.
|
||||||
|
if device is None or device.user_id != user.id:
|
||||||
|
raise HTTPException(404, "Device not found")
|
||||||
|
device.revoked_at = utcnow()
|
||||||
|
# Revoke this device's live refresh tokens.
|
||||||
|
await session.execute(
|
||||||
|
update(RefreshToken)
|
||||||
|
.where(RefreshToken.device_id == device.id, RefreshToken.revoked_at.is_(None))
|
||||||
|
.values(revoked_at=utcnow())
|
||||||
|
)
|
||||||
|
return None
|
||||||
0
backend/shonar/core/__init__.py
Normal file
0
backend/shonar/core/__init__.py
Normal file
138
backend/shonar/core/config.py
Normal file
138
backend/shonar/core/config.py
Normal file
|
|
@ -0,0 +1,138 @@
|
||||||
|
"""Application settings.
|
||||||
|
|
||||||
|
All configuration comes from environment variables (and optionally an
|
||||||
|
admin config file pointed at by SHONAR_CONFIG_FILE). API keys and other
|
||||||
|
secrets must NEVER be hard-coded.
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from functools import lru_cache
|
||||||
|
|
||||||
|
from pydantic import Field, field_validator
|
||||||
|
from pydantic_settings import BaseSettings, SettingsConfigDict
|
||||||
|
|
||||||
|
|
||||||
|
class Settings(BaseSettings):
|
||||||
|
model_config = SettingsConfigDict(
|
||||||
|
env_prefix="SHONAR_",
|
||||||
|
env_file=".env",
|
||||||
|
env_file_encoding="utf-8",
|
||||||
|
extra="ignore",
|
||||||
|
)
|
||||||
|
|
||||||
|
# --- Core -------------------------------------------------------------
|
||||||
|
app_name: str = "S.H.O.N.A.R."
|
||||||
|
debug: bool = False
|
||||||
|
# Secret used to sign access tokens. MUST be set in production.
|
||||||
|
secret_key: str = Field(default="", repr=False)
|
||||||
|
access_token_ttl_minutes: int = 15
|
||||||
|
refresh_token_ttl_days: int = 30
|
||||||
|
# Comma-separated list of allowed registration modes: "open", "invite"
|
||||||
|
allow_registration: bool = True
|
||||||
|
|
||||||
|
# --- Database -----------------------------------------------------------
|
||||||
|
database_url: str = "postgresql+asyncpg://shonar:shonar@localhost:5432/shonar"
|
||||||
|
db_pool_size: int = 5
|
||||||
|
db_max_overflow: int = 10
|
||||||
|
|
||||||
|
# --- Storage ------------------------------------------------------------
|
||||||
|
# "local" or "s3"
|
||||||
|
storage_backend: str = "local"
|
||||||
|
storage_path: str = "./data/storage"
|
||||||
|
# S3 (used when storage_backend == "s3")
|
||||||
|
s3_endpoint_url: str = ""
|
||||||
|
s3_bucket: str = ""
|
||||||
|
s3_region: str = "us-east-1"
|
||||||
|
s3_access_key_id: str = Field(default="", repr=False)
|
||||||
|
s3_secret_access_key: str = Field(default="", repr=False)
|
||||||
|
|
||||||
|
# --- Upload limits --------------------------------------------------------
|
||||||
|
max_upload_bytes: int = 2 * 1024 * 1024 * 1024 # 2 GiB
|
||||||
|
max_chunk_bytes: int = 16 * 1024 * 1024
|
||||||
|
allowed_audio_mime_types: list[str] = [
|
||||||
|
"audio/mp4",
|
||||||
|
"audio/m4a",
|
||||||
|
"audio/aac",
|
||||||
|
"audio/wav",
|
||||||
|
"audio/x-wav",
|
||||||
|
"audio/ogg",
|
||||||
|
"audio/opus",
|
||||||
|
"audio/webm",
|
||||||
|
"audio/mpeg",
|
||||||
|
]
|
||||||
|
|
||||||
|
# --- Worker / queue -------------------------------------------------------
|
||||||
|
redis_url: str = "redis://localhost:6379/0"
|
||||||
|
# "arq" = Redis transport (server deployments). "inline" = in-process
|
||||||
|
# asyncio runner (desktop bundled-lite engine: no Redis, one job at a
|
||||||
|
# time). DB rows are the queue of record either way.
|
||||||
|
queue_backend: str = "arq" # arq | inline
|
||||||
|
# Apply 'alembic upgrade head' at startup. The desktop bundled engine
|
||||||
|
# owns its SQLite file and must self-migrate; server deployments run
|
||||||
|
# migrations in their deploy flow instead.
|
||||||
|
auto_migrate: bool = False
|
||||||
|
|
||||||
|
# --- Audio processing -----------------------------------------------------
|
||||||
|
# Optional server-side conversion via FFmpeg. Off by default.
|
||||||
|
audio_conversion_enabled: bool = False
|
||||||
|
ffmpeg_bin: str = "ffmpeg"
|
||||||
|
|
||||||
|
# --- AI providers -----------------------------------------------------
|
||||||
|
# "none" disables all AI processing (recording/sync/playback still work).
|
||||||
|
transcription_provider: str = "none" # none | whisper_http | faster_whisper
|
||||||
|
transcription_model: str = "base"
|
||||||
|
transcription_base_url: str = "" # whisper-compatible HTTP server
|
||||||
|
transcription_api_key: str = Field(default="", repr=False)
|
||||||
|
|
||||||
|
llm_provider: str = "none" # none | openai_compat | ollama
|
||||||
|
llm_model: str = ""
|
||||||
|
llm_base_url: str = ""
|
||||||
|
llm_api_key: str = Field(default="", repr=False)
|
||||||
|
|
||||||
|
# --- Search ---------------------------------------------------------
|
||||||
|
search_backend: str = "postgres_fts" # postgres_fts (meilisearch: TODO)
|
||||||
|
|
||||||
|
# --- Retention --------------------------------------------------------
|
||||||
|
# Grace window before the sweep hard-deletes soft-deleted recordings
|
||||||
|
# and accounts (rows + stored files). Cancellations are valid until
|
||||||
|
# the sweep fires.
|
||||||
|
retention_grace_days: int = 30
|
||||||
|
|
||||||
|
# --- Misc -------------------------------------------------------------
|
||||||
|
rate_limit_auth: str = "10/minute"
|
||||||
|
rate_limit_default: str = "120/minute"
|
||||||
|
|
||||||
|
@field_validator("allowed_audio_mime_types", mode="before")
|
||||||
|
@classmethod
|
||||||
|
def _split_mime(cls, v): # noqa: ANN001, ANN206
|
||||||
|
if isinstance(v, str):
|
||||||
|
return [item.strip() for item in v.split(",") if item.strip()]
|
||||||
|
return v
|
||||||
|
|
||||||
|
@property
|
||||||
|
def ai_enabled(self) -> bool:
|
||||||
|
return self.transcription_provider != "none" or self.llm_provider != "none"
|
||||||
|
|
||||||
|
def validate_production(self) -> list[str]:
|
||||||
|
"""Return a list of configuration warnings (empty == OK)."""
|
||||||
|
warnings: list[str] = []
|
||||||
|
if not self.secret_key or len(self.secret_key) < 32:
|
||||||
|
warnings.append(
|
||||||
|
"SHONAR_SECRET_KEY is missing or shorter than 32 characters. "
|
||||||
|
"Set a strong random secret in production."
|
||||||
|
)
|
||||||
|
if self.storage_backend == "s3" and not self.s3_bucket:
|
||||||
|
warnings.append("storage_backend=s3 but SHONAR_S3_BUCKET is empty.")
|
||||||
|
if self.transcription_provider == "whisper_http" and not self.transcription_base_url:
|
||||||
|
warnings.append(
|
||||||
|
"transcription_provider=whisper_http requires SHONAR_TRANSCRIPTION_BASE_URL."
|
||||||
|
)
|
||||||
|
if self.llm_provider == "openai_compat" and not self.llm_base_url:
|
||||||
|
warnings.append("llm_provider=openai_compat requires SHONAR_LLM_BASE_URL.")
|
||||||
|
return warnings
|
||||||
|
|
||||||
|
|
||||||
|
@lru_cache
|
||||||
|
def get_settings() -> Settings:
|
||||||
|
return Settings()
|
||||||
19
backend/shonar/core/ratelimit.py
Normal file
19
backend/shonar/core/ratelimit.py
Normal file
|
|
@ -0,0 +1,19 @@
|
||||||
|
"""Rate limiting (slowapi). A single shared limiter instance so counters are
|
||||||
|
consistent across endpoints; limit strings come from settings/env."""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from slowapi import Limiter
|
||||||
|
from slowapi.util import get_remote_address
|
||||||
|
|
||||||
|
from shonar.core.config import get_settings
|
||||||
|
|
||||||
|
_settings = get_settings()
|
||||||
|
|
||||||
|
limiter = Limiter(
|
||||||
|
key_func=get_remote_address,
|
||||||
|
default_limits=[],
|
||||||
|
enabled=True,
|
||||||
|
)
|
||||||
|
|
||||||
|
auth_limit = limiter.limit(_settings.rate_limit_auth)
|
||||||
98
backend/shonar/core/security.py
Normal file
98
backend/shonar/core/security.py
Normal file
|
|
@ -0,0 +1,98 @@
|
||||||
|
"""Security primitives: password hashing, access/refresh tokens, rate limit
|
||||||
|
keying. Secrets come exclusively from settings/env."""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import hashlib
|
||||||
|
import hmac
|
||||||
|
import secrets
|
||||||
|
import uuid
|
||||||
|
from datetime import UTC, datetime, timedelta
|
||||||
|
from typing import Any
|
||||||
|
|
||||||
|
import jwt
|
||||||
|
from argon2 import PasswordHasher
|
||||||
|
from argon2.exceptions import InvalidHashError, VerificationError, VerifyMismatchError
|
||||||
|
|
||||||
|
from shonar.core.config import get_settings
|
||||||
|
|
||||||
|
_ph = PasswordHasher()
|
||||||
|
|
||||||
|
ACCESS_TOKEN_TYPE = "access"
|
||||||
|
REFRESH_TOKEN_TYPE = "refresh"
|
||||||
|
|
||||||
|
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
# Passwords
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
|
||||||
|
|
||||||
|
def hash_password(password: str) -> str:
|
||||||
|
return _ph.hash(password)
|
||||||
|
|
||||||
|
|
||||||
|
def verify_password(password_hash: str, password: str) -> bool:
|
||||||
|
try:
|
||||||
|
return _ph.verify(password_hash, password)
|
||||||
|
except (VerifyMismatchError, VerificationError, InvalidHashError):
|
||||||
|
return False
|
||||||
|
|
||||||
|
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
# Access tokens (JWT)
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
|
||||||
|
|
||||||
|
def create_access_token(user_id: uuid.UUID, device_id: uuid.UUID | None = None) -> tuple[str, int]:
|
||||||
|
"""Returns (token, ttl_seconds)."""
|
||||||
|
settings = get_settings()
|
||||||
|
ttl = settings.access_token_ttl_minutes * 60
|
||||||
|
now = datetime.now(UTC)
|
||||||
|
payload: dict[str, Any] = {
|
||||||
|
"sub": str(user_id),
|
||||||
|
"typ": ACCESS_TOKEN_TYPE,
|
||||||
|
"iat": now,
|
||||||
|
"exp": now + timedelta(seconds=ttl),
|
||||||
|
"jti": uuid.uuid4().hex,
|
||||||
|
}
|
||||||
|
if device_id is not None:
|
||||||
|
payload["dev"] = str(device_id)
|
||||||
|
token = jwt.encode(payload, settings.secret_key, algorithm="HS256")
|
||||||
|
return token, ttl
|
||||||
|
|
||||||
|
|
||||||
|
class TokenError(Exception):
|
||||||
|
"""Invalid or expired token."""
|
||||||
|
|
||||||
|
|
||||||
|
def decode_access_token(token: str) -> dict[str, Any]:
|
||||||
|
settings = get_settings()
|
||||||
|
try:
|
||||||
|
payload = jwt.decode(token, settings.secret_key, algorithms=["HS256"])
|
||||||
|
except jwt.PyJWTError as exc:
|
||||||
|
raise TokenError("invalid access token") from exc
|
||||||
|
if payload.get("typ") != ACCESS_TOKEN_TYPE:
|
||||||
|
raise TokenError("wrong token type")
|
||||||
|
return payload
|
||||||
|
|
||||||
|
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
# Refresh tokens (opaque, rotated, hashed at rest)
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
|
||||||
|
|
||||||
|
def generate_refresh_token() -> str:
|
||||||
|
"""Opaque high-entropy token; only its SHA-256 hash is ever stored."""
|
||||||
|
return secrets.token_urlsafe(48)
|
||||||
|
|
||||||
|
|
||||||
|
def hash_refresh_token(token: str) -> str:
|
||||||
|
return hashlib.sha256(token.encode("utf-8")).hexdigest()
|
||||||
|
|
||||||
|
|
||||||
|
def refresh_token_ttl() -> timedelta:
|
||||||
|
return timedelta(days=get_settings().refresh_token_ttl_days)
|
||||||
|
|
||||||
|
|
||||||
|
def constant_time_equals(a: str, b: str) -> bool:
|
||||||
|
return hmac.compare_digest(a, b)
|
||||||
3
backend/shonar/db/__init__.py
Normal file
3
backend/shonar/db/__init__.py
Normal file
|
|
@ -0,0 +1,3 @@
|
||||||
|
"""DB package. Importing this registers every model on Base.metadata."""
|
||||||
|
|
||||||
|
from shonar.db import models # noqa: F401
|
||||||
50
backend/shonar/db/base.py
Normal file
50
backend/shonar/db/base.py
Normal file
|
|
@ -0,0 +1,50 @@
|
||||||
|
"""Public SQLAlchemy model base."""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import uuid
|
||||||
|
from datetime import UTC, datetime
|
||||||
|
|
||||||
|
from sqlalchemy import DateTime, TypeDecorator, Uuid
|
||||||
|
from sqlalchemy.orm import DeclarativeBase, Mapped, mapped_column
|
||||||
|
|
||||||
|
|
||||||
|
class UTCDT(TypeDecorator):
|
||||||
|
"""Timezone-aware datetime that survives SQLite round-trips.
|
||||||
|
|
||||||
|
Postgres returns aware values (impl is a no-op there); SQLite has no
|
||||||
|
tz info and hands back naive datetimes, which crash comparisons with
|
||||||
|
``utcnow()``. Attach UTC on read whenever the driver lost it.
|
||||||
|
"""
|
||||||
|
|
||||||
|
impl = DateTime(timezone=True)
|
||||||
|
cache_ok = True
|
||||||
|
|
||||||
|
def process_result_value(self, value, dialect): # noqa: ANN001, ANN201
|
||||||
|
if value is not None and value.tzinfo is None:
|
||||||
|
return value.replace(tzinfo=UTC)
|
||||||
|
return value
|
||||||
|
|
||||||
|
|
||||||
|
def utcnow() -> datetime:
|
||||||
|
return datetime.now(UTC)
|
||||||
|
|
||||||
|
|
||||||
|
def new_uuid() -> uuid.UUID:
|
||||||
|
return uuid.uuid4()
|
||||||
|
|
||||||
|
|
||||||
|
class Base(DeclarativeBase):
|
||||||
|
pass
|
||||||
|
|
||||||
|
|
||||||
|
class PublicIdMixin:
|
||||||
|
"""Gives every table a UUID public id; integer PKs stay internal."""
|
||||||
|
|
||||||
|
id: Mapped[uuid.UUID] = mapped_column(Uuid, primary_key=True, default=new_uuid)
|
||||||
|
created_at: Mapped[datetime] = mapped_column(UTCDT(), default=utcnow, nullable=False)
|
||||||
|
# default= as well as onupdate=: SQLite enforces NOT NULL on insert
|
||||||
|
# (Postgres silently stored NULL for rows that never got an UPDATE).
|
||||||
|
updated_at: Mapped[datetime] = mapped_column(
|
||||||
|
UTCDT(), default=utcnow, onupdate=utcnow, nullable=False
|
||||||
|
)
|
||||||
33
backend/shonar/db/migrate.py
Normal file
33
backend/shonar/db/migrate.py
Normal file
|
|
@ -0,0 +1,33 @@
|
||||||
|
"""Programmatic schema migrations for the bundled desktop engine.
|
||||||
|
|
||||||
|
Server deployments run ``alembic upgrade head`` in their deploy flow;
|
||||||
|
the desktop app owns its SQLite file end-to-end, so the API applies
|
||||||
|
migrations itself at startup (``SHONAR_AUTO_MIGRATE=1``).
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import asyncio
|
||||||
|
import logging
|
||||||
|
from pathlib import Path
|
||||||
|
|
||||||
|
from alembic import command
|
||||||
|
from alembic.config import Config
|
||||||
|
|
||||||
|
logger = logging.getLogger("shonar.migrate")
|
||||||
|
|
||||||
|
# backend/ dir: migrations/ + alembic.ini live here relative to this file.
|
||||||
|
_BACKEND_DIR = Path(__file__).resolve().parents[2]
|
||||||
|
|
||||||
|
|
||||||
|
async def upgrade_head(database_url: str) -> None:
|
||||||
|
"""Run 'alembic upgrade head' for `database_url` without blocking the loop."""
|
||||||
|
|
||||||
|
def _run() -> None:
|
||||||
|
cfg = Config(str(_BACKEND_DIR / "alembic.ini"))
|
||||||
|
cfg.set_main_option("script_location", str(_BACKEND_DIR / "migrations"))
|
||||||
|
cfg.set_main_option("sqlalchemy.url", database_url)
|
||||||
|
command.upgrade(cfg, "head")
|
||||||
|
|
||||||
|
await asyncio.to_thread(_run)
|
||||||
|
logger.info("database schema migrated to head")
|
||||||
453
backend/shonar/db/models.py
Normal file
453
backend/shonar/db/models.py
Normal file
|
|
@ -0,0 +1,453 @@
|
||||||
|
"""All persistent models.
|
||||||
|
|
||||||
|
Design rules:
|
||||||
|
- Every client-visible identifier is a UUID (``PublicIdMixin.id``).
|
||||||
|
- Filesystem/storage paths are NEVER exposed to clients.
|
||||||
|
- Original uploads are immutable; processed audio lives in separate
|
||||||
|
``Asset`` rows and never replaces an original.
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import enum
|
||||||
|
import uuid
|
||||||
|
from datetime import datetime
|
||||||
|
|
||||||
|
# JSON that renders JSONB on PostgreSQL and plain JSON on SQLite (the
|
||||||
|
# desktop bundled-lite engine runs on SQLite; JSONB does not compile there).
|
||||||
|
from sqlalchemy import JSON as _JSON
|
||||||
|
from sqlalchemy import (
|
||||||
|
BigInteger,
|
||||||
|
Boolean,
|
||||||
|
Enum,
|
||||||
|
Float,
|
||||||
|
ForeignKey,
|
||||||
|
Index,
|
||||||
|
String,
|
||||||
|
Text,
|
||||||
|
UniqueConstraint,
|
||||||
|
)
|
||||||
|
from sqlalchemy.dialects.postgresql import JSONB
|
||||||
|
from sqlalchemy.orm import Mapped, mapped_column, relationship
|
||||||
|
|
||||||
|
from shonar.db.base import UTCDT, Base, PublicIdMixin, utcnow
|
||||||
|
|
||||||
|
JSONType = _JSON().with_variant(JSONB(), "postgresql")
|
||||||
|
|
||||||
|
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
# Users & auth
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
|
||||||
|
|
||||||
|
class User(Base, PublicIdMixin):
|
||||||
|
__tablename__ = "users"
|
||||||
|
|
||||||
|
email: Mapped[str] = mapped_column(String(320), unique=True, index=True, nullable=False)
|
||||||
|
password_hash: Mapped[str] = mapped_column(String(255), nullable=False)
|
||||||
|
display_name: Mapped[str | None] = mapped_column(String(120))
|
||||||
|
is_active: Mapped[bool] = mapped_column(Boolean, default=True, nullable=False)
|
||||||
|
# Feature switches the user controls (mirror of app settings, server truth
|
||||||
|
# for e.g. whether location metadata may be stored at all).
|
||||||
|
location_storage_enabled: Mapped[bool] = mapped_column(Boolean, default=False, nullable=False)
|
||||||
|
deleted_at: Mapped[datetime | None] = mapped_column(UTCDT())
|
||||||
|
|
||||||
|
devices: Mapped[list[Device]] = relationship(back_populates="user")
|
||||||
|
recordings: Mapped[list[Recording]] = relationship(back_populates="user")
|
||||||
|
|
||||||
|
|
||||||
|
class Device(Base, PublicIdMixin):
|
||||||
|
__tablename__ = "devices"
|
||||||
|
|
||||||
|
user_id: Mapped[uuid.UUID] = mapped_column(
|
||||||
|
ForeignKey("users.id", ondelete="CASCADE"), index=True, nullable=False
|
||||||
|
)
|
||||||
|
name: Mapped[str] = mapped_column(String(120), nullable=False)
|
||||||
|
platform: Mapped[str] = mapped_column(String(40), default="android", nullable=False)
|
||||||
|
last_seen_at: Mapped[datetime] = mapped_column(UTCDT(), default=utcnow)
|
||||||
|
revoked_at: Mapped[datetime | None] = mapped_column(UTCDT())
|
||||||
|
|
||||||
|
user: Mapped[User] = relationship(back_populates="devices")
|
||||||
|
|
||||||
|
|
||||||
|
class RefreshToken(Base, PublicIdMixin):
|
||||||
|
"""Rotating refresh tokens, stored hashed, grouped into families for
|
||||||
|
reuse detection. A refresh consumes one row and issues its replacement
|
||||||
|
with the same ``family``."""
|
||||||
|
|
||||||
|
__tablename__ = "refresh_tokens"
|
||||||
|
|
||||||
|
user_id: Mapped[uuid.UUID] = mapped_column(
|
||||||
|
ForeignKey("users.id", ondelete="CASCADE"), index=True, nullable=False
|
||||||
|
)
|
||||||
|
token_hash: Mapped[str] = mapped_column(String(128), unique=True, nullable=False)
|
||||||
|
family: Mapped[uuid.UUID] = mapped_column(index=True, nullable=False)
|
||||||
|
device_id: Mapped[uuid.UUID | None] = mapped_column(
|
||||||
|
ForeignKey("devices.id", ondelete="SET NULL")
|
||||||
|
)
|
||||||
|
expires_at: Mapped[datetime] = mapped_column(UTCDT(), nullable=False)
|
||||||
|
revoked_at: Mapped[datetime | None] = mapped_column(UTCDT())
|
||||||
|
replaced_by: Mapped[uuid.UUID | None] = mapped_column(
|
||||||
|
ForeignKey("refresh_tokens.id", ondelete="SET NULL")
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
# Recordings, assets, uploads
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
|
||||||
|
|
||||||
|
class ProcessingStatus(enum.StrEnum):
|
||||||
|
pending_upload = "pending_upload"
|
||||||
|
uploaded = "uploaded"
|
||||||
|
queued = "queued"
|
||||||
|
processing = "processing"
|
||||||
|
completed = "completed"
|
||||||
|
failed = "failed"
|
||||||
|
# No AI configured / requested — audio-only recording, fully usable.
|
||||||
|
ai_disabled = "ai_disabled"
|
||||||
|
|
||||||
|
|
||||||
|
class Recording(Base, PublicIdMixin):
|
||||||
|
__tablename__ = "recordings"
|
||||||
|
|
||||||
|
user_id: Mapped[uuid.UUID] = mapped_column(
|
||||||
|
ForeignKey("users.id", ondelete="CASCADE"), index=True, nullable=False
|
||||||
|
)
|
||||||
|
device_id: Mapped[uuid.UUID | None] = mapped_column(
|
||||||
|
ForeignKey("devices.id", ondelete="SET NULL")
|
||||||
|
)
|
||||||
|
# Client-generated idempotency id so retried uploads update, not duplicate.
|
||||||
|
client_recording_id: Mapped[str | None] = mapped_column(String(64), index=True)
|
||||||
|
|
||||||
|
title: Mapped[str] = mapped_column(String(300), nullable=False, default="Untitled recording")
|
||||||
|
recorded_at: Mapped[datetime] = mapped_column(UTCDT(), nullable=False)
|
||||||
|
duration_seconds: Mapped[float] = mapped_column(Float, default=0.0, nullable=False)
|
||||||
|
notes: Mapped[str | None] = mapped_column(Text)
|
||||||
|
# Location is stored ONLY when the user has explicitly enabled it.
|
||||||
|
latitude: Mapped[float | None] = mapped_column(Float)
|
||||||
|
longitude: Mapped[float | None] = mapped_column(Float)
|
||||||
|
location_accuracy_m: Mapped[float | None] = mapped_column(Float)
|
||||||
|
|
||||||
|
processing_status: Mapped[ProcessingStatus] = mapped_column(
|
||||||
|
Enum(
|
||||||
|
ProcessingStatus,
|
||||||
|
name="processing_status",
|
||||||
|
values_callable=lambda e: [m.value for m in e],
|
||||||
|
),
|
||||||
|
default=ProcessingStatus.pending_upload,
|
||||||
|
nullable=False,
|
||||||
|
index=True,
|
||||||
|
)
|
||||||
|
processing_error: Mapped[str | None] = mapped_column(Text)
|
||||||
|
# Effective faster-whisper model for this recording (e.g. "base").
|
||||||
|
# Set at finalize time from the per-recording override or the global
|
||||||
|
# default; NULL means "server default at processing time" (pre-model
|
||||||
|
# rows). Never rewritten: changing the default affects future rows only.
|
||||||
|
transcription_model: Mapped[str | None] = mapped_column(String(32))
|
||||||
|
# The original audio is the Asset row with kind=original for this
|
||||||
|
# recording (at most one, enforced by a partial unique index). Keeping
|
||||||
|
# the pointer one-directional avoids a recordings<->assets FK cycle.
|
||||||
|
deleted_at: Mapped[datetime | None] = mapped_column(UTCDT())
|
||||||
|
|
||||||
|
user: Mapped[User] = relationship(back_populates="recordings")
|
||||||
|
assets: Mapped[list[Asset]] = relationship(
|
||||||
|
back_populates="recording", cascade="all, delete-orphan"
|
||||||
|
)
|
||||||
|
transcripts: Mapped[list[Transcript]] = relationship(
|
||||||
|
back_populates="recording", cascade="all, delete-orphan"
|
||||||
|
)
|
||||||
|
summaries: Mapped[list[Summary]] = relationship(
|
||||||
|
back_populates="recording", cascade="all, delete-orphan"
|
||||||
|
)
|
||||||
|
tags: Mapped[list[Tag]] = relationship(secondary="recording_tags", back_populates="recordings")
|
||||||
|
|
||||||
|
__table_args__ = (
|
||||||
|
Index("ix_recordings_user_recorded", "user_id", "recorded_at"),
|
||||||
|
Index(
|
||||||
|
"ix_recordings_user_client_id",
|
||||||
|
"user_id",
|
||||||
|
"client_recording_id",
|
||||||
|
unique=True,
|
||||||
|
postgresql_where="client_recording_id IS NOT NULL",
|
||||||
|
),
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class AssetKind(enum.StrEnum):
|
||||||
|
original = "original" # user's uploaded file — NEVER modified
|
||||||
|
normalized = "normalized" # derivative for processing (ffmpeg)
|
||||||
|
export = "export" # generated export bundle
|
||||||
|
|
||||||
|
|
||||||
|
class Asset(Base, PublicIdMixin):
|
||||||
|
"""A stored binary object. ``storage_key`` is server-internal only."""
|
||||||
|
|
||||||
|
__tablename__ = "assets"
|
||||||
|
|
||||||
|
recording_id: Mapped[uuid.UUID | None] = mapped_column(
|
||||||
|
ForeignKey("recordings.id", ondelete="CASCADE"), index=True
|
||||||
|
)
|
||||||
|
user_id: Mapped[uuid.UUID] = mapped_column(
|
||||||
|
ForeignKey("users.id", ondelete="CASCADE"), index=True, nullable=False
|
||||||
|
)
|
||||||
|
kind: Mapped[AssetKind] = mapped_column(
|
||||||
|
Enum(AssetKind, name="asset_kind", values_callable=lambda e: [m.value for m in e]),
|
||||||
|
nullable=False,
|
||||||
|
)
|
||||||
|
storage_key: Mapped[str] = mapped_column(String(500), nullable=False)
|
||||||
|
mime_type: Mapped[str] = mapped_column(String(100), nullable=False)
|
||||||
|
size_bytes: Mapped[int] = mapped_column(BigInteger, nullable=False, default=0)
|
||||||
|
checksum_sha256: Mapped[str] = mapped_column(String(64), nullable=False)
|
||||||
|
|
||||||
|
recording: Mapped[Recording | None] = relationship(back_populates="assets")
|
||||||
|
|
||||||
|
__table_args__ = (
|
||||||
|
# At most one immutable "original" per recording.
|
||||||
|
Index(
|
||||||
|
"uq_assets_one_original",
|
||||||
|
"recording_id",
|
||||||
|
unique=True,
|
||||||
|
postgresql_where="kind = 'original'",
|
||||||
|
),
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class UploadSessionStatus(enum.StrEnum):
|
||||||
|
open = "open"
|
||||||
|
finalizing = "finalizing"
|
||||||
|
completed = "completed"
|
||||||
|
aborted = "aborted"
|
||||||
|
expired = "expired"
|
||||||
|
|
||||||
|
|
||||||
|
class UploadSession(Base, PublicIdMixin):
|
||||||
|
"""Chunked, resumable upload session."""
|
||||||
|
|
||||||
|
__tablename__ = "upload_sessions"
|
||||||
|
|
||||||
|
user_id: Mapped[uuid.UUID] = mapped_column(
|
||||||
|
ForeignKey("users.id", ondelete="CASCADE"), index=True, nullable=False
|
||||||
|
)
|
||||||
|
recording_id: Mapped[uuid.UUID | None] = mapped_column(
|
||||||
|
ForeignKey("recordings.id", ondelete="CASCADE")
|
||||||
|
)
|
||||||
|
client_recording_id: Mapped[str | None] = mapped_column(String(64))
|
||||||
|
title: Mapped[str | None] = mapped_column(String(300))
|
||||||
|
# Optional per-recording transcription model override chosen on the
|
||||||
|
# upload screen. Validated at finalize time; the finalize body wins when
|
||||||
|
# both specify one.
|
||||||
|
transcription_model: Mapped[str | None] = mapped_column(String(32))
|
||||||
|
declared_mime_type: Mapped[str] = mapped_column(String(100), nullable=False)
|
||||||
|
declared_size_bytes: Mapped[int] = mapped_column(BigInteger, nullable=False)
|
||||||
|
chunk_size_bytes: Mapped[int] = mapped_column(
|
||||||
|
BigInteger, nullable=False, default=8 * 1024 * 1024
|
||||||
|
)
|
||||||
|
status: Mapped[UploadSessionStatus] = mapped_column(
|
||||||
|
Enum(
|
||||||
|
UploadSessionStatus,
|
||||||
|
name="upload_session_status",
|
||||||
|
values_callable=lambda e: [m.value for m in e],
|
||||||
|
),
|
||||||
|
default=UploadSessionStatus.open,
|
||||||
|
nullable=False,
|
||||||
|
)
|
||||||
|
expires_at: Mapped[datetime] = mapped_column(UTCDT(), nullable=False)
|
||||||
|
completed_asset_id: Mapped[uuid.UUID | None] = mapped_column(
|
||||||
|
ForeignKey("assets.id", ondelete="SET NULL")
|
||||||
|
)
|
||||||
|
|
||||||
|
chunks: Mapped[list[UploadChunk]] = relationship(
|
||||||
|
back_populates="session", cascade="all, delete-orphan"
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class UploadChunk(Base):
|
||||||
|
__tablename__ = "upload_chunks"
|
||||||
|
|
||||||
|
id: Mapped[int] = mapped_column(primary_key=True)
|
||||||
|
session_id: Mapped[uuid.UUID] = mapped_column(
|
||||||
|
ForeignKey("upload_sessions.id", ondelete="CASCADE"), index=True, nullable=False
|
||||||
|
)
|
||||||
|
chunk_index: Mapped[int] = mapped_column(nullable=False)
|
||||||
|
size_bytes: Mapped[int] = mapped_column(BigInteger, nullable=False)
|
||||||
|
checksum_sha256: Mapped[str] = mapped_column(String(64), nullable=False)
|
||||||
|
storage_key: Mapped[str] = mapped_column(String(500), nullable=False)
|
||||||
|
created_at: Mapped[datetime] = mapped_column(UTCDT(), default=utcnow)
|
||||||
|
|
||||||
|
session: Mapped[UploadSession] = relationship(back_populates="chunks")
|
||||||
|
|
||||||
|
__table_args__ = (UniqueConstraint("session_id", "chunk_index"),)
|
||||||
|
|
||||||
|
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
# Transcript / summary / tags
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
|
||||||
|
|
||||||
|
class Transcript(Base, PublicIdMixin):
|
||||||
|
__tablename__ = "transcripts"
|
||||||
|
|
||||||
|
recording_id: Mapped[uuid.UUID] = mapped_column(
|
||||||
|
ForeignKey("recordings.id", ondelete="CASCADE"), index=True, nullable=False
|
||||||
|
)
|
||||||
|
# Versioned: regenerated transcripts supersede older rows; the newest
|
||||||
|
# non-superseded row is authoritative. ``edited_by_user`` rows win.
|
||||||
|
version: Mapped[int] = mapped_column(default=1, nullable=False)
|
||||||
|
superseded_at: Mapped[datetime | None] = mapped_column(UTCDT())
|
||||||
|
language: Mapped[str | None] = mapped_column(String(16))
|
||||||
|
provider: Mapped[str] = mapped_column(String(80), nullable=False, default="manual")
|
||||||
|
model: Mapped[str | None] = mapped_column(String(120))
|
||||||
|
# Full text, plus segments: [{"start": 0.0, "end": 2.5, "text": "...",
|
||||||
|
# "speaker": "S1"|null}, ...]
|
||||||
|
text: Mapped[str] = mapped_column(Text, nullable=False, default="")
|
||||||
|
segments: Mapped[list | None] = mapped_column(JSONType)
|
||||||
|
edited_by_user: Mapped[bool] = mapped_column(Boolean, default=False, nullable=False)
|
||||||
|
|
||||||
|
recording: Mapped[Recording] = relationship(back_populates="transcripts")
|
||||||
|
|
||||||
|
__table_args__ = (Index("ix_transcripts_recording_version", "recording_id", "version"),)
|
||||||
|
|
||||||
|
|
||||||
|
class Summary(Base, PublicIdMixin):
|
||||||
|
"""Structured AI summary — user-editable.
|
||||||
|
|
||||||
|
JSON shape of ``content``:
|
||||||
|
{
|
||||||
|
"short": str,
|
||||||
|
"detailed": str,
|
||||||
|
"key_points": [str],
|
||||||
|
"decisions": [str],
|
||||||
|
"action_items": [str],
|
||||||
|
"questions": [str]
|
||||||
|
}
|
||||||
|
"""
|
||||||
|
|
||||||
|
__tablename__ = "summaries"
|
||||||
|
|
||||||
|
recording_id: Mapped[uuid.UUID] = mapped_column(
|
||||||
|
ForeignKey("recordings.id", ondelete="CASCADE"), index=True, nullable=False
|
||||||
|
)
|
||||||
|
version: Mapped[int] = mapped_column(default=1, nullable=False)
|
||||||
|
superseded_at: Mapped[datetime | None] = mapped_column(UTCDT())
|
||||||
|
provider: Mapped[str] = mapped_column(String(80), nullable=False, default="manual")
|
||||||
|
model: Mapped[str | None] = mapped_column(String(120))
|
||||||
|
content: Mapped[dict] = mapped_column(JSONType, nullable=False, default=dict)
|
||||||
|
edited_by_user: Mapped[bool] = mapped_column(Boolean, default=False, nullable=False)
|
||||||
|
|
||||||
|
recording: Mapped[Recording] = relationship(back_populates="summaries")
|
||||||
|
|
||||||
|
|
||||||
|
class Tag(Base, PublicIdMixin):
|
||||||
|
__tablename__ = "tags"
|
||||||
|
|
||||||
|
user_id: Mapped[uuid.UUID] = mapped_column(
|
||||||
|
ForeignKey("users.id", ondelete="CASCADE"), index=True, nullable=False
|
||||||
|
)
|
||||||
|
name: Mapped[str] = mapped_column(String(80), nullable=False)
|
||||||
|
|
||||||
|
recordings: Mapped[list[Recording]] = relationship(
|
||||||
|
secondary="recording_tags", back_populates="tags"
|
||||||
|
)
|
||||||
|
|
||||||
|
__table_args__ = (UniqueConstraint("user_id", "name"),)
|
||||||
|
|
||||||
|
|
||||||
|
class RecordingTag(Base):
|
||||||
|
__tablename__ = "recording_tags"
|
||||||
|
|
||||||
|
recording_id: Mapped[uuid.UUID] = mapped_column(
|
||||||
|
ForeignKey("recordings.id", ondelete="CASCADE"), primary_key=True
|
||||||
|
)
|
||||||
|
tag_id: Mapped[uuid.UUID] = mapped_column(
|
||||||
|
ForeignKey("tags.id", ondelete="CASCADE"), primary_key=True
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
# Processing jobs
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
|
||||||
|
|
||||||
|
class JobStatus(enum.StrEnum):
|
||||||
|
queued = "queued"
|
||||||
|
running = "running"
|
||||||
|
succeeded = "succeeded"
|
||||||
|
failed = "failed"
|
||||||
|
skipped = "skipped" # e.g. AI disabled, or user-edited output protected
|
||||||
|
|
||||||
|
|
||||||
|
class JobType(enum.StrEnum):
|
||||||
|
normalize_audio = "normalize_audio"
|
||||||
|
transcribe = "transcribe"
|
||||||
|
summarize = "summarize"
|
||||||
|
|
||||||
|
|
||||||
|
class ProcessingJob(Base, PublicIdMixin):
|
||||||
|
"""One pipeline step for one recording. Idempotent: reruns overwrite
|
||||||
|
derived outputs (unless user-edited) and never touch originals."""
|
||||||
|
|
||||||
|
__tablename__ = "processing_jobs"
|
||||||
|
|
||||||
|
recording_id: Mapped[uuid.UUID] = mapped_column(
|
||||||
|
ForeignKey("recordings.id", ondelete="CASCADE"), index=True, nullable=False
|
||||||
|
)
|
||||||
|
job_type: Mapped[JobType] = mapped_column(
|
||||||
|
Enum(JobType, name="job_type", values_callable=lambda e: [m.value for m in e]),
|
||||||
|
nullable=False,
|
||||||
|
)
|
||||||
|
status: Mapped[JobStatus] = mapped_column(
|
||||||
|
Enum(JobStatus, name="job_status", values_callable=lambda e: [m.value for m in e]),
|
||||||
|
default=JobStatus.queued,
|
||||||
|
nullable=False,
|
||||||
|
index=True,
|
||||||
|
)
|
||||||
|
attempt: Mapped[int] = mapped_column(default=0, nullable=False)
|
||||||
|
max_attempts: Mapped[int] = mapped_column(default=3, nullable=False)
|
||||||
|
error: Mapped[str | None] = mapped_column(Text)
|
||||||
|
# Fine-grained phase for progress display (e.g. transcribe jobs report
|
||||||
|
# "loading-model" then "transcribing"). Nullable: older rows predate it.
|
||||||
|
# Never parsed by pipeline logic — display only.
|
||||||
|
stage: Mapped[str | None] = mapped_column(String(32))
|
||||||
|
# 0-100 work estimate within the current stage, when known.
|
||||||
|
progress: Mapped[int | None] = mapped_column()
|
||||||
|
started_at: Mapped[datetime | None] = mapped_column(UTCDT())
|
||||||
|
finished_at: Mapped[datetime | None] = mapped_column(UTCDT())
|
||||||
|
# Opaque arq task handle for observability.
|
||||||
|
task_handle: Mapped[str | None] = mapped_column(String(120))
|
||||||
|
|
||||||
|
__table_args__ = (Index("ix_processing_jobs_recording_type", "recording_id", "job_type"),)
|
||||||
|
|
||||||
|
|
||||||
|
class ExportJob(Base, PublicIdMixin):
|
||||||
|
__tablename__ = "export_jobs"
|
||||||
|
|
||||||
|
user_id: Mapped[uuid.UUID] = mapped_column(
|
||||||
|
ForeignKey("users.id", ondelete="CASCADE"), index=True, nullable=False
|
||||||
|
)
|
||||||
|
recording_id: Mapped[uuid.UUID] = mapped_column(
|
||||||
|
ForeignKey("recordings.id", ondelete="CASCADE"), index=True, nullable=False
|
||||||
|
)
|
||||||
|
# "audio", "transcript_txt", "notes_md", "bundle_zip"
|
||||||
|
export_type: Mapped[str] = mapped_column(String(40), nullable=False)
|
||||||
|
status: Mapped[JobStatus] = mapped_column(
|
||||||
|
Enum(JobStatus, name="export_job_status", values_callable=lambda e: [m.value for m in e]),
|
||||||
|
default=JobStatus.queued,
|
||||||
|
nullable=False,
|
||||||
|
)
|
||||||
|
asset_id: Mapped[uuid.UUID | None] = mapped_column(ForeignKey("assets.id", ondelete="SET NULL"))
|
||||||
|
error: Mapped[str | None] = mapped_column(Text)
|
||||||
|
|
||||||
|
|
||||||
|
# Ensure full-text search columns exist on Postgres (added via migration as
|
||||||
|
# tsvector generated columns; see migrations/versions/*_fts.py).
|
||||||
|
|
||||||
|
|
||||||
|
class AppSetting(Base, PublicIdMixin):
|
||||||
|
"""Server-wide settings editable at runtime (global transcription model
|
||||||
|
default, …). Single row per key; the desktop Settings page writes here."""
|
||||||
|
|
||||||
|
__tablename__ = "app_settings"
|
||||||
|
|
||||||
|
key: Mapped[str] = mapped_column(String(120), nullable=False, unique=True)
|
||||||
|
value: Mapped[dict] = mapped_column(JSONType, nullable=False, default=dict)
|
||||||
76
backend/shonar/db/session.py
Normal file
76
backend/shonar/db/session.py
Normal file
|
|
@ -0,0 +1,76 @@
|
||||||
|
"""Async SQLAlchemy engine/session management."""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from collections.abc import AsyncIterator
|
||||||
|
|
||||||
|
from sqlalchemy.ext.asyncio import AsyncSession, async_sessionmaker, create_async_engine
|
||||||
|
|
||||||
|
from shonar.core.config import get_settings
|
||||||
|
|
||||||
|
_engine = None
|
||||||
|
_session_factory: async_sessionmaker[AsyncSession] | None = None
|
||||||
|
|
||||||
|
|
||||||
|
def get_engine():
|
||||||
|
global _engine, _session_factory
|
||||||
|
if _engine is None:
|
||||||
|
settings = get_settings()
|
||||||
|
kwargs: dict = {"pool_pre_ping": True}
|
||||||
|
url = settings.database_url
|
||||||
|
# SQLite (tests, desktop bundled engine) does not support the pg
|
||||||
|
# pool sizing kwargs.
|
||||||
|
if url.startswith("sqlite"):
|
||||||
|
kwargs = {}
|
||||||
|
else:
|
||||||
|
kwargs.update(pool_size=settings.db_pool_size, max_overflow=settings.db_max_overflow)
|
||||||
|
_engine = create_async_engine(url, **kwargs)
|
||||||
|
if url.startswith("sqlite"):
|
||||||
|
_configure_sqlite(_engine)
|
||||||
|
_session_factory = async_sessionmaker(_engine, expire_on_commit=False)
|
||||||
|
return _engine
|
||||||
|
|
||||||
|
|
||||||
|
def _configure_sqlite(engine) -> None: # noqa: ANN001
|
||||||
|
"""Per-connection SQLite pragmas.
|
||||||
|
|
||||||
|
foreign_keys is OFF by default in SQLite; the schema leans on
|
||||||
|
ON DELETE CASCADE, so every connection must enable it. WAL lets the
|
||||||
|
inline-queue writer and API readers coexist without SQLITE_BUSY
|
||||||
|
storms on the desktop box.
|
||||||
|
"""
|
||||||
|
from sqlalchemy import event
|
||||||
|
|
||||||
|
@event.listens_for(engine.sync_engine, "connect")
|
||||||
|
def _pragmas(dbapi_conn, _record): # noqa: ANN001
|
||||||
|
cur = dbapi_conn.cursor()
|
||||||
|
cur.execute("PRAGMA foreign_keys=ON")
|
||||||
|
cur.execute("PRAGMA journal_mode=WAL")
|
||||||
|
cur.execute("PRAGMA busy_timeout=5000")
|
||||||
|
cur.close()
|
||||||
|
|
||||||
|
|
||||||
|
async def dispose_engine() -> None:
|
||||||
|
global _engine, _session_factory
|
||||||
|
if _engine is not None:
|
||||||
|
await _engine.dispose()
|
||||||
|
_engine = None
|
||||||
|
_session_factory = None
|
||||||
|
|
||||||
|
|
||||||
|
async def get_session() -> AsyncIterator[AsyncSession]:
|
||||||
|
"""FastAPI dependency yielding a database session."""
|
||||||
|
assert _session_factory is not None, "engine not initialised"
|
||||||
|
async with _session_factory() as session:
|
||||||
|
try:
|
||||||
|
yield session
|
||||||
|
await session.commit()
|
||||||
|
except Exception:
|
||||||
|
await session.rollback()
|
||||||
|
raise
|
||||||
|
|
||||||
|
|
||||||
|
def session_factory() -> async_sessionmaker[AsyncSession]:
|
||||||
|
"""Shareable session factory for the worker (outside requests)."""
|
||||||
|
assert _session_factory is not None, "engine not initialised"
|
||||||
|
return _session_factory
|
||||||
79
backend/shonar/main.py
Normal file
79
backend/shonar/main.py
Normal file
|
|
@ -0,0 +1,79 @@
|
||||||
|
"""S.H.O.N.A.R. FastAPI application entrypoint.
|
||||||
|
|
||||||
|
Run (dev): uvicorn shonar.main:app --reload --port 8000
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import logging
|
||||||
|
from contextlib import asynccontextmanager
|
||||||
|
|
||||||
|
from fastapi import FastAPI, Request
|
||||||
|
from fastapi.responses import JSONResponse
|
||||||
|
from slowapi import _rate_limit_exceeded_handler
|
||||||
|
from slowapi.errors import RateLimitExceeded
|
||||||
|
|
||||||
|
from shonar import __version__
|
||||||
|
from shonar.api.v1 import api_router
|
||||||
|
from shonar.core.config import get_settings
|
||||||
|
from shonar.core.ratelimit import limiter
|
||||||
|
from shonar.db.session import dispose_engine, get_engine
|
||||||
|
|
||||||
|
logging.basicConfig(level=logging.INFO)
|
||||||
|
logger = logging.getLogger("shonar")
|
||||||
|
|
||||||
|
|
||||||
|
@asynccontextmanager
|
||||||
|
async def lifespan(app: FastAPI):
|
||||||
|
settings = get_settings()
|
||||||
|
get_engine() # validate URL parses; connections are lazy
|
||||||
|
for warning in settings.validate_production():
|
||||||
|
logger.warning("CONFIG: %s", warning)
|
||||||
|
if settings.auto_migrate:
|
||||||
|
from shonar.db.migrate import upgrade_head
|
||||||
|
|
||||||
|
await upgrade_head(settings.database_url)
|
||||||
|
inline = settings.queue_backend.strip().lower() == "inline"
|
||||||
|
if inline:
|
||||||
|
from shonar.services import inline_queue
|
||||||
|
|
||||||
|
await inline_queue.start()
|
||||||
|
yield
|
||||||
|
if inline:
|
||||||
|
from shonar.services import inline_queue
|
||||||
|
|
||||||
|
await inline_queue.stop()
|
||||||
|
await dispose_engine()
|
||||||
|
|
||||||
|
|
||||||
|
app = FastAPI(
|
||||||
|
title="SHONAR API",
|
||||||
|
description=(
|
||||||
|
"SHONAR — Self-hosted Oral Notes and Audio Recorder. REST API. "
|
||||||
|
"All data stays on your server."
|
||||||
|
),
|
||||||
|
version=__version__,
|
||||||
|
lifespan=lifespan,
|
||||||
|
# Interactive docs at the conventional paths.
|
||||||
|
docs_url="/docs",
|
||||||
|
redoc_url="/redoc",
|
||||||
|
openapi_url="/openapi.json",
|
||||||
|
)
|
||||||
|
|
||||||
|
app.state.limiter = limiter
|
||||||
|
app.add_exception_handler(RateLimitExceeded, _rate_limit_exceeded_handler)
|
||||||
|
|
||||||
|
|
||||||
|
@app.exception_handler(Exception)
|
||||||
|
async def unhandled_exception_handler(request: Request, exc: Exception):
|
||||||
|
"""Safe error surface: log details server-side, never leak internals."""
|
||||||
|
logger.exception("Unhandled error on %s %s", request.method, request.url.path)
|
||||||
|
return JSONResponse(status_code=500, content={"detail": "Internal server error"})
|
||||||
|
|
||||||
|
|
||||||
|
app.include_router(api_router)
|
||||||
|
|
||||||
|
|
||||||
|
@app.get("/")
|
||||||
|
async def root():
|
||||||
|
return {"app": "SHONAR", "docs": "/docs", "health": "/api/v1/healthz"}
|
||||||
0
backend/shonar/services/__init__.py
Normal file
0
backend/shonar/services/__init__.py
Normal file
160
backend/shonar/services/ai/__init__.py
Normal file
160
backend/shonar/services/ai/__init__.py
Normal file
|
|
@ -0,0 +1,160 @@
|
||||||
|
"""AI provider interfaces (M7).
|
||||||
|
|
||||||
|
Two independent axes, both optional and both configured only through
|
||||||
|
environment variables (never hard-coded keys):
|
||||||
|
|
||||||
|
- transcription: none | whisper_http | faster_whisper
|
||||||
|
- LLM (summary/action items): none | openai_compat | ollama
|
||||||
|
|
||||||
|
"none" is a first-class choice: recording, sync, playback, and manual
|
||||||
|
transcripts work with no AI configured at all. The pipeline treats a
|
||||||
|
missing provider as "skip this stage", never as an error.
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from dataclasses import dataclass, field
|
||||||
|
from typing import Protocol
|
||||||
|
|
||||||
|
from shonar.core.config import Settings
|
||||||
|
|
||||||
|
|
||||||
|
class AIError(Exception):
|
||||||
|
"""Base for AI failures. Messages must be user-safe: they surface in
|
||||||
|
``processing_error`` and therefore on screens."""
|
||||||
|
|
||||||
|
|
||||||
|
class ProviderConfigError(AIError):
|
||||||
|
"""Persistent misconfiguration (bad credentials, unknown model, missing
|
||||||
|
dependency). Fails the job immediately — retrying cannot help."""
|
||||||
|
|
||||||
|
|
||||||
|
class ProviderTransientError(AIError):
|
||||||
|
"""May succeed on retry (timeouts, 429/5xx). The worker requeues these
|
||||||
|
up to the job's max attempts."""
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(frozen=True)
|
||||||
|
class Segment:
|
||||||
|
start: float
|
||||||
|
end: float
|
||||||
|
text: str
|
||||||
|
speaker: str | None = None
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(frozen=True)
|
||||||
|
class TranscriptResult:
|
||||||
|
text: str
|
||||||
|
language: str | None
|
||||||
|
segments: list[Segment] = field(default_factory=list)
|
||||||
|
model: str = ""
|
||||||
|
|
||||||
|
|
||||||
|
class TranscriptionProvider(Protocol):
|
||||||
|
name: str
|
||||||
|
|
||||||
|
async def transcribe(
|
||||||
|
self,
|
||||||
|
audio: bytes,
|
||||||
|
mime: str,
|
||||||
|
*,
|
||||||
|
language_hint: str | None = None,
|
||||||
|
on_progress=None, # optional Callable[[int], None], 0..99
|
||||||
|
) -> TranscriptResult: ...
|
||||||
|
|
||||||
|
|
||||||
|
SUMMARY_KEYS = ("short", "detailed", "key_points", "decisions", "action_items", "questions")
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(frozen=True)
|
||||||
|
class SummaryResult:
|
||||||
|
short: str = ""
|
||||||
|
detailed: str = ""
|
||||||
|
key_points: tuple[str, ...] = ()
|
||||||
|
decisions: tuple[str, ...] = ()
|
||||||
|
action_items: tuple[str, ...] = ()
|
||||||
|
questions: tuple[str, ...] = ()
|
||||||
|
model: str = ""
|
||||||
|
|
||||||
|
def to_dict(self) -> dict:
|
||||||
|
return {
|
||||||
|
"short": self.short,
|
||||||
|
"detailed": self.detailed,
|
||||||
|
"key_points": list(self.key_points),
|
||||||
|
"decisions": list(self.decisions),
|
||||||
|
"action_items": list(self.action_items),
|
||||||
|
"questions": list(self.questions),
|
||||||
|
}
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def from_dict(cls, raw: dict, model: str = "") -> SummaryResult:
|
||||||
|
def text(key: str) -> str:
|
||||||
|
v = raw.get(key)
|
||||||
|
return v if isinstance(v, str) else ""
|
||||||
|
|
||||||
|
def strs(key: str) -> tuple[str, ...]:
|
||||||
|
v = raw.get(key)
|
||||||
|
if not isinstance(v, list):
|
||||||
|
return ()
|
||||||
|
return tuple(s for s in v if isinstance(s, str) and s.strip())
|
||||||
|
|
||||||
|
return cls(
|
||||||
|
short=text("short"),
|
||||||
|
detailed=text("detailed"),
|
||||||
|
key_points=strs("key_points"),
|
||||||
|
decisions=strs("decisions"),
|
||||||
|
action_items=strs("action_items"),
|
||||||
|
questions=strs("questions"),
|
||||||
|
model=model,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class LlmProvider(Protocol):
|
||||||
|
name: str
|
||||||
|
|
||||||
|
async def summarize(self, transcript: str, *, title: str | None = None) -> SummaryResult: ...
|
||||||
|
|
||||||
|
|
||||||
|
def get_transcription_provider(settings: Settings) -> TranscriptionProvider | None:
|
||||||
|
"""None means "transcription stage skipped", never an error."""
|
||||||
|
kind = settings.transcription_provider.strip().lower()
|
||||||
|
if kind in ("", "none"):
|
||||||
|
return None
|
||||||
|
if kind == "whisper_http":
|
||||||
|
from shonar.services.ai.whisper_http import WhisperHttpProvider
|
||||||
|
|
||||||
|
return WhisperHttpProvider(
|
||||||
|
base_url=settings.transcription_base_url,
|
||||||
|
model=settings.transcription_model,
|
||||||
|
api_key=settings.transcription_api_key,
|
||||||
|
)
|
||||||
|
if kind == "faster_whisper":
|
||||||
|
from shonar.services.ai.faster_whisper import FasterWhisperProvider
|
||||||
|
|
||||||
|
return FasterWhisperProvider(model=settings.transcription_model)
|
||||||
|
raise ProviderConfigError(
|
||||||
|
f"Unknown transcription provider: {settings.transcription_provider!r}"
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def get_llm_provider(settings: Settings) -> LlmProvider | None:
|
||||||
|
"""None means "summary stage skipped", never an error."""
|
||||||
|
kind = settings.llm_provider.strip().lower()
|
||||||
|
if kind in ("", "none"):
|
||||||
|
return None
|
||||||
|
if kind == "openai_compat":
|
||||||
|
from shonar.services.ai.openai_compat import OpenAICompatProvider
|
||||||
|
|
||||||
|
return OpenAICompatProvider(
|
||||||
|
base_url=settings.llm_base_url,
|
||||||
|
model=settings.llm_model,
|
||||||
|
api_key=settings.llm_api_key,
|
||||||
|
)
|
||||||
|
if kind == "ollama":
|
||||||
|
from shonar.services.ai.ollama import OllamaProvider
|
||||||
|
|
||||||
|
return OllamaProvider(
|
||||||
|
base_url=settings.llm_base_url,
|
||||||
|
model=settings.llm_model,
|
||||||
|
)
|
||||||
|
raise ProviderConfigError(f"Unknown LLM provider: {settings.llm_provider!r}")
|
||||||
35
backend/shonar/services/ai/_llm.py
Normal file
35
backend/shonar/services/ai/_llm.py
Normal file
|
|
@ -0,0 +1,35 @@
|
||||||
|
"""Shared summary contract: one system prompt, one JSON shape.
|
||||||
|
|
||||||
|
The model replies with JSON only; partial replies are accepted and missing
|
||||||
|
keys default to empty (a terse-but-valid summary beats a failed job).
|
||||||
|
Transcript input is truncated to bound context — a way in, not a report.
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from shonar.services.ai import SummaryResult
|
||||||
|
|
||||||
|
SYSTEM_PROMPT = (
|
||||||
|
"You summarize voice recordings for the speaker's own later reference. "
|
||||||
|
"Reply with JSON only, exactly these keys: "
|
||||||
|
'{"short": "1-2 sentences", "detailed": "a faithful paragraph", '
|
||||||
|
'"key_points": [], "decisions": [], "action_items": [], "questions": []}. '
|
||||||
|
"Empty arrays when absent. Never invent names, dates, or commitments "
|
||||||
|
"not stated in the transcript."
|
||||||
|
)
|
||||||
|
|
||||||
|
MAX_TRANSCRIPT_CHARS = 12_000
|
||||||
|
|
||||||
|
|
||||||
|
def build_user_message(transcript: str, title: str | None) -> str:
|
||||||
|
text = transcript[:MAX_TRANSCRIPT_CHARS]
|
||||||
|
if len(transcript) > MAX_TRANSCRIPT_CHARS:
|
||||||
|
text += f"\n\n[truncated from {len(transcript)} chars]"
|
||||||
|
head = f'Title: "{title}"\n\n' if title else ""
|
||||||
|
return head + "Transcript:\n" + text
|
||||||
|
|
||||||
|
|
||||||
|
def parse_summary(data: object, model: str) -> SummaryResult:
|
||||||
|
if not isinstance(data, dict):
|
||||||
|
return SummaryResult(model=model)
|
||||||
|
return SummaryResult.from_dict(data, model=model)
|
||||||
108
backend/shonar/services/ai/faster_whisper.py
Normal file
108
backend/shonar/services/ai/faster_whisper.py
Normal file
|
|
@ -0,0 +1,108 @@
|
||||||
|
"""Local transcription via faster-whisper (optional dependency).
|
||||||
|
|
||||||
|
Runs fully on this machine: audio never leaves the server for this stage.
|
||||||
|
The import is lazy so the base install (and every test run) works without
|
||||||
|
the heavyweight dependency.
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import asyncio
|
||||||
|
import tempfile
|
||||||
|
import threading
|
||||||
|
from pathlib import Path
|
||||||
|
|
||||||
|
from shonar.services.ai import ProviderConfigError, Segment, TranscriptResult
|
||||||
|
from shonar.services.ai.model_registry import (
|
||||||
|
download_instructions,
|
||||||
|
is_model_downloaded,
|
||||||
|
validate_model_name,
|
||||||
|
)
|
||||||
|
|
||||||
|
# One loaded model per name, shared across jobs in this worker process.
|
||||||
|
# ctranslate2 inference is thread-safe; a lock serializes first-load only.
|
||||||
|
_MODEL_CACHE: dict[str, object] = {}
|
||||||
|
_MODEL_CACHE_LOCK = threading.Lock()
|
||||||
|
|
||||||
|
|
||||||
|
def _load_model(name: str) -> object:
|
||||||
|
from faster_whisper import WhisperModel
|
||||||
|
|
||||||
|
with _MODEL_CACHE_LOCK:
|
||||||
|
model = _MODEL_CACHE.get(name)
|
||||||
|
if model is None:
|
||||||
|
model = WhisperModel(name, device="auto")
|
||||||
|
_MODEL_CACHE[name] = model
|
||||||
|
return model
|
||||||
|
|
||||||
|
|
||||||
|
def ensure_model_available(name: str) -> None:
|
||||||
|
"""Fail fast with download instructions instead of triggering a surprise
|
||||||
|
multi-GB download inside a transcription job."""
|
||||||
|
if not is_model_downloaded(name):
|
||||||
|
raise ProviderConfigError(download_instructions(name))
|
||||||
|
|
||||||
|
|
||||||
|
class FasterWhisperProvider:
|
||||||
|
name = "faster_whisper"
|
||||||
|
|
||||||
|
def __init__(self, model: str = "base") -> None:
|
||||||
|
try:
|
||||||
|
import faster_whisper # noqa: F401
|
||||||
|
except ImportError as e:
|
||||||
|
raise ProviderConfigError(
|
||||||
|
"faster-whisper is not installed (pip install shonar-backend[faster-whisper])."
|
||||||
|
) from e
|
||||||
|
self.model = validate_model_name(model or "base")
|
||||||
|
|
||||||
|
async def transcribe(
|
||||||
|
self,
|
||||||
|
audio: bytes,
|
||||||
|
mime: str,
|
||||||
|
*,
|
||||||
|
language_hint: str | None = None,
|
||||||
|
on_progress=None, # Callable[[int], None] | None — 0..99 percent
|
||||||
|
) -> TranscriptResult:
|
||||||
|
# faster-whisper is blocking CPU work: keep it off the event loop.
|
||||||
|
return await asyncio.to_thread(self._run, audio, language_hint, on_progress)
|
||||||
|
|
||||||
|
def _run(
|
||||||
|
self,
|
||||||
|
audio: bytes,
|
||||||
|
language_hint: str | None,
|
||||||
|
on_progress=None,
|
||||||
|
) -> TranscriptResult:
|
||||||
|
path: Path | None = None
|
||||||
|
try:
|
||||||
|
with tempfile.NamedTemporaryFile(suffix=".m4a", delete=False) as f:
|
||||||
|
f.write(audio)
|
||||||
|
path = Path(f.name)
|
||||||
|
ensure_model_available(self.model)
|
||||||
|
model = _load_model(self.model)
|
||||||
|
segments_iter, info = model.transcribe( # type: ignore[union-attr]
|
||||||
|
str(path),
|
||||||
|
beam_size=5,
|
||||||
|
language=language_hint,
|
||||||
|
)
|
||||||
|
duration = float(getattr(info, "duration", 0.0) or 0.0)
|
||||||
|
segments = []
|
||||||
|
last_pct = -1
|
||||||
|
for s in segments_iter:
|
||||||
|
segments.append(Segment(start=s.start, end=s.end, text=s.text.strip()))
|
||||||
|
if on_progress is not None and duration > 0:
|
||||||
|
# Throttle: report only on whole-percent gains. Capped
|
||||||
|
# at 99 — the caller commits 100 when the row finishes.
|
||||||
|
pct = min(99, int(s.end / duration * 100))
|
||||||
|
if pct > last_pct:
|
||||||
|
last_pct = pct
|
||||||
|
on_progress(pct)
|
||||||
|
text = " ".join(s.text for s in segments).strip()
|
||||||
|
return TranscriptResult(
|
||||||
|
text=text,
|
||||||
|
language=getattr(info, "language", None),
|
||||||
|
segments=segments,
|
||||||
|
model=self.model,
|
||||||
|
)
|
||||||
|
finally:
|
||||||
|
if path is not None:
|
||||||
|
path.unlink(missing_ok=True)
|
||||||
165
backend/shonar/services/ai/model_registry.py
Normal file
165
backend/shonar/services/ai/model_registry.py
Normal file
|
|
@ -0,0 +1,165 @@
|
||||||
|
"""Transcription model registry (Stage 1).
|
||||||
|
|
||||||
|
The supported faster-whisper sizes, their display metadata, validation, and
|
||||||
|
local availability checks. Nothing here downloads anything: faster-whisper
|
||||||
|
fetches from HuggingFace on first use, so "downloaded" is answered by
|
||||||
|
inspecting the HF hub cache, and "available" additionally requires the
|
||||||
|
faster-whisper package itself.
|
||||||
|
|
||||||
|
No silent substitution anywhere: unknown names are rejected, missing
|
||||||
|
downloads fail fast with instructions.
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import os
|
||||||
|
from dataclasses import dataclass
|
||||||
|
from pathlib import Path
|
||||||
|
|
||||||
|
from sqlalchemy import select
|
||||||
|
from sqlalchemy.ext.asyncio import AsyncSession
|
||||||
|
|
||||||
|
from shonar.db.models import AppSetting
|
||||||
|
from shonar.services.ai import ProviderConfigError
|
||||||
|
|
||||||
|
DEFAULT_MODEL = "base"
|
||||||
|
DEFAULT_MODEL_KEY = "transcription.default_model"
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(frozen=True)
|
||||||
|
class TranscriptionModelInfo:
|
||||||
|
name: str
|
||||||
|
display_name: str
|
||||||
|
description: str
|
||||||
|
params: str
|
||||||
|
approx_memory: str
|
||||||
|
relative_speed: str
|
||||||
|
|
||||||
|
|
||||||
|
SUPPORTED_TRANSCRIPTION_MODELS: dict[str, TranscriptionModelInfo] = {
|
||||||
|
"tiny": TranscriptionModelInfo(
|
||||||
|
name="tiny",
|
||||||
|
display_name="Tiny",
|
||||||
|
description="Fastest and lightest. Good for quick drafts and slow machines.",
|
||||||
|
params="~39M",
|
||||||
|
approx_memory="~1 GB RAM",
|
||||||
|
relative_speed="~10x real-time (CPU)",
|
||||||
|
),
|
||||||
|
"base": TranscriptionModelInfo(
|
||||||
|
name="base",
|
||||||
|
display_name="Base (default)",
|
||||||
|
description="Balanced default. Works reasonably well on ordinary computers.",
|
||||||
|
params="~74M",
|
||||||
|
approx_memory="~1 GB RAM",
|
||||||
|
relative_speed="~7x real-time (CPU)",
|
||||||
|
),
|
||||||
|
"small": TranscriptionModelInfo(
|
||||||
|
name="small",
|
||||||
|
display_name="Small",
|
||||||
|
description="Better accuracy with higher resource usage.",
|
||||||
|
params="~244M",
|
||||||
|
approx_memory="~2 GB RAM",
|
||||||
|
relative_speed="~4x real-time (CPU)",
|
||||||
|
),
|
||||||
|
"medium": TranscriptionModelInfo(
|
||||||
|
name="medium",
|
||||||
|
display_name="Medium",
|
||||||
|
description="Higher accuracy and slower performance.",
|
||||||
|
params="~769M",
|
||||||
|
approx_memory="~5 GB RAM",
|
||||||
|
relative_speed="~2x real-time (CPU)",
|
||||||
|
),
|
||||||
|
"large-v3": TranscriptionModelInfo(
|
||||||
|
name="large-v3",
|
||||||
|
display_name="Large v3",
|
||||||
|
description="Highest accuracy and greatest resource requirements.",
|
||||||
|
params="~1.5B",
|
||||||
|
approx_memory="~10 GB RAM",
|
||||||
|
relative_speed="~1x real-time (CPU)",
|
||||||
|
),
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
def normalize_model_name(raw: str | None) -> str:
|
||||||
|
"""Case/whitespace-tolerant normalization. Never maps one model to another."""
|
||||||
|
return (raw or "").strip().lower()
|
||||||
|
|
||||||
|
|
||||||
|
def validate_model_name(raw: str | None) -> str:
|
||||||
|
"""Return the normalized name, or raise with a helpful message."""
|
||||||
|
name = normalize_model_name(raw)
|
||||||
|
if name in SUPPORTED_TRANSCRIPTION_MODELS:
|
||||||
|
return name
|
||||||
|
supported = ", ".join(sorted(SUPPORTED_TRANSCRIPTION_MODELS))
|
||||||
|
raise ProviderConfigError(
|
||||||
|
f"Unsupported transcription model {raw!r}. Supported models: {supported}. "
|
||||||
|
"Check the spelling — a different model is never substituted silently."
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _hf_hub_cache() -> Path:
|
||||||
|
try:
|
||||||
|
from huggingface_hub.constants import HF_HUB_CACHE
|
||||||
|
|
||||||
|
return Path(HF_HUB_CACHE)
|
||||||
|
except ImportError:
|
||||||
|
return Path(
|
||||||
|
os.environ.get("HF_HUB_CACHE", str(Path.home() / ".cache" / "huggingface" / "hub"))
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def is_model_downloaded(name: str) -> bool:
|
||||||
|
"""True when a non-empty faster-whisper snapshot for `name` sits in the
|
||||||
|
HuggingFace hub cache (repo Systran/faster-whisper-<name>)."""
|
||||||
|
repo_dir = _hf_hub_cache() / f"models--Systran--faster-whisper-{name}"
|
||||||
|
snapshots = repo_dir / "snapshots"
|
||||||
|
if not snapshots.is_dir():
|
||||||
|
return False
|
||||||
|
return any(s.is_dir() and any(s.iterdir()) for s in snapshots.iterdir())
|
||||||
|
|
||||||
|
|
||||||
|
def is_faster_whisper_installed() -> bool:
|
||||||
|
try:
|
||||||
|
import faster_whisper # noqa: F401
|
||||||
|
|
||||||
|
return True
|
||||||
|
except ImportError:
|
||||||
|
return False
|
||||||
|
|
||||||
|
|
||||||
|
def download_instructions(name: str) -> str:
|
||||||
|
return (
|
||||||
|
f'Model "{name}" is not downloaded. Download it with: '
|
||||||
|
f"POST /api/v1/models/{name}/download "
|
||||||
|
"(needs internet once), or run any transcription with that model selected — "
|
||||||
|
"faster-whisper fetches it from HuggingFace automatically."
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
async def get_global_default_model(session: AsyncSession) -> str:
|
||||||
|
row = await session.scalar(select(AppSetting).where(AppSetting.key == DEFAULT_MODEL_KEY))
|
||||||
|
if row is None:
|
||||||
|
return DEFAULT_MODEL
|
||||||
|
model = row.value.get("model") if isinstance(row.value, dict) else None
|
||||||
|
return model if model in SUPPORTED_TRANSCRIPTION_MODELS else DEFAULT_MODEL
|
||||||
|
|
||||||
|
|
||||||
|
async def set_global_default_model(session: AsyncSession, raw: str) -> str:
|
||||||
|
"""Validate + persist the global default. Affects future recordings only —
|
||||||
|
existing rows keep their saved model."""
|
||||||
|
name = validate_model_name(raw)
|
||||||
|
row = await session.scalar(select(AppSetting).where(AppSetting.key == DEFAULT_MODEL_KEY))
|
||||||
|
if row is None:
|
||||||
|
row = AppSetting(key=DEFAULT_MODEL_KEY, value={"model": name})
|
||||||
|
session.add(row)
|
||||||
|
else:
|
||||||
|
row.value = {"model": name}
|
||||||
|
await session.flush()
|
||||||
|
return name
|
||||||
|
|
||||||
|
|
||||||
|
def effective_model(recording_model: str | None, global_default: str) -> str:
|
||||||
|
"""Per-recording override wins; otherwise the global default."""
|
||||||
|
if recording_model and recording_model in SUPPORTED_TRANSCRIPTION_MODELS:
|
||||||
|
return recording_model
|
||||||
|
return global_default
|
||||||
78
backend/shonar/services/ai/ollama.py
Normal file
78
backend/shonar/services/ai/ollama.py
Normal file
|
|
@ -0,0 +1,78 @@
|
||||||
|
"""Summaries via a local Ollama server (`/api/chat`, JSON mode)."""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from collections.abc import AsyncIterator
|
||||||
|
from contextlib import asynccontextmanager
|
||||||
|
|
||||||
|
import httpx
|
||||||
|
|
||||||
|
from shonar.services.ai import ProviderConfigError, ProviderTransientError, SummaryResult
|
||||||
|
from shonar.services.ai._llm import SYSTEM_PROMPT, build_user_message, parse_summary
|
||||||
|
|
||||||
|
|
||||||
|
class OllamaProvider:
|
||||||
|
name = "ollama"
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
base_url: str,
|
||||||
|
model: str = "",
|
||||||
|
timeout_s: float = 900.0,
|
||||||
|
http_client: httpx.AsyncClient | None = None,
|
||||||
|
) -> None:
|
||||||
|
if not base_url.strip():
|
||||||
|
raise ProviderConfigError("ollama needs SHONAR_LLM_BASE_URL.")
|
||||||
|
if not model.strip():
|
||||||
|
raise ProviderConfigError("ollama needs SHONAR_LLM_MODEL.")
|
||||||
|
self.base_url = base_url.rstrip("/")
|
||||||
|
self.model = model
|
||||||
|
self.timeout_s = timeout_s
|
||||||
|
self.http_client = http_client
|
||||||
|
|
||||||
|
@asynccontextmanager
|
||||||
|
async def _client(self) -> AsyncIterator[httpx.AsyncClient]:
|
||||||
|
if self.http_client is not None:
|
||||||
|
yield self.http_client
|
||||||
|
else:
|
||||||
|
async with httpx.AsyncClient(timeout=self.timeout_s) as client:
|
||||||
|
yield client
|
||||||
|
|
||||||
|
async def summarize(self, transcript: str, *, title: str | None = None) -> SummaryResult:
|
||||||
|
payload = {
|
||||||
|
"model": self.model,
|
||||||
|
"stream": False,
|
||||||
|
"format": "json",
|
||||||
|
# qwen3-family models "think" by default: a long reasoning
|
||||||
|
# chain before the JSON answer, brutally slow on CPU and it
|
||||||
|
# does not improve the summary. Ask for the answer directly
|
||||||
|
# (ignored by non-thinking models).
|
||||||
|
"think": False,
|
||||||
|
"options": {"num_ctx": 8192},
|
||||||
|
"messages": [
|
||||||
|
{"role": "system", "content": SYSTEM_PROMPT},
|
||||||
|
{"role": "user", "content": build_user_message(transcript, title)},
|
||||||
|
],
|
||||||
|
}
|
||||||
|
try:
|
||||||
|
async with self._client() as client:
|
||||||
|
resp = await client.post(f"{self.base_url}/api/chat", json=payload)
|
||||||
|
except (httpx.TimeoutException, httpx.TransportError) as e:
|
||||||
|
raise ProviderTransientError(f"Ollama unreachable: {type(e).__name__}") from e
|
||||||
|
if resp.status_code == 404:
|
||||||
|
# Missing model and missing route both 404 here; both are
|
||||||
|
# configuration, not weather.
|
||||||
|
raise ProviderConfigError("Ollama has no such model or route (HTTP 404).")
|
||||||
|
if resp.status_code != 200:
|
||||||
|
raise ProviderTransientError(f"Summarization failed (HTTP {resp.status_code}).")
|
||||||
|
try:
|
||||||
|
content = resp.json()["message"]["content"]
|
||||||
|
except (ValueError, KeyError, TypeError) as e:
|
||||||
|
raise ProviderTransientError("Ollama sent an unreadable reply.") from e
|
||||||
|
import json as _json
|
||||||
|
|
||||||
|
try:
|
||||||
|
data = _json.loads(content)
|
||||||
|
except ValueError as e:
|
||||||
|
raise ProviderTransientError("Ollama reply was not JSON.") from e
|
||||||
|
return parse_summary(data, self.model)
|
||||||
82
backend/shonar/services/ai/openai_compat.py
Normal file
82
backend/shonar/services/ai/openai_compat.py
Normal file
|
|
@ -0,0 +1,82 @@
|
||||||
|
"""Summaries via any OpenAI-compatible chat endpoint (self-hosted
|
||||||
|
vLLM/llama.cpp server, commercial API, …) with JSON mode.
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from collections.abc import AsyncIterator
|
||||||
|
from contextlib import asynccontextmanager
|
||||||
|
|
||||||
|
import httpx
|
||||||
|
|
||||||
|
from shonar.services.ai import ProviderConfigError, ProviderTransientError, SummaryResult
|
||||||
|
from shonar.services.ai._llm import SYSTEM_PROMPT, build_user_message, parse_summary
|
||||||
|
|
||||||
|
|
||||||
|
class OpenAICompatProvider:
|
||||||
|
name = "openai_compat"
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
base_url: str,
|
||||||
|
model: str = "",
|
||||||
|
api_key: str = "",
|
||||||
|
timeout_s: float = 180.0,
|
||||||
|
http_client: httpx.AsyncClient | None = None,
|
||||||
|
) -> None:
|
||||||
|
if not base_url.strip():
|
||||||
|
raise ProviderConfigError("openai_compat needs SHONAR_LLM_BASE_URL.")
|
||||||
|
if not model.strip():
|
||||||
|
raise ProviderConfigError("openai_compat needs SHONAR_LLM_MODEL.")
|
||||||
|
self.base_url = base_url.rstrip("/")
|
||||||
|
self.model = model
|
||||||
|
self.api_key = api_key
|
||||||
|
self.timeout_s = timeout_s
|
||||||
|
self.http_client = http_client
|
||||||
|
|
||||||
|
@asynccontextmanager
|
||||||
|
async def _client(self) -> AsyncIterator[httpx.AsyncClient]:
|
||||||
|
if self.http_client is not None:
|
||||||
|
yield self.http_client
|
||||||
|
else:
|
||||||
|
async with httpx.AsyncClient(timeout=self.timeout_s) as client:
|
||||||
|
yield client
|
||||||
|
|
||||||
|
async def summarize(self, transcript: str, *, title: str | None = None) -> SummaryResult:
|
||||||
|
headers = (
|
||||||
|
{"Authorization": f"Bearer {self.api_key}"} if self.api_key else {}
|
||||||
|
)
|
||||||
|
payload = {
|
||||||
|
"model": self.model,
|
||||||
|
"temperature": 0.2,
|
||||||
|
"response_format": {"type": "json_object"},
|
||||||
|
"messages": [
|
||||||
|
{"role": "system", "content": SYSTEM_PROMPT},
|
||||||
|
{"role": "user", "content": build_user_message(transcript, title)},
|
||||||
|
],
|
||||||
|
}
|
||||||
|
try:
|
||||||
|
async with self._client() as client:
|
||||||
|
resp = await client.post(
|
||||||
|
f"{self.base_url}/v1/chat/completions", headers=headers, json=payload
|
||||||
|
)
|
||||||
|
except (httpx.TimeoutException, httpx.TransportError) as e:
|
||||||
|
raise ProviderTransientError(f"LLM unreachable: {type(e).__name__}") from e
|
||||||
|
if resp.status_code in (401, 403, 404):
|
||||||
|
raise ProviderConfigError(f"LLM refused the request (HTTP {resp.status_code}).")
|
||||||
|
if resp.status_code == 429 or resp.status_code >= 500:
|
||||||
|
raise ProviderTransientError(f"LLM busy (HTTP {resp.status_code}).")
|
||||||
|
if resp.status_code != 200:
|
||||||
|
raise ProviderTransientError(f"Summarization failed (HTTP {resp.status_code}).")
|
||||||
|
try:
|
||||||
|
body = resp.json()
|
||||||
|
content = body["choices"][0]["message"]["content"]
|
||||||
|
except (ValueError, KeyError, IndexError, TypeError) as e:
|
||||||
|
raise ProviderTransientError("LLM sent an unreadable reply.") from e
|
||||||
|
import json as _json
|
||||||
|
|
||||||
|
try:
|
||||||
|
data = _json.loads(content)
|
||||||
|
except ValueError as e:
|
||||||
|
raise ProviderTransientError("LLM reply was not JSON.") from e
|
||||||
|
return parse_summary(data, self.model)
|
||||||
128
backend/shonar/services/ai/whisper_http.py
Normal file
128
backend/shonar/services/ai/whisper_http.py
Normal file
|
|
@ -0,0 +1,128 @@
|
||||||
|
"""Transcription via any OpenAI-compatible `/v1/audio/transcriptions`
|
||||||
|
endpoint (self-hosted whisper.cpp server, commercial Whisper API, …).
|
||||||
|
Sends `verbose_json` so segment timings come back with the text.
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from collections.abc import AsyncIterator
|
||||||
|
from contextlib import asynccontextmanager
|
||||||
|
|
||||||
|
import httpx
|
||||||
|
|
||||||
|
from shonar.services.ai import (
|
||||||
|
ProviderConfigError,
|
||||||
|
ProviderTransientError,
|
||||||
|
Segment,
|
||||||
|
TranscriptResult,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class WhisperHttpProvider:
|
||||||
|
name = "whisper_http"
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
base_url: str,
|
||||||
|
model: str = "base",
|
||||||
|
api_key: str = "",
|
||||||
|
timeout_s: float = 300.0,
|
||||||
|
http_client: httpx.AsyncClient | None = None,
|
||||||
|
) -> None:
|
||||||
|
if not base_url.strip():
|
||||||
|
raise ProviderConfigError(
|
||||||
|
"whisper_http needs SHONAR_TRANSCRIPTION_BASE_URL."
|
||||||
|
)
|
||||||
|
self.base_url = base_url.rstrip("/")
|
||||||
|
self.model = model
|
||||||
|
self.api_key = api_key
|
||||||
|
self.timeout_s = timeout_s
|
||||||
|
self.http_client = http_client
|
||||||
|
|
||||||
|
@asynccontextmanager
|
||||||
|
async def _client(self) -> AsyncIterator[httpx.AsyncClient]:
|
||||||
|
if self.http_client is not None:
|
||||||
|
yield self.http_client
|
||||||
|
else:
|
||||||
|
async with httpx.AsyncClient(timeout=self.timeout_s) as client:
|
||||||
|
yield client
|
||||||
|
|
||||||
|
async def transcribe(
|
||||||
|
self,
|
||||||
|
audio: bytes,
|
||||||
|
mime: str,
|
||||||
|
*,
|
||||||
|
language_hint: str | None = None,
|
||||||
|
on_progress=None, # accepted for protocol parity; not reported
|
||||||
|
) -> TranscriptResult:
|
||||||
|
headers = (
|
||||||
|
{"Authorization": f"Bearer {self.api_key}"} if self.api_key else {}
|
||||||
|
)
|
||||||
|
data: dict[str, str] = {"model": self.model, "response_format": "verbose_json"}
|
||||||
|
if language_hint:
|
||||||
|
data["language"] = language_hint
|
||||||
|
files = {"file": (f"audio.{_ext(mime)}", audio, mime or "application/octet-stream")}
|
||||||
|
try:
|
||||||
|
async with self._client() as client:
|
||||||
|
resp = await client.post(
|
||||||
|
f"{self.base_url}/v1/audio/transcriptions",
|
||||||
|
headers=headers,
|
||||||
|
data=data,
|
||||||
|
files=files,
|
||||||
|
)
|
||||||
|
except (httpx.TimeoutException, httpx.TransportError) as e:
|
||||||
|
raise ProviderTransientError(
|
||||||
|
f"Transcription service unreachable: {type(e).__name__}"
|
||||||
|
) from e
|
||||||
|
if resp.status_code in (401, 403, 404):
|
||||||
|
raise ProviderConfigError(
|
||||||
|
f"Transcription service refused the request (HTTP {resp.status_code})."
|
||||||
|
)
|
||||||
|
if resp.status_code == 429 or resp.status_code >= 500:
|
||||||
|
raise ProviderTransientError(
|
||||||
|
f"Transcription service busy (HTTP {resp.status_code})."
|
||||||
|
)
|
||||||
|
if resp.status_code != 200:
|
||||||
|
raise ProviderTransientError(
|
||||||
|
f"Transcription failed (HTTP {resp.status_code})."
|
||||||
|
)
|
||||||
|
try:
|
||||||
|
body = resp.json()
|
||||||
|
except ValueError as e:
|
||||||
|
raise ProviderTransientError("Transcription service sent no JSON.") from e
|
||||||
|
segments = []
|
||||||
|
raw_segs = body.get("segments")
|
||||||
|
if isinstance(raw_segs, list):
|
||||||
|
for s in raw_segs:
|
||||||
|
if not isinstance(s, dict):
|
||||||
|
continue
|
||||||
|
try:
|
||||||
|
segments.append(
|
||||||
|
Segment(
|
||||||
|
start=float(s.get("start", 0.0)),
|
||||||
|
end=float(s.get("end", 0.0)),
|
||||||
|
text=str(s.get("text", "")),
|
||||||
|
)
|
||||||
|
)
|
||||||
|
except (TypeError, ValueError):
|
||||||
|
continue
|
||||||
|
text = body.get("text")
|
||||||
|
return TranscriptResult(
|
||||||
|
text=text if isinstance(text, str) else "",
|
||||||
|
language=body.get("language") if isinstance(body.get("language"), str) else None,
|
||||||
|
segments=segments,
|
||||||
|
model=self.model,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _ext(mime: str) -> str:
|
||||||
|
return {
|
||||||
|
"audio/mp4": "m4a",
|
||||||
|
"audio/m4a": "m4a",
|
||||||
|
"audio/wav": "wav",
|
||||||
|
"audio/x-wav": "wav",
|
||||||
|
"audio/ogg": "ogg",
|
||||||
|
"audio/opus": "ogg",
|
||||||
|
"audio/webm": "webm",
|
||||||
|
"audio/mpeg": "mp3",
|
||||||
|
}.get(mime.lower().split(";")[0].strip(), "bin")
|
||||||
167
backend/shonar/services/auth.py
Normal file
167
backend/shonar/services/auth.py
Normal file
|
|
@ -0,0 +1,167 @@
|
||||||
|
"""Authentication service: registration, login, rotating refresh tokens with
|
||||||
|
reuse detection, logout, account deletion."""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import uuid
|
||||||
|
from datetime import UTC, datetime
|
||||||
|
|
||||||
|
from sqlalchemy import select, update
|
||||||
|
from sqlalchemy.ext.asyncio import AsyncSession
|
||||||
|
|
||||||
|
from shonar.core.security import (
|
||||||
|
generate_refresh_token,
|
||||||
|
hash_password,
|
||||||
|
hash_refresh_token,
|
||||||
|
refresh_token_ttl,
|
||||||
|
verify_password,
|
||||||
|
)
|
||||||
|
from shonar.db.models import Device, RefreshToken, User, utcnow
|
||||||
|
|
||||||
|
|
||||||
|
class AuthError(Exception):
|
||||||
|
"""Safe, client-displayable auth failure (never leaks which half failed
|
||||||
|
beyond what the flow requires)."""
|
||||||
|
|
||||||
|
def __init__(self, message: str, status_code: int = 401):
|
||||||
|
super().__init__(message)
|
||||||
|
self.message = message
|
||||||
|
self.status_code = status_code
|
||||||
|
|
||||||
|
|
||||||
|
async def register_user(
|
||||||
|
session: AsyncSession, email: str, password: str, display_name: str | None
|
||||||
|
) -> User:
|
||||||
|
email = email.strip().lower()
|
||||||
|
existing = await session.scalar(select(User).where(User.email == email))
|
||||||
|
if existing is not None:
|
||||||
|
# Use a generic message; do not reveal whether the account exists in
|
||||||
|
# flows where that matters. For self-hosted registration the UX cost
|
||||||
|
# of "email already registered" is acceptable and helpful.
|
||||||
|
raise AuthError("An account with this email already exists.", 409)
|
||||||
|
user = User(
|
||||||
|
email=email,
|
||||||
|
password_hash=hash_password(password),
|
||||||
|
display_name=display_name,
|
||||||
|
)
|
||||||
|
session.add(user)
|
||||||
|
await session.flush()
|
||||||
|
return user
|
||||||
|
|
||||||
|
|
||||||
|
async def issue_refresh_token(
|
||||||
|
session: AsyncSession,
|
||||||
|
user_id: uuid.UUID,
|
||||||
|
family: uuid.UUID | None,
|
||||||
|
device_id: uuid.UUID | None,
|
||||||
|
) -> tuple[str, RefreshToken]:
|
||||||
|
token = generate_refresh_token()
|
||||||
|
rt = RefreshToken(
|
||||||
|
user_id=user_id,
|
||||||
|
token_hash=hash_refresh_token(token),
|
||||||
|
family=family or uuid.uuid4(),
|
||||||
|
device_id=device_id,
|
||||||
|
expires_at=datetime.now(UTC) + refresh_token_ttl(),
|
||||||
|
)
|
||||||
|
session.add(rt)
|
||||||
|
await session.flush()
|
||||||
|
return token, rt
|
||||||
|
|
||||||
|
|
||||||
|
async def login(
|
||||||
|
session: AsyncSession,
|
||||||
|
email: str,
|
||||||
|
password: str,
|
||||||
|
device_name: str | None,
|
||||||
|
platform: str,
|
||||||
|
) -> tuple[User, str, Device]:
|
||||||
|
"""Returns (user, refresh_token, device). Raises AuthError safely."""
|
||||||
|
email = email.strip().lower()
|
||||||
|
user = await session.scalar(select(User).where(User.email == email))
|
||||||
|
if user is None or user.deleted_at is not None or not user.is_active:
|
||||||
|
raise AuthError("Invalid email or password.")
|
||||||
|
if not verify_password(user.password_hash, password):
|
||||||
|
raise AuthError("Invalid email or password.")
|
||||||
|
|
||||||
|
device = Device(user_id=user.id, name=device_name or "Android device", platform=platform)
|
||||||
|
session.add(device)
|
||||||
|
await session.flush()
|
||||||
|
|
||||||
|
refresh_token, _ = await issue_refresh_token(session, user.id, None, device.id)
|
||||||
|
return user, refresh_token, device
|
||||||
|
|
||||||
|
|
||||||
|
async def rotate_refresh_token(
|
||||||
|
session: AsyncSession, presented_token: str
|
||||||
|
) -> tuple[User, str, uuid.UUID | None]:
|
||||||
|
"""Consume a refresh token and issue a replacement in the same family.
|
||||||
|
|
||||||
|
Reuse detection: presenting an already-consumed/revoked token revokes the
|
||||||
|
entire family (an attacker's stolen token dies along with the real one).
|
||||||
|
"""
|
||||||
|
token_hash = hash_refresh_token(presented_token)
|
||||||
|
rt = await session.scalar(select(RefreshToken).where(RefreshToken.token_hash == token_hash))
|
||||||
|
now = utcnow()
|
||||||
|
|
||||||
|
if rt is None:
|
||||||
|
raise AuthError("Invalid refresh token.")
|
||||||
|
|
||||||
|
if rt.revoked_at is not None or rt.replaced_by is not None:
|
||||||
|
# REUSE DETECTED — revoke the whole family. Commit BEFORE raising:
|
||||||
|
# the request's transaction would otherwise roll back on the 401 and
|
||||||
|
# silently undo the security-revocation.
|
||||||
|
await session.execute(
|
||||||
|
update(RefreshToken)
|
||||||
|
.where(RefreshToken.family == rt.family, RefreshToken.revoked_at.is_(None))
|
||||||
|
.values(revoked_at=now)
|
||||||
|
)
|
||||||
|
await session.commit()
|
||||||
|
raise AuthError("Refresh token reuse detected. Please log in again.", 401)
|
||||||
|
|
||||||
|
if rt.expires_at < now:
|
||||||
|
raise AuthError("Refresh token expired.", 401)
|
||||||
|
|
||||||
|
user = await session.get(User, rt.user_id)
|
||||||
|
if user is None or user.deleted_at is not None or not user.is_active:
|
||||||
|
raise AuthError("Account unavailable.", 401)
|
||||||
|
|
||||||
|
new_token, new_rt = await issue_refresh_token(session, user.id, rt.family, rt.device_id)
|
||||||
|
rt.revoked_at = now
|
||||||
|
rt.replaced_by = new_rt.id
|
||||||
|
|
||||||
|
if rt.device_id is not None:
|
||||||
|
device = await session.get(Device, rt.device_id)
|
||||||
|
if device is not None:
|
||||||
|
device.last_seen_at = now
|
||||||
|
await session.flush()
|
||||||
|
return user, new_token, rt.device_id
|
||||||
|
|
||||||
|
|
||||||
|
async def logout(session: AsyncSession, presented_token: str) -> None:
|
||||||
|
"""Revoke the presented token's whole family (logs the device out)."""
|
||||||
|
token_hash = hash_refresh_token(presented_token)
|
||||||
|
rt = await session.scalar(select(RefreshToken).where(RefreshToken.token_hash == token_hash))
|
||||||
|
if rt is None:
|
||||||
|
return # idempotent
|
||||||
|
await session.execute(
|
||||||
|
update(RefreshToken)
|
||||||
|
.where(RefreshToken.family == rt.family, RefreshToken.revoked_at.is_(None))
|
||||||
|
.values(revoked_at=utcnow())
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
async def delete_account(session: AsyncSession, user: User, password: str) -> None:
|
||||||
|
if not verify_password(user.password_hash, password):
|
||||||
|
raise AuthError("Invalid password.", 403)
|
||||||
|
now = utcnow()
|
||||||
|
user.deleted_at = now
|
||||||
|
user.is_active = False
|
||||||
|
# Revoke every refresh token for the user.
|
||||||
|
await session.execute(
|
||||||
|
update(RefreshToken)
|
||||||
|
.where(RefreshToken.user_id == user.id, RefreshToken.revoked_at.is_(None))
|
||||||
|
.values(revoked_at=now)
|
||||||
|
)
|
||||||
|
# NOTE: hard deletion of rows/files is performed by the retention sweep
|
||||||
|
# (services/retention.py) so an accidental deletion can be cancelled
|
||||||
|
# within the grace window (see docs/security.md).
|
||||||
251
backend/shonar/services/exports.py
Normal file
251
backend/shonar/services/exports.py
Normal file
|
|
@ -0,0 +1,251 @@
|
||||||
|
"""Exports (M9): audio, transcript txt, notes markdown, bundle zip.
|
||||||
|
|
||||||
|
Synchronous generation — every artifact is small enough (text, or one
|
||||||
|
audio file) that a background job adds failure modes, not speed. Each
|
||||||
|
successful export records an ExportJob row and stores the produced bytes
|
||||||
|
as an Asset(kind=export) so the audit trail exists; the response is the
|
||||||
|
file itself (no separate download-asset round trip).
|
||||||
|
|
||||||
|
Formats:
|
||||||
|
audio — the original upload, byte-identical, original mime/extension
|
||||||
|
txt — current transcript text
|
||||||
|
md — notes.md: title, metadata, notes, summary sections, transcript
|
||||||
|
zip — bundle: original audio + transcript.txt + notes.md
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import io
|
||||||
|
import uuid
|
||||||
|
import zipfile
|
||||||
|
from dataclasses import dataclass
|
||||||
|
|
||||||
|
from sqlalchemy import select
|
||||||
|
from sqlalchemy.ext.asyncio import AsyncSession
|
||||||
|
|
||||||
|
from shonar.db.models import (
|
||||||
|
Asset,
|
||||||
|
AssetKind,
|
||||||
|
ExportJob,
|
||||||
|
JobStatus,
|
||||||
|
Recording,
|
||||||
|
Summary,
|
||||||
|
Transcript,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class ExportError(Exception):
|
||||||
|
def __init__(self, status_code: int, message: str):
|
||||||
|
super().__init__(message)
|
||||||
|
self.status_code = status_code
|
||||||
|
self.message = message
|
||||||
|
|
||||||
|
|
||||||
|
EXPORT_FORMATS = ("audio", "txt", "md", "zip")
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class ExportResult:
|
||||||
|
filename: str
|
||||||
|
mime_type: str
|
||||||
|
data: bytes
|
||||||
|
|
||||||
|
|
||||||
|
def _safe_filename(rec: Recording) -> str:
|
||||||
|
"""Slug the title; fall back to the recorded_at stamp. Never leaks ids."""
|
||||||
|
base = "".join(
|
||||||
|
c if (c.isalnum() or c in "-_ ") else " " for c in (rec.title or "")
|
||||||
|
).strip()
|
||||||
|
if not base:
|
||||||
|
base = f"recording-{rec.recorded_at:%Y%m%d-%H%M%S}"
|
||||||
|
return base[:120]
|
||||||
|
|
||||||
|
|
||||||
|
async def _current_transcript(session: AsyncSession, rec_id: uuid.UUID) -> Transcript | None:
|
||||||
|
return await session.scalar(
|
||||||
|
select(Transcript)
|
||||||
|
.where(Transcript.recording_id == rec_id, Transcript.superseded_at.is_(None))
|
||||||
|
.order_by(Transcript.version.desc())
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
async def _current_summary(session: AsyncSession, rec_id: uuid.UUID) -> Summary | None:
|
||||||
|
return await session.scalar(
|
||||||
|
select(Summary)
|
||||||
|
.where(Summary.recording_id == rec_id, Summary.superseded_at.is_(None))
|
||||||
|
.order_by(Summary.version.desc())
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
async def _original_asset(session: AsyncSession, rec_id: uuid.UUID) -> Asset | None:
|
||||||
|
return await session.scalar(
|
||||||
|
select(Asset).where(Asset.recording_id == rec_id, Asset.kind == AssetKind.original)
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _summary_md(summary: Summary | None) -> str:
|
||||||
|
if summary is None:
|
||||||
|
return ""
|
||||||
|
c = summary.content or {}
|
||||||
|
out = ["## Summary\n"]
|
||||||
|
if c.get("short"):
|
||||||
|
out.append(f"{c['short']}\n")
|
||||||
|
if c.get("detailed"):
|
||||||
|
out.append(f"### Detailed\n\n{c['detailed']}\n")
|
||||||
|
for key, header in (
|
||||||
|
("key_points", "Key Points"),
|
||||||
|
("decisions", "Decisions"),
|
||||||
|
("action_items", "Action Items"),
|
||||||
|
("questions", "Questions"),
|
||||||
|
):
|
||||||
|
items = c.get(key)
|
||||||
|
if isinstance(items, list) and items:
|
||||||
|
out.append(f"### {header}\n")
|
||||||
|
out.extend(f"- {x}" for x in items)
|
||||||
|
out.append("")
|
||||||
|
return "\n".join(out)
|
||||||
|
|
||||||
|
|
||||||
|
def _transcript_md(t: Transcript | None) -> str:
|
||||||
|
if t is None:
|
||||||
|
return ""
|
||||||
|
lines = ["## Transcript\n"]
|
||||||
|
segs = [s for s in (t.segments or []) if isinstance(s, dict)]
|
||||||
|
if segs:
|
||||||
|
for s in segs:
|
||||||
|
start = float(s.get("start", 0.0))
|
||||||
|
stamp = f"{int(start // 60):02d}:{start % 60:04.1f}"
|
||||||
|
speaker = f"**{s['speaker']}**: " if s.get("speaker") else ""
|
||||||
|
lines.append(f"- `[{stamp}]` {speaker}{s.get('text', '').strip()}")
|
||||||
|
else:
|
||||||
|
lines.append(t.text or "")
|
||||||
|
return "\n".join(lines) + "\n"
|
||||||
|
|
||||||
|
|
||||||
|
def _notes_md(
|
||||||
|
rec: Recording, t: Transcript | None, summary: Summary | None,
|
||||||
|
tag_names: list[str] | None = None,
|
||||||
|
) -> str:
|
||||||
|
parts = [
|
||||||
|
f"# {rec.title}\n",
|
||||||
|
f"- Recorded: {rec.recorded_at:%Y-%m-%d %H:%M} UTC",
|
||||||
|
f"- Duration: {rec.duration_seconds:.1f}s",
|
||||||
|
]
|
||||||
|
if tag_names:
|
||||||
|
parts.append("- Tags: " + ", ".join(f"`{x}`" for x in tag_names))
|
||||||
|
parts.append("")
|
||||||
|
if rec.notes:
|
||||||
|
parts.append(f"## Notes\n\n{rec.notes}\n")
|
||||||
|
s = _summary_md(summary)
|
||||||
|
if s:
|
||||||
|
parts.append(s)
|
||||||
|
tr = _transcript_md(t)
|
||||||
|
if tr:
|
||||||
|
parts.append(tr)
|
||||||
|
return "\n".join(parts)
|
||||||
|
|
||||||
|
|
||||||
|
def _zip(files: list[tuple[str, bytes]]) -> bytes:
|
||||||
|
buf = io.BytesIO()
|
||||||
|
with zipfile.ZipFile(buf, "w", zipfile.ZIP_DEFLATED) as z:
|
||||||
|
for name, data in files:
|
||||||
|
z.writestr(name, data)
|
||||||
|
return buf.getvalue()
|
||||||
|
|
||||||
|
|
||||||
|
async def build_export(
|
||||||
|
session: AsyncSession, rec: Recording, export_format: str
|
||||||
|
) -> ExportResult:
|
||||||
|
"""Build one export artifact for an owned, non-deleted recording."""
|
||||||
|
from shonar.storage import get_storage
|
||||||
|
|
||||||
|
if export_format not in EXPORT_FORMATS:
|
||||||
|
raise ExportError(422, f"Unknown export format. Use one of: {', '.join(EXPORT_FORMATS)}")
|
||||||
|
|
||||||
|
original = await _original_asset(session, rec.id)
|
||||||
|
t = await _current_transcript(session, rec.id)
|
||||||
|
summary = await _current_summary(session, rec.id)
|
||||||
|
from shonar.services.search import tag_names
|
||||||
|
|
||||||
|
tags = await tag_names(session, rec.id)
|
||||||
|
base = _safe_filename(rec)
|
||||||
|
|
||||||
|
if export_format == "audio":
|
||||||
|
if original is None:
|
||||||
|
raise ExportError(404, "No audio stored for this recording")
|
||||||
|
data = await get_storage().get(original.storage_key)
|
||||||
|
ext = original.storage_key[original.storage_key.rfind(".") :]
|
||||||
|
return ExportResult(filename=f"{base}{ext}", mime_type=original.mime_type, data=data)
|
||||||
|
|
||||||
|
if export_format == "txt":
|
||||||
|
if t is None:
|
||||||
|
raise ExportError(404, "No transcript yet — transcribe first")
|
||||||
|
return ExportResult(
|
||||||
|
filename=f"{base}.txt", mime_type="text/plain; charset=utf-8",
|
||||||
|
data=(t.text or "").encode("utf-8"),
|
||||||
|
)
|
||||||
|
|
||||||
|
if export_format == "md":
|
||||||
|
if t is None and summary is None and not rec.notes:
|
||||||
|
raise ExportError(
|
||||||
|
404, "Nothing to export — this recording has no notes, transcript, or summary"
|
||||||
|
)
|
||||||
|
return ExportResult(
|
||||||
|
filename=f"{base}.md", mime_type="text/markdown; charset=utf-8",
|
||||||
|
data=_notes_md(rec, t, summary, tags).encode("utf-8"),
|
||||||
|
)
|
||||||
|
|
||||||
|
# zip bundle: whatever exists, always at least the audio when present.
|
||||||
|
if original is None and t is None and summary is None and not rec.notes:
|
||||||
|
raise ExportError(404, "Nothing to export for this recording")
|
||||||
|
files: list[tuple[str, bytes]] = []
|
||||||
|
if original is not None:
|
||||||
|
audio = await get_storage().get(original.storage_key)
|
||||||
|
ext = original.storage_key[original.storage_key.rfind(".") :]
|
||||||
|
files.append((f"{base}{ext}", audio))
|
||||||
|
if t is not None:
|
||||||
|
files.append(("transcript.txt", (t.text or "").encode("utf-8")))
|
||||||
|
files.append(("notes.md", _notes_md(rec, t, summary, tags).encode("utf-8")))
|
||||||
|
return ExportResult(
|
||||||
|
filename=f"{base}.zip", mime_type="application/zip", data=_zip(files)
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
async def record_export(
|
||||||
|
session: AsyncSession, user_id: uuid.UUID, rec: Recording,
|
||||||
|
export_format: str, result: ExportResult,
|
||||||
|
) -> None:
|
||||||
|
"""Persist the audit trail: ExportJob(succeeded) + Asset(kind=export).
|
||||||
|
|
||||||
|
Best-effort storage of the artifact bytes; a storage failure never
|
||||||
|
fails the download the user already received.
|
||||||
|
"""
|
||||||
|
from shonar.storage import get_storage
|
||||||
|
|
||||||
|
job = ExportJob(
|
||||||
|
user_id=user_id,
|
||||||
|
recording_id=rec.id,
|
||||||
|
export_type=export_format,
|
||||||
|
status=JobStatus.succeeded,
|
||||||
|
)
|
||||||
|
session.add(job)
|
||||||
|
try:
|
||||||
|
key = f"exports/{user_id}/{rec.id}/{export_format}-{uuid.uuid4().hex}"
|
||||||
|
await get_storage().put(key, result.data)
|
||||||
|
import hashlib
|
||||||
|
|
||||||
|
asset = Asset(
|
||||||
|
recording_id=rec.id,
|
||||||
|
user_id=user_id,
|
||||||
|
kind=AssetKind.export,
|
||||||
|
storage_key=key,
|
||||||
|
mime_type=result.mime_type,
|
||||||
|
size_bytes=len(result.data),
|
||||||
|
checksum_sha256=hashlib.sha256(result.data).hexdigest(),
|
||||||
|
)
|
||||||
|
session.add(asset)
|
||||||
|
await session.flush()
|
||||||
|
job.asset_id = asset.id
|
||||||
|
except Exception: # noqa: BLE001 — audit copy is best-effort
|
||||||
|
pass
|
||||||
|
await session.flush()
|
||||||
141
backend/shonar/services/inline_queue.py
Normal file
141
backend/shonar/services/inline_queue.py
Normal file
|
|
@ -0,0 +1,141 @@
|
||||||
|
"""In-process job runner (desktop bundled-lite engine).
|
||||||
|
|
||||||
|
``queue_backend=inline`` replaces the arq/Redis transport with a single
|
||||||
|
asyncio consumer inside the uvicorn process: DB rows stay the source of
|
||||||
|
truth (ProcessingJob), this just runs the work. One job at a time —
|
||||||
|
local faster-whisper is multi-GB per pass, mirroring the worker's
|
||||||
|
``max_jobs=1`` rule.
|
||||||
|
|
||||||
|
Transient failures retry in-process with backoff up to MAX_TRIES (the
|
||||||
|
same budget arq gives via ``max_tries``); the final failure is recorded
|
||||||
|
by the task body itself (see processing._fail).
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import asyncio
|
||||||
|
import contextlib
|
||||||
|
import logging
|
||||||
|
|
||||||
|
from shonar.db.models import JobType
|
||||||
|
from shonar.services import processing
|
||||||
|
|
||||||
|
logger = logging.getLogger("shonar.inline_queue")
|
||||||
|
|
||||||
|
RETRY_DELAY_SECONDS = 5.0
|
||||||
|
# Desktop engine: hard-delete expired soft-deletes once the app has been up
|
||||||
|
# for a day, then daily. Delayed first run keeps startup snappy.
|
||||||
|
RETENTION_INTERVAL_SECONDS = 24 * 3600.0
|
||||||
|
|
||||||
|
_queue: asyncio.Queue[tuple[str, str, int]] | None = None
|
||||||
|
_consumer: asyncio.Task | None = None
|
||||||
|
_retention: asyncio.Task | None = None
|
||||||
|
|
||||||
|
|
||||||
|
def _get_queue() -> asyncio.Queue[tuple[str, str, int]]:
|
||||||
|
global _queue
|
||||||
|
if _queue is None:
|
||||||
|
_queue = asyncio.Queue()
|
||||||
|
return _queue
|
||||||
|
|
||||||
|
|
||||||
|
async def start() -> None:
|
||||||
|
"""Start the consumer and re-run anything the DB says is pending."""
|
||||||
|
global _consumer, _retention
|
||||||
|
_get_queue()
|
||||||
|
# Inline engine = single process: any row still marked `running` at
|
||||||
|
# startup is a corpse from the previous process (the worker died with
|
||||||
|
# it). sweep_stale's 2h live-worker grace — correct for multi-worker
|
||||||
|
# arq deployments — would starve these jobs, so reclaim them first.
|
||||||
|
from sqlalchemy import update
|
||||||
|
|
||||||
|
from shonar.db.models import JobStatus, ProcessingJob
|
||||||
|
from shonar.db.session import session_factory
|
||||||
|
|
||||||
|
async with session_factory()() as s:
|
||||||
|
await s.execute(
|
||||||
|
update(ProcessingJob)
|
||||||
|
.where(ProcessingJob.status == JobStatus.running)
|
||||||
|
.values(status=JobStatus.queued, started_at=None)
|
||||||
|
)
|
||||||
|
await s.commit()
|
||||||
|
if _consumer is None or _consumer.done():
|
||||||
|
_consumer = asyncio.create_task(_consume(), name="shonar-inline-queue")
|
||||||
|
if _retention is None or _retention.done():
|
||||||
|
_retention = asyncio.create_task(_retention_loop(), name="shonar-retention")
|
||||||
|
# Crash recovery: queued rows (and orphaned running rows requeued by
|
||||||
|
# the sweep) go back on the in-process queue.
|
||||||
|
count = await processing.sweep_stale()
|
||||||
|
if count:
|
||||||
|
logger.info("inline queue startup sweep requeued %d jobs", count)
|
||||||
|
|
||||||
|
|
||||||
|
async def _retention_loop() -> None:
|
||||||
|
"""Daily hard-delete sweep for the desktop engine (no arq cron here).
|
||||||
|
|
||||||
|
First pass a few minutes after start (the app may only run for hours
|
||||||
|
at a time, so a full-day initial sleep could starve the sweep), then
|
||||||
|
once a day while running.
|
||||||
|
"""
|
||||||
|
from shonar.services import retention
|
||||||
|
|
||||||
|
await asyncio.sleep(120.0)
|
||||||
|
while True:
|
||||||
|
try:
|
||||||
|
purged = await retention.sweep_deleted()
|
||||||
|
if purged["recordings"] or purged["users"]:
|
||||||
|
logger.info("inline retention sweep: %s", purged)
|
||||||
|
except asyncio.CancelledError:
|
||||||
|
raise
|
||||||
|
except Exception: # noqa: BLE001 — the loop must survive any failure
|
||||||
|
logger.exception("inline retention sweep failed")
|
||||||
|
await asyncio.sleep(RETENTION_INTERVAL_SECONDS)
|
||||||
|
|
||||||
|
|
||||||
|
async def stop() -> None:
|
||||||
|
global _consumer, _retention
|
||||||
|
if _consumer is not None:
|
||||||
|
_consumer.cancel()
|
||||||
|
with contextlib.suppress(BaseException): # noqa: BLE001 — shutdown is best-effort
|
||||||
|
await _consumer
|
||||||
|
_consumer = None
|
||||||
|
if _retention is not None:
|
||||||
|
_retention.cancel()
|
||||||
|
with contextlib.suppress(BaseException): # noqa: BLE001
|
||||||
|
await _retention
|
||||||
|
_retention = None
|
||||||
|
|
||||||
|
|
||||||
|
async def enqueue(job_type: JobType, recording_id: str) -> None:
|
||||||
|
await _get_queue().put((job_type.value, str(recording_id), 1))
|
||||||
|
|
||||||
|
|
||||||
|
async def _consume() -> None:
|
||||||
|
q = _get_queue()
|
||||||
|
while True:
|
||||||
|
job_value, recording_id, attempt = await q.get()
|
||||||
|
ctx = {"job_try": attempt}
|
||||||
|
try:
|
||||||
|
if job_value == JobType.transcribe.value:
|
||||||
|
await processing.run_transcribe(ctx, recording_id)
|
||||||
|
else:
|
||||||
|
await processing.run_summarize(ctx, recording_id)
|
||||||
|
except processing.ProviderTransientError as e:
|
||||||
|
if attempt < processing.MAX_TRIES:
|
||||||
|
logger.warning(
|
||||||
|
"inline job %s:%s transient failure (%s); retry %d/%d",
|
||||||
|
job_value, recording_id, e, attempt + 1, processing.MAX_TRIES,
|
||||||
|
)
|
||||||
|
await asyncio.sleep(RETRY_DELAY_SECONDS)
|
||||||
|
await q.put((job_value, recording_id, attempt + 1))
|
||||||
|
else:
|
||||||
|
logger.error(
|
||||||
|
"inline job %s:%s failed after %d tries: %s",
|
||||||
|
job_value, recording_id, attempt, e,
|
||||||
|
)
|
||||||
|
except asyncio.CancelledError:
|
||||||
|
raise
|
||||||
|
except Exception: # noqa: BLE001 — consumer must survive any task crash
|
||||||
|
logger.exception("inline job %s:%s crashed", job_value, recording_id)
|
||||||
|
finally:
|
||||||
|
q.task_done()
|
||||||
74
backend/shonar/services/media.py
Normal file
74
backend/shonar/services/media.py
Normal file
|
|
@ -0,0 +1,74 @@
|
||||||
|
"""Audio format validation: declared MIME type vs magic bytes.
|
||||||
|
|
||||||
|
Families decouple declared MIME aliases (audio/mp4 vs audio/m4a) from what
|
||||||
|
the bytes actually are. The original file is stored exactly as uploaded —
|
||||||
|
validation never rewrites it.
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
FAMILY_BY_MIME = {
|
||||||
|
"audio/mp4": "mp4",
|
||||||
|
"audio/m4a": "mp4",
|
||||||
|
"audio/aac": "aac",
|
||||||
|
"audio/wav": "wav",
|
||||||
|
"audio/x-wav": "wav",
|
||||||
|
"audio/ogg": "ogg",
|
||||||
|
"audio/opus": "ogg",
|
||||||
|
"audio/webm": "webm",
|
||||||
|
"audio/mpeg": "mpeg",
|
||||||
|
}
|
||||||
|
|
||||||
|
EXT_BY_FAMILY = {
|
||||||
|
"mp4": ".m4a",
|
||||||
|
"aac": ".aac",
|
||||||
|
"wav": ".wav",
|
||||||
|
"ogg": ".ogg",
|
||||||
|
"webm": ".webm",
|
||||||
|
"mpeg": ".mp3",
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
def sniff_audio_family(data: bytes) -> str | None:
|
||||||
|
"""Return the audio family from magic bytes, or None if unrecognized."""
|
||||||
|
if len(data) >= 12:
|
||||||
|
if data[4:8] == b"ftyp":
|
||||||
|
return "mp4"
|
||||||
|
if data[:4] == b"RIFF" and data[8:12] == b"WAVE":
|
||||||
|
return "wav"
|
||||||
|
if data[:4] == b"OggS":
|
||||||
|
return "ogg"
|
||||||
|
if data[:4] == b"\x1a\x45\xdf\xa3":
|
||||||
|
return "webm"
|
||||||
|
if data[:3] == b"ID3":
|
||||||
|
return "mpeg"
|
||||||
|
if len(data) >= 2 and data[0] == 0xFF and (data[1] & 0xF6) == 0xF0:
|
||||||
|
# ADTS frame sync: AAC (also accepted as mpeg-family audio)
|
||||||
|
return "aac"
|
||||||
|
return None
|
||||||
|
|
||||||
|
|
||||||
|
def declared_family(mime_type: str) -> str | None:
|
||||||
|
return FAMILY_BY_MIME.get(mime_type.lower().split(";")[0].strip())
|
||||||
|
|
||||||
|
|
||||||
|
def is_compatible(mime_type: str, data: bytes) -> bool:
|
||||||
|
"""True when the declared MIME matches the sniffed magic bytes.
|
||||||
|
|
||||||
|
mpeg and aac are treated as one family: Android records AAC in ADTS or
|
||||||
|
in MP4 containers and MIME reporting around these is inconsistent.
|
||||||
|
"""
|
||||||
|
declared = declared_family(mime_type)
|
||||||
|
sniffed = sniff_audio_family(data)
|
||||||
|
if declared is None or sniffed is None:
|
||||||
|
return False
|
||||||
|
if {declared, sniffed} == {"mpeg", "aac"}:
|
||||||
|
return True
|
||||||
|
return declared == sniffed
|
||||||
|
|
||||||
|
|
||||||
|
def extension_for(mime_type: str, data: bytes) -> str:
|
||||||
|
sniffed = sniff_audio_family(data)
|
||||||
|
if sniffed is not None:
|
||||||
|
return EXT_BY_FAMILY[sniffed]
|
||||||
|
return EXT_BY_FAMILY.get(declared_family(mime_type) or "", ".bin")
|
||||||
591
backend/shonar/services/processing.py
Normal file
591
backend/shonar/services/processing.py
Normal file
|
|
@ -0,0 +1,591 @@
|
||||||
|
"""AI processing pipeline (M7): uploaded -> transcribed -> summarized.
|
||||||
|
|
||||||
|
State lives in the database (ProcessingJob rows); arq/redis is transport
|
||||||
|
only. That ordering is deliberate: if redis is down, uploads still succeed
|
||||||
|
(the jobs sit queued) and the worker sweep picks them up. The only race is
|
||||||
|
a task running before the API transaction commits (recording invisible) —
|
||||||
|
missing rows are transient failures, so arq retry absorbs it.
|
||||||
|
|
||||||
|
Entry points:
|
||||||
|
- ``enqueue_for_recording``: called from upload finalize AND re-callable
|
||||||
|
later (manual transcripts in M8 re-enter here). Idempotent.
|
||||||
|
- ``run_transcribe`` / ``run_summarize``: arq task bodies. ``ctx`` is an
|
||||||
|
arq context in production and a plain dict in tests; only
|
||||||
|
``ctx.get("job_try", 1)`` is read.
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import asyncio
|
||||||
|
import logging
|
||||||
|
import uuid
|
||||||
|
from datetime import timedelta
|
||||||
|
|
||||||
|
from sqlalchemy import func, select
|
||||||
|
from sqlalchemy.ext.asyncio import AsyncSession
|
||||||
|
|
||||||
|
from shonar.core.config import get_settings
|
||||||
|
from shonar.db.models import (
|
||||||
|
Asset,
|
||||||
|
AssetKind,
|
||||||
|
JobStatus,
|
||||||
|
JobType,
|
||||||
|
ProcessingJob,
|
||||||
|
ProcessingStatus,
|
||||||
|
Recording,
|
||||||
|
Summary,
|
||||||
|
Transcript,
|
||||||
|
utcnow,
|
||||||
|
)
|
||||||
|
from shonar.services.ai import (
|
||||||
|
AIError,
|
||||||
|
ProviderTransientError,
|
||||||
|
get_llm_provider,
|
||||||
|
get_transcription_provider,
|
||||||
|
)
|
||||||
|
from shonar.storage import get_storage
|
||||||
|
|
||||||
|
logger = logging.getLogger("shonar.processing")
|
||||||
|
|
||||||
|
|
||||||
|
def _accepts_on_progress(provider) -> bool:
|
||||||
|
"""True when the provider's transcribe() accepts on_progress=."""
|
||||||
|
import inspect
|
||||||
|
|
||||||
|
try:
|
||||||
|
return "on_progress" in inspect.signature(provider.transcribe).parameters
|
||||||
|
except (TypeError, ValueError): # builtins / exotic callables
|
||||||
|
return False
|
||||||
|
|
||||||
|
|
||||||
|
def _thread_progress_reporter(job_id):
|
||||||
|
"""Callback safe to invoke from a worker thread (faster-whisper runs
|
||||||
|
via asyncio.to_thread): schedules a tiny session update on the loop.
|
||||||
|
|
||||||
|
Progress writes are best-effort display state — failures are logged
|
||||||
|
and swallowed, never allowed to disturb the transcription itself."""
|
||||||
|
loop = asyncio.get_running_loop()
|
||||||
|
|
||||||
|
def report(pct: int) -> None:
|
||||||
|
async def _write() -> None:
|
||||||
|
from sqlalchemy import update
|
||||||
|
|
||||||
|
from shonar.db.session import session_factory
|
||||||
|
|
||||||
|
try:
|
||||||
|
async with session_factory()() as s:
|
||||||
|
await s.execute(
|
||||||
|
update(ProcessingJob)
|
||||||
|
.where(ProcessingJob.id == job_id)
|
||||||
|
.values(progress=int(pct))
|
||||||
|
)
|
||||||
|
await s.commit()
|
||||||
|
except Exception: # pragma: no cover - display state only
|
||||||
|
logger.debug("progress write failed for job %s", job_id, exc_info=True)
|
||||||
|
|
||||||
|
try:
|
||||||
|
asyncio.run_coroutine_threadsafe(_write(), loop)
|
||||||
|
except RuntimeError: # loop already gone (shutdown race)
|
||||||
|
logger.debug("progress dropped for job %s (loop gone)", job_id)
|
||||||
|
|
||||||
|
return report
|
||||||
|
|
||||||
|
MAX_TRIES = 3
|
||||||
|
|
||||||
|
# A `running` job younger than this is treated as live work, not a crash
|
||||||
|
# relic — long transcriptions legitimately outrun the 5-minute sweep, and
|
||||||
|
# re-enqueuing them mid-flight caused a copy storm (one whisper per copy).
|
||||||
|
STALE_RUNNING_AFTER = timedelta(hours=2)
|
||||||
|
|
||||||
|
|
||||||
|
# --- entry -----------------------------------------------------------------
|
||||||
|
|
||||||
|
|
||||||
|
async def enqueue_for_recording(session: AsyncSession, rec: Recording) -> list[JobType]:
|
||||||
|
"""Queue whatever AI stages apply. Safe to call repeatedly: completed
|
||||||
|
work is never redone, failed work is reset for another attempt."""
|
||||||
|
if rec.deleted_at is not None:
|
||||||
|
return []
|
||||||
|
settings = get_settings()
|
||||||
|
tprov = get_transcription_provider(settings)
|
||||||
|
lprov = get_llm_provider(settings)
|
||||||
|
queued: list[JobType] = []
|
||||||
|
if tprov is not None and await reset_or_create(session, rec, JobType.transcribe):
|
||||||
|
queued.append(JobType.transcribe)
|
||||||
|
if (
|
||||||
|
lprov is not None
|
||||||
|
and await latest_transcript_text(session, rec.id) is not None
|
||||||
|
and await reset_or_create(session, rec, JobType.summarize)
|
||||||
|
):
|
||||||
|
queued.append(JobType.summarize)
|
||||||
|
if tprov is None and lprov is None:
|
||||||
|
if rec.processing_status != ProcessingStatus.ai_disabled:
|
||||||
|
rec.processing_status = ProcessingStatus.ai_disabled
|
||||||
|
rec.processing_error = None
|
||||||
|
elif queued and rec.processing_status not in (
|
||||||
|
ProcessingStatus.processing,
|
||||||
|
ProcessingStatus.completed,
|
||||||
|
):
|
||||||
|
rec.processing_status = ProcessingStatus.queued
|
||||||
|
rec.processing_error = None
|
||||||
|
await session.flush()
|
||||||
|
for jt in queued:
|
||||||
|
await transport_enqueue(jt, rec.id)
|
||||||
|
return queued
|
||||||
|
|
||||||
|
|
||||||
|
async def reset_or_create(
|
||||||
|
session: AsyncSession, rec: Recording, job_type: JobType
|
||||||
|
) -> bool:
|
||||||
|
"""Ensure a queued job row. Returns True when (re)queued now: new rows,
|
||||||
|
plus failed/skipped rows (a re-upload or re-entry deserves another
|
||||||
|
attempt). Queued/running/succeeded rows are left alone."""
|
||||||
|
existing = await session.scalar(
|
||||||
|
select(ProcessingJob)
|
||||||
|
.where(
|
||||||
|
ProcessingJob.recording_id == rec.id,
|
||||||
|
ProcessingJob.job_type == job_type,
|
||||||
|
)
|
||||||
|
.order_by(ProcessingJob.id.desc())
|
||||||
|
)
|
||||||
|
if existing is None:
|
||||||
|
session.add(ProcessingJob(recording_id=rec.id, job_type=job_type))
|
||||||
|
return True
|
||||||
|
if existing.status in (JobStatus.queued, JobStatus.running, JobStatus.succeeded):
|
||||||
|
return False
|
||||||
|
existing.status = JobStatus.queued
|
||||||
|
existing.attempt = 0
|
||||||
|
existing.error = None
|
||||||
|
existing.stage = None
|
||||||
|
existing.progress = None
|
||||||
|
existing.started_at = None
|
||||||
|
existing.finished_at = None
|
||||||
|
return True
|
||||||
|
|
||||||
|
|
||||||
|
async def latest_transcript_text(
|
||||||
|
session: AsyncSession, recording_id: uuid.UUID
|
||||||
|
) -> str | None:
|
||||||
|
"""Newest non-superseded transcript; user-edited rows win over newer
|
||||||
|
machine rows (an edit is a verdict, not a draft)."""
|
||||||
|
rows = (
|
||||||
|
await session.scalars(
|
||||||
|
select(Transcript)
|
||||||
|
.where(
|
||||||
|
Transcript.recording_id == recording_id,
|
||||||
|
Transcript.superseded_at.is_(None),
|
||||||
|
)
|
||||||
|
.order_by(Transcript.version.desc())
|
||||||
|
)
|
||||||
|
).all()
|
||||||
|
if not rows:
|
||||||
|
return None
|
||||||
|
for r in rows:
|
||||||
|
if r.edited_by_user and r.text.strip():
|
||||||
|
return r.text
|
||||||
|
text = rows[0].text
|
||||||
|
return text if text.strip() else None
|
||||||
|
|
||||||
|
|
||||||
|
async def transport_enqueue(job_type: JobType, recording_id: uuid.UUID) -> None:
|
||||||
|
"""Best-effort trigger for the configured backend ('arq' or 'inline').
|
||||||
|
Failure only logs: the DB rows are the real queue and the sweep picks
|
||||||
|
up anything the transport missed."""
|
||||||
|
if get_settings().queue_backend.strip().lower() == "inline":
|
||||||
|
from shonar.services import inline_queue
|
||||||
|
|
||||||
|
await inline_queue.enqueue(job_type, str(recording_id))
|
||||||
|
return
|
||||||
|
from arq import create_pool
|
||||||
|
from arq.connections import RedisSettings
|
||||||
|
from arq.constants import result_key_prefix
|
||||||
|
|
||||||
|
job_id = f"{job_type.value}:{recording_id}"
|
||||||
|
try:
|
||||||
|
pool = await create_pool(RedisSettings.from_dsn(get_settings().redis_url))
|
||||||
|
try:
|
||||||
|
# Deterministic _job_id: arq refuses a duplicate while a copy of
|
||||||
|
# (job_type, recording) is queued or running, so sweep pokes and
|
||||||
|
# retried enqueues can never pile up concurrent copies of the
|
||||||
|
# same work.
|
||||||
|
#
|
||||||
|
# arq dedupes on job key OR result key, and failed runs leave a
|
||||||
|
# result behind for keep_result days — that would silently swallow
|
||||||
|
# deliberate re-runs of previously-finished work. Drop only the
|
||||||
|
# stale RESULT key first (never the job key: that is exactly what
|
||||||
|
# shields a live copy from duplicate enqueue).
|
||||||
|
await pool.delete(result_key_prefix + job_id)
|
||||||
|
await pool.enqueue_job(
|
||||||
|
"run_" + job_type.value,
|
||||||
|
str(recording_id),
|
||||||
|
_job_id=job_id,
|
||||||
|
)
|
||||||
|
finally:
|
||||||
|
await pool.aclose()
|
||||||
|
except Exception as e: # noqa: BLE001 — transport must never break uploads
|
||||||
|
logger.warning("arq enqueue failed (%s); worker sweep will pick it up", e)
|
||||||
|
|
||||||
|
|
||||||
|
# --- tasks -----------------------------------------------------------------
|
||||||
|
|
||||||
|
|
||||||
|
def _job_try(ctx: dict) -> int:
|
||||||
|
try:
|
||||||
|
return int(ctx.get("job_try", 1))
|
||||||
|
except (TypeError, ValueError):
|
||||||
|
return 1
|
||||||
|
|
||||||
|
|
||||||
|
async def _load(session: AsyncSession, recording_id: str) -> Recording | None:
|
||||||
|
try:
|
||||||
|
rid = uuid.UUID(recording_id)
|
||||||
|
except ValueError:
|
||||||
|
return None
|
||||||
|
return await session.get(Recording, rid)
|
||||||
|
|
||||||
|
|
||||||
|
async def _job(
|
||||||
|
session: AsyncSession, rec: Recording, job_type: JobType
|
||||||
|
) -> ProcessingJob:
|
||||||
|
job = await session.scalar(
|
||||||
|
select(ProcessingJob)
|
||||||
|
.where(
|
||||||
|
ProcessingJob.recording_id == rec.id,
|
||||||
|
ProcessingJob.job_type == job_type,
|
||||||
|
)
|
||||||
|
.order_by(ProcessingJob.id.desc())
|
||||||
|
)
|
||||||
|
if job is None:
|
||||||
|
job = ProcessingJob(recording_id=rec.id, job_type=job_type)
|
||||||
|
session.add(job)
|
||||||
|
await session.flush()
|
||||||
|
return job
|
||||||
|
|
||||||
|
|
||||||
|
async def _fail(
|
||||||
|
session: AsyncSession,
|
||||||
|
rec: Recording,
|
||||||
|
job: ProcessingJob,
|
||||||
|
message: str,
|
||||||
|
ctx: dict,
|
||||||
|
exc: AIError | None = None,
|
||||||
|
) -> None:
|
||||||
|
"""Config errors fail now; transient errors fail only on the last try
|
||||||
|
(returning normally), otherwise they raise for arq retry."""
|
||||||
|
transient = exc is None or isinstance(exc, ProviderTransientError)
|
||||||
|
if transient and _job_try(ctx) < MAX_TRIES:
|
||||||
|
# Running state was committed before the long phase; put the row
|
||||||
|
# back to queued for the retry and persist that (a raise no longer
|
||||||
|
# rolls the pre-phase commit back).
|
||||||
|
job.status = JobStatus.queued
|
||||||
|
job.attempt = _job_try(ctx)
|
||||||
|
job.stage = None
|
||||||
|
await session.commit()
|
||||||
|
raise ProviderTransientError(message)
|
||||||
|
job.status = JobStatus.failed
|
||||||
|
job.error = message
|
||||||
|
job.stage = None
|
||||||
|
job.finished_at = utcnow()
|
||||||
|
if job.job_type == JobType.summarize and await latest_transcript_text(
|
||||||
|
session, rec.id
|
||||||
|
):
|
||||||
|
# The transcript is usable; a summary timeout must not mark the
|
||||||
|
# whole recording failed (the summary can be re-run separately).
|
||||||
|
rec.processing_status = ProcessingStatus.completed
|
||||||
|
rec.processing_error = f"Summary failed: {message}"
|
||||||
|
else:
|
||||||
|
rec.processing_status = ProcessingStatus.failed
|
||||||
|
rec.processing_error = message
|
||||||
|
await session.flush()
|
||||||
|
|
||||||
|
|
||||||
|
async def run_transcribe(ctx: dict, recording_id: str) -> None:
|
||||||
|
"""Transcribe the original audio; chain into summarization when an LLM
|
||||||
|
is configured."""
|
||||||
|
from shonar.db.session import session_factory
|
||||||
|
|
||||||
|
settings = get_settings()
|
||||||
|
async with session_factory()() as session:
|
||||||
|
rec = await _load(session, recording_id)
|
||||||
|
if rec is None or rec.deleted_at is not None:
|
||||||
|
# Finalize race: the API transaction may not have committed yet.
|
||||||
|
raise ProviderTransientError("Recording not ready; retrying.")
|
||||||
|
job = await _job(session, rec, JobType.transcribe)
|
||||||
|
provider = get_transcription_provider(settings)
|
||||||
|
if provider is None:
|
||||||
|
job.status = JobStatus.skipped
|
||||||
|
await session.flush()
|
||||||
|
await session.commit()
|
||||||
|
await _maybe_chain_summarize(session, rec)
|
||||||
|
return
|
||||||
|
# The recording's saved model is authoritative: a per-recording
|
||||||
|
# override wins, otherwise the global default in force at finalize
|
||||||
|
# time (changing the default never rewrites history).
|
||||||
|
from shonar.services.ai.model_registry import (
|
||||||
|
effective_model,
|
||||||
|
get_global_default_model,
|
||||||
|
)
|
||||||
|
|
||||||
|
model = effective_model(rec.transcription_model, await get_global_default_model(session))
|
||||||
|
if provider.name == "faster_whisper":
|
||||||
|
from shonar.services.ai.faster_whisper import FasterWhisperProvider
|
||||||
|
|
||||||
|
try:
|
||||||
|
provider = FasterWhisperProvider(model=model)
|
||||||
|
except AIError as e:
|
||||||
|
await _fail(session, rec, job, str(e), ctx, e)
|
||||||
|
await session.commit()
|
||||||
|
return
|
||||||
|
original = await session.scalar(
|
||||||
|
select(Asset).where(
|
||||||
|
Asset.recording_id == rec.id, Asset.kind == AssetKind.original
|
||||||
|
)
|
||||||
|
)
|
||||||
|
if original is None:
|
||||||
|
raise ProviderTransientError("Audio not ready; retrying.")
|
||||||
|
job.status = JobStatus.running
|
||||||
|
job.attempt = _job_try(ctx)
|
||||||
|
job.started_at = utcnow()
|
||||||
|
job.stage = "loading-model"
|
||||||
|
job.progress = None
|
||||||
|
rec.processing_status = ProcessingStatus.processing
|
||||||
|
rec.processing_error = None
|
||||||
|
# Commit before the long CPU phase: an open transaction is invisible
|
||||||
|
# to other readers (and on SQLite it locks out the progress writer).
|
||||||
|
await session.commit()
|
||||||
|
try:
|
||||||
|
audio = await get_storage().get(original.storage_key)
|
||||||
|
job.stage = "transcribing"
|
||||||
|
await session.commit()
|
||||||
|
kwargs = {}
|
||||||
|
if _accepts_on_progress(provider):
|
||||||
|
kwargs["on_progress"] = _thread_progress_reporter(job.id)
|
||||||
|
result = await provider.transcribe(audio, original.mime_type, **kwargs)
|
||||||
|
except AIError as e:
|
||||||
|
await _fail(session, rec, job, str(e), ctx, e)
|
||||||
|
await session.commit()
|
||||||
|
return
|
||||||
|
await store_transcript(
|
||||||
|
session, rec, result.text, result.segments, result.language,
|
||||||
|
provider.name, getattr(result, "model", ""),
|
||||||
|
)
|
||||||
|
job.status = JobStatus.succeeded
|
||||||
|
job.stage = None
|
||||||
|
job.progress = 100
|
||||||
|
job.finished_at = utcnow()
|
||||||
|
await session.flush()
|
||||||
|
await _maybe_chain_summarize(session, rec)
|
||||||
|
await session.commit()
|
||||||
|
|
||||||
|
|
||||||
|
async def _maybe_chain_summarize(session: AsyncSession, rec: Recording) -> None:
|
||||||
|
"""After transcription (or a skip): summarize when possible, else finish."""
|
||||||
|
settings = get_settings()
|
||||||
|
if get_llm_provider(settings) is None:
|
||||||
|
if rec.processing_status != ProcessingStatus.completed:
|
||||||
|
rec.processing_status = ProcessingStatus.completed
|
||||||
|
rec.processing_error = None
|
||||||
|
await session.flush()
|
||||||
|
return
|
||||||
|
if await latest_transcript_text(session, rec.id) is None:
|
||||||
|
# LLM configured but nothing to summarize (e.g. empty transcript).
|
||||||
|
existing = await session.scalar(
|
||||||
|
select(ProcessingJob).where(
|
||||||
|
ProcessingJob.recording_id == rec.id,
|
||||||
|
ProcessingJob.job_type == JobType.summarize,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
if existing is not None and existing.status == JobStatus.queued:
|
||||||
|
existing.status = JobStatus.skipped
|
||||||
|
if rec.processing_status != ProcessingStatus.completed:
|
||||||
|
rec.processing_status = ProcessingStatus.completed
|
||||||
|
await session.flush()
|
||||||
|
return
|
||||||
|
if await reset_or_create(session, rec, JobType.summarize):
|
||||||
|
await transport_enqueue(JobType.summarize, rec.id)
|
||||||
|
# Status stays `processing` until the summarize task lands.
|
||||||
|
|
||||||
|
|
||||||
|
async def run_summarize(ctx: dict, recording_id: str) -> None:
|
||||||
|
"""Summarize the latest transcript into the structured summary shape."""
|
||||||
|
from shonar.db.session import session_factory
|
||||||
|
|
||||||
|
settings = get_settings()
|
||||||
|
async with session_factory()() as session:
|
||||||
|
rec = await _load(session, recording_id)
|
||||||
|
if rec is None or rec.deleted_at is not None:
|
||||||
|
raise ProviderTransientError("Recording not ready; retrying.")
|
||||||
|
job = await _job(session, rec, JobType.summarize)
|
||||||
|
provider = get_llm_provider(settings)
|
||||||
|
if provider is None:
|
||||||
|
job.status = JobStatus.skipped
|
||||||
|
await session.flush()
|
||||||
|
await session.commit()
|
||||||
|
return
|
||||||
|
text = await latest_transcript_text(session, rec.id)
|
||||||
|
if text is None:
|
||||||
|
job.status = JobStatus.skipped
|
||||||
|
await session.flush()
|
||||||
|
if rec.processing_status != ProcessingStatus.completed:
|
||||||
|
rec.processing_status = ProcessingStatus.completed
|
||||||
|
await session.commit()
|
||||||
|
return
|
||||||
|
job.status = JobStatus.running
|
||||||
|
job.attempt = _job_try(ctx)
|
||||||
|
job.started_at = utcnow()
|
||||||
|
job.stage = "summarizing"
|
||||||
|
rec.processing_status = ProcessingStatus.processing
|
||||||
|
await session.commit() # visible before the long LLM call
|
||||||
|
try:
|
||||||
|
result = await provider.summarize(text, title=rec.title)
|
||||||
|
except AIError as e:
|
||||||
|
await _fail(session, rec, job, str(e), ctx, e)
|
||||||
|
await session.commit()
|
||||||
|
return
|
||||||
|
await store_summary(session, rec, result.to_dict(), provider.name, result.model)
|
||||||
|
job.status = JobStatus.succeeded
|
||||||
|
job.stage = None
|
||||||
|
job.progress = 100
|
||||||
|
job.finished_at = utcnow()
|
||||||
|
rec.processing_status = ProcessingStatus.completed
|
||||||
|
rec.processing_error = None
|
||||||
|
await session.flush()
|
||||||
|
await session.commit()
|
||||||
|
|
||||||
|
|
||||||
|
async def store_transcript(
|
||||||
|
session: AsyncSession,
|
||||||
|
rec: Recording,
|
||||||
|
text: str,
|
||||||
|
segments: list,
|
||||||
|
language: str | None,
|
||||||
|
provider_name: str,
|
||||||
|
model: str,
|
||||||
|
) -> None:
|
||||||
|
"""Insert a new auto version; a newest user-edited row wins instead and
|
||||||
|
nothing is inserted (edits are verdicts)."""
|
||||||
|
existing = (
|
||||||
|
await session.scalars(
|
||||||
|
select(Transcript)
|
||||||
|
.where(
|
||||||
|
Transcript.recording_id == rec.id,
|
||||||
|
Transcript.superseded_at.is_(None),
|
||||||
|
)
|
||||||
|
.order_by(Transcript.version.desc())
|
||||||
|
)
|
||||||
|
).all()
|
||||||
|
if existing and existing[0].edited_by_user:
|
||||||
|
return
|
||||||
|
now = utcnow()
|
||||||
|
max_version = await session.scalar(
|
||||||
|
select(func.max(Transcript.version)).where(Transcript.recording_id == rec.id)
|
||||||
|
)
|
||||||
|
for row in existing:
|
||||||
|
row.superseded_at = now
|
||||||
|
session.add(
|
||||||
|
Transcript(
|
||||||
|
recording_id=rec.id,
|
||||||
|
version=(max_version or 0) + 1,
|
||||||
|
language=language,
|
||||||
|
provider=provider_name,
|
||||||
|
model=model,
|
||||||
|
text=text,
|
||||||
|
segments=[
|
||||||
|
{"start": s.start, "end": s.end, "text": s.text, "speaker": s.speaker}
|
||||||
|
for s in segments
|
||||||
|
],
|
||||||
|
edited_by_user=False,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
await session.flush()
|
||||||
|
|
||||||
|
|
||||||
|
async def store_summary(
|
||||||
|
session: AsyncSession,
|
||||||
|
rec: Recording,
|
||||||
|
content: dict,
|
||||||
|
provider_name: str,
|
||||||
|
model: str,
|
||||||
|
) -> None:
|
||||||
|
existing = (
|
||||||
|
await session.scalars(
|
||||||
|
select(Summary)
|
||||||
|
.where(
|
||||||
|
Summary.recording_id == rec.id,
|
||||||
|
Summary.superseded_at.is_(None),
|
||||||
|
)
|
||||||
|
.order_by(Summary.version.desc())
|
||||||
|
)
|
||||||
|
).all()
|
||||||
|
if existing and existing[0].edited_by_user:
|
||||||
|
return
|
||||||
|
now = utcnow()
|
||||||
|
max_version = await session.scalar(
|
||||||
|
select(func.max(Summary.version)).where(Summary.recording_id == rec.id)
|
||||||
|
)
|
||||||
|
for row in existing:
|
||||||
|
row.superseded_at = now
|
||||||
|
session.add(
|
||||||
|
Summary(
|
||||||
|
recording_id=rec.id,
|
||||||
|
version=(max_version or 0) + 1,
|
||||||
|
provider=provider_name,
|
||||||
|
model=model,
|
||||||
|
content=content,
|
||||||
|
edited_by_user=False,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
await session.flush()
|
||||||
|
|
||||||
|
|
||||||
|
async def sweep_stale(limit: int = 100) -> int:
|
||||||
|
"""Crash recovery + transport-loss backstop: requeue jobs stuck running
|
||||||
|
or sitting queued, newest catastrophe first. Returns jobs re-enqueued."""
|
||||||
|
from shonar.db.session import session_factory
|
||||||
|
|
||||||
|
count = 0
|
||||||
|
async with session_factory()() as session:
|
||||||
|
rows = (
|
||||||
|
await session.scalars(
|
||||||
|
select(ProcessingJob)
|
||||||
|
.where(ProcessingJob.status.in_([JobStatus.queued, JobStatus.running]))
|
||||||
|
.order_by(ProcessingJob.id.desc())
|
||||||
|
.limit(limit)
|
||||||
|
)
|
||||||
|
).all()
|
||||||
|
for job in rows:
|
||||||
|
rec = await session.get(Recording, job.recording_id)
|
||||||
|
if rec is None or rec.deleted_at is not None:
|
||||||
|
job.status = JobStatus.skipped
|
||||||
|
continue
|
||||||
|
if job.status == JobStatus.running:
|
||||||
|
if job.attempt >= MAX_TRIES:
|
||||||
|
job.status = JobStatus.failed
|
||||||
|
job.error = "Worker died too many times."
|
||||||
|
job.finished_at = utcnow()
|
||||||
|
continue
|
||||||
|
# Live work, not a crash relic: a running job that started
|
||||||
|
# recently belongs to a worker still chewing it (long audio
|
||||||
|
# outruns the sweep cadence). Only orphaned runs — no
|
||||||
|
# started_at, or older than the grace window — get requeued.
|
||||||
|
if job.started_at is not None and (
|
||||||
|
utcnow() - job.started_at
|
||||||
|
) < STALE_RUNNING_AFTER:
|
||||||
|
continue
|
||||||
|
job.status = JobStatus.queued
|
||||||
|
job.error = None
|
||||||
|
count += 1
|
||||||
|
await session.commit()
|
||||||
|
# Transport outside the transaction: rows are the queue, this just pokes.
|
||||||
|
async with session_factory()() as session:
|
||||||
|
rows = (
|
||||||
|
await session.scalars(
|
||||||
|
select(ProcessingJob)
|
||||||
|
.where(ProcessingJob.status == JobStatus.queued)
|
||||||
|
.order_by(ProcessingJob.id.desc())
|
||||||
|
.limit(limit)
|
||||||
|
)
|
||||||
|
).all()
|
||||||
|
for job in rows:
|
||||||
|
await transport_enqueue(job.job_type, job.recording_id)
|
||||||
|
return count
|
||||||
101
backend/shonar/services/retention.py
Normal file
101
backend/shonar/services/retention.py
Normal file
|
|
@ -0,0 +1,101 @@
|
||||||
|
"""Retention sweep (M9): hard-delete what passed its grace window.
|
||||||
|
|
||||||
|
Two policies, both driven by ``deleted_at``:
|
||||||
|
|
||||||
|
* **Recordings** soft-deleted more than ``retention_grace_days`` ago have
|
||||||
|
their rows removed (cascade cleans transcripts/summaries/jobs/tags) and
|
||||||
|
every stored asset file deleted best-effort.
|
||||||
|
* **Accounts** deleted more than ``retention_grace_days`` ago are hard
|
||||||
|
deleted (user cascade takes their recordings/assets/devices/tokens);
|
||||||
|
their storage files are collected the same way.
|
||||||
|
|
||||||
|
Running inside the grace window is a no-op, so an accidental delete stays
|
||||||
|
cancellable until the sweep actually fires. The sweep is idempotent and
|
||||||
|
safe to run on any cadence (worker cron + inline-queue timer).
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import contextlib
|
||||||
|
import logging
|
||||||
|
from datetime import timedelta
|
||||||
|
|
||||||
|
from sqlalchemy import delete as sql_delete
|
||||||
|
from sqlalchemy import select
|
||||||
|
|
||||||
|
from shonar.core.config import get_settings
|
||||||
|
from shonar.db.models import Asset, Recording, User, utcnow
|
||||||
|
from shonar.db.session import session_factory
|
||||||
|
from shonar.storage import get_storage
|
||||||
|
|
||||||
|
logger = logging.getLogger("shonar.retention")
|
||||||
|
|
||||||
|
|
||||||
|
async def sweep_deleted(limit: int = 200) -> dict[str, int]:
|
||||||
|
"""Hard-purge expired recordings and accounts. Returns counts."""
|
||||||
|
grace = timedelta(days=get_settings().retention_grace_days)
|
||||||
|
cutoff = utcnow() - grace
|
||||||
|
purged = {"recordings": 0, "users": 0, "files": 0}
|
||||||
|
storage = get_storage()
|
||||||
|
|
||||||
|
async with session_factory()() as session:
|
||||||
|
# --- recordings (skip rows whose account is also expiring: the
|
||||||
|
# user cascade below collects their files in one pass) ---
|
||||||
|
expiring_users = select(User.id).where(
|
||||||
|
User.deleted_at.is_not(None), User.deleted_at < cutoff
|
||||||
|
)
|
||||||
|
recs = list(
|
||||||
|
await session.scalars(
|
||||||
|
select(Recording)
|
||||||
|
.where(
|
||||||
|
Recording.deleted_at.is_not(None),
|
||||||
|
Recording.deleted_at < cutoff,
|
||||||
|
Recording.user_id.notin_(expiring_users),
|
||||||
|
)
|
||||||
|
.limit(limit)
|
||||||
|
)
|
||||||
|
)
|
||||||
|
for rec in recs:
|
||||||
|
assets = list(
|
||||||
|
await session.scalars(select(Asset).where(Asset.recording_id == rec.id))
|
||||||
|
)
|
||||||
|
await session.delete(rec)
|
||||||
|
await session.flush()
|
||||||
|
for a in assets:
|
||||||
|
with contextlib.suppress(Exception): # best effort; DB row is gone
|
||||||
|
await storage.delete(a.storage_key)
|
||||||
|
purged["files"] += 1
|
||||||
|
purged["recordings"] += 1
|
||||||
|
|
||||||
|
# --- accounts ---
|
||||||
|
users = list(
|
||||||
|
await session.scalars(
|
||||||
|
select(User)
|
||||||
|
.where(User.deleted_at.is_not(None), User.deleted_at < cutoff)
|
||||||
|
.limit(limit)
|
||||||
|
)
|
||||||
|
)
|
||||||
|
for user in users:
|
||||||
|
assets = list(
|
||||||
|
await session.scalars(select(Asset).where(Asset.user_id == user.id))
|
||||||
|
)
|
||||||
|
keys = [a.storage_key for a in assets]
|
||||||
|
# DB-level delete: the ORM would null the NOT NULL FKs of the
|
||||||
|
# user's recordings before the ON DELETE CASCADE could fire.
|
||||||
|
await session.execute(sql_delete(User).where(User.id == user.id))
|
||||||
|
await session.flush()
|
||||||
|
for key in keys:
|
||||||
|
with contextlib.suppress(Exception):
|
||||||
|
await storage.delete(key)
|
||||||
|
purged["files"] += 1
|
||||||
|
purged["users"] += 1
|
||||||
|
|
||||||
|
await session.commit()
|
||||||
|
|
||||||
|
if purged["recordings"] or purged["users"]:
|
||||||
|
logger.info(
|
||||||
|
"retention sweep purged %d recordings, %d accounts, %d files (grace %dd)",
|
||||||
|
purged["recordings"], purged["users"], purged["files"],
|
||||||
|
get_settings().retention_grace_days,
|
||||||
|
)
|
||||||
|
return purged
|
||||||
292
backend/shonar/services/search.py
Normal file
292
backend/shonar/services/search.py
Normal file
|
|
@ -0,0 +1,292 @@
|
||||||
|
"""Full-text search across recordings (M9).
|
||||||
|
|
||||||
|
Postgres uses the tsvector columns from migration ``fts0000000001``
|
||||||
|
(title/notes, transcript text, summary content, tag names). SQLite (the
|
||||||
|
desktop bundled-lite engine) falls back to a substring scan; libraries
|
||||||
|
there are single-user and small, and the endpoint contract is identical.
|
||||||
|
|
||||||
|
This module is the whole SearchBackend seam — a Meilisearch/OpenSearch
|
||||||
|
implementation would replace it, not the callers.
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import uuid
|
||||||
|
from dataclasses import dataclass
|
||||||
|
|
||||||
|
from sqlalchemy import select, text
|
||||||
|
from sqlalchemy.ext.asyncio import AsyncSession
|
||||||
|
|
||||||
|
from shonar.db.models import Recording, RecordingTag, Summary, Tag, Transcript
|
||||||
|
|
||||||
|
# ``scope`` values accepted by the endpoint.
|
||||||
|
SCOPES = ("all", "title", "notes", "transcript", "summary", "tag")
|
||||||
|
|
||||||
|
# Hit ordering: a title hit outranks a transcript hit.
|
||||||
|
_FIELD_RANK = {"title": 0, "tag": 1, "notes": 2, "summary": 3, "transcript": 4}
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class SearchHit:
|
||||||
|
recording: Recording
|
||||||
|
field: str # where the best match landed
|
||||||
|
snippet: str
|
||||||
|
|
||||||
|
|
||||||
|
def _snippet_around(value: str, at: int, width: int = 160) -> str:
|
||||||
|
"""~``width`` chars centred on ``at``, word-bounded, with ellipses."""
|
||||||
|
half = width // 2
|
||||||
|
start = max(0, at - half)
|
||||||
|
end = min(len(value), at + half)
|
||||||
|
if start > 0:
|
||||||
|
start = value.rfind(" ", 0, start) + 1 or start
|
||||||
|
if end < len(value):
|
||||||
|
nxt = value.find(" ", end)
|
||||||
|
end = nxt if nxt != -1 else end
|
||||||
|
prefix = "…" if start > 0 else ""
|
||||||
|
suffix = "…" if end < len(value) else ""
|
||||||
|
return f"{prefix}{value[start:end].strip()}{suffix}"
|
||||||
|
|
||||||
|
|
||||||
|
async def tag_names(session: AsyncSession, recording_id: uuid.UUID) -> list[str]:
|
||||||
|
return list(
|
||||||
|
await session.scalars(
|
||||||
|
select(Tag.name)
|
||||||
|
.join(RecordingTag, RecordingTag.tag_id == Tag.id)
|
||||||
|
.where(RecordingTag.recording_id == recording_id)
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
async def _current_transcript_text(session: AsyncSession, recording_id: uuid.UUID) -> str:
|
||||||
|
row = await session.scalar(
|
||||||
|
select(Transcript.text)
|
||||||
|
.where(Transcript.recording_id == recording_id, Transcript.superseded_at.is_(None))
|
||||||
|
.order_by(Transcript.version.desc())
|
||||||
|
)
|
||||||
|
return row or ""
|
||||||
|
|
||||||
|
|
||||||
|
def _summary_flat(content: dict | None) -> str:
|
||||||
|
if not content:
|
||||||
|
return ""
|
||||||
|
parts = [str(content.get("short", "")), str(content.get("detailed", ""))]
|
||||||
|
for key in ("key_points", "decisions", "action_items", "questions"):
|
||||||
|
v = content.get(key)
|
||||||
|
if isinstance(v, list):
|
||||||
|
parts.extend(str(x) for x in v)
|
||||||
|
return " ".join(p for p in parts if p)
|
||||||
|
|
||||||
|
|
||||||
|
# --- SQLite fallback ----------------------------------------------------------
|
||||||
|
|
||||||
|
|
||||||
|
async def _search_sqlite(
|
||||||
|
session: AsyncSession, user_id: uuid.UUID, q: str, scope: str,
|
||||||
|
limit: int, offset: int,
|
||||||
|
) -> tuple[list[SearchHit], int]:
|
||||||
|
needles = [t for t in q.lower().split() if t]
|
||||||
|
if not needles:
|
||||||
|
return [], 0
|
||||||
|
recs = list(
|
||||||
|
await session.scalars(
|
||||||
|
select(Recording)
|
||||||
|
.where(Recording.user_id == user_id, Recording.deleted_at.is_(None))
|
||||||
|
.order_by(Recording.recorded_at.desc())
|
||||||
|
)
|
||||||
|
)
|
||||||
|
hits: list[SearchHit] = []
|
||||||
|
for rec in recs:
|
||||||
|
fields: list[tuple[str, str]] = []
|
||||||
|
if scope in ("all", "title"):
|
||||||
|
fields.append(("title", rec.title or ""))
|
||||||
|
if scope in ("all", "notes"):
|
||||||
|
fields.append(("notes", rec.notes or ""))
|
||||||
|
if scope in ("all", "tag"):
|
||||||
|
fields.append(("tag", " ".join(await tag_names(session, rec.id))))
|
||||||
|
if scope in ("all", "transcript"):
|
||||||
|
fields.append(("transcript", await _current_transcript_text(session, rec.id)))
|
||||||
|
if scope in ("all", "summary"):
|
||||||
|
content = await session.scalar(
|
||||||
|
select(Summary.content)
|
||||||
|
.where(Summary.recording_id == rec.id, Summary.superseded_at.is_(None))
|
||||||
|
.order_by(Summary.version.desc())
|
||||||
|
)
|
||||||
|
fields.append(("summary", _summary_flat(content)))
|
||||||
|
# A field matches when EVERY needle appears in it (AND semantics,
|
||||||
|
# matching plainto_tsquery on the Postgres side).
|
||||||
|
best: SearchHit | None = None
|
||||||
|
for field, value in fields:
|
||||||
|
low = value.lower()
|
||||||
|
if not all(n in low for n in needles):
|
||||||
|
continue
|
||||||
|
at = low.find(needles[0])
|
||||||
|
cand = SearchHit(
|
||||||
|
recording=rec, field=field, snippet=_snippet_around(value, max(at, 0))
|
||||||
|
)
|
||||||
|
if best is None or _FIELD_RANK[field] < _FIELD_RANK[best.field]:
|
||||||
|
best = cand
|
||||||
|
if best is not None:
|
||||||
|
hits.append(best)
|
||||||
|
hits.sort(key=lambda h: (_FIELD_RANK[h.field], h.recording.recorded_at), reverse=False)
|
||||||
|
return hits[offset : offset + limit], len(hits)
|
||||||
|
|
||||||
|
|
||||||
|
# --- PostgreSQL tsvector path -------------------------------------------------
|
||||||
|
|
||||||
|
# CTE ``q`` carries the parsed tsquery so it is computed once. Notes are not
|
||||||
|
# in the recordings search_vector weights the way we want headlines, so notes
|
||||||
|
# match by substring like the SQLite path (title/notes share the vector; the
|
||||||
|
# field classifier prefers 'title' when the vector hits).
|
||||||
|
_PG_MATCHES = """
|
||||||
|
WITH q AS (SELECT plainto_tsquery('simple', :q) AS ts),
|
||||||
|
lt AS (
|
||||||
|
SELECT DISTINCT ON (recording_id) recording_id, search_vector, text
|
||||||
|
FROM transcripts WHERE superseded_at IS NULL
|
||||||
|
ORDER BY recording_id, version DESC
|
||||||
|
),
|
||||||
|
ls AS (
|
||||||
|
SELECT DISTINCT ON (recording_id) recording_id, content, search_vector
|
||||||
|
FROM summaries WHERE superseded_at IS NULL
|
||||||
|
ORDER BY recording_id, version DESC
|
||||||
|
),
|
||||||
|
tm AS (
|
||||||
|
SELECT DISTINCT rt.recording_id,
|
||||||
|
ts_headline('simple', t.name, q.ts,
|
||||||
|
'StartSel=,StopSel=,MaxFragments=0') AS snip
|
||||||
|
FROM recording_tags rt
|
||||||
|
JOIN tags t ON t.id = rt.tag_id AND t.user_id = :uid, q
|
||||||
|
),
|
||||||
|
matched AS (
|
||||||
|
SELECT r.id AS id,
|
||||||
|
CASE
|
||||||
|
WHEN :want_title AND r.search_vector @@ q.ts
|
||||||
|
AND coalesce(r.title, '') <> ''
|
||||||
|
THEN 'title'
|
||||||
|
WHEN :want_tag AND tm.recording_id IS NOT NULL THEN 'tag'
|
||||||
|
WHEN :want_notes AND coalesce(r.notes, '') ILIKE '%' || :raw || '%' THEN 'notes'
|
||||||
|
WHEN :want_summary AND ls.search_vector @@ q.ts THEN 'summary'
|
||||||
|
WHEN :want_transcript AND lt.search_vector @@ q.ts THEN 'transcript'
|
||||||
|
END AS field,
|
||||||
|
CASE
|
||||||
|
WHEN :want_title AND r.search_vector @@ q.ts
|
||||||
|
AND coalesce(r.title, '') <> ''
|
||||||
|
THEN ts_headline('simple', r.title, q.ts,
|
||||||
|
'StartSel=,StopSel=,MaxFragments=0,MaxWords=25')
|
||||||
|
WHEN :want_tag AND tm.recording_id IS NOT NULL THEN tm.snip
|
||||||
|
WHEN :want_notes AND coalesce(r.notes, '') ILIKE '%' || :raw || '%'
|
||||||
|
THEN left(r.notes, 200)
|
||||||
|
WHEN :want_summary AND ls.search_vector @@ q.ts
|
||||||
|
THEN ts_headline('simple',
|
||||||
|
coalesce(ls.content->>'short', '') || ' ' || coalesce(ls.content->>'detailed', ''),
|
||||||
|
q.ts, 'StartSel=,StopSel=,MaxFragments=1,MinWords=10,MaxWords=25')
|
||||||
|
WHEN :want_transcript AND lt.search_vector @@ q.ts
|
||||||
|
THEN ts_headline('simple', coalesce(lt.text, ''), q.ts,
|
||||||
|
'StartSel=,StopSel=,MaxFragments=1,MinWords=10,MaxWords=25')
|
||||||
|
END AS snippet
|
||||||
|
FROM recordings r
|
||||||
|
CROSS JOIN q
|
||||||
|
LEFT JOIN lt ON lt.recording_id = r.id
|
||||||
|
LEFT JOIN ls ON ls.recording_id = r.id
|
||||||
|
LEFT JOIN tm ON tm.recording_id = r.id
|
||||||
|
WHERE r.user_id = :uid AND r.deleted_at IS NULL
|
||||||
|
)
|
||||||
|
SELECT id, field, snippet FROM matched
|
||||||
|
WHERE field IS NOT NULL
|
||||||
|
ORDER BY CASE field WHEN 'title' THEN 0 WHEN 'tag' THEN 1 WHEN 'notes' THEN 2
|
||||||
|
WHEN 'summary' THEN 3 ELSE 4 END,
|
||||||
|
id
|
||||||
|
LIMIT :limit OFFSET :offset
|
||||||
|
"""
|
||||||
|
|
||||||
|
_PG_COUNT = """
|
||||||
|
WITH q AS (SELECT plainto_tsquery('simple', :q) AS ts),
|
||||||
|
lt AS (
|
||||||
|
SELECT DISTINCT ON (recording_id) recording_id, search_vector
|
||||||
|
FROM transcripts WHERE superseded_at IS NULL
|
||||||
|
),
|
||||||
|
ls AS (
|
||||||
|
SELECT DISTINCT ON (recording_id) recording_id, search_vector
|
||||||
|
FROM summaries WHERE superseded_at IS NULL
|
||||||
|
),
|
||||||
|
tm AS (
|
||||||
|
SELECT DISTINCT rt.recording_id
|
||||||
|
FROM recording_tags rt
|
||||||
|
JOIN tags t ON t.id = rt.tag_id AND t.user_id = :uid, q
|
||||||
|
WHERE t.search_vector @@ q.ts
|
||||||
|
)
|
||||||
|
SELECT count(*)
|
||||||
|
FROM recordings r
|
||||||
|
CROSS JOIN q
|
||||||
|
LEFT JOIN lt ON lt.recording_id = r.id
|
||||||
|
LEFT JOIN ls ON ls.recording_id = r.id
|
||||||
|
LEFT JOIN tm ON tm.recording_id = r.id
|
||||||
|
WHERE r.user_id = :uid AND r.deleted_at IS NULL
|
||||||
|
AND (
|
||||||
|
(:want_title AND r.search_vector @@ q.ts) OR
|
||||||
|
(:want_transcript AND lt.search_vector @@ q.ts) OR
|
||||||
|
(:want_summary AND ls.search_vector @@ q.ts) OR
|
||||||
|
(:want_tag AND tm.recording_id IS NOT NULL) OR
|
||||||
|
(:want_notes AND coalesce(r.notes, '') ILIKE '%' || :raw || '%')
|
||||||
|
)
|
||||||
|
"""
|
||||||
|
|
||||||
|
|
||||||
|
async def _search_postgres(
|
||||||
|
session: AsyncSession, user_id: uuid.UUID, q: str, scope: str,
|
||||||
|
limit: int, offset: int,
|
||||||
|
) -> tuple[list[SearchHit], int]:
|
||||||
|
# NOTE: the title field matches anything the recordings vector hits
|
||||||
|
# (title + notes); notes-only hits surface under 'title' headlines from
|
||||||
|
# the title text. Acceptable precision tradeoff for a GIN-indexed path.
|
||||||
|
params = {
|
||||||
|
"uid": user_id,
|
||||||
|
"q": q,
|
||||||
|
"raw": q,
|
||||||
|
"limit": limit,
|
||||||
|
"offset": offset,
|
||||||
|
"want_title": scope in ("all", "title"),
|
||||||
|
"want_transcript": scope in ("all", "transcript"),
|
||||||
|
"want_summary": scope in ("all", "summary"),
|
||||||
|
"want_tag": scope in ("all", "tag"),
|
||||||
|
"want_notes": scope in ("all", "notes"),
|
||||||
|
}
|
||||||
|
rows = (await session.execute(text(_PG_MATCHES), params)).all()
|
||||||
|
total = await session.scalar(text(_PG_COUNT), params) or 0
|
||||||
|
if not rows:
|
||||||
|
return [], total
|
||||||
|
ids = [r[0] for r in rows]
|
||||||
|
by_id = {
|
||||||
|
rec.id: rec
|
||||||
|
for rec in (
|
||||||
|
await session.scalars(select(Recording).where(Recording.id.in_(ids)))
|
||||||
|
).all()
|
||||||
|
}
|
||||||
|
hits = [
|
||||||
|
SearchHit(recording=by_id[r.id], field=r.field, snippet=r.snippet or "")
|
||||||
|
for r in rows
|
||||||
|
if r.id in by_id
|
||||||
|
]
|
||||||
|
return hits, total
|
||||||
|
|
||||||
|
|
||||||
|
# --- public API ----------------------------------------------------------------
|
||||||
|
|
||||||
|
|
||||||
|
async def search_recordings(
|
||||||
|
session: AsyncSession,
|
||||||
|
user_id: uuid.UUID,
|
||||||
|
q: str,
|
||||||
|
*,
|
||||||
|
scope: str = "all",
|
||||||
|
limit: int = 20,
|
||||||
|
offset: int = 0,
|
||||||
|
) -> tuple[list[SearchHit], int]:
|
||||||
|
"""Returns (page of hits ordered by field rank, total). Empty q → no hits."""
|
||||||
|
q = q.strip()
|
||||||
|
if not q:
|
||||||
|
return [], 0
|
||||||
|
dialect = session.bind.dialect.name if session.bind else "sqlite"
|
||||||
|
if dialect == "postgresql":
|
||||||
|
return await _search_postgres(session, user_id, q, scope, limit, offset)
|
||||||
|
return await _search_sqlite(session, user_id, q, scope, limit, offset)
|
||||||
342
backend/shonar/services/uploads.py
Normal file
342
backend/shonar/services/uploads.py
Normal file
|
|
@ -0,0 +1,342 @@
|
||||||
|
"""Chunked, resumable upload sessions.
|
||||||
|
|
||||||
|
Flow:
|
||||||
|
1. POST /uploads -> session (uuid), chunk size, expiry
|
||||||
|
2. PUT /uploads/{id}/chunks/{n} (idempotent per index; GET status lists
|
||||||
|
received indexes so clients resume)
|
||||||
|
3. POST /uploads/{id}/finalize -> validates size + magic bytes,
|
||||||
|
assembles the object, creates the
|
||||||
|
immutable original Asset and the
|
||||||
|
Recording (or updates the existing
|
||||||
|
recording for a retried
|
||||||
|
client_recording_id).
|
||||||
|
|
||||||
|
Storage keys are server-generated UUID paths; clients never see them.
|
||||||
|
Originals are immutable: finalize never overwrites an existing original.
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import contextlib
|
||||||
|
import hashlib
|
||||||
|
import uuid
|
||||||
|
from datetime import datetime, timedelta
|
||||||
|
|
||||||
|
from sqlalchemy import select
|
||||||
|
from sqlalchemy.ext.asyncio import AsyncSession
|
||||||
|
|
||||||
|
from shonar.core.config import get_settings
|
||||||
|
from shonar.db.models import (
|
||||||
|
Asset,
|
||||||
|
AssetKind,
|
||||||
|
ProcessingStatus,
|
||||||
|
Recording,
|
||||||
|
UploadChunk,
|
||||||
|
UploadSession,
|
||||||
|
UploadSessionStatus,
|
||||||
|
utcnow,
|
||||||
|
)
|
||||||
|
from shonar.services.media import is_compatible
|
||||||
|
from shonar.storage import get_storage
|
||||||
|
|
||||||
|
SESSION_TTL = timedelta(hours=24)
|
||||||
|
|
||||||
|
|
||||||
|
class UploadError(Exception):
|
||||||
|
def __init__(self, message: str, status_code: int = 400):
|
||||||
|
super().__init__(message)
|
||||||
|
self.message = message
|
||||||
|
self.status_code = status_code
|
||||||
|
|
||||||
|
|
||||||
|
async def create_session(
|
||||||
|
session: AsyncSession,
|
||||||
|
user_id: uuid.UUID,
|
||||||
|
declared_mime_type: str,
|
||||||
|
declared_size_bytes: int,
|
||||||
|
title: str | None,
|
||||||
|
client_recording_id: str | None,
|
||||||
|
transcription_model: str | None = None,
|
||||||
|
) -> UploadSession:
|
||||||
|
settings = get_settings()
|
||||||
|
if declared_size_bytes <= 0 or declared_size_bytes > settings.max_upload_bytes:
|
||||||
|
raise UploadError(
|
||||||
|
f"Declared size must be between 1 and {settings.max_upload_bytes} bytes.", 413
|
||||||
|
)
|
||||||
|
allowed = [m.lower() for m in settings.allowed_audio_mime_types]
|
||||||
|
if declared_mime_type.lower() not in allowed:
|
||||||
|
raise UploadError(f"MIME type not allowed. Allowed: {', '.join(allowed)}", 415)
|
||||||
|
|
||||||
|
us = UploadSession(
|
||||||
|
user_id=user_id,
|
||||||
|
client_recording_id=client_recording_id,
|
||||||
|
title=title,
|
||||||
|
declared_mime_type=declared_mime_type.lower(),
|
||||||
|
declared_size_bytes=declared_size_bytes,
|
||||||
|
chunk_size_bytes=settings.max_chunk_bytes,
|
||||||
|
expires_at=utcnow() + SESSION_TTL,
|
||||||
|
transcription_model=(transcription_model.strip() if transcription_model else None),
|
||||||
|
)
|
||||||
|
session.add(us)
|
||||||
|
await session.flush()
|
||||||
|
return us
|
||||||
|
|
||||||
|
|
||||||
|
async def get_owned_session(
|
||||||
|
session: AsyncSession, user_id: uuid.UUID, session_id: uuid.UUID
|
||||||
|
) -> UploadSession:
|
||||||
|
us = await session.get(UploadSession, session_id)
|
||||||
|
if us is None or us.user_id != user_id:
|
||||||
|
raise UploadError("Upload session not found.", 404)
|
||||||
|
if us.status == UploadSessionStatus.expired or (
|
||||||
|
us.status != UploadSessionStatus.completed and us.expires_at < utcnow()
|
||||||
|
):
|
||||||
|
us.status = UploadSessionStatus.expired
|
||||||
|
raise UploadError("Upload session expired. Start a new upload.", 410)
|
||||||
|
return us
|
||||||
|
|
||||||
|
|
||||||
|
async def put_chunk(
|
||||||
|
session: AsyncSession,
|
||||||
|
user_id: uuid.UUID,
|
||||||
|
session_id: uuid.UUID,
|
||||||
|
chunk_index: int,
|
||||||
|
data: bytes,
|
||||||
|
checksum_sha256: str | None,
|
||||||
|
) -> UploadChunk:
|
||||||
|
settings = get_settings()
|
||||||
|
us = await get_owned_session(session, user_id, session_id)
|
||||||
|
if us.status != UploadSessionStatus.open:
|
||||||
|
raise UploadError(f"Session is {us.status.value}; cannot accept chunks.", 409)
|
||||||
|
if chunk_index < 0:
|
||||||
|
raise UploadError("Chunk index must be >= 0.", 400)
|
||||||
|
if not data:
|
||||||
|
raise UploadError("Empty chunk.", 400)
|
||||||
|
if len(data) > settings.max_chunk_bytes:
|
||||||
|
raise UploadError(f"Chunk exceeds max size {settings.max_chunk_bytes}.", 413)
|
||||||
|
if checksum_sha256 and hashlib.sha256(data).hexdigest() != checksum_sha256.lower():
|
||||||
|
raise UploadError("Chunk checksum mismatch.", 422)
|
||||||
|
|
||||||
|
existing = await session.scalar(
|
||||||
|
select(UploadChunk).where(
|
||||||
|
UploadChunk.session_id == us.id, UploadChunk.chunk_index == chunk_index
|
||||||
|
)
|
||||||
|
)
|
||||||
|
if existing is not None:
|
||||||
|
# Idempotent retry: same index re-sent replaces the stored bytes.
|
||||||
|
if existing.size_bytes != len(data):
|
||||||
|
await get_storage().delete(existing.storage_key)
|
||||||
|
existing.size_bytes = len(data)
|
||||||
|
existing.checksum_sha256 = hashlib.sha256(data).hexdigest()
|
||||||
|
await get_storage().put(existing.storage_key, data)
|
||||||
|
await session.flush()
|
||||||
|
return existing
|
||||||
|
|
||||||
|
key = f"uploads/{us.id}/{chunk_index:06d}.part"
|
||||||
|
await get_storage().put(key, data)
|
||||||
|
chunk = UploadChunk(
|
||||||
|
session_id=us.id,
|
||||||
|
chunk_index=chunk_index,
|
||||||
|
size_bytes=len(data),
|
||||||
|
checksum_sha256=hashlib.sha256(data).hexdigest(),
|
||||||
|
storage_key=key,
|
||||||
|
)
|
||||||
|
session.add(chunk)
|
||||||
|
await session.flush()
|
||||||
|
return chunk
|
||||||
|
|
||||||
|
|
||||||
|
async def received_indexes(
|
||||||
|
session: AsyncSession, user_id: uuid.UUID, session_id: uuid.UUID
|
||||||
|
) -> list[int]:
|
||||||
|
us = await get_owned_session(session, user_id, session_id)
|
||||||
|
rows = await session.scalars(
|
||||||
|
select(UploadChunk.chunk_index).where(UploadChunk.session_id == us.id)
|
||||||
|
)
|
||||||
|
return sorted(rows)
|
||||||
|
|
||||||
|
|
||||||
|
async def finalize(
|
||||||
|
session: AsyncSession,
|
||||||
|
user_id: uuid.UUID,
|
||||||
|
session_id: uuid.UUID,
|
||||||
|
*,
|
||||||
|
recorded_at: datetime | None,
|
||||||
|
duration_seconds: float,
|
||||||
|
latitude: float | None = None,
|
||||||
|
longitude: float | None = None,
|
||||||
|
location_accuracy_m: float | None = None,
|
||||||
|
notes: str | None = None,
|
||||||
|
transcription_model: str | None = None,
|
||||||
|
) -> tuple[UploadSession, Recording, Asset]:
|
||||||
|
"""Assemble chunks, validate, store the immutable original, and create or
|
||||||
|
update the recording. Idempotent per client_recording_id."""
|
||||||
|
from shonar.services.ai import ProviderConfigError
|
||||||
|
from shonar.services.ai.model_registry import (
|
||||||
|
effective_model,
|
||||||
|
get_global_default_model,
|
||||||
|
validate_model_name,
|
||||||
|
)
|
||||||
|
|
||||||
|
us = await get_owned_session(session, user_id, session_id)
|
||||||
|
# Resolve + validate the transcription model before touching audio.
|
||||||
|
# Finalize body wins over the session's upload-screen choice; an
|
||||||
|
# explicit override updates history, the default never rewrites it.
|
||||||
|
override_raw = transcription_model or us.transcription_model
|
||||||
|
override: str | None = None
|
||||||
|
if override_raw:
|
||||||
|
try:
|
||||||
|
override = validate_model_name(override_raw)
|
||||||
|
except ProviderConfigError as e:
|
||||||
|
raise UploadError(str(e), 422) from None
|
||||||
|
model = effective_model(override, await get_global_default_model(session))
|
||||||
|
if us.status == UploadSessionStatus.completed and us.completed_asset_id:
|
||||||
|
# Already finalized: return existing recording (retry-safe client).
|
||||||
|
asset = await session.get(Asset, us.completed_asset_id)
|
||||||
|
rec = await session.scalar(
|
||||||
|
select(Recording).where(Recording.id == asset.recording_id)
|
||||||
|
)
|
||||||
|
if asset and rec:
|
||||||
|
return us, rec, asset
|
||||||
|
if us.status != UploadSessionStatus.open:
|
||||||
|
raise UploadError(f"Session is {us.status.value}.", 409)
|
||||||
|
|
||||||
|
chunks = list(
|
||||||
|
await session.scalars(
|
||||||
|
select(UploadChunk).where(UploadChunk.session_id == us.id).order_by(
|
||||||
|
UploadChunk.chunk_index
|
||||||
|
)
|
||||||
|
)
|
||||||
|
)
|
||||||
|
total = sum(c.size_bytes for c in chunks)
|
||||||
|
if total != us.declared_size_bytes:
|
||||||
|
raise UploadError(
|
||||||
|
f"Size mismatch: received {total} of declared {us.declared_size_bytes} bytes. "
|
||||||
|
"Upload missing chunks and retry.",
|
||||||
|
422,
|
||||||
|
)
|
||||||
|
expected_indexes = list(range(len(chunks)))
|
||||||
|
if [c.chunk_index for c in chunks] != expected_indexes:
|
||||||
|
raise UploadError("Chunk sequence has gaps. Upload missing chunks and retry.", 422)
|
||||||
|
|
||||||
|
storage = get_storage()
|
||||||
|
# Validate magic bytes from the first chunk.
|
||||||
|
first = await storage.get(chunks[0].storage_key)
|
||||||
|
if not is_compatible(us.declared_mime_type, first):
|
||||||
|
raise UploadError(
|
||||||
|
"File contents do not match the declared audio MIME type.", 415
|
||||||
|
)
|
||||||
|
|
||||||
|
# Idempotency: same client_recording_id => update existing recording.
|
||||||
|
recording: Recording | None = None
|
||||||
|
if us.client_recording_id:
|
||||||
|
recording = await session.scalar(
|
||||||
|
select(Recording).where(
|
||||||
|
Recording.user_id == user_id,
|
||||||
|
Recording.client_recording_id == us.client_recording_id,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
# NOTE(original-immutability): if the recording already has an original
|
||||||
|
# asset we do NOT replace it; a re-upload with the same client id after
|
||||||
|
# local edits updates metadata only, and the new audio is rejected as a
|
||||||
|
# duplicate (the existing original is returned instead).
|
||||||
|
if recording is not None:
|
||||||
|
existing_original = await session.scalar(
|
||||||
|
select(Asset).where(
|
||||||
|
Asset.recording_id == recording.id, Asset.kind == AssetKind.original
|
||||||
|
)
|
||||||
|
)
|
||||||
|
if existing_original is not None:
|
||||||
|
us.status = UploadSessionStatus.completed
|
||||||
|
us.completed_asset_id = existing_original.id
|
||||||
|
recording.title = us.title or recording.title
|
||||||
|
if override is not None:
|
||||||
|
# Explicit re-choice replaces the saved model; the default
|
||||||
|
# never rewrites history.
|
||||||
|
recording.transcription_model = override
|
||||||
|
await session.flush()
|
||||||
|
from shonar.services import processing as _processing
|
||||||
|
|
||||||
|
# Same audio, maybe new metadata — and a failed pipeline deserves
|
||||||
|
# another attempt. Idempotent: completed work is never redone.
|
||||||
|
await _processing.enqueue_for_recording(session, recording)
|
||||||
|
return us, recording, existing_original
|
||||||
|
|
||||||
|
# Assemble into the final object (streamed per chunk to bound memory).
|
||||||
|
from shonar.services.media import extension_for
|
||||||
|
|
||||||
|
digest = hashlib.sha256()
|
||||||
|
parts: list[bytes] = []
|
||||||
|
for c in chunks:
|
||||||
|
data = await storage.get(c.storage_key)
|
||||||
|
digest.update(data)
|
||||||
|
parts.append(data)
|
||||||
|
checksum = digest.hexdigest()
|
||||||
|
ext = extension_for(us.declared_mime_type, parts[0])
|
||||||
|
final_key = f"recordings/{user_id}/{uuid.uuid4()}{ext}"
|
||||||
|
blob = b"".join(parts)
|
||||||
|
await storage.put(final_key, blob)
|
||||||
|
|
||||||
|
asset = Asset(
|
||||||
|
recording_id=recording.id if recording else None,
|
||||||
|
user_id=user_id,
|
||||||
|
kind=AssetKind.original,
|
||||||
|
storage_key=final_key,
|
||||||
|
mime_type=us.declared_mime_type,
|
||||||
|
size_bytes=total,
|
||||||
|
checksum_sha256=checksum,
|
||||||
|
)
|
||||||
|
session.add(asset)
|
||||||
|
|
||||||
|
if recording is None:
|
||||||
|
recording = Recording(
|
||||||
|
user_id=user_id,
|
||||||
|
client_recording_id=us.client_recording_id,
|
||||||
|
title=us.title or "Untitled recording",
|
||||||
|
recorded_at=recorded_at or utcnow(),
|
||||||
|
duration_seconds=duration_seconds,
|
||||||
|
notes=notes,
|
||||||
|
latitude=latitude,
|
||||||
|
longitude=longitude,
|
||||||
|
location_accuracy_m=location_accuracy_m,
|
||||||
|
processing_status=ProcessingStatus.uploaded,
|
||||||
|
transcription_model=model,
|
||||||
|
)
|
||||||
|
session.add(recording)
|
||||||
|
await session.flush()
|
||||||
|
asset.recording_id = recording.id
|
||||||
|
else:
|
||||||
|
recording.title = us.title or recording.title
|
||||||
|
recording.duration_seconds = duration_seconds or recording.duration_seconds
|
||||||
|
recording.processing_status = ProcessingStatus.uploaded
|
||||||
|
recording.processing_error = None
|
||||||
|
|
||||||
|
us.status = UploadSessionStatus.completed
|
||||||
|
us.completed_asset_id = asset.id
|
||||||
|
|
||||||
|
# Clean up chunk parts (the assembled object is the source of truth).
|
||||||
|
for c in chunks: # best-effort cleanup of chunk parts
|
||||||
|
with contextlib.suppress(Exception):
|
||||||
|
await storage.delete(c.storage_key)
|
||||||
|
await session.flush()
|
||||||
|
from shonar.services import processing as _processing
|
||||||
|
|
||||||
|
# New audio on disk: queue whatever AI stages apply (none configured =
|
||||||
|
# ai_disabled, never an error).
|
||||||
|
await _processing.enqueue_for_recording(session, recording)
|
||||||
|
return us, recording, asset
|
||||||
|
|
||||||
|
|
||||||
|
async def abort(session: AsyncSession, user_id: uuid.UUID, session_id: uuid.UUID) -> None:
|
||||||
|
us = await get_owned_session(session, user_id, session_id)
|
||||||
|
if us.status in (UploadSessionStatus.completed, UploadSessionStatus.aborted):
|
||||||
|
return
|
||||||
|
storage = get_storage()
|
||||||
|
chunks = list(
|
||||||
|
await session.scalars(select(UploadChunk).where(UploadChunk.session_id == us.id))
|
||||||
|
)
|
||||||
|
for c in chunks: # best-effort cleanup
|
||||||
|
with contextlib.suppress(Exception):
|
||||||
|
await storage.delete(c.storage_key)
|
||||||
|
us.status = UploadSessionStatus.aborted
|
||||||
179
backend/shonar/storage/__init__.py
Normal file
179
backend/shonar/storage/__init__.py
Normal file
|
|
@ -0,0 +1,179 @@
|
||||||
|
"""Storage abstraction.
|
||||||
|
|
||||||
|
Backends:
|
||||||
|
- ``LocalStorage``: filesystem under ``SHONAR_STORAGE_PATH`` (dev + default).
|
||||||
|
- ``S3Storage``: any S3-compatible object store (extra: ``pip install
|
||||||
|
shonar-backend[s3]``).
|
||||||
|
|
||||||
|
Storage keys are server-internal and validated against path traversal:
|
||||||
|
they are UUID-based by construction and never derived from user input.
|
||||||
|
At-rest encryption is NOT implemented; see docs/security.md for the
|
||||||
|
documented optional design.
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import asyncio
|
||||||
|
import shutil
|
||||||
|
from pathlib import Path, PurePosixPath
|
||||||
|
from typing import Protocol
|
||||||
|
|
||||||
|
from shonar.core.config import get_settings
|
||||||
|
|
||||||
|
_KEY_CHARSET = set("abcdefghijklmnopqrstuvwxyz0123456789-./")
|
||||||
|
|
||||||
|
|
||||||
|
class StorageError(Exception):
|
||||||
|
pass
|
||||||
|
|
||||||
|
|
||||||
|
def validate_storage_key(key: str) -> str:
|
||||||
|
"""Reject anything that could escape the storage root."""
|
||||||
|
if not key or len(key) > 500:
|
||||||
|
raise StorageError("invalid storage key")
|
||||||
|
pure = PurePosixPath(key)
|
||||||
|
if pure.is_absolute() or ".." in pure.parts:
|
||||||
|
raise StorageError("invalid storage key")
|
||||||
|
if not set(key) <= _KEY_CHARSET:
|
||||||
|
raise StorageError("invalid storage key")
|
||||||
|
return key
|
||||||
|
|
||||||
|
|
||||||
|
class StorageBackend(Protocol):
|
||||||
|
async def put(self, key: str, data: bytes) -> int: ...
|
||||||
|
async def put_file(self, key: str, src_path: Path) -> int: ...
|
||||||
|
async def get(self, key: str) -> bytes: ...
|
||||||
|
async def open_path(self, key: str) -> Path | None:
|
||||||
|
"""Local file path if the backend can provide one, else None."""
|
||||||
|
...
|
||||||
|
|
||||||
|
async def delete(self, key: str) -> None: ...
|
||||||
|
async def exists(self, key: str) -> bool: ...
|
||||||
|
|
||||||
|
|
||||||
|
class LocalStorage:
|
||||||
|
def __init__(self, root: str | Path):
|
||||||
|
self.root = Path(root).resolve()
|
||||||
|
self.root.mkdir(parents=True, exist_ok=True)
|
||||||
|
|
||||||
|
def _path(self, key: str) -> Path:
|
||||||
|
validate_storage_key(key)
|
||||||
|
path = (self.root / key).resolve()
|
||||||
|
# Defense in depth: resolved path must stay under root.
|
||||||
|
if not path.is_relative_to(self.root):
|
||||||
|
raise StorageError("storage key escapes storage root")
|
||||||
|
return path
|
||||||
|
|
||||||
|
async def put(self, key: str, data: bytes) -> int:
|
||||||
|
path = self._path(key)
|
||||||
|
path.parent.mkdir(parents=True, exist_ok=True)
|
||||||
|
await asyncio.to_thread(path.write_bytes, data)
|
||||||
|
return len(data)
|
||||||
|
|
||||||
|
async def put_file(self, key: str, src_path: Path) -> int:
|
||||||
|
path = self._path(key)
|
||||||
|
path.parent.mkdir(parents=True, exist_ok=True)
|
||||||
|
await asyncio.to_thread(shutil.copyfile, src_path, path)
|
||||||
|
return path.stat().st_size
|
||||||
|
|
||||||
|
async def get(self, key: str) -> bytes:
|
||||||
|
path = self._path(key)
|
||||||
|
if not path.exists():
|
||||||
|
raise StorageError("object not found")
|
||||||
|
return await asyncio.to_thread(path.read_bytes)
|
||||||
|
|
||||||
|
async def open_path(self, key: str) -> Path | None:
|
||||||
|
path = self._path(key)
|
||||||
|
return path if path.exists() else None
|
||||||
|
|
||||||
|
async def delete(self, key: str) -> None:
|
||||||
|
path = self._path(key)
|
||||||
|
if path.exists():
|
||||||
|
await asyncio.to_thread(path.unlink)
|
||||||
|
|
||||||
|
async def exists(self, key: str) -> bool:
|
||||||
|
return self._path(key).exists()
|
||||||
|
|
||||||
|
|
||||||
|
class S3Storage: # pragma: no cover - requires boto3 + a real/mini endpoint
|
||||||
|
def __init__(
|
||||||
|
self, endpoint_url: str, bucket: str, region: str, access_key: str, secret_key: str
|
||||||
|
):
|
||||||
|
import boto3 # optional extra
|
||||||
|
|
||||||
|
self.bucket = bucket
|
||||||
|
self.s3 = boto3.client(
|
||||||
|
"s3",
|
||||||
|
endpoint_url=endpoint_url or None,
|
||||||
|
region_name=region,
|
||||||
|
aws_access_key_id=access_key or None,
|
||||||
|
aws_secret_access_key=secret_key or None,
|
||||||
|
)
|
||||||
|
|
||||||
|
async def put(self, key: str, data: bytes) -> int:
|
||||||
|
validate_storage_key(key)
|
||||||
|
await asyncio.to_thread(self.s3.put_object, Bucket=self.bucket, Key=key, Body=data)
|
||||||
|
return len(data)
|
||||||
|
|
||||||
|
async def put_file(self, key: str, src_path: Path) -> int:
|
||||||
|
validate_storage_key(key)
|
||||||
|
await asyncio.to_thread(self.s3.upload_file, str(src_path), self.bucket, key)
|
||||||
|
return src_path.stat().st_size
|
||||||
|
|
||||||
|
async def get(self, key: str) -> bytes:
|
||||||
|
validate_storage_key(key)
|
||||||
|
|
||||||
|
def _get() -> bytes:
|
||||||
|
obj = self.s3.get_object(Bucket=self.bucket, Key=key)
|
||||||
|
return obj["Body"].read()
|
||||||
|
|
||||||
|
return await asyncio.to_thread(_get)
|
||||||
|
|
||||||
|
async def open_path(self, key: str) -> Path | None:
|
||||||
|
return None # callers must use get()/streaming
|
||||||
|
|
||||||
|
async def delete(self, key: str) -> None:
|
||||||
|
validate_storage_key(key)
|
||||||
|
await asyncio.to_thread(self.s3.delete_object, Bucket=self.bucket, Key=key)
|
||||||
|
|
||||||
|
async def exists(self, key: str) -> bool:
|
||||||
|
validate_storage_key(key)
|
||||||
|
|
||||||
|
def _head() -> bool:
|
||||||
|
from botocore.exceptions import ClientError
|
||||||
|
|
||||||
|
try:
|
||||||
|
self.s3.head_object(Bucket=self.bucket, Key=key)
|
||||||
|
return True
|
||||||
|
except ClientError:
|
||||||
|
return False
|
||||||
|
|
||||||
|
return await asyncio.to_thread(_head)
|
||||||
|
|
||||||
|
|
||||||
|
_backend: StorageBackend | None = None
|
||||||
|
|
||||||
|
|
||||||
|
def get_storage() -> StorageBackend:
|
||||||
|
global _backend
|
||||||
|
if _backend is None:
|
||||||
|
settings = get_settings()
|
||||||
|
if settings.storage_backend == "local":
|
||||||
|
_backend = LocalStorage(settings.storage_path)
|
||||||
|
elif settings.storage_backend == "s3":
|
||||||
|
_backend = S3Storage(
|
||||||
|
settings.s3_endpoint_url,
|
||||||
|
settings.s3_bucket,
|
||||||
|
settings.s3_region,
|
||||||
|
settings.s3_access_key_id,
|
||||||
|
settings.s3_secret_access_key,
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
raise StorageError(f"unknown storage backend: {settings.storage_backend}")
|
||||||
|
return _backend
|
||||||
|
|
||||||
|
|
||||||
|
def set_storage(backend: StorageBackend | None) -> None:
|
||||||
|
"""Test seam."""
|
||||||
|
global _backend
|
||||||
|
_backend = backend
|
||||||
79
backend/shonar/worker.py
Normal file
79
backend/shonar/worker.py
Normal file
|
|
@ -0,0 +1,79 @@
|
||||||
|
"""arq worker entrypoint (M7): runs the AI pipeline tasks.
|
||||||
|
|
||||||
|
Run: arq shonar.worker.WorkerSettings
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import logging
|
||||||
|
|
||||||
|
from arq import cron
|
||||||
|
from arq.connections import RedisSettings
|
||||||
|
|
||||||
|
from shonar.core.config import get_settings
|
||||||
|
from shonar.db.session import dispose_engine, get_engine
|
||||||
|
from shonar.services import processing
|
||||||
|
|
||||||
|
logger = logging.getLogger("shonar.worker")
|
||||||
|
|
||||||
|
|
||||||
|
async def run_transcribe(ctx: dict, recording_id: str) -> None:
|
||||||
|
await processing.run_transcribe(ctx, recording_id)
|
||||||
|
|
||||||
|
|
||||||
|
async def run_summarize(ctx: dict, recording_id: str) -> None:
|
||||||
|
await processing.run_summarize(ctx, recording_id)
|
||||||
|
|
||||||
|
|
||||||
|
async def sweep(ctx: dict) -> None: # noqa: ARG001 — arq cron signature
|
||||||
|
count = await processing.sweep_stale()
|
||||||
|
if count:
|
||||||
|
logger.info("sweep re-enqueued %d stale jobs", count)
|
||||||
|
|
||||||
|
|
||||||
|
async def retention_sweep(ctx: dict) -> None: # noqa: ARG001 — arq cron signature
|
||||||
|
from shonar.services import retention
|
||||||
|
|
||||||
|
purged = await retention.sweep_deleted()
|
||||||
|
if purged["recordings"] or purged["users"]:
|
||||||
|
logger.info("retention sweep: %s", purged)
|
||||||
|
|
||||||
|
|
||||||
|
async def startup(ctx: dict) -> None:
|
||||||
|
get_engine()
|
||||||
|
# Crash recovery before accepting new work: jobs stuck `running` and
|
||||||
|
# queued rows the transport missed go back through arq.
|
||||||
|
count = await processing.sweep_stale()
|
||||||
|
if count:
|
||||||
|
logger.info("startup sweep re-enqueued %d stale jobs", count)
|
||||||
|
|
||||||
|
|
||||||
|
async def shutdown(ctx: dict) -> None: # noqa: ARG001
|
||||||
|
await dispose_engine()
|
||||||
|
|
||||||
|
|
||||||
|
def _redis() -> RedisSettings:
|
||||||
|
return RedisSettings.from_dsn(get_settings().redis_url)
|
||||||
|
|
||||||
|
|
||||||
|
class WorkerSettings:
|
||||||
|
functions = [run_transcribe, run_summarize, sweep, retention_sweep]
|
||||||
|
# Transport-loss backstop beyond the startup sweep: anything still
|
||||||
|
# queued (missed enqueue, dead worker between runs) goes back through
|
||||||
|
# arq every 5 minutes. Rows are the queue; this just pokes.
|
||||||
|
cron_jobs = [
|
||||||
|
cron(sweep, minute={0, 5, 10, 15, 20, 25, 30, 35, 40, 45, 50, 55}),
|
||||||
|
# Hard-delete expired soft-deletes once a day (3:17 local, off-peak).
|
||||||
|
cron(retention_sweep, hour=3, minute=17),
|
||||||
|
]
|
||||||
|
on_startup = startup
|
||||||
|
on_shutdown = shutdown
|
||||||
|
redis_settings = _redis()
|
||||||
|
# Retry budget for transient provider failures; the tasks themselves
|
||||||
|
# mark jobs failed on the last try (see processing.MAX_TRIES).
|
||||||
|
max_tries = 3
|
||||||
|
# One job per worker process at a time. Local faster-whisper is
|
||||||
|
# multi-GB per concurrent pass; default max_jobs=10 let one worker run
|
||||||
|
# ~8 transcriptions at once and OOM-swap the box. Scale by running more
|
||||||
|
# worker processes, never by raising this.
|
||||||
|
max_jobs = 1
|
||||||
0
backend/tests/__init__.py
Normal file
0
backend/tests/__init__.py
Normal file
113
backend/tests/conftest.py
Normal file
113
backend/tests/conftest.py
Normal file
|
|
@ -0,0 +1,113 @@
|
||||||
|
"""Pytest fixtures.
|
||||||
|
|
||||||
|
Tests run against a real PostgreSQL (deploy/docker-compose.dev.yml) using a
|
||||||
|
dedicated ``shonar_test`` database, plus a temporary local storage root.
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import os
|
||||||
|
import tempfile
|
||||||
|
from collections.abc import AsyncIterator
|
||||||
|
from pathlib import Path
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
import pytest_asyncio
|
||||||
|
|
||||||
|
# Configure env BEFORE importing the app so Settings picks it up.
|
||||||
|
TEST_DB = os.environ.get(
|
||||||
|
"SHONAR_TEST_DATABASE_URL",
|
||||||
|
"postgresql+asyncpg://shonar:shonar@localhost:5432/shonar_test",
|
||||||
|
)
|
||||||
|
os.environ["SHONAR_DATABASE_URL"] = TEST_DB
|
||||||
|
os.environ["SHONAR_SECRET_KEY"] = "test-secret-key-0123456789abcdef0123456789abcdef"
|
||||||
|
os.environ["SHONAR_STORAGE_BACKEND"] = "local"
|
||||||
|
# Effectively disable the auth rate limit under test (dedicated tests cover
|
||||||
|
# the limiter behaviour itself).
|
||||||
|
os.environ["SHONAR_RATE_LIMIT_AUTH"] = "10000/minute"
|
||||||
|
|
||||||
|
_tmp_storage = tempfile.mkdtemp(prefix="shonar-test-storage-")
|
||||||
|
os.environ["SHONAR_STORAGE_PATH"] = _tmp_storage
|
||||||
|
|
||||||
|
|
||||||
|
@pytest_asyncio.fixture(scope="session", loop_scope="session")
|
||||||
|
async def _setup_db() -> AsyncIterator[None]:
|
||||||
|
from shonar.db import models # noqa: F401
|
||||||
|
from shonar.db.base import Base
|
||||||
|
from shonar.db.session import dispose_engine, get_engine
|
||||||
|
|
||||||
|
engine = get_engine()
|
||||||
|
async with engine.begin() as conn:
|
||||||
|
await conn.run_sync(Base.metadata.drop_all)
|
||||||
|
await conn.run_sync(Base.metadata.create_all)
|
||||||
|
# The tsvector generated columns live only in migration
|
||||||
|
# fts0000000001 (not in the ORM models), so create_all misses them.
|
||||||
|
# Apply the same DDL the migration applies (M9 search needs them).
|
||||||
|
if engine.dialect.name == "postgresql":
|
||||||
|
# Reuse the real migration's DDL (not a copy) via a sync
|
||||||
|
# MigrationContext — op.execute() is synchronous there.
|
||||||
|
import importlib.util
|
||||||
|
from pathlib import Path
|
||||||
|
|
||||||
|
def _apply(sync_conn):
|
||||||
|
from alembic.migration import MigrationContext
|
||||||
|
from alembic.operations import Operations
|
||||||
|
|
||||||
|
spec = importlib.util.spec_from_file_location(
|
||||||
|
"fts_migration",
|
||||||
|
Path(__file__).resolve().parents[1]
|
||||||
|
/ "migrations/versions/fts0000000001_fts_columns.py",
|
||||||
|
)
|
||||||
|
assert spec is not None and spec.loader is not None
|
||||||
|
mod = importlib.util.module_from_spec(spec)
|
||||||
|
spec.loader.exec_module(mod)
|
||||||
|
ctx = MigrationContext.configure(sync_conn)
|
||||||
|
with Operations.context(ctx):
|
||||||
|
mod.upgrade()
|
||||||
|
|
||||||
|
await conn.run_sync(_apply)
|
||||||
|
yield
|
||||||
|
await dispose_engine()
|
||||||
|
|
||||||
|
|
||||||
|
@pytest_asyncio.fixture(loop_scope="session", autouse=True)
|
||||||
|
async def clean_db(_setup_db: None) -> AsyncIterator[None]:
|
||||||
|
"""Truncate between tests for isolation."""
|
||||||
|
yield
|
||||||
|
from sqlalchemy import delete, text
|
||||||
|
|
||||||
|
from shonar.db.base import Base
|
||||||
|
from shonar.db.session import _session_factory, get_engine # type: ignore[attr-defined]
|
||||||
|
|
||||||
|
assert _session_factory is not None
|
||||||
|
async with _session_factory() as s:
|
||||||
|
if get_engine().dialect.name == "postgresql":
|
||||||
|
await s.execute(
|
||||||
|
text(
|
||||||
|
"TRUNCATE users, devices, refresh_tokens, recordings, assets, "
|
||||||
|
"upload_sessions, upload_chunks, transcripts, summaries, tags, "
|
||||||
|
"recording_tags, processing_jobs, export_jobs, app_settings "
|
||||||
|
"RESTART IDENTITY CASCADE"
|
||||||
|
)
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
# SQLite: delete child-first (reverse dependency order).
|
||||||
|
for table in reversed(Base.metadata.sorted_tables):
|
||||||
|
await s.execute(delete(table))
|
||||||
|
await s.commit()
|
||||||
|
|
||||||
|
|
||||||
|
@pytest_asyncio.fixture(loop_scope="session")
|
||||||
|
async def client(_setup_db) -> AsyncIterator:
|
||||||
|
from httpx import ASGITransport, AsyncClient
|
||||||
|
|
||||||
|
from shonar.main import app
|
||||||
|
|
||||||
|
transport = ASGITransport(app=app)
|
||||||
|
async with AsyncClient(transport=transport, base_url="http://test") as c:
|
||||||
|
yield c
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.fixture()
|
||||||
|
def storage_root() -> Path:
|
||||||
|
return Path(_tmp_storage)
|
||||||
184
backend/tests/test_ai_adapters.py
Normal file
184
backend/tests/test_ai_adapters.py
Normal file
|
|
@ -0,0 +1,184 @@
|
||||||
|
"""Adapter wire-protocol tests (M7): JSON shapes in/out, error mapping.
|
||||||
|
|
||||||
|
HTTP adapters take an injectable httpx client; faster-whisper is an
|
||||||
|
optional dependency and is only exercised when installed.
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import importlib.util
|
||||||
|
|
||||||
|
import httpx
|
||||||
|
import pytest
|
||||||
|
|
||||||
|
from shonar.core.config import Settings
|
||||||
|
from shonar.services import ai
|
||||||
|
from shonar.services.ai import ProviderConfigError, SummaryResult
|
||||||
|
from shonar.services.ai.ollama import OllamaProvider
|
||||||
|
from shonar.services.ai.openai_compat import OpenAICompatProvider
|
||||||
|
from shonar.services.ai.whisper_http import WhisperHttpProvider
|
||||||
|
|
||||||
|
|
||||||
|
def mock_client(handler) -> httpx.AsyncClient:
|
||||||
|
return httpx.AsyncClient(transport=httpx.MockTransport(handler))
|
||||||
|
|
||||||
|
|
||||||
|
# --- factories -----------------------------------------------------------------
|
||||||
|
|
||||||
|
|
||||||
|
def test_factories_none_means_skip():
|
||||||
|
s = Settings(transcription_provider="none", llm_provider="none")
|
||||||
|
assert ai.get_transcription_provider(s) is None
|
||||||
|
assert ai.get_llm_provider(s) is None
|
||||||
|
|
||||||
|
|
||||||
|
def test_factories_unknown_names_raise_config_error():
|
||||||
|
s = Settings(transcription_provider="whisper-9k")
|
||||||
|
with pytest.raises(ProviderConfigError):
|
||||||
|
ai.get_transcription_provider(s)
|
||||||
|
s = Settings(llm_provider="clippy")
|
||||||
|
with pytest.raises(ProviderConfigError):
|
||||||
|
ai.get_llm_provider(s)
|
||||||
|
|
||||||
|
|
||||||
|
def test_factories_missing_urls_raise_config_error():
|
||||||
|
s = Settings(transcription_provider="whisper_http", transcription_base_url="")
|
||||||
|
with pytest.raises(ProviderConfigError):
|
||||||
|
ai.get_transcription_provider(s)
|
||||||
|
s = Settings(llm_provider="openai_compat", llm_base_url="http://x", llm_model="")
|
||||||
|
with pytest.raises(ProviderConfigError):
|
||||||
|
ai.get_llm_provider(s)
|
||||||
|
|
||||||
|
|
||||||
|
def test_faster_whisper_missing_dep_is_config_error():
|
||||||
|
if importlib.util.find_spec("faster_whisper") is not None:
|
||||||
|
pytest.skip("faster-whisper installed; missing-dep path N/A")
|
||||||
|
# Constructor validates eagerly so misconfiguration fails at startup,
|
||||||
|
# not on the first recording.
|
||||||
|
with pytest.raises(ProviderConfigError):
|
||||||
|
ai.get_transcription_provider(Settings(transcription_provider="faster_whisper"))
|
||||||
|
|
||||||
|
|
||||||
|
# --- whisper_http -----------------------------------------------------------------
|
||||||
|
|
||||||
|
|
||||||
|
def whisper_ok(request: httpx.Request) -> httpx.Response:
|
||||||
|
assert request.url.path == "/v1/audio/transcriptions"
|
||||||
|
assert request.method == "POST"
|
||||||
|
return httpx.Response(200, json={
|
||||||
|
"text": "hello world",
|
||||||
|
"language": "en",
|
||||||
|
"segments": [{"start": 0.0, "end": 1.2, "text": "hello world"}],
|
||||||
|
})
|
||||||
|
|
||||||
|
|
||||||
|
async def test_whisper_http_happy_path():
|
||||||
|
p = WhisperHttpProvider("http://stt:8000", model="small",
|
||||||
|
http_client=mock_client(whisper_ok))
|
||||||
|
res = await p.transcribe(b"\x00" * 16, "audio/wav")
|
||||||
|
assert res.text == "hello world"
|
||||||
|
assert res.language == "en"
|
||||||
|
assert [(s.start, s.end, s.text) for s in res.segments] == [(0.0, 1.2, "hello world")]
|
||||||
|
assert res.model == "small"
|
||||||
|
|
||||||
|
|
||||||
|
async def test_whisper_http_401_is_config_error():
|
||||||
|
async def denied(request: httpx.Request) -> httpx.Response:
|
||||||
|
return httpx.Response(401, json={"detail": "nope"})
|
||||||
|
p = WhisperHttpProvider("http://stt:8000", http_client=mock_client(denied))
|
||||||
|
with pytest.raises(ProviderConfigError):
|
||||||
|
await p.transcribe(b"\x00" * 16, "audio/wav")
|
||||||
|
|
||||||
|
|
||||||
|
async def test_whisper_http_503_is_transient():
|
||||||
|
async def busy(request: httpx.Request) -> httpx.Response:
|
||||||
|
return httpx.Response(503, text="overloaded")
|
||||||
|
from shonar.services.ai import ProviderTransientError
|
||||||
|
|
||||||
|
p = WhisperHttpProvider("http://stt:8000", http_client=mock_client(busy))
|
||||||
|
with pytest.raises(ProviderTransientError):
|
||||||
|
await p.transcribe(b"\x00" * 16, "audio/wav")
|
||||||
|
|
||||||
|
|
||||||
|
# --- openai_compat ------------------------------------------------------------------
|
||||||
|
|
||||||
|
|
||||||
|
def chat_ok(request: httpx.Request) -> httpx.Response:
|
||||||
|
assert request.url.path == "/v1/chat/completions"
|
||||||
|
body = {
|
||||||
|
"short": "Standup.",
|
||||||
|
"detailed": "The team met.",
|
||||||
|
"key_points": ["a", "b"],
|
||||||
|
"decisions": ["ship"],
|
||||||
|
"action_items": [{"not": "a string"}, "call ana"],
|
||||||
|
"questions": [],
|
||||||
|
"extra_key": "ignored",
|
||||||
|
}
|
||||||
|
import json as _json
|
||||||
|
|
||||||
|
return httpx.Response(200, json={"choices": [{"message": {"content": _json.dumps(body)}}]})
|
||||||
|
|
||||||
|
|
||||||
|
async def test_openai_compat_parses_and_sanitizes():
|
||||||
|
p = OpenAICompatProvider("http://llm:8000", model="qwen",
|
||||||
|
http_client=mock_client(chat_ok))
|
||||||
|
res = await p.summarize("a very long meeting transcript", title="Standup")
|
||||||
|
assert isinstance(res, SummaryResult)
|
||||||
|
assert res.short == "Standup."
|
||||||
|
assert res.action_items == ("call ana",) # non-strings dropped
|
||||||
|
assert res.model == "qwen"
|
||||||
|
|
||||||
|
|
||||||
|
async def test_openai_compat_partial_json_gets_defaults():
|
||||||
|
import json as _json
|
||||||
|
|
||||||
|
async def partial(request: httpx.Request) -> httpx.Response:
|
||||||
|
return httpx.Response(200, json={
|
||||||
|
"choices": [{"message": {"content": _json.dumps({"short": "Hi."})}}]
|
||||||
|
})
|
||||||
|
p = OpenAICompatProvider("http://llm:8000", model="qwen",
|
||||||
|
http_client=mock_client(partial))
|
||||||
|
res = await p.summarize("hi")
|
||||||
|
assert res.short == "Hi."
|
||||||
|
assert res.detailed == "" and res.key_points == ()
|
||||||
|
|
||||||
|
|
||||||
|
async def test_openai_compat_non_json_is_transient():
|
||||||
|
async def garbage(request: httpx.Request) -> httpx.Response:
|
||||||
|
return httpx.Response(200, json={
|
||||||
|
"choices": [{"message": {"content": "Sure! Here it is..."}}]
|
||||||
|
})
|
||||||
|
from shonar.services.ai import ProviderTransientError
|
||||||
|
|
||||||
|
p = OpenAICompatProvider("http://llm:8000", model="qwen",
|
||||||
|
http_client=mock_client(garbage))
|
||||||
|
with pytest.raises(ProviderTransientError):
|
||||||
|
await p.summarize("hi")
|
||||||
|
|
||||||
|
|
||||||
|
# --- ollama ---------------------------------------------------------------------------
|
||||||
|
|
||||||
|
|
||||||
|
async def test_ollama_happy_path():
|
||||||
|
import json as _json
|
||||||
|
|
||||||
|
async def ok(request: httpx.Request) -> httpx.Response:
|
||||||
|
assert request.url.path == "/api/chat"
|
||||||
|
payload = _json.loads(request.content.decode())
|
||||||
|
assert payload["format"] == "json" and payload["stream"] is False
|
||||||
|
return httpx.Response(200, json={
|
||||||
|
"message": {"content": _json.dumps({"short": "S.", "detailed": "D."})}
|
||||||
|
})
|
||||||
|
p = OllamaProvider("http://ollama:11434", model="llama3",
|
||||||
|
http_client=mock_client(ok))
|
||||||
|
res = await p.summarize("meeting notes")
|
||||||
|
assert res.short == "S." and res.model == "llama3"
|
||||||
|
|
||||||
|
|
||||||
|
async def test_ollama_404_is_config_error():
|
||||||
|
async def missing(request: httpx.Request) -> httpx.Response:
|
||||||
|
return httpx.Response(404, text="model not found")
|
||||||
|
p = OllamaProvider("http://ollama:11434", model="nope",
|
||||||
|
http_client=mock_client(missing))
|
||||||
|
with pytest.raises(ProviderConfigError):
|
||||||
|
await p.summarize("hi")
|
||||||
410
backend/tests/test_ai_pipeline.py
Normal file
410
backend/tests/test_ai_pipeline.py
Normal file
|
|
@ -0,0 +1,410 @@
|
||||||
|
"""AI pipeline tests (M7): status flow, versioning, failure modes, endpoints.
|
||||||
|
|
||||||
|
Provider fakes stand in for real STT/LLM services (no network, no model
|
||||||
|
downloads); adapter wire-protocol tests live in test_ai_adapters.py.
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import struct
|
||||||
|
import uuid
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
from sqlalchemy import func, select
|
||||||
|
|
||||||
|
from shonar.db import session as db_session
|
||||||
|
from shonar.db.models import (
|
||||||
|
JobStatus,
|
||||||
|
JobType,
|
||||||
|
ProcessingJob,
|
||||||
|
Transcript,
|
||||||
|
)
|
||||||
|
from shonar.services import processing
|
||||||
|
from shonar.services.ai import (
|
||||||
|
ProviderConfigError,
|
||||||
|
ProviderTransientError,
|
||||||
|
Segment,
|
||||||
|
SummaryResult,
|
||||||
|
TranscriptResult,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def wav_bytes(payload_len: int = 64) -> bytes:
|
||||||
|
data = bytes(range(payload_len % 256)) * (payload_len // 256 + 1)
|
||||||
|
data = data[:payload_len]
|
||||||
|
header = (
|
||||||
|
b"RIFF" + struct.pack("<I", 36 + len(data)) + b"WAVE"
|
||||||
|
+ b"fmt " + struct.pack("<IHHIIHH", 16, 1, 1, 8000, 8000, 1, 8)
|
||||||
|
+ b"data" + struct.pack("<I", len(data))
|
||||||
|
)
|
||||||
|
return header + data
|
||||||
|
|
||||||
|
|
||||||
|
AUTH = {"email": "m7@example.com", "password": "m7-test-passw0rd-123"}
|
||||||
|
|
||||||
|
|
||||||
|
async def user_tokens(client, email=AUTH["email"], password=AUTH["password"]):
|
||||||
|
r = await client.post("/api/v1/auth/register", json={"email": email, "password": password})
|
||||||
|
assert r.status_code == 201, r.text
|
||||||
|
return r.json()["access_token"]
|
||||||
|
|
||||||
|
|
||||||
|
async def upload_recording(client, token, client_id=None, title="M7 standup"):
|
||||||
|
h = {"Authorization": f"Bearer {token}"}
|
||||||
|
data = wav_bytes()
|
||||||
|
r = await client.post(
|
||||||
|
"/api/v1/uploads",
|
||||||
|
json={"declared_mime_type": "audio/wav", "declared_size_bytes": len(data),
|
||||||
|
"client_recording_id": client_id, "title": title},
|
||||||
|
headers=h,
|
||||||
|
)
|
||||||
|
assert r.status_code == 201, r.text
|
||||||
|
sid = r.json()["id"]
|
||||||
|
r = await client.put(f"/api/v1/uploads/{sid}/chunks/0", content=data,
|
||||||
|
headers={**h, "content-type": "application/octet-stream"})
|
||||||
|
assert r.status_code == 201, r.text
|
||||||
|
r = await client.post(f"/api/v1/uploads/{sid}/finalize",
|
||||||
|
json={"duration_seconds": 5.0}, headers=h)
|
||||||
|
assert r.status_code == 201, r.text
|
||||||
|
return r.json()
|
||||||
|
|
||||||
|
|
||||||
|
class FakeTranscriber:
|
||||||
|
name = "fake-stt"
|
||||||
|
|
||||||
|
def __init__(self, text="hello world from the meeting", fail=None):
|
||||||
|
self.text = text
|
||||||
|
self.fail = fail
|
||||||
|
self.calls = 0
|
||||||
|
|
||||||
|
async def transcribe(self, audio, mime, *, language_hint=None):
|
||||||
|
self.calls += 1
|
||||||
|
assert len(audio) > 0 and mime == "audio/wav"
|
||||||
|
if self.fail is not None:
|
||||||
|
raise self.fail
|
||||||
|
return TranscriptResult(
|
||||||
|
text=self.text, language="en",
|
||||||
|
segments=[Segment(0.0, 1.0, self.text)], model="fake-stt-1",
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class FakeLlm:
|
||||||
|
name = "fake-llm"
|
||||||
|
|
||||||
|
def __init__(self, fail=None):
|
||||||
|
self.fail = fail
|
||||||
|
self.seen = []
|
||||||
|
|
||||||
|
async def summarize(self, transcript, *, title=None):
|
||||||
|
self.seen.append(transcript)
|
||||||
|
if self.fail is not None:
|
||||||
|
raise self.fail
|
||||||
|
return SummaryResult(
|
||||||
|
short="Standup happened.", detailed="The team met and spoke.",
|
||||||
|
key_points=("a",), decisions=(), action_items=("ship it",),
|
||||||
|
questions=(), model="fake-llm-1",
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
_UNSET = object()
|
||||||
|
|
||||||
|
|
||||||
|
def use_fakes(monkeypatch, stt=_UNSET, llm=_UNSET):
|
||||||
|
tprov = FakeTranscriber() if stt is _UNSET else stt
|
||||||
|
lprov = FakeLlm() if llm is _UNSET else llm
|
||||||
|
monkeypatch.setattr(processing, "get_transcription_provider", lambda settings: tprov)
|
||||||
|
monkeypatch.setattr(processing, "get_llm_provider", lambda settings: lprov)
|
||||||
|
|
||||||
|
|
||||||
|
async def jobs_for(recording_id):
|
||||||
|
async with db_session._session_factory() as s:
|
||||||
|
rows = (await s.scalars(
|
||||||
|
select(ProcessingJob).where(ProcessingJob.recording_id == uuid.UUID(recording_id))
|
||||||
|
)).all()
|
||||||
|
return {j.job_type: j for j in rows}
|
||||||
|
|
||||||
|
|
||||||
|
# --- no AI configured --------------------------------------------------------
|
||||||
|
|
||||||
|
|
||||||
|
async def test_none_configured_marks_ai_disabled(client):
|
||||||
|
token = await user_tokens(client)
|
||||||
|
rec = await upload_recording(client, token)
|
||||||
|
assert rec["processing_status"] == "ai_disabled"
|
||||||
|
assert await jobs_for(rec["id"]) == {}
|
||||||
|
# Status endpoints: nothing there yet, but the recording exists.
|
||||||
|
h = {"Authorization": f"Bearer {token}"}
|
||||||
|
r = await client.get(f"/api/v1/recordings/{rec['id']}/transcript", headers=h)
|
||||||
|
assert r.status_code == 404
|
||||||
|
r = await client.get(f"/api/v1/recordings/{rec['id']}/summary", headers=h)
|
||||||
|
assert r.status_code == 404
|
||||||
|
r = await client.get(f"/api/v1/recordings/{rec['id']}/jobs", headers=h)
|
||||||
|
assert r.status_code == 200 and r.json() == []
|
||||||
|
|
||||||
|
|
||||||
|
# --- happy path ----------------------------------------------------------------
|
||||||
|
|
||||||
|
|
||||||
|
async def test_full_pipeline_transcribe_then_summarize(client, monkeypatch):
|
||||||
|
stt, llm = FakeTranscriber(), FakeLlm()
|
||||||
|
use_fakes(monkeypatch, stt, llm)
|
||||||
|
token = await user_tokens(client)
|
||||||
|
rec = await upload_recording(client, token, client_id="m7-full-1")
|
||||||
|
assert rec["processing_status"] == "queued"
|
||||||
|
jobs = await jobs_for(rec["id"])
|
||||||
|
assert set(jobs) == {JobType.transcribe}
|
||||||
|
assert jobs[JobType.transcribe].status == JobStatus.queued
|
||||||
|
|
||||||
|
await processing.run_transcribe({}, rec["id"])
|
||||||
|
jobs = await jobs_for(rec["id"])
|
||||||
|
assert jobs[JobType.transcribe].status == JobStatus.succeeded
|
||||||
|
assert set(jobs) == {JobType.transcribe, JobType.summarize}
|
||||||
|
h = {"Authorization": f"Bearer {token}"}
|
||||||
|
r = await client.get(f"/api/v1/recordings/{rec['id']}", headers=h)
|
||||||
|
assert r.json()["processing_status"] == "processing"
|
||||||
|
|
||||||
|
t = await client.get(f"/api/v1/recordings/{rec['id']}/transcript", headers=h)
|
||||||
|
assert t.status_code == 200, t.text
|
||||||
|
body = t.json()
|
||||||
|
assert body["text"] == "hello world from the meeting"
|
||||||
|
assert body["version"] == 1 and body["provider"] == "fake-stt"
|
||||||
|
assert body["segments"][0]["text"].startswith("hello")
|
||||||
|
assert llm.seen == []
|
||||||
|
|
||||||
|
await processing.run_summarize({}, rec["id"])
|
||||||
|
jobs = await jobs_for(rec["id"])
|
||||||
|
assert jobs[JobType.summarize].status == JobStatus.succeeded
|
||||||
|
s = await client.get(f"/api/v1/recordings/{rec['id']}/summary", headers=h)
|
||||||
|
assert s.status_code == 200, s.text
|
||||||
|
content = s.json()["content"]
|
||||||
|
assert content["short"] == "Standup happened."
|
||||||
|
assert content["action_items"] == ["ship it"]
|
||||||
|
assert set(content) == {"short", "detailed", "key_points", "decisions",
|
||||||
|
"action_items", "questions"}
|
||||||
|
r = await client.get(f"/api/v1/recordings/{rec['id']}", headers=h)
|
||||||
|
assert r.json()["processing_status"] == "completed"
|
||||||
|
assert llm.seen == ["hello world from the meeting"]
|
||||||
|
|
||||||
|
jobs_rows = await client.get(f"/api/v1/recordings/{rec['id']}/jobs", headers=h)
|
||||||
|
assert jobs_rows.status_code == 200
|
||||||
|
assert {j["job_type"] for j in jobs_rows.json()} == {"transcribe", "summarize"}
|
||||||
|
|
||||||
|
|
||||||
|
async def test_enqueue_is_idempotent(client, monkeypatch):
|
||||||
|
use_fakes(monkeypatch, FakeTranscriber(), FakeLlm())
|
||||||
|
token = await user_tokens(client)
|
||||||
|
rec = await upload_recording(client, token, client_id="m7-idem-1")
|
||||||
|
await processing.run_transcribe({}, rec["id"])
|
||||||
|
await processing.run_summarize({}, rec["id"])
|
||||||
|
async with db_session._session_factory() as s:
|
||||||
|
from shonar.db.models import Recording
|
||||||
|
|
||||||
|
row = await s.get(Recording, uuid.UUID(rec["id"]))
|
||||||
|
queued = await processing.enqueue_for_recording(s, row)
|
||||||
|
assert queued == []
|
||||||
|
n = await s.scalar(
|
||||||
|
select(func.count(ProcessingJob.id)).where(
|
||||||
|
ProcessingJob.recording_id == uuid.UUID(rec["id"]))
|
||||||
|
)
|
||||||
|
assert n == 2
|
||||||
|
await s.commit()
|
||||||
|
|
||||||
|
|
||||||
|
# --- failure modes -------------------------------------------------------------
|
||||||
|
|
||||||
|
|
||||||
|
async def test_transient_failure_retries_then_fails(client, monkeypatch):
|
||||||
|
stt = FakeTranscriber(fail=ProviderTransientError("stt down"))
|
||||||
|
use_fakes(monkeypatch, stt, FakeLlm())
|
||||||
|
token = await user_tokens(client)
|
||||||
|
rec = await upload_recording(client, token, client_id="m7-fail-1")
|
||||||
|
# First tries raise for arq retry; the attempt rolls back with the
|
||||||
|
# transaction, so the row still reads queued (arq tracks the tries).
|
||||||
|
with pytest.raises(ProviderTransientError):
|
||||||
|
await processing.run_transcribe({"job_try": 1}, rec["id"])
|
||||||
|
jobs = await jobs_for(rec["id"])
|
||||||
|
assert jobs[JobType.transcribe].status == JobStatus.queued
|
||||||
|
# Last try marks the job (and recording) failed with a safe message.
|
||||||
|
await processing.run_transcribe({"job_try": 3}, rec["id"])
|
||||||
|
jobs = await jobs_for(rec["id"])
|
||||||
|
assert jobs[JobType.transcribe].status == JobStatus.failed
|
||||||
|
assert jobs[JobType.transcribe].error == "stt down"
|
||||||
|
h = {"Authorization": f"Bearer {token}"}
|
||||||
|
r = await client.get(f"/api/v1/recordings/{rec['id']}", headers=h)
|
||||||
|
assert r.json()["processing_status"] == "failed"
|
||||||
|
assert r.json()["processing_error"] == "stt down"
|
||||||
|
|
||||||
|
|
||||||
|
async def test_config_error_fails_fast(client, monkeypatch):
|
||||||
|
stt = FakeTranscriber(fail=ProviderConfigError("bad credentials"))
|
||||||
|
use_fakes(monkeypatch, stt, FakeLlm())
|
||||||
|
token = await user_tokens(client)
|
||||||
|
rec = await upload_recording(client, token, client_id="m7-cfg-1")
|
||||||
|
await processing.run_transcribe({"job_try": 1}, rec["id"]) # no raise
|
||||||
|
jobs = await jobs_for(rec["id"])
|
||||||
|
assert jobs[JobType.transcribe].status == JobStatus.failed
|
||||||
|
assert stt.calls == 1
|
||||||
|
|
||||||
|
|
||||||
|
# --- versioning -----------------------------------------------------------------
|
||||||
|
|
||||||
|
|
||||||
|
async def test_rerun_supersedes_auto_but_not_user_edits(client, monkeypatch):
|
||||||
|
use_fakes(monkeypatch, FakeTranscriber("v1 text"), FakeLlm())
|
||||||
|
token = await user_tokens(client)
|
||||||
|
rec = await upload_recording(client, token, client_id="m7-ver-1")
|
||||||
|
await processing.run_transcribe({}, rec["id"])
|
||||||
|
rid = uuid.UUID(rec["id"])
|
||||||
|
|
||||||
|
async def versions():
|
||||||
|
async with db_session._session_factory() as s:
|
||||||
|
rows = (await s.scalars(
|
||||||
|
select(Transcript).where(Transcript.recording_id == rid)
|
||||||
|
.order_by(Transcript.version))).all()
|
||||||
|
return [(r.version, r.text, r.superseded_at is not None, r.edited_by_user)
|
||||||
|
for r in rows]
|
||||||
|
|
||||||
|
assert await versions() == [(1, "v1 text", False, False)]
|
||||||
|
|
||||||
|
# Queue another transcription run manually (re-entry path).
|
||||||
|
async with db_session._session_factory() as s:
|
||||||
|
from shonar.db.models import Recording
|
||||||
|
|
||||||
|
row = await s.get(Recording, rid)
|
||||||
|
await processing.enqueue_for_recording(s, row)
|
||||||
|
await s.commit()
|
||||||
|
use_fakes(monkeypatch, FakeTranscriber("v2 text"), FakeLlm())
|
||||||
|
await processing.run_transcribe({}, rec["id"])
|
||||||
|
assert await versions() == [(1, "v1 text", True, False), (2, "v2 text", False, False)]
|
||||||
|
|
||||||
|
# A user edit wins: the next auto run inserts nothing.
|
||||||
|
async with db_session._session_factory() as s:
|
||||||
|
v2 = await s.scalar(
|
||||||
|
select(Transcript).where(Transcript.recording_id == rid,
|
||||||
|
Transcript.version == 2))
|
||||||
|
v2.edited_by_user = True
|
||||||
|
v2.text = "user corrected text"
|
||||||
|
from shonar.db.models import Recording
|
||||||
|
|
||||||
|
row = await s.get(Recording, rid)
|
||||||
|
await processing.enqueue_for_recording(s, row)
|
||||||
|
await s.commit()
|
||||||
|
use_fakes(monkeypatch, FakeTranscriber("v3 text"), FakeLlm())
|
||||||
|
await processing.run_transcribe({}, rec["id"])
|
||||||
|
got = await versions()
|
||||||
|
assert len(got) == 2 and got[1][1] == "user corrected text"
|
||||||
|
|
||||||
|
|
||||||
|
# --- llm-only and skips ----------------------------------------------------------
|
||||||
|
|
||||||
|
|
||||||
|
async def test_llm_only_without_transcript_queues_nothing(client, monkeypatch):
|
||||||
|
use_fakes(monkeypatch, None, FakeLlm())
|
||||||
|
token = await user_tokens(client)
|
||||||
|
rec = await upload_recording(client, token, client_id="m7-llm-1")
|
||||||
|
assert rec["processing_status"] == "uploaded"
|
||||||
|
assert await jobs_for(rec["id"]) == {}
|
||||||
|
|
||||||
|
|
||||||
|
async def test_llm_only_with_manual_transcript_summarizes(client, monkeypatch):
|
||||||
|
use_fakes(monkeypatch, None, FakeLlm())
|
||||||
|
token = await user_tokens(client)
|
||||||
|
rec = await upload_recording(client, token, client_id="m7-llm-2")
|
||||||
|
rid = uuid.UUID(rec["id"])
|
||||||
|
async with db_session._session_factory() as s:
|
||||||
|
from shonar.db.models import Recording
|
||||||
|
|
||||||
|
row = await s.get(Recording, rid)
|
||||||
|
s.add(Transcript(recording_id=rid, version=1, provider="manual",
|
||||||
|
text="handwritten notes", edited_by_user=True))
|
||||||
|
await s.flush()
|
||||||
|
queued = await processing.enqueue_for_recording(s, row)
|
||||||
|
assert queued == [JobType.summarize]
|
||||||
|
await s.commit()
|
||||||
|
await processing.run_summarize({}, rec["id"])
|
||||||
|
h = {"Authorization": f"Bearer {token}"}
|
||||||
|
s = await client.get(f"/api/v1/recordings/{rec['id']}/summary", headers=h)
|
||||||
|
assert s.status_code == 200
|
||||||
|
r = await client.get(f"/api/v1/recordings/{rec['id']}", headers=h)
|
||||||
|
assert r.json()["processing_status"] == "completed"
|
||||||
|
|
||||||
|
|
||||||
|
# --- sweep ------------------------------------------------------------------------
|
||||||
|
|
||||||
|
|
||||||
|
async def test_sweep_requeues_stale_jobs(client, monkeypatch):
|
||||||
|
use_fakes(monkeypatch, FakeTranscriber(), FakeLlm())
|
||||||
|
token = await user_tokens(client)
|
||||||
|
rec = await upload_recording(client, token, client_id="m7-sweep-1")
|
||||||
|
rid = uuid.UUID(rec["id"])
|
||||||
|
async with db_session._session_factory() as s:
|
||||||
|
from shonar.db.models import Recording
|
||||||
|
|
||||||
|
row = await s.get(Recording, rid)
|
||||||
|
# Simulate a crashed worker + a lost transport, respectively.
|
||||||
|
await processing.reset_or_create(s, row, JobType.transcribe)
|
||||||
|
jobs = await jobs_for(rec["id"])
|
||||||
|
jobs[JobType.transcribe].status = JobStatus.running
|
||||||
|
await s.commit()
|
||||||
|
# reset_or_create leaves running rows alone, so force the second shape:
|
||||||
|
async with db_session._session_factory() as s:
|
||||||
|
extra = ProcessingJob(recording_id=rid, job_type=JobType.summarize,
|
||||||
|
status=JobStatus.queued)
|
||||||
|
s.add(extra)
|
||||||
|
await s.commit()
|
||||||
|
count = await processing.sweep_stale()
|
||||||
|
assert count == 2
|
||||||
|
jobs = await jobs_for(rec["id"])
|
||||||
|
assert jobs[JobType.transcribe].status == JobStatus.queued
|
||||||
|
assert jobs[JobType.summarize].status == JobStatus.queued
|
||||||
|
|
||||||
|
|
||||||
|
async def test_sweep_leaves_live_running_jobs_alone(client, monkeypatch):
|
||||||
|
"""Regression: sweep re-enqueued every in-flight transcription every
|
||||||
|
5 minutes; with local whisper each copy ran concurrently and one worker
|
||||||
|
OOM-swap-swelled the box. A running job started recently is live work."""
|
||||||
|
from datetime import UTC, datetime, timedelta
|
||||||
|
|
||||||
|
use_fakes(monkeypatch, FakeTranscriber(), FakeLlm())
|
||||||
|
token = await user_tokens(client)
|
||||||
|
rec = await upload_recording(client, token, client_id="m7-sweep-live")
|
||||||
|
rid = uuid.UUID(rec["id"])
|
||||||
|
async with db_session._session_factory() as s:
|
||||||
|
from shonar.db.models import Recording
|
||||||
|
|
||||||
|
row = await s.get(Recording, rid)
|
||||||
|
await processing.reset_or_create(s, row, JobType.transcribe)
|
||||||
|
live = (await s.scalars(
|
||||||
|
select(ProcessingJob).where(
|
||||||
|
ProcessingJob.recording_id == rid,
|
||||||
|
ProcessingJob.job_type == JobType.transcribe)
|
||||||
|
)).one()
|
||||||
|
live.status = JobStatus.running
|
||||||
|
live.started_at = datetime.now(UTC) # worker is chewing it right now
|
||||||
|
stale = ProcessingJob(recording_id=rid, job_type=JobType.summarize,
|
||||||
|
status=JobStatus.running, attempt=1)
|
||||||
|
stale.started_at = datetime.now(UTC) - timedelta(
|
||||||
|
hours=processing.STALE_RUNNING_AFTER.total_seconds() / 3600 + 1)
|
||||||
|
s.add(stale)
|
||||||
|
await s.commit()
|
||||||
|
count = await processing.sweep_stale()
|
||||||
|
assert count == 1 # only the orphaned old run
|
||||||
|
jobs = await jobs_for(rec["id"])
|
||||||
|
assert jobs[JobType.transcribe].status == JobStatus.running
|
||||||
|
assert jobs[JobType.summarize].status == JobStatus.queued
|
||||||
|
|
||||||
|
|
||||||
|
# --- ownership ---------------------------------------------------------------------
|
||||||
|
|
||||||
|
|
||||||
|
async def test_ai_endpoints_enforce_ownership(client, monkeypatch):
|
||||||
|
use_fakes(monkeypatch, FakeTranscriber(), FakeLlm())
|
||||||
|
token = await user_tokens(client)
|
||||||
|
rec = await upload_recording(client, token, client_id="m7-own-1")
|
||||||
|
await processing.run_transcribe({}, rec["id"])
|
||||||
|
other = await user_tokens(client, email="m7-other@example.com",
|
||||||
|
password="m7-test-passw0rd-456")
|
||||||
|
h = {"Authorization": f"Bearer {other}"}
|
||||||
|
for path in ("transcript", "summary", "jobs"):
|
||||||
|
r = await client.get(f"/api/v1/recordings/{rec['id']}/{path}", headers=h)
|
||||||
|
assert r.status_code == 404, path
|
||||||
110
backend/tests/test_auth.py
Normal file
110
backend/tests/test_auth.py
Normal file
|
|
@ -0,0 +1,110 @@
|
||||||
|
"""Auth flow tests: register, login, refresh rotation + reuse detection,
|
||||||
|
logout, protected access, account guards."""
|
||||||
|
|
||||||
|
AUTH = {"email": "test@example.com", "password": "correct-horse-battery"}
|
||||||
|
|
||||||
|
|
||||||
|
async def register(client, email=AUTH["email"], password=AUTH["password"]):
|
||||||
|
return await client.post(
|
||||||
|
"/api/v1/auth/register",
|
||||||
|
json={"email": email, "password": password, "display_name": "Tester"},
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
async def test_register_returns_token_pair(client):
|
||||||
|
r = await register(client)
|
||||||
|
assert r.status_code == 201, r.text
|
||||||
|
body = r.json()
|
||||||
|
assert body["token_type"] == "bearer"
|
||||||
|
assert body["expires_in"] == 15 * 60
|
||||||
|
assert body["access_token"] and body["refresh_token"]
|
||||||
|
|
||||||
|
|
||||||
|
async def test_register_rejects_duplicate_email(client):
|
||||||
|
assert (await register(client)).status_code == 201
|
||||||
|
r = await register(client)
|
||||||
|
assert r.status_code == 409
|
||||||
|
|
||||||
|
|
||||||
|
async def test_register_rejects_weak_password(client):
|
||||||
|
r = await client.post(
|
||||||
|
"/api/v1/auth/register", json={"email": "x@example.com", "password": "short"}
|
||||||
|
)
|
||||||
|
assert r.status_code == 422
|
||||||
|
|
||||||
|
|
||||||
|
async def test_login_success_and_failure(client):
|
||||||
|
await register(client)
|
||||||
|
r = await client.post(
|
||||||
|
"/api/v1/auth/login", json={"email": AUTH["email"], "password": AUTH["password"]}
|
||||||
|
)
|
||||||
|
assert r.status_code == 200
|
||||||
|
assert r.json()["device_id"]
|
||||||
|
|
||||||
|
bad = await client.post("/api/v1/auth/login", json={"email": AUTH["email"], "password": "***"})
|
||||||
|
assert bad.status_code == 401
|
||||||
|
# Same generic message either way (no user enumeration via password check)
|
||||||
|
bad2 = await client.post(
|
||||||
|
"/api/v1/auth/login", json={"email": "nobody@example.com", "password": "***"}
|
||||||
|
)
|
||||||
|
assert bad2.status_code == 401
|
||||||
|
assert bad.json()["detail"] == bad2.json()["detail"]
|
||||||
|
|
||||||
|
|
||||||
|
async def test_me_requires_valid_token(client):
|
||||||
|
r = await client.get("/api/v1/auth/me")
|
||||||
|
assert r.status_code == 401
|
||||||
|
tok = (await register(client)).json()["access_token"]
|
||||||
|
r = await client.get("/api/v1/auth/me", headers={"Authorization": f"Bearer {tok}"})
|
||||||
|
assert r.status_code == 200
|
||||||
|
assert r.json()["email"] == AUTH["email"]
|
||||||
|
|
||||||
|
|
||||||
|
async def test_refresh_rotates_and_detects_reuse(client):
|
||||||
|
tok = (await register(client)).json()
|
||||||
|
old_refresh = tok["refresh_token"]
|
||||||
|
|
||||||
|
r = await client.post("/api/v1/auth/refresh", json={"refresh_token": old_refresh})
|
||||||
|
assert r.status_code == 200
|
||||||
|
new_refresh = r.json()["refresh_token"]
|
||||||
|
assert new_refresh != old_refresh
|
||||||
|
|
||||||
|
# Old token is dead; reuse kills the whole family.
|
||||||
|
r2 = await client.post("/api/v1/auth/refresh", json={"refresh_token": old_refresh})
|
||||||
|
assert r2.status_code == 401
|
||||||
|
|
||||||
|
# The replacement is also revoked (family revocation).
|
||||||
|
r3 = await client.post("/api/v1/auth/refresh", json={"refresh_token": new_refresh})
|
||||||
|
assert r3.status_code == 401
|
||||||
|
|
||||||
|
|
||||||
|
async def test_logout_revokes_refresh(client):
|
||||||
|
tok = (await register(client)).json()
|
||||||
|
r = await client.post("/api/v1/auth/logout", json={"refresh_token": tok["refresh_token"]})
|
||||||
|
assert r.status_code == 204
|
||||||
|
r2 = await client.post("/api/v1/auth/refresh", json={"refresh_token": tok["refresh_token"]})
|
||||||
|
assert r2.status_code == 401
|
||||||
|
|
||||||
|
|
||||||
|
async def test_access_token_with_refresh_type_rejected(client):
|
||||||
|
tok = (await register(client)).json()
|
||||||
|
r = await client.get(
|
||||||
|
"/api/v1/auth/me", headers={"Authorization": f"Bearer {tok['refresh_token']}"}
|
||||||
|
)
|
||||||
|
assert r.status_code == 401
|
||||||
|
|
||||||
|
|
||||||
|
async def test_delete_account_requires_password(client):
|
||||||
|
tok = (await register(client)).json()
|
||||||
|
h = {"Authorization": f"Bearer {tok['access_token']}"}
|
||||||
|
r = await client.post("/api/v1/auth/delete-account", json={"password": "***"}, headers=h)
|
||||||
|
assert r.status_code == 403
|
||||||
|
r = await client.post(
|
||||||
|
"/api/v1/auth/delete-account", json={"password": AUTH["password"]}, headers=h
|
||||||
|
)
|
||||||
|
assert r.status_code == 202
|
||||||
|
# Deleted account can no longer log in.
|
||||||
|
r = await client.post(
|
||||||
|
"/api/v1/auth/login", json={"email": AUTH["email"], "password": AUTH["password"]}
|
||||||
|
)
|
||||||
|
assert r.status_code == 401
|
||||||
29
backend/tests/test_health.py
Normal file
29
backend/tests/test_health.py
Normal file
|
|
@ -0,0 +1,29 @@
|
||||||
|
"""Health-check and system-status tests."""
|
||||||
|
|
||||||
|
|
||||||
|
async def test_healthz(client):
|
||||||
|
r = await client.get("/api/v1/healthz")
|
||||||
|
assert r.status_code == 200
|
||||||
|
body = r.json()
|
||||||
|
assert body["status"] == "ok"
|
||||||
|
assert body["uptime_seconds"] >= 0
|
||||||
|
|
||||||
|
|
||||||
|
async def test_readyz_with_db(client):
|
||||||
|
r = await client.get("/api/v1/readyz")
|
||||||
|
assert r.status_code == 200
|
||||||
|
assert r.json() == {"status": "ok", "database": True}
|
||||||
|
|
||||||
|
|
||||||
|
async def test_system_status_reports_ai_state_honestly(client):
|
||||||
|
r = await client.get("/api/v1/system/status")
|
||||||
|
assert r.status_code == 200
|
||||||
|
body = r.json()
|
||||||
|
# Default test config: AI disabled, no external calls.
|
||||||
|
assert body["ai"]["transcription_enabled"] is False
|
||||||
|
assert body["ai"]["llm_enabled"] is False
|
||||||
|
assert body["ai"]["external_ai_in_use"] is False
|
||||||
|
# No secrets ever present in the public status payload.
|
||||||
|
body_text = r.text.lower()
|
||||||
|
for leak in ("api_key", "password", "secret_key"):
|
||||||
|
assert leak not in body_text
|
||||||
153
backend/tests/test_inline_queue.py
Normal file
153
backend/tests/test_inline_queue.py
Normal file
|
|
@ -0,0 +1,153 @@
|
||||||
|
"""Inline queue tests (desktop bundled-lite engine): with
|
||||||
|
``queue_backend=inline`` an upload runs transcribe→summarize automatically
|
||||||
|
in-process — no arq, no Redis. Providers are fakes (see test_ai_pipeline).
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import asyncio
|
||||||
|
import uuid
|
||||||
|
|
||||||
|
from sqlalchemy import select
|
||||||
|
|
||||||
|
from shonar.core.config import get_settings
|
||||||
|
from shonar.db import session as db_session
|
||||||
|
from shonar.db.models import JobStatus, JobType, ProcessingJob
|
||||||
|
from shonar.services import inline_queue, processing
|
||||||
|
|
||||||
|
from .test_ai_pipeline import FakeLlm, FakeTranscriber, upload_recording, use_fakes, user_tokens
|
||||||
|
|
||||||
|
|
||||||
|
def use_inline(monkeypatch):
|
||||||
|
"""Force queue_backend=inline for everything resolved via processing."""
|
||||||
|
inline_settings = get_settings().model_copy(update={"queue_backend": "inline"})
|
||||||
|
monkeypatch.setattr(processing, "get_settings", lambda: inline_settings)
|
||||||
|
|
||||||
|
|
||||||
|
async def _wait_terminal(recording_id: str, want: int = 2, timeout: float = 20.0):
|
||||||
|
"""Poll the DB until `want` jobs reach a terminal state."""
|
||||||
|
deadline = asyncio.get_running_loop().time() + timeout
|
||||||
|
rows: list[ProcessingJob] = []
|
||||||
|
while asyncio.get_running_loop().time() < deadline:
|
||||||
|
async with db_session._session_factory() as s:
|
||||||
|
rows = list(
|
||||||
|
(
|
||||||
|
await s.scalars(
|
||||||
|
select(ProcessingJob).where(
|
||||||
|
ProcessingJob.recording_id == uuid.UUID(recording_id)
|
||||||
|
)
|
||||||
|
)
|
||||||
|
).all()
|
||||||
|
)
|
||||||
|
if sum(
|
||||||
|
1
|
||||||
|
for j in rows
|
||||||
|
if j.status in (JobStatus.succeeded, JobStatus.failed, JobStatus.skipped)
|
||||||
|
) >= want:
|
||||||
|
return {j.job_type: j.status for j in rows}
|
||||||
|
await asyncio.sleep(0.1)
|
||||||
|
raise AssertionError(
|
||||||
|
f"jobs did not reach terminal state in {timeout}s: "
|
||||||
|
f"{[(j.job_type, j.status) for j in rows]}"
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
async def test_inline_backend_processes_upload_without_arq(client, monkeypatch):
|
||||||
|
use_fakes(monkeypatch, FakeTranscriber(), FakeLlm())
|
||||||
|
use_inline(monkeypatch)
|
||||||
|
await inline_queue.start()
|
||||||
|
try:
|
||||||
|
token = await user_tokens(client, email="inline@shonar.dev")
|
||||||
|
rec = await upload_recording(client, token, client_id="inline-1")
|
||||||
|
|
||||||
|
statuses = await _wait_terminal(rec["id"], want=2)
|
||||||
|
assert statuses[JobType.transcribe] == JobStatus.succeeded
|
||||||
|
assert statuses[JobType.summarize] == JobStatus.succeeded
|
||||||
|
|
||||||
|
h = {"Authorization": f"Bearer {token}"}
|
||||||
|
r = await client.get(f"/api/v1/recordings/{rec['id']}", headers=h)
|
||||||
|
assert r.json()["processing_status"] == "completed"
|
||||||
|
t = await client.get(f"/api/v1/recordings/{rec['id']}/transcript", headers=h)
|
||||||
|
assert t.json()["text"] == "hello world from the meeting"
|
||||||
|
s = await client.get(f"/api/v1/recordings/{rec['id']}/summary", headers=h)
|
||||||
|
assert s.json()["content"]["short"] == "Standup happened."
|
||||||
|
finally:
|
||||||
|
await inline_queue.stop()
|
||||||
|
|
||||||
|
|
||||||
|
async def test_inline_retries_transient_failure(client, monkeypatch):
|
||||||
|
stt = FakeTranscriber()
|
||||||
|
state = {"tries": 0}
|
||||||
|
original = stt.transcribe
|
||||||
|
|
||||||
|
async def flaky(audio, mime, *, language_hint=None):
|
||||||
|
state["tries"] += 1
|
||||||
|
if state["tries"] == 1:
|
||||||
|
raise processing.ProviderTransientError("temporary hiccup")
|
||||||
|
return await original(audio, mime, language_hint=language_hint)
|
||||||
|
|
||||||
|
stt.transcribe = flaky
|
||||||
|
use_fakes(monkeypatch, stt, FakeLlm())
|
||||||
|
use_inline(monkeypatch)
|
||||||
|
monkeypatch.setattr(inline_queue, "RETRY_DELAY_SECONDS", 0.05)
|
||||||
|
await inline_queue.start()
|
||||||
|
try:
|
||||||
|
token = await user_tokens(client, email="inline2@shonar.dev")
|
||||||
|
rec = await upload_recording(client, token, client_id="inline-2")
|
||||||
|
statuses = await _wait_terminal(rec["id"], want=2)
|
||||||
|
assert statuses[JobType.transcribe] == JobStatus.succeeded
|
||||||
|
assert state["tries"] >= 2 # the retry actually happened
|
||||||
|
finally:
|
||||||
|
await inline_queue.stop()
|
||||||
|
|
||||||
|
|
||||||
|
async def test_reprocess_reruns_in_place_with_model(client, monkeypatch):
|
||||||
|
"""POST /reprocess?job=transcribe&model= re-runs the pipeline on the
|
||||||
|
SAME recording (no re-upload) and persists the model override."""
|
||||||
|
stt = FakeTranscriber()
|
||||||
|
use_fakes(monkeypatch, stt, FakeLlm())
|
||||||
|
use_inline(monkeypatch)
|
||||||
|
await inline_queue.start()
|
||||||
|
try:
|
||||||
|
token = await user_tokens(client, email="reproc@shonar.dev")
|
||||||
|
h = {"Authorization": f"Bearer {token}"}
|
||||||
|
rec = await upload_recording(client, token, client_id="reproc-1")
|
||||||
|
await _wait_terminal(rec["id"], want=2)
|
||||||
|
calls_before = stt.calls
|
||||||
|
|
||||||
|
r = await client.post(
|
||||||
|
f"/api/v1/recordings/{rec['id']}/reprocess?job=transcribe&model=small",
|
||||||
|
headers=h,
|
||||||
|
)
|
||||||
|
assert r.status_code == 200, r.text
|
||||||
|
statuses = await _wait_terminal(rec["id"], want=2)
|
||||||
|
assert statuses[JobType.transcribe] == JobStatus.succeeded
|
||||||
|
assert stt.calls == calls_before + 1 # re-ran, no new recording
|
||||||
|
|
||||||
|
# Same recording id, model override persisted on the row.
|
||||||
|
r = await client.get(f"/api/v1/recordings/{rec['id']}", headers=h)
|
||||||
|
assert r.json()["id"] == rec["id"]
|
||||||
|
assert r.json()["transcription_model"] == "small"
|
||||||
|
finally:
|
||||||
|
await inline_queue.stop()
|
||||||
|
|
||||||
|
|
||||||
|
async def test_reprocess_validation(client, monkeypatch):
|
||||||
|
token = await user_tokens(client, email="reproc2@shonar.dev")
|
||||||
|
h = {"Authorization": f"Bearer {token}"}
|
||||||
|
rec = await upload_recording(client, token, client_id="reproc-2")
|
||||||
|
# Nothing transcribed yet -> summarize refuses.
|
||||||
|
r = await client.post(
|
||||||
|
f"/api/v1/recordings/{rec['id']}/reprocess?job=summarize", headers=h)
|
||||||
|
assert r.status_code == 409
|
||||||
|
# Bad model name -> 422 (never silently substituted).
|
||||||
|
r = await client.post(
|
||||||
|
f"/api/v1/recordings/{rec['id']}/reprocess?job=transcribe&model=not-a-model",
|
||||||
|
headers=h)
|
||||||
|
assert r.status_code == 422
|
||||||
|
# Not the owner -> 404.
|
||||||
|
other = await user_tokens(client, email="reproc3@shonar.dev")
|
||||||
|
r = await client.post(
|
||||||
|
f"/api/v1/recordings/{rec['id']}/reprocess?job=transcribe",
|
||||||
|
headers={"Authorization": f"Bearer {other}"})
|
||||||
|
assert r.status_code == 404
|
||||||
304
backend/tests/test_m9.py
Normal file
304
backend/tests/test_m9.py
Normal file
|
|
@ -0,0 +1,304 @@
|
||||||
|
"""M9 tests: search, exports, retention sweep.
|
||||||
|
|
||||||
|
Search runs against the real backend dialect (Postgres in CI/dev, SQLite
|
||||||
|
via SHONAR_TEST_DATABASE_URL) — both paths share the endpoint contract.
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import io
|
||||||
|
import zipfile
|
||||||
|
|
||||||
|
from tests.test_recordings import auth, user_tokens, wav_bytes
|
||||||
|
|
||||||
|
|
||||||
|
async def make_recording(client, token, title, notes=None, tags=None, recorded_at=None):
|
||||||
|
h = await auth(token)
|
||||||
|
audio = wav_bytes()
|
||||||
|
body = {"declared_mime_type": "audio/wav", "declared_size_bytes": len(audio), "title": title}
|
||||||
|
r = await client.post("/api/v1/uploads", json=body, headers=h)
|
||||||
|
sid = r.json()["id"]
|
||||||
|
await client.put(
|
||||||
|
f"/api/v1/uploads/{sid}/chunks/0", content=audio,
|
||||||
|
headers={**h, "content-type": "application/octet-stream"},
|
||||||
|
)
|
||||||
|
fin = {"duration_seconds": 5.0}
|
||||||
|
if recorded_at:
|
||||||
|
fin["recorded_at"] = recorded_at
|
||||||
|
if notes:
|
||||||
|
fin["notes"] = notes
|
||||||
|
r = await client.post(f"/api/v1/uploads/{sid}/finalize", json=fin, headers=h)
|
||||||
|
assert r.status_code == 201, r.text
|
||||||
|
rec = r.json()
|
||||||
|
if tags:
|
||||||
|
r = await client.patch(
|
||||||
|
f"/api/v1/recordings/{rec['id']}", json={"tags": tags}, headers=h
|
||||||
|
)
|
||||||
|
assert r.status_code == 200
|
||||||
|
return rec["id"]
|
||||||
|
|
||||||
|
|
||||||
|
async def add_transcript(rec_id: str, text: str, segments=None):
|
||||||
|
"""Insert a transcript row directly (no AI provider under test)."""
|
||||||
|
from shonar.db.models import Transcript
|
||||||
|
from shonar.db.session import session_factory
|
||||||
|
|
||||||
|
async with session_factory()() as s:
|
||||||
|
s.add(Transcript(recording_id=rec_id, text=text, segments=segments, provider="test"))
|
||||||
|
await s.commit()
|
||||||
|
|
||||||
|
|
||||||
|
async def add_summary(rec_id: str, content: dict):
|
||||||
|
from shonar.db.models import Summary
|
||||||
|
from shonar.db.session import session_factory
|
||||||
|
|
||||||
|
async with session_factory()() as s:
|
||||||
|
s.add(Summary(recording_id=rec_id, content=content, provider="test"))
|
||||||
|
await s.commit()
|
||||||
|
|
||||||
|
|
||||||
|
# --- search -------------------------------------------------------------------
|
||||||
|
|
||||||
|
|
||||||
|
async def test_search_title_and_transcript(client):
|
||||||
|
token = await user_tokens(client, email="m9s1@example.com")
|
||||||
|
rid_t = await make_recording(client, token, "Quarterly budget review")
|
||||||
|
rid_x = await make_recording(client, token, "Grocery list")
|
||||||
|
await add_transcript(rid_x, "remember to buy kale chips and quinoa tonight")
|
||||||
|
|
||||||
|
r = await client.get("/api/v1/search", params={"q": "budget"}, headers=await auth(token))
|
||||||
|
assert r.status_code == 200, r.text
|
||||||
|
body = r.json()
|
||||||
|
assert body["total"] == 1
|
||||||
|
assert body["items"][0]["id"] == rid_t
|
||||||
|
assert body["items"][0]["field"] == "title"
|
||||||
|
|
||||||
|
r = await client.get("/api/v1/search", params={"q": "quinoa"}, headers=await auth(token))
|
||||||
|
body = r.json()
|
||||||
|
assert body["total"] == 1
|
||||||
|
assert body["items"][0]["id"] == rid_x
|
||||||
|
assert body["items"][0]["field"] == "transcript"
|
||||||
|
assert "quinoa" in body["items"][0]["snippet"].lower()
|
||||||
|
|
||||||
|
|
||||||
|
async def test_search_scope_and_tag(client):
|
||||||
|
token = await user_tokens(client, email="m9s2@example.com")
|
||||||
|
rid = await make_recording(client, token, "Standup", tags=["daily"])
|
||||||
|
await add_transcript(rid, "we discussed the daily standup format")
|
||||||
|
|
||||||
|
# tag scope finds by tag name
|
||||||
|
r = await client.get(
|
||||||
|
"/api/v1/search", params={"q": "daily", "scope": "tag"}, headers=await auth(token)
|
||||||
|
)
|
||||||
|
assert r.json()["total"] == 1
|
||||||
|
assert r.json()["items"][0]["field"] == "tag"
|
||||||
|
|
||||||
|
# a transcript-only word does NOT match under scope=title
|
||||||
|
r = await client.get(
|
||||||
|
"/api/v1/search", params={"q": "discussed", "scope": "title"}, headers=await auth(token)
|
||||||
|
)
|
||||||
|
assert r.json()["total"] == 0
|
||||||
|
r = await client.get(
|
||||||
|
"/api/v1/search", params={"q": "discussed", "scope": "transcript"},
|
||||||
|
headers=await auth(token),
|
||||||
|
)
|
||||||
|
assert r.json()["total"] == 1
|
||||||
|
|
||||||
|
|
||||||
|
async def test_search_isolation_and_deleted(client):
|
||||||
|
token_a = await user_tokens(client, email="m9s3a@example.com")
|
||||||
|
token_b = await user_tokens(client, email="m9s3b@example.com")
|
||||||
|
rid = await make_recording(client, token_a, "secret sauce recipe")
|
||||||
|
|
||||||
|
h_b = await auth(token_b)
|
||||||
|
r = await client.get("/api/v1/search", params={"q": "secret"}, headers=h_b)
|
||||||
|
assert r.json()["total"] == 0 # other users' data invisible
|
||||||
|
|
||||||
|
# soft-deleted rows drop out of search
|
||||||
|
await client.delete(f"/api/v1/recordings/{rid}", headers=await auth(token_a))
|
||||||
|
r = await client.get(
|
||||||
|
"/api/v1/search", params={"q": "secret"}, headers=await auth(token_a)
|
||||||
|
)
|
||||||
|
assert r.json()["total"] == 0
|
||||||
|
|
||||||
|
|
||||||
|
async def test_search_summary_and_notes(client):
|
||||||
|
token = await user_tokens(client, email="m9s4@example.com")
|
||||||
|
rid = await make_recording(client, token, "Meeting", notes="bring the projector cable")
|
||||||
|
await add_summary(rid, {"short": "sprint retro", "action_items": ["fix flaky test"]})
|
||||||
|
|
||||||
|
r = await client.get("/api/v1/search", params={"q": "projector"}, headers=await auth(token))
|
||||||
|
assert r.json()["total"] == 1
|
||||||
|
r = await client.get("/api/v1/search", params={"q": "flaky"}, headers=await auth(token))
|
||||||
|
assert r.json()["total"] == 1
|
||||||
|
|
||||||
|
|
||||||
|
# --- list filters ---------------------------------------------------------------
|
||||||
|
|
||||||
|
|
||||||
|
async def test_list_filters_tag_status_date(client):
|
||||||
|
token = await user_tokens(client, email="m9f1@example.com")
|
||||||
|
await make_recording(client, token, "Old one", tags=["keep"],
|
||||||
|
recorded_at="2020-01-01T10:00:00Z")
|
||||||
|
rid_new = await make_recording(client, token, "New one", tags=["keep"],
|
||||||
|
recorded_at="2026-01-01T10:00:00Z")
|
||||||
|
h = await auth(token)
|
||||||
|
|
||||||
|
r = await client.get("/api/v1/recordings", params={"tag": "keep"}, headers=h)
|
||||||
|
assert r.json()["total"] == 2
|
||||||
|
r = await client.get(
|
||||||
|
"/api/v1/recordings", params={"tag": "keep", "from_date": "2025-06-01T00:00:00Z"},
|
||||||
|
headers=h,
|
||||||
|
)
|
||||||
|
body = r.json()
|
||||||
|
assert body["total"] == 1 and body["items"][0]["id"] == rid_new
|
||||||
|
|
||||||
|
r = await client.get("/api/v1/recordings", params={"status": "ai_disabled"}, headers=h)
|
||||||
|
assert r.json()["total"] == 2
|
||||||
|
r = await client.get("/api/v1/recordings", params={"status": "bogus"}, headers=h)
|
||||||
|
assert r.status_code == 422
|
||||||
|
r = await client.get("/api/v1/recordings", params={"tag": "nope"}, headers=h)
|
||||||
|
assert r.json()["total"] == 0
|
||||||
|
|
||||||
|
|
||||||
|
# --- exports --------------------------------------------------------------------
|
||||||
|
|
||||||
|
|
||||||
|
async def test_export_formats(client):
|
||||||
|
token = await user_tokens(client, email="m9e1@example.com")
|
||||||
|
rid = await make_recording(client, token, "Retro & Planning", notes="retro notes here",
|
||||||
|
tags=["team"])
|
||||||
|
await add_transcript(
|
||||||
|
rid, "first segment second segment",
|
||||||
|
segments=[{"start": 0.0, "end": 2.5, "text": "first segment", "speaker": None},
|
||||||
|
{"start": 2.5, "end": 5.0, "text": "second segment", "speaker": "S1"}],
|
||||||
|
)
|
||||||
|
await add_summary(rid, {"short": "one line", "action_items": ["do the thing"]})
|
||||||
|
h = await auth(token)
|
||||||
|
|
||||||
|
r = await client.get(f"/api/v1/recordings/{rid}/export", params={"fmt": "txt"}, headers=h)
|
||||||
|
assert r.status_code == 200
|
||||||
|
assert r.text == "first segment second segment"
|
||||||
|
assert "attachment" in r.headers["content-disposition"]
|
||||||
|
assert ".txt" in r.headers["content-disposition"]
|
||||||
|
|
||||||
|
r = await client.get(f"/api/v1/recordings/{rid}/export", params={"fmt": "md"}, headers=h)
|
||||||
|
assert r.status_code == 200
|
||||||
|
assert "# Retro & Planning" in r.text
|
||||||
|
assert "do the thing" in r.text
|
||||||
|
assert "`[00:02.5]` **S1**: second segment" in r.text
|
||||||
|
assert "`team`" in r.text
|
||||||
|
|
||||||
|
r = await client.get(f"/api/v1/recordings/{rid}/export", params={"fmt": "zip"}, headers=h)
|
||||||
|
assert r.status_code == 200
|
||||||
|
assert r.headers["content-type"] == "application/zip"
|
||||||
|
with zipfile.ZipFile(io.BytesIO(r.content)) as z:
|
||||||
|
names = z.namelist()
|
||||||
|
assert "transcript.txt" in names and "notes.md" in names
|
||||||
|
assert any(n.endswith(".wav") for n in names)
|
||||||
|
|
||||||
|
r = await client.get(f"/api/v1/recordings/{rid}/export", params={"fmt": "audio"}, headers=h)
|
||||||
|
assert r.status_code == 200
|
||||||
|
assert r.content.startswith(b"RIFF")
|
||||||
|
|
||||||
|
|
||||||
|
async def test_export_missing_and_foreign(client):
|
||||||
|
token = await user_tokens(client, email="m9e2@example.com")
|
||||||
|
rid = await make_recording(client, token, "Bare") # no transcript/summary/notes
|
||||||
|
token2 = await user_tokens(client, email="m9e2b@example.com")
|
||||||
|
|
||||||
|
r = await client.get(f"/api/v1/recordings/{rid}/export", params={"fmt": "txt"},
|
||||||
|
headers=await auth(token))
|
||||||
|
assert r.status_code == 404 # no transcript yet
|
||||||
|
r = await client.get(f"/api/v1/recordings/{rid}/export", params={"fmt": "md"},
|
||||||
|
headers=await auth(token))
|
||||||
|
assert r.status_code == 404
|
||||||
|
# zip still works with just the audio
|
||||||
|
r = await client.get(f"/api/v1/recordings/{rid}/export", params={"fmt": "zip"},
|
||||||
|
headers=await auth(token))
|
||||||
|
assert r.status_code == 200
|
||||||
|
|
||||||
|
r = await client.get(f"/api/v1/recordings/{rid}/export", params={"fmt": "audio"},
|
||||||
|
headers=await auth(token2))
|
||||||
|
assert r.status_code == 404 # not yours
|
||||||
|
|
||||||
|
|
||||||
|
# --- retention sweep --------------------------------------------------------------
|
||||||
|
|
||||||
|
|
||||||
|
async def test_retention_sweep_purges_expired(client):
|
||||||
|
from datetime import timedelta
|
||||||
|
|
||||||
|
from shonar.db.models import Asset, Recording, utcnow
|
||||||
|
from shonar.db.session import session_factory
|
||||||
|
from shonar.services import retention
|
||||||
|
|
||||||
|
token = await user_tokens(client, email="m9r1@example.com")
|
||||||
|
rid = await make_recording(client, token, "Doomed")
|
||||||
|
h = await auth(token)
|
||||||
|
|
||||||
|
# Soft-delete, then push deleted_at past the grace window directly.
|
||||||
|
r = await client.delete(f"/api/v1/recordings/{rid}", headers=h)
|
||||||
|
assert r.status_code == 204
|
||||||
|
async with session_factory()() as s:
|
||||||
|
rec = await s.get(Recording, rid)
|
||||||
|
rec.deleted_at = utcnow() - timedelta(days=31)
|
||||||
|
# remember the storage key before the row vanishes
|
||||||
|
from shonar.db.models import Asset
|
||||||
|
|
||||||
|
asset = (
|
||||||
|
await s.execute(
|
||||||
|
Asset.__table__.select().where(Asset.recording_id == rec.id) # noqa: SLF001
|
||||||
|
)
|
||||||
|
).first()
|
||||||
|
storage_key = asset.storage_key if asset else None
|
||||||
|
await s.commit()
|
||||||
|
assert storage_key is not None
|
||||||
|
|
||||||
|
purged = await retention.sweep_deleted()
|
||||||
|
assert purged["recordings"] == 1
|
||||||
|
assert purged["files"] >= 1
|
||||||
|
|
||||||
|
async with session_factory()() as s:
|
||||||
|
assert await s.get(Recording, rid) is None
|
||||||
|
from shonar.storage import get_storage
|
||||||
|
|
||||||
|
assert not await get_storage().exists(storage_key) # file gone too
|
||||||
|
|
||||||
|
# Within the window, nothing is purged (cancellable).
|
||||||
|
rid2 = await make_recording(client, token, "Fresh delete")
|
||||||
|
await client.delete(f"/api/v1/recordings/{rid2}", headers=h)
|
||||||
|
purged = await retention.sweep_deleted()
|
||||||
|
assert purged["recordings"] == 0
|
||||||
|
async with session_factory()() as s:
|
||||||
|
assert await s.get(Recording, rid2) is not None
|
||||||
|
|
||||||
|
|
||||||
|
async def test_retention_sweep_account(client):
|
||||||
|
from datetime import timedelta
|
||||||
|
|
||||||
|
from shonar.db.models import Recording, utcnow
|
||||||
|
from shonar.db.session import session_factory
|
||||||
|
from shonar.services import retention
|
||||||
|
|
||||||
|
token = await user_tokens(client, email="m9r2@example.com")
|
||||||
|
rid = await make_recording(client, token, "Gone with user")
|
||||||
|
|
||||||
|
async with session_factory()() as s:
|
||||||
|
user = await s.scalar(select_user("m9r2@example.com"))
|
||||||
|
user.deleted_at = utcnow() - timedelta(days=31)
|
||||||
|
await s.commit()
|
||||||
|
|
||||||
|
purged = await retention.sweep_deleted()
|
||||||
|
assert purged["users"] == 1
|
||||||
|
async with session_factory()() as s:
|
||||||
|
assert await s.get(Recording, rid) is None # cascade
|
||||||
|
assert await s.scalar(select_user("m9r2@example.com")) is None
|
||||||
|
|
||||||
|
|
||||||
|
def select_user(email: str):
|
||||||
|
from sqlalchemy import select
|
||||||
|
|
||||||
|
from shonar.db.models import User
|
||||||
|
|
||||||
|
return select(User).where(User.email == email)
|
||||||
245
backend/tests/test_models.py
Normal file
245
backend/tests/test_models.py
Normal file
|
|
@ -0,0 +1,245 @@
|
||||||
|
"""Stage 1: per-recording transcription models + models API.
|
||||||
|
|
||||||
|
Global default (base) with optional per-recording overrides; the worker
|
||||||
|
uses the exact saved model; history is never rewritten; unknown models
|
||||||
|
are rejected; missing downloads fail fast with instructions.
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import struct
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
from sqlalchemy import select
|
||||||
|
|
||||||
|
from shonar.db import session as db_session
|
||||||
|
from shonar.db.models import Recording
|
||||||
|
from shonar.services import processing
|
||||||
|
|
||||||
|
|
||||||
|
def wav_bytes(payload_len: int = 64) -> bytes:
|
||||||
|
data = bytes(range(payload_len % 256)) * (payload_len // 256 + 1)
|
||||||
|
data = data[:payload_len]
|
||||||
|
return (
|
||||||
|
b"RIFF" + struct.pack("<I", 36 + len(data)) + b"WAVE"
|
||||||
|
+ b"fmt " + struct.pack("<IHHIIHH", 16, 1, 1, 8000, 8000, 1, 8)
|
||||||
|
+ b"data" + struct.pack("<I", len(data))
|
||||||
|
) + data
|
||||||
|
|
||||||
|
|
||||||
|
AUTH = {"email": "models@example.com", "password": "models-test-passw0rd-123"}
|
||||||
|
|
||||||
|
|
||||||
|
async def user_tokens(client):
|
||||||
|
r = await client.post(
|
||||||
|
"/api/v1/auth/register",
|
||||||
|
json={"email": AUTH["email"], "password": AUTH["password"]},
|
||||||
|
)
|
||||||
|
assert r.status_code == 201, r.text
|
||||||
|
return r.json()["access_token"]
|
||||||
|
|
||||||
|
|
||||||
|
async def upload_recording(client, token, client_id, session_model=None, finalize_model=None):
|
||||||
|
h = {"Authorization": f"Bearer {token}"}
|
||||||
|
data = wav_bytes()
|
||||||
|
body = {"declared_mime_type": "audio/wav", "declared_size_bytes": len(data),
|
||||||
|
"client_recording_id": client_id}
|
||||||
|
if session_model is not None:
|
||||||
|
body["transcription_model"] = session_model
|
||||||
|
r = await client.post("/api/v1/uploads", json=body, headers=h)
|
||||||
|
assert r.status_code == 201, r.text
|
||||||
|
sid = r.json()["id"]
|
||||||
|
r = await client.put(f"/api/v1/uploads/{sid}/chunks/0", content=data,
|
||||||
|
headers={**h, "content-type": "application/octet-stream"})
|
||||||
|
assert r.status_code == 201, r.text
|
||||||
|
fbody = {"duration_seconds": 5.0}
|
||||||
|
if finalize_model is not None:
|
||||||
|
fbody["transcription_model"] = finalize_model
|
||||||
|
r = await client.post(f"/api/v1/uploads/{sid}/finalize", json=fbody, headers=h)
|
||||||
|
return r
|
||||||
|
|
||||||
|
|
||||||
|
async def test_models_endpoint_lists_five(client):
|
||||||
|
token = await user_tokens(client)
|
||||||
|
r = await client.get("/api/v1/models", headers={"Authorization": f"Bearer {token}"})
|
||||||
|
assert r.status_code == 200, r.text
|
||||||
|
body = r.json()
|
||||||
|
assert body["default_model"] == "base"
|
||||||
|
assert [m["name"] for m in body["models"]] == ["tiny", "base", "small", "medium", "large-v3"]
|
||||||
|
base = next(m for m in body["models"] if m["name"] == "base")
|
||||||
|
assert base["is_default"] is True
|
||||||
|
assert "ordinary computers" in base["description"]
|
||||||
|
for m in body["models"]:
|
||||||
|
assert isinstance(m["downloaded"], bool)
|
||||||
|
assert isinstance(m["available"], bool)
|
||||||
|
|
||||||
|
|
||||||
|
async def test_default_get_put_validation(client):
|
||||||
|
token = await user_tokens(client)
|
||||||
|
h = {"Authorization": f"Bearer {token}"}
|
||||||
|
r = await client.get("/api/v1/models/default", headers=h)
|
||||||
|
assert r.json() == {"default_model": "base"}
|
||||||
|
r = await client.put("/api/v1/models/default", json={"model": "small"}, headers=h)
|
||||||
|
assert r.status_code == 200, r.text
|
||||||
|
assert r.json() == {"default_model": "small"}
|
||||||
|
r = await client.put("/api/v1/models/default", json={"model": "xxl-turbo"}, headers=h)
|
||||||
|
assert r.status_code == 422
|
||||||
|
assert "Supported models" in r.text
|
||||||
|
# rejected change did not stick
|
||||||
|
r = await client.get("/api/v1/models/default", headers=h)
|
||||||
|
assert r.json() == {"default_model": "small"}
|
||||||
|
|
||||||
|
|
||||||
|
async def test_finalize_override_and_default_history(client):
|
||||||
|
token = await user_tokens(client)
|
||||||
|
h = {"Authorization": f"Bearer {token}"}
|
||||||
|
r = await upload_recording(client, token, "m-override-1", finalize_model="small")
|
||||||
|
assert r.status_code == 201, r.text
|
||||||
|
assert r.json()["transcription_model"] == "small"
|
||||||
|
|
||||||
|
r = await upload_recording(client, token, "m-default-1")
|
||||||
|
assert r.json()["transcription_model"] == "base"
|
||||||
|
|
||||||
|
# Changing the default affects future rows only.
|
||||||
|
r = await client.put("/api/v1/models/default", json={"model": "tiny"}, headers=h)
|
||||||
|
assert r.status_code == 200
|
||||||
|
r = await upload_recording(client, token, "m-default-2")
|
||||||
|
assert r.json()["transcription_model"] == "tiny"
|
||||||
|
|
||||||
|
async with db_session._session_factory() as s:
|
||||||
|
got = dict((await s.execute(
|
||||||
|
select(Recording.client_recording_id, Recording.transcription_model)
|
||||||
|
.where(Recording.client_recording_id.in_(
|
||||||
|
["m-override-1", "m-default-1", "m-default-2"])))
|
||||||
|
).all())
|
||||||
|
assert got == {"m-override-1": "small", "m-default-1": "base", "m-default-2": "tiny"}
|
||||||
|
|
||||||
|
|
||||||
|
async def test_session_override_and_finalize_wins(client):
|
||||||
|
token = await user_tokens(client)
|
||||||
|
r = await upload_recording(client, token, "m-sess-1", session_model="small")
|
||||||
|
assert r.status_code == 201, r.text
|
||||||
|
assert r.json()["transcription_model"] == "small"
|
||||||
|
|
||||||
|
r = await upload_recording(client, token, "m-sess-2", session_model="small",
|
||||||
|
finalize_model="tiny")
|
||||||
|
assert r.status_code == 201, r.text
|
||||||
|
assert r.json()["transcription_model"] == "tiny"
|
||||||
|
|
||||||
|
|
||||||
|
async def test_finalize_invalid_override_rejected(client):
|
||||||
|
token = await user_tokens(client)
|
||||||
|
r = await upload_recording(client, token, "m-bad-1", finalize_model="xxl-turbo")
|
||||||
|
assert r.status_code == 422
|
||||||
|
assert "Supported models" in r.text
|
||||||
|
|
||||||
|
|
||||||
|
class SpyTranscriber:
|
||||||
|
"""Stands in for FasterWhisperProvider; records the model it was built with."""
|
||||||
|
|
||||||
|
name = "faster_whisper"
|
||||||
|
seen_models: list = []
|
||||||
|
|
||||||
|
def __init__(self, model: str = "base"):
|
||||||
|
type(self).seen_models.append(model)
|
||||||
|
self.model = model
|
||||||
|
|
||||||
|
async def transcribe(self, audio, mime, *, language_hint=None):
|
||||||
|
from shonar.services.ai import Segment, TranscriptResult
|
||||||
|
|
||||||
|
return TranscriptResult(
|
||||||
|
text="spy text", language="en",
|
||||||
|
segments=[Segment(0.0, 1.0, "spy text")], model=self.model,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.fixture(autouse=True)
|
||||||
|
def _clear_spy():
|
||||||
|
SpyTranscriber.seen_models = []
|
||||||
|
yield
|
||||||
|
SpyTranscriber.seen_models = []
|
||||||
|
|
||||||
|
|
||||||
|
async def _run_with_spy(monkeypatch, recording_id):
|
||||||
|
import shonar.services.ai.faster_whisper as fw
|
||||||
|
|
||||||
|
# The worker rebuilds a faster_whisper provider from the saved model;
|
||||||
|
# the late import inside run_transcribe picks up this spy.
|
||||||
|
monkeypatch.setattr(fw, "FasterWhisperProvider", SpyTranscriber)
|
||||||
|
monkeypatch.setattr(
|
||||||
|
processing, "get_transcription_provider",
|
||||||
|
lambda settings: SpyTranscriber(model="ignored"),
|
||||||
|
)
|
||||||
|
await processing.run_transcribe({}, recording_id)
|
||||||
|
|
||||||
|
|
||||||
|
async def test_worker_uses_saved_model(client, monkeypatch):
|
||||||
|
token = await user_tokens(client)
|
||||||
|
rec = (await upload_recording(client, token, "m-work-1", finalize_model="small")).json()
|
||||||
|
await _run_with_spy(monkeypatch, rec["id"])
|
||||||
|
# Last construction is the worker rebuild from the saved model.
|
||||||
|
assert SpyTranscriber.seen_models[-1] == "small"
|
||||||
|
|
||||||
|
h = {"Authorization": f"Bearer {token}"}
|
||||||
|
r = await client.get(f"/api/v1/recordings/{rec['id']}/transcript", headers=h)
|
||||||
|
assert r.status_code == 200
|
||||||
|
assert r.json()["model"] == "small"
|
||||||
|
|
||||||
|
r = await client.get(f"/api/v1/recordings/{rec['id']}/jobs", headers=h)
|
||||||
|
job = next(j for j in r.json() if j["job_type"] == "transcribe")
|
||||||
|
assert job["status"] == "succeeded"
|
||||||
|
assert job["stage"] is None
|
||||||
|
assert job["progress"] == 100
|
||||||
|
|
||||||
|
|
||||||
|
async def test_worker_keeps_history_after_default_change(client, monkeypatch):
|
||||||
|
token = await user_tokens(client)
|
||||||
|
h = {"Authorization": f"Bearer {token}"}
|
||||||
|
rec = (await upload_recording(client, token, "m-work-2")).json()
|
||||||
|
assert rec["transcription_model"] == "base"
|
||||||
|
r = await client.put("/api/v1/models/default", json={"model": "tiny"}, headers=h)
|
||||||
|
assert r.status_code == 200
|
||||||
|
await _run_with_spy(monkeypatch, rec["id"])
|
||||||
|
assert SpyTranscriber.seen_models[-1] == "base"
|
||||||
|
|
||||||
|
|
||||||
|
async def test_unavailable_model_fails_with_instructions(client, monkeypatch):
|
||||||
|
import shonar.services.ai.faster_whisper as fw
|
||||||
|
from shonar.services.ai.faster_whisper import FasterWhisperProvider
|
||||||
|
|
||||||
|
# Real provider class, but the model is (simulated) not downloaded.
|
||||||
|
monkeypatch.setattr(fw, "is_model_downloaded", lambda name: False)
|
||||||
|
monkeypatch.setattr(
|
||||||
|
processing, "get_transcription_provider",
|
||||||
|
lambda settings: FasterWhisperProvider(model="base"),
|
||||||
|
)
|
||||||
|
token = await user_tokens(client)
|
||||||
|
h = {"Authorization": f"Bearer {token}"}
|
||||||
|
rec = (await upload_recording(client, token, "m-work-3")).json()
|
||||||
|
await processing.run_transcribe({}, rec["id"])
|
||||||
|
|
||||||
|
r = await client.get(f"/api/v1/recordings/{rec['id']}", headers=h)
|
||||||
|
assert r.json()["processing_status"] == "failed"
|
||||||
|
assert "not downloaded" in (r.json()["processing_error"] or "")
|
||||||
|
r = await client.get(f"/api/v1/recordings/{rec['id']}/jobs", headers=h)
|
||||||
|
job = next(j for j in r.json() if j["job_type"] == "transcribe")
|
||||||
|
assert job["status"] == "failed"
|
||||||
|
assert "download" in (job["error"] or "").lower()
|
||||||
|
|
||||||
|
|
||||||
|
async def test_reupload_explicit_override_updates_model(client):
|
||||||
|
token = await user_tokens(client)
|
||||||
|
r = await upload_recording(client, token, "m-reup-1", finalize_model="small")
|
||||||
|
assert r.json()["transcription_model"] == "small"
|
||||||
|
# Same client id, new explicit choice: model history moves with it.
|
||||||
|
r = await upload_recording(client, token, "m-reup-1", finalize_model="tiny")
|
||||||
|
assert r.status_code == 201, r.text
|
||||||
|
assert r.json()["transcription_model"] == "tiny"
|
||||||
|
# Same client id, no choice: history untouched.
|
||||||
|
r = await upload_recording(client, token, "m-reup-1")
|
||||||
|
assert r.json()["transcription_model"] == "tiny"
|
||||||
|
|
||||||
|
|
||||||
|
async def test_models_need_auth(client):
|
||||||
|
r = await client.get("/api/v1/models")
|
||||||
|
assert r.status_code in (401, 403)
|
||||||
35
backend/tests/test_provider_info.py
Normal file
35
backend/tests/test_provider_info.py
Normal file
|
|
@ -0,0 +1,35 @@
|
||||||
|
"""provider-info handshake tests (P2/P3 client probing depends on this)."""
|
||||||
|
|
||||||
|
|
||||||
|
async def test_provider_info_identifies_shonar(client):
|
||||||
|
r = await client.get("/api/v1/provider-info")
|
||||||
|
assert r.status_code == 200
|
||||||
|
body = r.json()
|
||||||
|
assert body["kind"] == "shonar"
|
||||||
|
assert body["api_version"] == "v1"
|
||||||
|
caps = body["capabilities"]
|
||||||
|
assert caps["chunked_upload"] is True
|
||||||
|
assert caps["account_deletion"] is True
|
||||||
|
# test env configures no AI providers -> flags must be false
|
||||||
|
assert caps["server_transcription"] is False
|
||||||
|
assert caps["server_summary"] is False
|
||||||
|
assert body["storage_backend"] in ("local", "s3")
|
||||||
|
|
||||||
|
|
||||||
|
async def test_provider_info_leaks_no_paths_or_secrets(client):
|
||||||
|
r = await client.get("/api/v1/provider-info")
|
||||||
|
body = r.json()
|
||||||
|
# every string value must be a bare identifier, never a filesystem path
|
||||||
|
def walk(v):
|
||||||
|
if isinstance(v, str):
|
||||||
|
assert not v.startswith("/"), f"path-like value in handshake: {v!r}"
|
||||||
|
low = v.lower()
|
||||||
|
for banned in ("secret", "token", "password", "key"):
|
||||||
|
assert banned not in low
|
||||||
|
elif isinstance(v, dict):
|
||||||
|
for x in v.values():
|
||||||
|
walk(x)
|
||||||
|
elif isinstance(v, list):
|
||||||
|
for x in v:
|
||||||
|
walk(x)
|
||||||
|
walk(body)
|
||||||
264
backend/tests/test_recordings.py
Normal file
264
backend/tests/test_recordings.py
Normal file
|
|
@ -0,0 +1,264 @@
|
||||||
|
"""Upload sessions + recordings CRUD tests (M2).
|
||||||
|
|
||||||
|
Uses real WAV magic bytes; validation is byte-level, so fakes would be
|
||||||
|
testing the wrong thing.
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import struct
|
||||||
|
import uuid
|
||||||
|
|
||||||
|
|
||||||
|
def wav_bytes(payload_len: int = 64) -> bytes:
|
||||||
|
data = bytes(range(payload_len % 256)) * (payload_len // 256 + 1)
|
||||||
|
data = data[:payload_len]
|
||||||
|
header = (
|
||||||
|
b"RIFF" + struct.pack("<I", 36 + len(data)) + b"WAVE"
|
||||||
|
+ b"fmt " + struct.pack("<IHHIIHH", 16, 1, 1, 8000, 8000, 1, 8)
|
||||||
|
+ b"data" + struct.pack("<I", len(data))
|
||||||
|
)
|
||||||
|
return header + data
|
||||||
|
|
||||||
|
|
||||||
|
def mp4_bytes() -> bytes:
|
||||||
|
return b"\x00\x00\x00 ftypM4A " + b"\x00" * 64
|
||||||
|
|
||||||
|
|
||||||
|
AUTH = {"email": "m2@example.com",
|
||||||
|
"password": "m2-test-" + "passw0rd-123"}
|
||||||
|
|
||||||
|
|
||||||
|
async def user_tokens(client, email=AUTH["email"], password=AUTH["password"]):
|
||||||
|
r = await client.post("/api/v1/auth/register", json={"email": email, "password": password})
|
||||||
|
assert r.status_code == 201, r.text
|
||||||
|
return r.json()["access_token"]
|
||||||
|
|
||||||
|
|
||||||
|
async def auth(token: str) -> dict:
|
||||||
|
return {"Authorization": f"Bearer {token}"}
|
||||||
|
|
||||||
|
|
||||||
|
async def upload_full(client, token: str, data: bytes, mime="audio/wav",
|
||||||
|
client_id=None, title=None):
|
||||||
|
h = await auth(token)
|
||||||
|
r = await client.post(
|
||||||
|
"/api/v1/uploads",
|
||||||
|
json={"declared_mime_type": mime, "declared_size_bytes": len(data),
|
||||||
|
"client_recording_id": client_id, "title": title},
|
||||||
|
headers=h,
|
||||||
|
)
|
||||||
|
assert r.status_code == 201, r.text
|
||||||
|
sid = r.json()["id"]
|
||||||
|
r = await client.put(f"/api/v1/uploads/{sid}/chunks/0", content=data,
|
||||||
|
headers={**h, "content-type": "application/octet-stream"})
|
||||||
|
assert r.status_code == 201, r.text
|
||||||
|
r = await client.post(
|
||||||
|
f"/api/v1/uploads/{sid}/finalize",
|
||||||
|
json={"duration_seconds": 12.5},
|
||||||
|
headers=h,
|
||||||
|
)
|
||||||
|
return sid, r
|
||||||
|
|
||||||
|
|
||||||
|
# --- upload session ----------------------------------------------------------
|
||||||
|
|
||||||
|
|
||||||
|
async def test_upload_happy_path_creates_recording(client):
|
||||||
|
token = await user_tokens(client)
|
||||||
|
data = wav_bytes()
|
||||||
|
sid, r = await upload_full(client, token, data, title="Standup")
|
||||||
|
assert r.status_code == 201, r.text
|
||||||
|
rec = r.json()
|
||||||
|
assert rec["title"] == "Standup"
|
||||||
|
assert rec["has_audio"] is True
|
||||||
|
# M7: with no AI configured the pipeline marks audio-only explicitly.
|
||||||
|
assert rec["processing_status"] == "ai_disabled"
|
||||||
|
assert rec["duration_seconds"] == 12.5
|
||||||
|
# No storage keys or internals leak.
|
||||||
|
assert "storage" not in r.text and "key" not in r.text.lower().replace("chunk", "")
|
||||||
|
|
||||||
|
|
||||||
|
async def test_upload_rejects_bad_mime_declared(client):
|
||||||
|
token = await user_tokens(client)
|
||||||
|
h = await auth(token)
|
||||||
|
r = await client.post("/api/v1/uploads",
|
||||||
|
json={"declared_mime_type": "application/x-msdownload",
|
||||||
|
"declared_size_bytes": 100}, headers=h)
|
||||||
|
assert r.status_code == 415
|
||||||
|
|
||||||
|
|
||||||
|
async def test_upload_rejects_oversize(client):
|
||||||
|
token = await user_tokens(client)
|
||||||
|
h = await auth(token)
|
||||||
|
r = await client.post("/api/v1/uploads",
|
||||||
|
json={"declared_mime_type": "audio/wav",
|
||||||
|
"declared_size_bytes": 5 * 1024**3}, headers=h)
|
||||||
|
assert r.status_code == 413
|
||||||
|
|
||||||
|
|
||||||
|
async def test_finalize_rejects_bytes_not_matching_mime(client):
|
||||||
|
token = await user_tokens(client)
|
||||||
|
data = mp4_bytes()
|
||||||
|
_sid, r = await upload_full(client, token, data, mime="audio/wav")
|
||||||
|
assert r.status_code == 415
|
||||||
|
|
||||||
|
|
||||||
|
async def test_finalize_rejects_size_mismatch(client):
|
||||||
|
token = await user_tokens(client)
|
||||||
|
h = await auth(token)
|
||||||
|
data = wav_bytes()
|
||||||
|
r = await client.post("/api/v1/uploads",
|
||||||
|
json={"declared_mime_type": "audio/wav",
|
||||||
|
"declared_size_bytes": len(data) + 10}, headers=h)
|
||||||
|
sid = r.json()["id"]
|
||||||
|
await client.put(f"/api/v1/uploads/{sid}/chunks/0", content=data,
|
||||||
|
headers={**h, "content-type": "application/octet-stream"})
|
||||||
|
r = await client.post(f"/api/v1/uploads/{sid}/finalize", json={}, headers=h)
|
||||||
|
assert r.status_code == 422
|
||||||
|
detail = r.json()["detail"].lower()
|
||||||
|
assert "missing chunks" in detail or "size mismatch" in detail
|
||||||
|
|
||||||
|
|
||||||
|
async def test_chunk_resume_status_and_idempotency(client):
|
||||||
|
token = await user_tokens(client)
|
||||||
|
h = await auth(token)
|
||||||
|
data = wav_bytes()
|
||||||
|
r = await client.post("/api/v1/uploads",
|
||||||
|
json={"declared_mime_type": "audio/wav",
|
||||||
|
"declared_size_bytes": len(data)}, headers=h)
|
||||||
|
sid = r.json()["id"]
|
||||||
|
r = await client.get(f"/api/v1/uploads/{sid}", headers=h)
|
||||||
|
assert r.status_code == 200
|
||||||
|
assert r.json()["received_chunk_indexes"] == []
|
||||||
|
hdr = {**h, "content-type": "application/octet-stream"}
|
||||||
|
await client.put(f"/api/v1/uploads/{sid}/chunks/0", content=data, headers=hdr)
|
||||||
|
# Duplicate PUT of chunk 0 (retry) must not duplicate or corrupt.
|
||||||
|
await client.put(f"/api/v1/uploads/{sid}/chunks/0", content=data, headers=hdr)
|
||||||
|
r = await client.get(f"/api/v1/uploads/{sid}", headers=h)
|
||||||
|
assert r.json()["received_chunk_indexes"] == [0]
|
||||||
|
r = await client.post(f"/api/v1/uploads/{sid}/finalize", json={}, headers=h)
|
||||||
|
assert r.status_code == 201
|
||||||
|
|
||||||
|
|
||||||
|
async def test_chunk_checksum_enforced(client):
|
||||||
|
token = await user_tokens(client)
|
||||||
|
h = await auth(token)
|
||||||
|
data = wav_bytes()
|
||||||
|
r = await client.post("/api/v1/uploads",
|
||||||
|
json={"declared_mime_type": "audio/wav",
|
||||||
|
"declared_size_bytes": len(data)}, headers=h)
|
||||||
|
sid = r.json()["id"]
|
||||||
|
r = await client.put(f"/api/v1/uploads/{sid}/chunks/0", content=data,
|
||||||
|
headers={**h, "content-type": "application/octet-stream",
|
||||||
|
"x-chunk-sha256": "0" * 64})
|
||||||
|
assert r.status_code == 422
|
||||||
|
|
||||||
|
|
||||||
|
async def test_finalize_idempotent_per_client_recording_id(client):
|
||||||
|
token = await user_tokens(client)
|
||||||
|
cid = str(uuid.uuid4())
|
||||||
|
_sid1, r1 = await upload_full(client, token, wav_bytes(64), client_id=cid)
|
||||||
|
_sid2, r2 = await upload_full(client, token, wav_bytes(64), client_id=cid, title="Renamed")
|
||||||
|
assert r1.status_code == 201 and r2.status_code == 201
|
||||||
|
# Same recording id, original preserved, metadata updated.
|
||||||
|
assert r1.json()["id"] == r2.json()["id"]
|
||||||
|
assert r2.json()["title"] == "Renamed"
|
||||||
|
|
||||||
|
|
||||||
|
# --- ownership ---------------------------------------------------------------
|
||||||
|
|
||||||
|
|
||||||
|
async def test_cross_user_isolation(client):
|
||||||
|
ta = await user_tokens(client, "a@example.com")
|
||||||
|
tb = await user_tokens(client, "b@example.com")
|
||||||
|
_sid, r = await upload_full(client, ta, wav_bytes())
|
||||||
|
rec_id = r.json()["id"]
|
||||||
|
r = await client.get(f"/api/v1/recordings/{rec_id}", headers=await auth(tb))
|
||||||
|
assert r.status_code == 404
|
||||||
|
r = await client.get(f"/api/v1/recordings/{rec_id}/audio", headers=await auth(tb))
|
||||||
|
assert r.status_code == 404
|
||||||
|
r = await client.get("/api/v1/recordings", headers=await auth(tb))
|
||||||
|
assert r.json()["total"] == 0
|
||||||
|
|
||||||
|
|
||||||
|
async def test_uploads_require_auth(client):
|
||||||
|
r = await client.get("/api/v1/recordings")
|
||||||
|
assert r.status_code == 401
|
||||||
|
|
||||||
|
|
||||||
|
# --- recordings CRUD -----------------------------------------------------------
|
||||||
|
|
||||||
|
|
||||||
|
async def test_update_metadata_and_tags(client):
|
||||||
|
token = await user_tokens(client)
|
||||||
|
h = await auth(token)
|
||||||
|
_sid, r = await upload_full(client, token, wav_bytes())
|
||||||
|
rec_id = r.json()["id"]
|
||||||
|
r = await client.patch(f"/api/v1/recordings/{rec_id}",
|
||||||
|
json={"title": "Sync meeting", "notes": "n1",
|
||||||
|
"tags": ["Work", " meeting ", "work"]}, headers=h)
|
||||||
|
assert r.status_code == 200
|
||||||
|
body = r.json()
|
||||||
|
assert body["title"] == "Sync meeting"
|
||||||
|
assert body["tags"] == ["meeting", "work"] # normalized, deduped, sorted
|
||||||
|
# listing shows same
|
||||||
|
r = await client.get("/api/v1/recordings", headers=h)
|
||||||
|
assert r.json()["total"] == 1
|
||||||
|
assert r.json()["items"][0]["tags"] == ["meeting", "work"]
|
||||||
|
|
||||||
|
|
||||||
|
async def test_location_dropped_without_consent(client):
|
||||||
|
token = await user_tokens(client)
|
||||||
|
h = await auth(token)
|
||||||
|
_sid, r = await upload_full(client, token, wav_bytes())
|
||||||
|
rec_id = r.json()["id"]
|
||||||
|
r = await client.patch(f"/api/v1/recordings/{rec_id}",
|
||||||
|
json={"latitude": 41.8, "longitude": -87.6}, headers=h)
|
||||||
|
assert r.json()["latitude"] is None
|
||||||
|
# enable consent
|
||||||
|
await client.patch("/api/v1/users/me", json={"location_storage_enabled": True}, headers=h)
|
||||||
|
_sid2, r2 = await upload_full(client, token, wav_bytes(128), client_id=str(uuid.uuid4()))
|
||||||
|
rid2 = r2.json()["id"]
|
||||||
|
r = await client.patch(f"/api/v1/recordings/{rid2}",
|
||||||
|
json={"latitude": 41.8, "longitude": -87.6}, headers=h)
|
||||||
|
assert r.json()["latitude"] == 41.8
|
||||||
|
|
||||||
|
|
||||||
|
async def test_soft_then_purge_delete(client, storage_root):
|
||||||
|
token = await user_tokens(client)
|
||||||
|
h = await auth(token)
|
||||||
|
_sid, r = await upload_full(client, token, wav_bytes())
|
||||||
|
rec_id = r.json()["id"]
|
||||||
|
|
||||||
|
files_before = list(storage_root.rglob("*"))
|
||||||
|
assert any(p.is_file() for p in files_before)
|
||||||
|
|
||||||
|
r = await client.delete(f"/api/v1/recordings/{rec_id}", headers=h)
|
||||||
|
assert r.status_code == 204
|
||||||
|
r = await client.get(f"/api/v1/recordings/{rec_id}", headers=h)
|
||||||
|
assert r.status_code == 404
|
||||||
|
|
||||||
|
# Purge deletes rows AND stored files.
|
||||||
|
token2 = await user_tokens(client, "p2@example.com")
|
||||||
|
h2 = await auth(token2)
|
||||||
|
_sid, r = await upload_full(client, token2, wav_bytes(96))
|
||||||
|
rec2 = r.json()["id"]
|
||||||
|
r = await client.delete(f"/api/v1/recordings/{rec2}?purge=true", headers=h2)
|
||||||
|
assert r.status_code == 204
|
||||||
|
remaining = [p for p in storage_root.rglob("*") if p.is_file() and f"{rec2}" in str(p)]
|
||||||
|
assert remaining == []
|
||||||
|
|
||||||
|
|
||||||
|
async def test_download_audio_roundtrip(client):
|
||||||
|
token = await user_tokens(client)
|
||||||
|
h = await auth(token)
|
||||||
|
data = wav_bytes(128)
|
||||||
|
_sid, r = await upload_full(client, token, data)
|
||||||
|
rec_id = r.json()["id"]
|
||||||
|
r = await client.get(f"/api/v1/recordings/{rec_id}/audio", headers=h)
|
||||||
|
assert r.status_code == 200
|
||||||
|
assert r.content == data
|
||||||
|
assert r.headers["content-type"] == "audio/wav"
|
||||||
|
assert "attachment" in r.headers["content-disposition"]
|
||||||
|
assert "no-store" in r.headers["cache-control"]
|
||||||
189
backend/tests/test_transcript_edits.py
Normal file
189
backend/tests/test_transcript_edits.py
Normal file
|
|
@ -0,0 +1,189 @@
|
||||||
|
"""M8: user edits to transcript/summary via PUT endpoints.
|
||||||
|
|
||||||
|
Edits create a new version with edited_by_user=True; the AI pipeline
|
||||||
|
must not overwrite them afterwards.
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import struct
|
||||||
|
import uuid
|
||||||
|
|
||||||
|
from sqlalchemy import select
|
||||||
|
|
||||||
|
from shonar.db import session as db_session
|
||||||
|
from shonar.db.models import Transcript
|
||||||
|
from shonar.services import processing
|
||||||
|
|
||||||
|
AUTH = {"email": "m8@example.com", "password": "m8-test-passw0rd-123"}
|
||||||
|
|
||||||
|
|
||||||
|
async def user_tokens(client):
|
||||||
|
r = await client.post(
|
||||||
|
"/api/v1/auth/register",
|
||||||
|
json={"email": AUTH["email"], "password": AUTH["password"]},
|
||||||
|
)
|
||||||
|
assert r.status_code == 201, r.text
|
||||||
|
return r.json()["access_token"]
|
||||||
|
|
||||||
|
|
||||||
|
def wav_bytes(payload_len: int = 64) -> bytes:
|
||||||
|
data = bytes(range(payload_len % 256)) * (payload_len // 256 + 1)
|
||||||
|
data = data[:payload_len]
|
||||||
|
header = (
|
||||||
|
b"RIFF" + struct.pack("<I", 36 + len(data)) + b"WAVE"
|
||||||
|
+ b"fmt " + struct.pack("<IHHIIHH", 16, 1, 1, 8000, 8000, 1, 8)
|
||||||
|
+ b"data" + struct.pack("<I", len(data))
|
||||||
|
)
|
||||||
|
return header + data
|
||||||
|
|
||||||
|
|
||||||
|
async def upload_recording(client, token, client_id=None):
|
||||||
|
h = {"Authorization": f"Bearer {token}"}
|
||||||
|
data = wav_bytes()
|
||||||
|
r = await client.post(
|
||||||
|
"/api/v1/uploads",
|
||||||
|
json={"declared_mime_type": "audio/wav", "declared_size_bytes": len(data),
|
||||||
|
"client_recording_id": client_id},
|
||||||
|
headers=h,
|
||||||
|
)
|
||||||
|
assert r.status_code == 201, r.text
|
||||||
|
sid = r.json()["id"]
|
||||||
|
r = await client.put(f"/api/v1/uploads/{sid}/chunks/0", content=data,
|
||||||
|
headers={**h, "content-type": "application/octet-stream"})
|
||||||
|
assert r.status_code == 201, r.text
|
||||||
|
r = await client.post(f"/api/v1/uploads/{sid}/finalize",
|
||||||
|
json={"duration_seconds": 5.0}, headers=h)
|
||||||
|
assert r.status_code == 201, r.text
|
||||||
|
return r.json()
|
||||||
|
|
||||||
|
|
||||||
|
async def test_put_transcript_creates_user_version(client):
|
||||||
|
token = await user_tokens(client)
|
||||||
|
rec = await upload_recording(client, token, client_id="m8-t-1")
|
||||||
|
h = {"Authorization": f"Bearer {token}"}
|
||||||
|
rid = rec["id"]
|
||||||
|
|
||||||
|
r = await client.put(f"/api/v1/recordings/{rid}/transcript",
|
||||||
|
json={"text": "hello corrected"}, headers=h)
|
||||||
|
assert r.status_code == 200, r.text
|
||||||
|
body = r.json()
|
||||||
|
assert body["version"] == 1
|
||||||
|
assert body["text"] == "hello corrected"
|
||||||
|
assert body["edited_by_user"] is True
|
||||||
|
assert body["provider"] == "user"
|
||||||
|
|
||||||
|
r = await client.get(f"/api/v1/recordings/{rid}/transcript", headers=h)
|
||||||
|
assert r.status_code == 200
|
||||||
|
assert r.json()["text"] == "hello corrected"
|
||||||
|
|
||||||
|
# Second edit supersedes the first.
|
||||||
|
r = await client.put(f"/api/v1/recordings/{rid}/transcript",
|
||||||
|
json={"text": "second pass",
|
||||||
|
"segments": [{"start": 0.0, "end": 1.0, "text": "second pass"}]},
|
||||||
|
headers=h)
|
||||||
|
assert r.status_code == 200, r.text
|
||||||
|
assert r.json()["version"] == 2
|
||||||
|
assert r.json()["segments"][0]["text"] == "second pass"
|
||||||
|
|
||||||
|
async with db_session._session_factory() as s:
|
||||||
|
rows = (await s.scalars(
|
||||||
|
select(Transcript).where(Transcript.recording_id == uuid.UUID(rid))
|
||||||
|
.order_by(Transcript.version))).all()
|
||||||
|
assert [(r.version, r.superseded_at is not None, r.edited_by_user) for r in rows] == [
|
||||||
|
(1, True, True), (2, False, True)]
|
||||||
|
|
||||||
|
|
||||||
|
async def test_put_summary_creates_user_version(client):
|
||||||
|
token = await user_tokens(client)
|
||||||
|
rec = await upload_recording(client, token, client_id="m8-s-1")
|
||||||
|
h = {"Authorization": f"Bearer {token}"}
|
||||||
|
rid = rec["id"]
|
||||||
|
|
||||||
|
content = {"short": "s", "action_items": ["ship it"]}
|
||||||
|
r = await client.put(f"/api/v1/recordings/{rid}/summary",
|
||||||
|
json={"content": content}, headers=h)
|
||||||
|
assert r.status_code == 200, r.text
|
||||||
|
body = r.json()
|
||||||
|
assert body["version"] == 1
|
||||||
|
assert body["content"] == content
|
||||||
|
assert body["edited_by_user"] is True
|
||||||
|
|
||||||
|
r = await client.get(f"/api/v1/recordings/{rid}/summary", headers=h)
|
||||||
|
assert r.status_code == 200
|
||||||
|
assert r.json()["content"]["action_items"] == ["ship it"]
|
||||||
|
|
||||||
|
|
||||||
|
async def test_user_edit_survives_auto_pipeline(client, monkeypatch):
|
||||||
|
"""An auto transcribe run after a user edit inserts nothing."""
|
||||||
|
from shonar.services.ai import Segment as AiSegment
|
||||||
|
from shonar.services.ai import SummaryResult, TranscriptResult
|
||||||
|
|
||||||
|
class FakeTranscriber:
|
||||||
|
name = "fake-stt"
|
||||||
|
|
||||||
|
async def transcribe(self, audio, mime, *, language_hint=None):
|
||||||
|
return TranscriptResult(
|
||||||
|
text="auto text", language="en",
|
||||||
|
segments=[AiSegment(0.0, 1.0, "auto text")], model="fake-stt-1",
|
||||||
|
)
|
||||||
|
|
||||||
|
class FakeLlm:
|
||||||
|
name = "fake-llm"
|
||||||
|
|
||||||
|
async def summarize(self, transcript, *, title=None):
|
||||||
|
return SummaryResult(short="s", detailed="d", key_points=(),
|
||||||
|
decisions=(), action_items=(), questions=(),
|
||||||
|
model="fake-llm-1")
|
||||||
|
|
||||||
|
monkeypatch.setattr(
|
||||||
|
processing, "get_transcription_provider", lambda settings: FakeTranscriber())
|
||||||
|
monkeypatch.setattr(processing, "get_llm_provider", lambda settings: FakeLlm())
|
||||||
|
|
||||||
|
token = await user_tokens(client)
|
||||||
|
rec = await upload_recording(client, token, client_id="m8-t-2")
|
||||||
|
h = {"Authorization": f"Bearer {token}"}
|
||||||
|
rid = rec["id"]
|
||||||
|
|
||||||
|
r = await client.put(f"/api/v1/recordings/{rid}/transcript",
|
||||||
|
json={"text": "user verdict"}, headers=h)
|
||||||
|
assert r.status_code == 200, r.text
|
||||||
|
|
||||||
|
monkeypatch.setattr(
|
||||||
|
processing, "get_transcription_provider", lambda settings: FakeTranscriber())
|
||||||
|
monkeypatch.setattr(processing, "get_llm_provider", lambda settings: FakeLlm())
|
||||||
|
await processing.run_transcribe({}, rid)
|
||||||
|
|
||||||
|
r = await client.get(f"/api/v1/recordings/{rid}/transcript", headers=h)
|
||||||
|
assert r.status_code == 200
|
||||||
|
assert r.json()["text"] == "user verdict"
|
||||||
|
|
||||||
|
|
||||||
|
async def test_edit_endpoints_enforce_ownership(client):
|
||||||
|
token = await user_tokens(client)
|
||||||
|
rec = await upload_recording(client, token, client_id="m8-t-3")
|
||||||
|
rid = rec["id"]
|
||||||
|
|
||||||
|
r = await client.post("/api/v1/auth/register",
|
||||||
|
json={"email": "m8-other@example.com", "password": "m8-other-passw0rd-1"})
|
||||||
|
assert r.status_code == 201, r.text
|
||||||
|
other = r.json()["access_token"]
|
||||||
|
h2 = {"Authorization": f"Bearer {other}"}
|
||||||
|
|
||||||
|
r = await client.put(f"/api/v1/recordings/{rid}/transcript",
|
||||||
|
json={"text": "hijack"}, headers=h2)
|
||||||
|
assert r.status_code == 404
|
||||||
|
r = await client.put(f"/api/v1/recordings/{rid}/summary",
|
||||||
|
json={"content": {}}, headers=h2)
|
||||||
|
assert r.status_code == 404
|
||||||
|
|
||||||
|
|
||||||
|
async def test_edit_validation(client):
|
||||||
|
token = await user_tokens(client)
|
||||||
|
rec = await upload_recording(client, token, client_id="m8-t-4")
|
||||||
|
h = {"Authorization": f"Bearer {token}"}
|
||||||
|
rid = rec["id"]
|
||||||
|
|
||||||
|
r = await client.put(f"/api/v1/recordings/{rid}/transcript",
|
||||||
|
json={"text": ""}, headers=h)
|
||||||
|
assert r.status_code == 422
|
||||||
60
desktop-dev.sh
Executable file
60
desktop-dev.sh
Executable file
|
|
@ -0,0 +1,60 @@
|
||||||
|
#!/usr/bin/env bash
|
||||||
|
# SHONAR desktop dev loop: runs the app and auto-restarts it when any
|
||||||
|
# Kotlin/Gradle source changes (inotifywait). Engine runs in dev mode
|
||||||
|
# (SHONAR_DEV=1 -> uvicorn --reload), so backend edits apply live too.
|
||||||
|
#
|
||||||
|
# Stop with Ctrl-C (or pkill -f desktop-dev.sh).
|
||||||
|
set -uo pipefail
|
||||||
|
cd "$(dirname "$0")"
|
||||||
|
|
||||||
|
export SHONAR_DEV=1
|
||||||
|
export DISPLAY="${DISPLAY:-:0}"
|
||||||
|
|
||||||
|
FLAG="$(pwd)/.dev-restart-flag"
|
||||||
|
|
||||||
|
# Watch targets resolved to absolute paths, non-existent dirs skipped at
|
||||||
|
# start but retried each cycle (inotifywait -r tolerates a dir appearing).
|
||||||
|
watch_targets() {
|
||||||
|
local t=()
|
||||||
|
for d in "$PWD/app/src/main/kotlin" \
|
||||||
|
"$PWD/shared/com/shonar" \
|
||||||
|
"$PWD/backend/shonar"; do
|
||||||
|
[ -d "$d" ] && t+=("$d")
|
||||||
|
done
|
||||||
|
printf '%s\n' "${t[@]}"
|
||||||
|
}
|
||||||
|
|
||||||
|
pkill -f "com.shonar.desktop.MainKt" 2>/dev/null || true
|
||||||
|
rm -f "$FLAG"
|
||||||
|
sleep 1
|
||||||
|
|
||||||
|
(
|
||||||
|
while true; do
|
||||||
|
mapfile -t TARGETS < <(watch_targets)
|
||||||
|
[ ${#TARGETS[@]} -eq 0 ] && { sleep 5; continue; }
|
||||||
|
if inotifywait -qq -e close_write,move --include '\.kt$|\.kts$' \
|
||||||
|
-r "${TARGETS[@]}"; then
|
||||||
|
echo "[dev] source change -> restarting app…"
|
||||||
|
touch "$FLAG"
|
||||||
|
pkill -f "com.shonar.desktop.MainKt" 2>/dev/null || true
|
||||||
|
sleep 2 # debounce bursts (editor writes can fire several events)
|
||||||
|
else
|
||||||
|
sleep 5 # targets vanished; retry with a fresh list
|
||||||
|
fi
|
||||||
|
done
|
||||||
|
) &
|
||||||
|
WATCHER=$!
|
||||||
|
trap 'kill $WATCHER 2>/dev/null || true; rm -f "$FLAG"' EXIT
|
||||||
|
|
||||||
|
while true; do
|
||||||
|
rm -f "$FLAG"
|
||||||
|
./gradlew run --offline
|
||||||
|
if [ -f "$FLAG" ]; then
|
||||||
|
continue # killed by a source edit: rebuild now
|
||||||
|
fi
|
||||||
|
rm -f "$FLAG"
|
||||||
|
# App closed cleanly or build failed — wait for the next edit.
|
||||||
|
mapfile -t TARGETS < <(watch_targets)
|
||||||
|
[ ${#TARGETS[@]} -eq 0 ] && { sleep 5; continue; }
|
||||||
|
inotifywait -qq -e close_write,move --include '\.kt$|\.kts$' -r "${TARGETS[@]}" || sleep 5
|
||||||
|
done
|
||||||
BIN
gradle/wrapper/gradle-wrapper.jar
vendored
Normal file
BIN
gradle/wrapper/gradle-wrapper.jar
vendored
Normal file
Binary file not shown.
7
gradle/wrapper/gradle-wrapper.properties
vendored
Normal file
7
gradle/wrapper/gradle-wrapper.properties
vendored
Normal file
|
|
@ -0,0 +1,7 @@
|
||||||
|
distributionBase=GRADLE_USER_HOME
|
||||||
|
distributionPath=wrapper/dists
|
||||||
|
distributionUrl=https\://services.gradle.org/distributions/gradle-8.11.1-bin.zip
|
||||||
|
networkTimeout=10000
|
||||||
|
validateDistributionUrl=true
|
||||||
|
zipStoreBase=GRADLE_USER_HOME
|
||||||
|
zipStorePath=wrapper/dists
|
||||||
251
gradlew
vendored
Executable file
251
gradlew
vendored
Executable file
|
|
@ -0,0 +1,251 @@
|
||||||
|
#!/bin/sh
|
||||||
|
|
||||||
|
#
|
||||||
|
# Copyright © 2015-2021 the original authors.
|
||||||
|
#
|
||||||
|
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||||
|
# you may not use this file except in compliance with the License.
|
||||||
|
# You may obtain a copy of the License at
|
||||||
|
#
|
||||||
|
# https://www.apache.org/licenses/LICENSE-2.0
|
||||||
|
#
|
||||||
|
# Unless required by applicable law or agreed to in writing, software
|
||||||
|
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||||
|
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||||
|
# See the License for the specific language governing permissions and
|
||||||
|
# limitations under the License.
|
||||||
|
#
|
||||||
|
# SPDX-License-Identifier: Apache-2.0
|
||||||
|
#
|
||||||
|
|
||||||
|
##############################################################################
|
||||||
|
#
|
||||||
|
# Gradle start up script for POSIX generated by Gradle.
|
||||||
|
#
|
||||||
|
# Important for running:
|
||||||
|
#
|
||||||
|
# (1) You need a POSIX-compliant shell to run this script. If your /bin/sh is
|
||||||
|
# noncompliant, but you have some other compliant shell such as ksh or
|
||||||
|
# bash, then to run this script, type that shell name before the whole
|
||||||
|
# command line, like:
|
||||||
|
#
|
||||||
|
# ksh Gradle
|
||||||
|
#
|
||||||
|
# Busybox and similar reduced shells will NOT work, because this script
|
||||||
|
# requires all of these POSIX shell features:
|
||||||
|
# * functions;
|
||||||
|
# * expansions «$var», «${var}», «${var:-default}», «${var+SET}»,
|
||||||
|
# «${var#prefix}», «${var%suffix}», and «$( cmd )»;
|
||||||
|
# * compound commands having a testable exit status, especially «case»;
|
||||||
|
# * various built-in commands including «command», «set», and «ulimit».
|
||||||
|
#
|
||||||
|
# Important for patching:
|
||||||
|
#
|
||||||
|
# (2) This script targets any POSIX shell, so it avoids extensions provided
|
||||||
|
# by Bash, Ksh, etc; in particular arrays are avoided.
|
||||||
|
#
|
||||||
|
# The "traditional" practice of packing multiple parameters into a
|
||||||
|
# space-separated string is a well documented source of bugs and security
|
||||||
|
# problems, so this is (mostly) avoided, by progressively accumulating
|
||||||
|
# options in "$@", and eventually passing that to Java.
|
||||||
|
#
|
||||||
|
# Where the inherited environment variables (DEFAULT_JVM_OPTS, JAVA_OPTS,
|
||||||
|
# and GRADLE_OPTS) rely on word-splitting, this is performed explicitly;
|
||||||
|
# see the in-line comments for details.
|
||||||
|
#
|
||||||
|
# There are tweaks for specific operating systems such as AIX, CygWin,
|
||||||
|
# Darwin, MinGW, and NonStop.
|
||||||
|
#
|
||||||
|
# (3) This script is generated from the Groovy template
|
||||||
|
# https://github.com/gradle/gradle/blob/HEAD/platforms/jvm/plugins-application/src/main/resources/org/gradle/api/internal/plugins/unixStartScript.txt
|
||||||
|
# within the Gradle project.
|
||||||
|
#
|
||||||
|
# You can find Gradle at https://github.com/gradle/gradle/.
|
||||||
|
#
|
||||||
|
##############################################################################
|
||||||
|
|
||||||
|
# Attempt to set APP_HOME
|
||||||
|
|
||||||
|
# Resolve links: $0 may be a link
|
||||||
|
app_path=$0
|
||||||
|
|
||||||
|
# Need this for daisy-chained symlinks.
|
||||||
|
while
|
||||||
|
APP_HOME=${app_path%"${app_path##*/}"} # leaves a trailing /; empty if no leading path
|
||||||
|
[ -h "$app_path" ]
|
||||||
|
do
|
||||||
|
ls=$( ls -ld "$app_path" )
|
||||||
|
link=${ls#*' -> '}
|
||||||
|
case $link in #(
|
||||||
|
/*) app_path=$link ;; #(
|
||||||
|
*) app_path=$APP_HOME$link ;;
|
||||||
|
esac
|
||||||
|
done
|
||||||
|
|
||||||
|
# This is normally unused
|
||||||
|
# shellcheck disable=SC2034
|
||||||
|
APP_BASE_NAME=${0##*/}
|
||||||
|
# Discard cd standard output in case $CDPATH is set (https://github.com/gradle/gradle/issues/25036)
|
||||||
|
APP_HOME=$( cd -P "${APP_HOME:-./}" > /dev/null && printf '%s\n' "$PWD" ) || exit
|
||||||
|
|
||||||
|
# Use the maximum available, or set MAX_FD != -1 to use that value.
|
||||||
|
MAX_FD=maximum
|
||||||
|
|
||||||
|
warn () {
|
||||||
|
echo "$*"
|
||||||
|
} >&2
|
||||||
|
|
||||||
|
die () {
|
||||||
|
echo
|
||||||
|
echo "$*"
|
||||||
|
echo
|
||||||
|
exit 1
|
||||||
|
} >&2
|
||||||
|
|
||||||
|
# OS specific support (must be 'true' or 'false').
|
||||||
|
cygwin=false
|
||||||
|
msys=false
|
||||||
|
darwin=false
|
||||||
|
nonstop=false
|
||||||
|
case "$( uname )" in #(
|
||||||
|
CYGWIN* ) cygwin=true ;; #(
|
||||||
|
Darwin* ) darwin=true ;; #(
|
||||||
|
MSYS* | MINGW* ) msys=true ;; #(
|
||||||
|
NONSTOP* ) nonstop=true ;;
|
||||||
|
esac
|
||||||
|
|
||||||
|
CLASSPATH="\\\"\\\""
|
||||||
|
|
||||||
|
|
||||||
|
# Determine the Java command to use to start the JVM.
|
||||||
|
if [ -n "$JAVA_HOME" ] ; then
|
||||||
|
if [ -x "$JAVA_HOME/jre/sh/java" ] ; then
|
||||||
|
# IBM's JDK on AIX uses strange locations for the executables
|
||||||
|
JAVACMD=$JAVA_HOME/jre/sh/java
|
||||||
|
else
|
||||||
|
JAVACMD=$JAVA_HOME/bin/java
|
||||||
|
fi
|
||||||
|
if [ ! -x "$JAVACMD" ] ; then
|
||||||
|
die "ERROR: JAVA_HOME is set to an invalid directory: $JAVA_HOME
|
||||||
|
|
||||||
|
Please set the JAVA_HOME variable in your environment to match the
|
||||||
|
location of your Java installation."
|
||||||
|
fi
|
||||||
|
else
|
||||||
|
JAVACMD=java
|
||||||
|
if ! command -v java >/dev/null 2>&1
|
||||||
|
then
|
||||||
|
die "ERROR: JAVA_HOME is not set and no 'java' command could be found in your PATH.
|
||||||
|
|
||||||
|
Please set the JAVA_HOME variable in your environment to match the
|
||||||
|
location of your Java installation."
|
||||||
|
fi
|
||||||
|
fi
|
||||||
|
|
||||||
|
# Increase the maximum file descriptors if we can.
|
||||||
|
if ! "$cygwin" && ! "$darwin" && ! "$nonstop" ; then
|
||||||
|
case $MAX_FD in #(
|
||||||
|
max*)
|
||||||
|
# In POSIX sh, ulimit -H is undefined. That's why the result is checked to see if it worked.
|
||||||
|
# shellcheck disable=SC2039,SC3045
|
||||||
|
MAX_FD=$( ulimit -H -n ) ||
|
||||||
|
warn "Could not query maximum file descriptor limit"
|
||||||
|
esac
|
||||||
|
case $MAX_FD in #(
|
||||||
|
'' | soft) :;; #(
|
||||||
|
*)
|
||||||
|
# In POSIX sh, ulimit -n is undefined. That's why the result is checked to see if it worked.
|
||||||
|
# shellcheck disable=SC2039,SC3045
|
||||||
|
ulimit -n "$MAX_FD" ||
|
||||||
|
warn "Could not set maximum file descriptor limit to $MAX_FD"
|
||||||
|
esac
|
||||||
|
fi
|
||||||
|
|
||||||
|
# Collect all arguments for the java command, stacking in reverse order:
|
||||||
|
# * args from the command line
|
||||||
|
# * the main class name
|
||||||
|
# * -classpath
|
||||||
|
# * -D...appname settings
|
||||||
|
# * --module-path (only if needed)
|
||||||
|
# * DEFAULT_JVM_OPTS, JAVA_OPTS, and GRADLE_OPTS environment variables.
|
||||||
|
|
||||||
|
# For Cygwin or MSYS, switch paths to Windows format before running java
|
||||||
|
if "$cygwin" || "$msys" ; then
|
||||||
|
APP_HOME=$( cygpath --path --mixed "$APP_HOME" )
|
||||||
|
CLASSPATH=$( cygpath --path --mixed "$CLASSPATH" )
|
||||||
|
|
||||||
|
JAVACMD=$( cygpath --unix "$JAVACMD" )
|
||||||
|
|
||||||
|
# Now convert the arguments - kludge to limit ourselves to /bin/sh
|
||||||
|
for arg do
|
||||||
|
if
|
||||||
|
case $arg in #(
|
||||||
|
-*) false ;; # don't mess with options #(
|
||||||
|
/?*) t=${arg#/} t=/${t%%/*} # looks like a POSIX filepath
|
||||||
|
[ -e "$t" ] ;; #(
|
||||||
|
*) false ;;
|
||||||
|
esac
|
||||||
|
then
|
||||||
|
arg=$( cygpath --path --ignore --mixed "$arg" )
|
||||||
|
fi
|
||||||
|
# Roll the args list around exactly as many times as the number of
|
||||||
|
# args, so each arg winds up back in the position where it started, but
|
||||||
|
# possibly modified.
|
||||||
|
#
|
||||||
|
# NB: a `for` loop captures its iteration list before it begins, so
|
||||||
|
# changing the positional parameters here affects neither the number of
|
||||||
|
# iterations, nor the values presented in `arg`.
|
||||||
|
shift # remove old arg
|
||||||
|
set -- "$@" "$arg" # push replacement arg
|
||||||
|
done
|
||||||
|
fi
|
||||||
|
|
||||||
|
|
||||||
|
# Add default JVM options here. You can also use JAVA_OPTS and GRADLE_OPTS to pass JVM options to this script.
|
||||||
|
DEFAULT_JVM_OPTS='"-Xmx64m" "-Xms64m"'
|
||||||
|
|
||||||
|
# Collect all arguments for the java command:
|
||||||
|
# * DEFAULT_JVM_OPTS, JAVA_OPTS, and optsEnvironmentVar are not allowed to contain shell fragments,
|
||||||
|
# and any embedded shellness will be escaped.
|
||||||
|
# * For example: A user cannot expect ${Hostname} to be expanded, as it is an environment variable and will be
|
||||||
|
# treated as '${Hostname}' itself on the command line.
|
||||||
|
|
||||||
|
set -- \
|
||||||
|
"-Dorg.gradle.appname=$APP_BASE_NAME" \
|
||||||
|
-classpath "$CLASSPATH" \
|
||||||
|
-jar "$APP_HOME/gradle/wrapper/gradle-wrapper.jar" \
|
||||||
|
"$@"
|
||||||
|
|
||||||
|
# Stop when "xargs" is not available.
|
||||||
|
if ! command -v xargs >/dev/null 2>&1
|
||||||
|
then
|
||||||
|
die "xargs is not available"
|
||||||
|
fi
|
||||||
|
|
||||||
|
# Use "xargs" to parse quoted args.
|
||||||
|
#
|
||||||
|
# With -n1 it outputs one arg per line, with the quotes and backslashes removed.
|
||||||
|
#
|
||||||
|
# In Bash we could simply go:
|
||||||
|
#
|
||||||
|
# readarray ARGS < <( xargs -n1 <<<"$var" ) &&
|
||||||
|
# set -- "${ARGS[@]}" "$@"
|
||||||
|
#
|
||||||
|
# but POSIX shell has neither arrays nor command substitution, so instead we
|
||||||
|
# post-process each arg (as a line of input to sed) to backslash-escape any
|
||||||
|
# character that might be a shell metacharacter, then use eval to reverse
|
||||||
|
# that process (while maintaining the separation between arguments), and wrap
|
||||||
|
# the whole thing up as a single "set" statement.
|
||||||
|
#
|
||||||
|
# This will of course break if any of these variables contains a newline or
|
||||||
|
# an unmatched quote.
|
||||||
|
#
|
||||||
|
|
||||||
|
eval "set -- $(
|
||||||
|
printf '%s\n' "$DEFAULT_JVM_OPTS $JAVA_OPTS $GRADLE_OPTS" |
|
||||||
|
xargs -n1 |
|
||||||
|
sed ' s~[^-[:alnum:]+,./:=@_]~\\&~g; ' |
|
||||||
|
tr '\n' ' '
|
||||||
|
)" '"$@"'
|
||||||
|
|
||||||
|
exec "$JAVACMD" "$@"
|
||||||
11
launch-desktop.sh
Executable file
11
launch-desktop.sh
Executable file
|
|
@ -0,0 +1,11 @@
|
||||||
|
#!/usr/bin/env bash
|
||||||
|
# Launch the SHONAR desktop app (packaged binary, no Gradle needed).
|
||||||
|
set -euo pipefail
|
||||||
|
cd "$(dirname "$0")"
|
||||||
|
BIN="app/build/compose/binaries/main/app/shonar-desktop/bin/shonar-desktop"
|
||||||
|
if [ ! -x "$BIN" ]; then
|
||||||
|
echo "Desktop build not found. Run: ./gradlew :app:createDistributable" >&2
|
||||||
|
exit 1
|
||||||
|
fi
|
||||||
|
export DISPLAY="${DISPLAY:-:0}"
|
||||||
|
exec "$BIN"
|
||||||
17
settings.gradle.kts
Normal file
17
settings.gradle.kts
Normal file
|
|
@ -0,0 +1,17 @@
|
||||||
|
pluginManagement {
|
||||||
|
repositories {
|
||||||
|
mavenCentral()
|
||||||
|
gradlePluginPortal()
|
||||||
|
google()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
dependencyResolutionManagement {
|
||||||
|
repositories {
|
||||||
|
mavenCentral()
|
||||||
|
google()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
rootProject.name = "shonar-desktop"
|
||||||
|
include(":app")
|
||||||
176
shared/com/shonar/MainActivity.kt
Normal file
176
shared/com/shonar/MainActivity.kt
Normal file
|
|
@ -0,0 +1,176 @@
|
||||||
|
package com.shonar
|
||||||
|
|
||||||
|
import android.content.Intent
|
||||||
|
import android.os.Bundle
|
||||||
|
import androidx.activity.ComponentActivity
|
||||||
|
import androidx.activity.compose.setContent
|
||||||
|
import androidx.activity.result.contract.ActivityResultContracts
|
||||||
|
import androidx.compose.foundation.layout.fillMaxSize
|
||||||
|
import androidx.compose.material3.MaterialTheme
|
||||||
|
import androidx.compose.material3.Surface
|
||||||
|
import androidx.compose.runtime.collectAsState
|
||||||
|
import androidx.compose.runtime.getValue
|
||||||
|
import androidx.compose.runtime.mutableStateOf
|
||||||
|
import androidx.compose.runtime.remember
|
||||||
|
import androidx.compose.runtime.setValue
|
||||||
|
import androidx.compose.ui.Modifier
|
||||||
|
import androidx.navigation.compose.NavHost
|
||||||
|
import androidx.navigation.compose.composable
|
||||||
|
import androidx.navigation.compose.rememberNavController
|
||||||
|
import androidx.navigation.navArgument
|
||||||
|
import com.shonar.recording.RecordingFiles
|
||||||
|
import com.shonar.recording.RecordingService
|
||||||
|
import com.shonar.recording.RenamePendingHolder
|
||||||
|
import com.shonar.ui.detail.DetailsScreen
|
||||||
|
import com.shonar.ui.folder.FolderBrowserScreen
|
||||||
|
import com.shonar.ui.home.HomeScreen
|
||||||
|
import com.shonar.ui.provider.ProviderSelectionScreen
|
||||||
|
import com.shonar.ui.provider.StorageScreen
|
||||||
|
import com.shonar.ui.rename.RenameAfterSaveDialog
|
||||||
|
import com.shonar.ui.settings.SettingsScreen
|
||||||
|
import com.shonar.ui.theme.ShonarTheme
|
||||||
|
|
||||||
|
class MainActivity : ComponentActivity() {
|
||||||
|
|
||||||
|
private val notifPermLauncher =
|
||||||
|
registerForActivityResult(ActivityResultContracts.RequestPermission()) { }
|
||||||
|
|
||||||
|
/** Recording id awaiting rename after a service SAVE. Null = no sheet. */
|
||||||
|
private val pendingRename = mutableStateOf<String?>(null)
|
||||||
|
|
||||||
|
override fun onCreate(savedInstanceState: Bundle?) {
|
||||||
|
super.onCreate(savedInstanceState)
|
||||||
|
if (android.os.Build.VERSION.SDK_INT >= 33 &&
|
||||||
|
androidx.core.content.ContextCompat.checkSelfPermission(
|
||||||
|
this, android.Manifest.permission.POST_NOTIFICATIONS,
|
||||||
|
) != android.content.pm.PackageManager.PERMISSION_GRANTED
|
||||||
|
) {
|
||||||
|
runCatching { notifPermLauncher.launch(android.Manifest.permission.POST_NOTIFICATIONS) }
|
||||||
|
}
|
||||||
|
routeIntent(intent)
|
||||||
|
setContent {
|
||||||
|
ShonarTheme {
|
||||||
|
val nav = rememberNavController()
|
||||||
|
Surface(
|
||||||
|
modifier = Modifier.fillMaxSize(),
|
||||||
|
color = MaterialTheme.colorScheme.background,
|
||||||
|
) {
|
||||||
|
NavHost(navController = nav, startDestination = "home") {
|
||||||
|
composable("home") {
|
||||||
|
HomeScreen(
|
||||||
|
onOpenSettings = { nav.navigate("settings") },
|
||||||
|
onOpenStorage = { nav.navigate("storage") },
|
||||||
|
onOpenDetail = { id -> nav.navigate("detail/$id") },
|
||||||
|
)
|
||||||
|
}
|
||||||
|
composable(
|
||||||
|
"detail/{id}",
|
||||||
|
arguments = listOf(navArgument("id") { type = androidx.navigation.NavType.StringType }),
|
||||||
|
) { backStack ->
|
||||||
|
DetailsScreen(
|
||||||
|
recordingId = backStack.arguments?.getString("id").orEmpty(),
|
||||||
|
onBack = { nav.popBackStack() },
|
||||||
|
)
|
||||||
|
}
|
||||||
|
composable("storage") {
|
||||||
|
StorageScreen(
|
||||||
|
onBack = { nav.popBackStack() },
|
||||||
|
onSwitchProvider = { nav.navigate("provider") },
|
||||||
|
)
|
||||||
|
}
|
||||||
|
composable("provider") { entry ->
|
||||||
|
val pick by entry.savedStateHandle
|
||||||
|
.getStateFlow<String?>("picked_path", null)
|
||||||
|
.collectAsState()
|
||||||
|
ProviderSelectionScreen(
|
||||||
|
onDone = { nav.popBackStack() },
|
||||||
|
onBack = { nav.popBackStack() },
|
||||||
|
browsePick = pick,
|
||||||
|
onPickConsumed = {
|
||||||
|
entry.savedStateHandle.remove<String>("picked_path")
|
||||||
|
},
|
||||||
|
onBrowse = { nav.navigate("folderBrowser") },
|
||||||
|
)
|
||||||
|
}
|
||||||
|
composable("folderBrowser") {
|
||||||
|
FolderBrowserScreen(
|
||||||
|
onPick = { path ->
|
||||||
|
nav.previousBackStackEntry
|
||||||
|
?.savedStateHandle
|
||||||
|
?.set("picked_path", path)
|
||||||
|
nav.popBackStack()
|
||||||
|
},
|
||||||
|
onBack = { nav.popBackStack() },
|
||||||
|
)
|
||||||
|
}
|
||||||
|
composable("settings") { SettingsScreen(onBack = { nav.popBackStack() }) }
|
||||||
|
}
|
||||||
|
|
||||||
|
// Post-save rename sheet: opens automatically after Save.
|
||||||
|
// Cancel keeps the file under its default name.
|
||||||
|
val renameId = pendingRename.value
|
||||||
|
if (renameId != null) {
|
||||||
|
val app = applicationContext as ShonarApplication
|
||||||
|
val recordings by app.recordingRepository.recordings
|
||||||
|
.collectAsState(initial = emptyList())
|
||||||
|
val row = recordings.firstOrNull { it.id == renameId }
|
||||||
|
val defaultName = row?.title
|
||||||
|
?: RecordingFiles.defaultDisplayName(System.currentTimeMillis())
|
||||||
|
var err by androidx.compose.runtime.remember(renameId) {
|
||||||
|
androidx.compose.runtime.mutableStateOf<String?>(null)
|
||||||
|
}
|
||||||
|
val vm: com.shonar.ui.rename.RenameAfterSaveViewModel =
|
||||||
|
androidx.lifecycle.viewmodel.compose.viewModel(
|
||||||
|
factory = object : androidx.lifecycle.ViewModelProvider.Factory {
|
||||||
|
@Suppress("UNCHECKED_CAST")
|
||||||
|
override fun <T : androidx.lifecycle.ViewModel> create(
|
||||||
|
modelClass: Class<T>,
|
||||||
|
): T = com.shonar.ui.rename.RenameAfterSaveViewModel(app) as T
|
||||||
|
},
|
||||||
|
)
|
||||||
|
RenameAfterSaveDialog(
|
||||||
|
defaultName = defaultName,
|
||||||
|
error = err,
|
||||||
|
onConfirm = { name ->
|
||||||
|
val stem = RecordingFiles.sanitizeStem(name)
|
||||||
|
if (stem == null) {
|
||||||
|
err = "Enter a valid name."
|
||||||
|
} else {
|
||||||
|
vm.confirm(renameId, stem) { pendingRename.value = null }
|
||||||
|
}
|
||||||
|
},
|
||||||
|
onCancel = { pendingRename.value = null },
|
||||||
|
)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
override fun onNewIntent(intent: Intent) {
|
||||||
|
super.onNewIntent(intent)
|
||||||
|
setIntent(intent)
|
||||||
|
routeIntent(intent)
|
||||||
|
}
|
||||||
|
|
||||||
|
private fun routeIntent(intent: Intent?) {
|
||||||
|
// Widget without RECORD_AUDIO routes here to request permission first.
|
||||||
|
if (intent?.action == RecordingService.ACTION_START) {
|
||||||
|
if (androidx.core.content.ContextCompat.checkSelfPermission(
|
||||||
|
this, android.Manifest.permission.RECORD_AUDIO,
|
||||||
|
) == android.content.pm.PackageManager.PERMISSION_GRANTED
|
||||||
|
) {
|
||||||
|
RecordingService.command(this, RecordingService.ACTION_START)
|
||||||
|
}
|
||||||
|
intent.action = Intent.ACTION_MAIN
|
||||||
|
}
|
||||||
|
val renameId = intent?.takeIf { it.action == RecordingService.ACTION_RENAME }
|
||||||
|
?.getStringExtra(RecordingService.EXTRA_RECORDING_ID)
|
||||||
|
?: RenamePendingHolder.pendingRenameId
|
||||||
|
if (intent?.action == RecordingService.ACTION_RENAME) {
|
||||||
|
RenamePendingHolder.pendingRenameId = null
|
||||||
|
intent.action = Intent.ACTION_MAIN
|
||||||
|
}
|
||||||
|
if (renameId != null) pendingRename.value = renameId
|
||||||
|
}
|
||||||
|
}
|
||||||
128
shared/com/shonar/ShonarApplication.kt
Normal file
128
shared/com/shonar/ShonarApplication.kt
Normal file
|
|
@ -0,0 +1,128 @@
|
||||||
|
package com.shonar
|
||||||
|
|
||||||
|
import android.app.Application
|
||||||
|
import androidx.room.Room
|
||||||
|
import com.shonar.recording.RecordingRepository
|
||||||
|
import com.shonar.recording.ShonarDatabase
|
||||||
|
import com.shonar.settings.DataStoreSettingsStore
|
||||||
|
import com.shonar.settings.SecureSettingsStore
|
||||||
|
import com.shonar.settings.SettingsManager
|
||||||
|
import java.util.concurrent.atomic.AtomicBoolean
|
||||||
|
import kotlinx.coroutines.CoroutineScope
|
||||||
|
import kotlinx.coroutines.Dispatchers
|
||||||
|
import kotlinx.coroutines.SupervisorJob
|
||||||
|
import kotlinx.coroutines.flow.collect
|
||||||
|
import kotlinx.coroutines.launch
|
||||||
|
|
||||||
|
class ShonarApplication : Application() {
|
||||||
|
|
||||||
|
private val appScope = CoroutineScope(SupervisorJob() + Dispatchers.IO)
|
||||||
|
|
||||||
|
val database: ShonarDatabase by lazy {
|
||||||
|
Room.databaseBuilder(this, ShonarDatabase::class.java, "shonar.db")
|
||||||
|
.addMigrations(
|
||||||
|
com.shonar.recording.MIGRATION_1_2,
|
||||||
|
com.shonar.recording.MIGRATION_2_3,
|
||||||
|
)
|
||||||
|
.build()
|
||||||
|
}
|
||||||
|
|
||||||
|
val recordingRepository: RecordingRepository by lazy {
|
||||||
|
RecordingRepository(filesDir, database.recordingDao()) {
|
||||||
|
// New recordings belong in the picked folder (local-only
|
||||||
|
// provider), or the app-private library by default / for
|
||||||
|
// server-backed setups. Never throws — null means default.
|
||||||
|
runCatching {
|
||||||
|
settingsManager.ensureLoaded()
|
||||||
|
val pid = settingsManager.string(
|
||||||
|
com.shonar.settings.BuiltInSettings.PROVIDER_ID
|
||||||
|
).ifBlank { "local-only" }
|
||||||
|
(providerRegistry.provider(pid) as? com.shonar.provider.LocalOnlyProvider)
|
||||||
|
?.currentRoot()
|
||||||
|
}.getOrNull()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
val secureStore: SecureSettingsStore by lazy { SecureSettingsStore(this) }
|
||||||
|
|
||||||
|
val settingsManager: SettingsManager by lazy {
|
||||||
|
SettingsManager(
|
||||||
|
store = DataStoreSettingsStore(this),
|
||||||
|
secureStore = secureStore,
|
||||||
|
)
|
||||||
|
}
|
||||||
|
|
||||||
|
/** P5: trust-on-first-use pins (secure store) + in-memory cache. */
|
||||||
|
val tofu: com.shonar.provider.TofuManager by lazy {
|
||||||
|
com.shonar.provider.TofuManager(
|
||||||
|
com.shonar.provider.TofuStore(secureStore)
|
||||||
|
)
|
||||||
|
}
|
||||||
|
|
||||||
|
/**
|
||||||
|
* HTTP debug logging follows the `log_http_bodies` setting. Interceptors
|
||||||
|
* run on OkHttp threads that cannot suspend, so the flag is cached here
|
||||||
|
* and refreshed whenever settings change — toggling needs no restart.
|
||||||
|
*/
|
||||||
|
private val bodyLogging = AtomicBoolean(false)
|
||||||
|
|
||||||
|
/** Server providers (P1 local-only, P3 custom SHONAR, P4 Nextcloud + sync folder). */
|
||||||
|
val providerRegistry: com.shonar.provider.ProviderRegistry by lazy {
|
||||||
|
com.shonar.provider.ProviderRegistry.withDefaults(
|
||||||
|
appFilesDir = filesDir,
|
||||||
|
secureStore = secureStore,
|
||||||
|
plainStore = DataStoreSettingsStore(this),
|
||||||
|
tlsPolicy = com.shonar.provider.TlsPolicy(
|
||||||
|
tofu = tofu,
|
||||||
|
bodiesEnabled = bodyLogging::get,
|
||||||
|
sink = { msg -> android.util.Log.d("ShonarNet", msg) },
|
||||||
|
),
|
||||||
|
)
|
||||||
|
}
|
||||||
|
|
||||||
|
override fun onCreate() {
|
||||||
|
super.onCreate()
|
||||||
|
appScope.launch {
|
||||||
|
settingsManager.ensureLoaded()
|
||||||
|
// TOFU pins must be in cache before any TLS handshake needs them.
|
||||||
|
tofu.refresh()
|
||||||
|
// Point new recordings at the picked folder (if any).
|
||||||
|
runCatching { recordingRepository.refreshRoot() }
|
||||||
|
syncBodyLoggingFlag()
|
||||||
|
scheduleSync()
|
||||||
|
settingsManager.valuesChanged.collect {
|
||||||
|
syncBodyLoggingFlag()
|
||||||
|
scheduleSync()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
/** M5: steady-state periodic drain plus an immediate drain at startup
|
||||||
|
* and whenever sync constraints change. Re-running on unrelated
|
||||||
|
* settings edits only refreshes the periodic schedule (cheap UPDATE),
|
||||||
|
* never a drain. */
|
||||||
|
private var lastSyncFlags: Pair<Boolean, Boolean>? = null
|
||||||
|
|
||||||
|
private suspend fun scheduleSync() {
|
||||||
|
val flags = wifiOnly() to chargingOnly()
|
||||||
|
com.shonar.recording.SyncScheduler.ensurePeriodic(this, flags.first, flags.second)
|
||||||
|
if (lastSyncFlags == null || lastSyncFlags != flags) {
|
||||||
|
com.shonar.recording.SyncScheduler.requestNow(this, flags.first, flags.second)
|
||||||
|
}
|
||||||
|
lastSyncFlags = flags
|
||||||
|
}
|
||||||
|
|
||||||
|
private suspend fun wifiOnly(): Boolean = runCatching {
|
||||||
|
settingsManager.bool(com.shonar.settings.BuiltInSettings.WIFI_ONLY_UPLOAD)
|
||||||
|
}.getOrDefault(true)
|
||||||
|
|
||||||
|
private suspend fun chargingOnly(): Boolean = runCatching {
|
||||||
|
settingsManager.bool(com.shonar.settings.BuiltInSettings.CHARGING_ONLY_UPLOAD)
|
||||||
|
}.getOrDefault(false)
|
||||||
|
|
||||||
|
private suspend fun syncBodyLoggingFlag() {
|
||||||
|
bodyLogging.set(
|
||||||
|
runCatching { settingsManager.bool("log_http_bodies") }.getOrDefault(false)
|
||||||
|
)
|
||||||
|
}
|
||||||
|
}
|
||||||
748
shared/com/shonar/provider/CustomShonarProvider.kt
Normal file
748
shared/com/shonar/provider/CustomShonarProvider.kt
Normal file
|
|
@ -0,0 +1,748 @@
|
||||||
|
package com.shonar.provider
|
||||||
|
|
||||||
|
import java.io.File
|
||||||
|
import java.security.MessageDigest
|
||||||
|
import java.time.Instant
|
||||||
|
import java.util.concurrent.TimeUnit
|
||||||
|
import kotlinx.coroutines.Dispatchers
|
||||||
|
import kotlinx.coroutines.ensureActive
|
||||||
|
import kotlinx.coroutines.flow.MutableStateFlow
|
||||||
|
import kotlinx.coroutines.flow.StateFlow
|
||||||
|
import kotlinx.coroutines.sync.Mutex
|
||||||
|
import kotlinx.coroutines.sync.withLock
|
||||||
|
import kotlinx.coroutines.withContext
|
||||||
|
import okhttp3.MediaType.Companion.toMediaType
|
||||||
|
import okhttp3.OkHttpClient
|
||||||
|
import okhttp3.Request
|
||||||
|
import okhttp3.RequestBody.Companion.toRequestBody
|
||||||
|
import okhttp3.Response
|
||||||
|
|
||||||
|
/**
|
||||||
|
* P3: SHONAR-backend provider (this repo's FastAPI server, M1 auth + M2
|
||||||
|
* uploads). Speaks only documented endpoints:
|
||||||
|
*
|
||||||
|
* - GET /api/v1/provider-info (probe; no credentials)
|
||||||
|
* - POST /api/v1/auth/login {email, password, device_name, platform}
|
||||||
|
* - POST /api/v1/auth/refresh {refresh_token} (rotating; reuse-detected)
|
||||||
|
* - POST /api/v1/auth/logout {refresh_token}
|
||||||
|
* - GET /api/v1/auth/me (token validation)
|
||||||
|
* - POST /api/v1/auth/delete-account {password} (via [deleteAccount])
|
||||||
|
* - POST /api/v1/uploads (create session)
|
||||||
|
* - GET /api/v1/uploads/{id} (resume: received indexes)
|
||||||
|
* - PUT /api/v1/uploads/{id}/chunks/{n} (+ X-Chunk-Sha256)
|
||||||
|
* - POST /api/v1/uploads/{id}/finalize
|
||||||
|
* - GET /api/v1/recordings?limit&offset&sort&order
|
||||||
|
* - GET /api/v1/recordings/{id}/audio
|
||||||
|
* - DELETE /api/v1/recordings/{id}?purge=true
|
||||||
|
*
|
||||||
|
* Token discipline: the backend rotates refresh tokens with reuse
|
||||||
|
* detection, so the stored pair is overwritten on every login AND every
|
||||||
|
* refresh, and concurrent 401s serialize on [refreshMutex] — two parallel
|
||||||
|
* refreshes would look like token reuse and burn the whole family.
|
||||||
|
*
|
||||||
|
* Sidecars: file cache under [sidecarRoot] keyed by remote recording id
|
||||||
|
* (same layout as LocalOnlyProvider). AI content (transcript/summary/jobs)
|
||||||
|
* goes through the real M7 endpoints instead — see [fetchTranscript] and
|
||||||
|
* friends, consumed by the M8 details screen.
|
||||||
|
*
|
||||||
|
* Cancellation safety: a cancelled upload leaves an open server session
|
||||||
|
* with some chunks stored — invisible until finalize, resumable via the
|
||||||
|
* status endpoint, and idempotent per draft id through
|
||||||
|
* `client_recording_id`. Nothing half-visible ever appears in listings.
|
||||||
|
*/
|
||||||
|
class CustomShonarProvider(
|
||||||
|
private val auth: ShonarAuthStore,
|
||||||
|
private val sidecarRoot: File,
|
||||||
|
private val client: OkHttpClient = defaultClient(),
|
||||||
|
private val handshake: ShonarHandshake = ShonarHandshake(),
|
||||||
|
/** Identifies this client to the server (login device_name/platform). */
|
||||||
|
private val deviceName: String = "SHONAR Android",
|
||||||
|
private val platform: String = "android",
|
||||||
|
) : ShonarProvider {
|
||||||
|
|
||||||
|
override val descriptor = ProviderDescriptor(
|
||||||
|
id = ProviderRegistry.CUSTOM_SHONAR_ID,
|
||||||
|
displayName = "Custom SHONAR server",
|
||||||
|
capabilities = setOf(
|
||||||
|
ProviderDescriptor.Capability.CHUNKED_UPLOAD,
|
||||||
|
ProviderDescriptor.Capability.ACCOUNT_DELETION,
|
||||||
|
),
|
||||||
|
)
|
||||||
|
|
||||||
|
private val _authState = MutableStateFlow(AuthState.DISCONNECTED)
|
||||||
|
override val authState: StateFlow<AuthState> = _authState
|
||||||
|
|
||||||
|
/** Base origin, e.g. https://shonar.example.com. Set by connect/reconnect. */
|
||||||
|
private var origin: String? = null
|
||||||
|
|
||||||
|
private val refreshMutex = Mutex()
|
||||||
|
|
||||||
|
// ---- lifecycle ---------------------------------------------------------
|
||||||
|
|
||||||
|
override suspend fun probe(baseUrl: ServerUrl): ProbeResult = handshake.probe(baseUrl)
|
||||||
|
|
||||||
|
override suspend fun connect(credential: ProviderCredential) = withContext(Dispatchers.IO) {
|
||||||
|
when (credential) {
|
||||||
|
is ProviderCredential.ShonarLogin -> login(
|
||||||
|
credential.serverUrl.origin, credential.email, credential.password
|
||||||
|
)
|
||||||
|
is ProviderCredential.OAuthTokens -> resumeWithTokens(credential)
|
||||||
|
else -> throw ProviderError.InvalidUrl(
|
||||||
|
"Custom SHONAR server needs an email + password login"
|
||||||
|
)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
override suspend fun reconnect(): AuthState = withContext(Dispatchers.IO) {
|
||||||
|
val saved = auth.load()
|
||||||
|
if (saved == null || saved.refreshToken.isBlank()) {
|
||||||
|
_authState.value = AuthState.DISCONNECTED
|
||||||
|
return@withContext _authState.value
|
||||||
|
}
|
||||||
|
origin = saved.baseUrl
|
||||||
|
val base = saved.baseUrl
|
||||||
|
try {
|
||||||
|
executeAuthed(base) { token -> get(base, "/api/v1/auth/me", token) }.use { resp ->
|
||||||
|
if (resp.code != 200) throw ProviderError.AuthExpired()
|
||||||
|
}
|
||||||
|
_authState.value = AuthState.CONNECTED
|
||||||
|
} catch (e: ProviderError.AuthExpired) {
|
||||||
|
// Refresh already failed inside executeAuthed: tokens are dead.
|
||||||
|
auth.clearTokens()
|
||||||
|
_authState.value = AuthState.EXPIRED
|
||||||
|
} catch (e: ProviderError) {
|
||||||
|
_authState.value = AuthState.OFFLINE
|
||||||
|
}
|
||||||
|
_authState.value
|
||||||
|
}
|
||||||
|
|
||||||
|
override suspend fun disconnect(revokeOnServer: Boolean) = withContext(Dispatchers.IO) {
|
||||||
|
if (revokeOnServer) {
|
||||||
|
// Best effort: local state is cleared even if revoke fails.
|
||||||
|
runCatching {
|
||||||
|
val saved = auth.load()
|
||||||
|
if (saved != null && saved.refreshToken.isNotBlank()) {
|
||||||
|
postUnauthed(
|
||||||
|
saved.baseUrl, "/api/v1/auth/logout",
|
||||||
|
"""{"refresh_token":${jsonStr(saved.refreshToken)}}""",
|
||||||
|
).close()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
auth.clearTokens()
|
||||||
|
_authState.value = AuthState.DISCONNECTED
|
||||||
|
}
|
||||||
|
|
||||||
|
override suspend fun deleteAccountAndData() {
|
||||||
|
// No password is available here (never stored), so server-side
|
||||||
|
// account purge needs the explicit [deleteAccount] call below.
|
||||||
|
// This path revokes the session and wipes everything local.
|
||||||
|
withContext(Dispatchers.IO) {
|
||||||
|
disconnect(revokeOnServer = true)
|
||||||
|
sidecarRoot.deleteRecursively()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
/**
|
||||||
|
* Full server-side account purge (backend: 30-day grace, then hard
|
||||||
|
* delete). Needs the password because it is never stored — call this
|
||||||
|
* from a confirmation screen that asks for it once.
|
||||||
|
*/
|
||||||
|
suspend fun deleteAccount(password: String) = withContext(Dispatchers.IO) {
|
||||||
|
val base = origin ?: auth.load()?.baseUrl ?: throw ProviderError.NotConnected()
|
||||||
|
executeAuthed(base) { token ->
|
||||||
|
post(base, "/api/v1/auth/delete-account", token,
|
||||||
|
"""{"password":${jsonStr(password)}}""")
|
||||||
|
}.use { resp ->
|
||||||
|
if (resp.code != 202 && resp.code != 204) {
|
||||||
|
throw ProviderError.Transient("Account deletion refused (HTTP ${resp.code})")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
auth.clearAll()
|
||||||
|
origin = null
|
||||||
|
sidecarRoot.deleteRecursively()
|
||||||
|
_authState.value = AuthState.DISCONNECTED
|
||||||
|
}
|
||||||
|
|
||||||
|
// ---- storage -----------------------------------------------------------
|
||||||
|
|
||||||
|
override suspend fun upload(draft: RecordingDraft, onProgress: (Float) -> Unit): RemoteRef =
|
||||||
|
withContext(Dispatchers.IO) {
|
||||||
|
ensureConnected()
|
||||||
|
val base = currentOrigin()
|
||||||
|
val total = draft.sizeBytes.coerceAtLeast(1)
|
||||||
|
|
||||||
|
val (sessionId, chunkSize) = createSession(base, draft)
|
||||||
|
val received = uploadStatus(base, sessionId).toMutableSet()
|
||||||
|
|
||||||
|
var sent = 0L
|
||||||
|
// Account progress for already-present chunks so resume continues
|
||||||
|
// the bar instead of restarting it.
|
||||||
|
var idx = 0
|
||||||
|
val chunkCount = ((draft.sizeBytes + chunkSize - 1) / chunkSize).toInt().coerceAtLeast(1)
|
||||||
|
var resumedBytes = 0L
|
||||||
|
// Estimate resumed bytes from the received set (last chunk may be short).
|
||||||
|
for (i in received) {
|
||||||
|
resumedBytes += if (i < chunkCount - 1) chunkSize.toLong()
|
||||||
|
else (draft.sizeBytes - chunkSize * (chunkCount - 1)).coerceAtLeast(0)
|
||||||
|
}
|
||||||
|
sent = resumedBytes
|
||||||
|
if (sent > 0) onProgress((sent.toFloat() / total).coerceIn(0f, 1f))
|
||||||
|
|
||||||
|
while (idx < chunkCount) {
|
||||||
|
// Cooperative cancellation between chunks; a cancelled
|
||||||
|
// upload leaves a resumable server session, never a
|
||||||
|
// half-visible object.
|
||||||
|
ensureActive()
|
||||||
|
if (idx !in received) {
|
||||||
|
val slice = readSlice(draft.sourceFile, idx.toLong() * chunkSize, chunkSize)
|
||||||
|
putChunk(base, sessionId, idx, slice)
|
||||||
|
sent += slice.size
|
||||||
|
onProgress((sent.toFloat() / total).coerceIn(0f, 1f))
|
||||||
|
}
|
||||||
|
idx++
|
||||||
|
}
|
||||||
|
|
||||||
|
val recordingId = finalize(base, sessionId, draft)
|
||||||
|
onProgress(1f)
|
||||||
|
RemoteRef(ProviderRegistry.CUSTOM_SHONAR_ID, recordingId, etag = null, sizeBytes = draft.sizeBytes)
|
||||||
|
}
|
||||||
|
|
||||||
|
override suspend fun download(ref: RemoteRef, dest: File, onProgress: (Float) -> Unit) =
|
||||||
|
withContext(Dispatchers.IO) {
|
||||||
|
ensureConnected()
|
||||||
|
val base = currentOrigin()
|
||||||
|
executeAuthed(base) { token -> get(base, "/api/v1/recordings/${ref.key}/audio", token) }.use { resp ->
|
||||||
|
when (resp.code) {
|
||||||
|
200 -> {
|
||||||
|
val body = resp.body ?: throw ProviderError.Transient("Empty download response")
|
||||||
|
val total = body.contentLength().takeIf { it > 0 } ?: -1
|
||||||
|
dest.parentFile?.mkdirs()
|
||||||
|
body.byteStream().use { input ->
|
||||||
|
dest.outputStream().use { output ->
|
||||||
|
val buf = ByteArray(64 * 1024)
|
||||||
|
var written = 0L
|
||||||
|
var lastReported = -1f
|
||||||
|
while (true) {
|
||||||
|
val n = input.read(buf)
|
||||||
|
if (n < 0) break
|
||||||
|
output.write(buf, 0, n)
|
||||||
|
written += n
|
||||||
|
if (total > 0) {
|
||||||
|
val p = (written.toFloat() / total).coerceIn(0f, 1f)
|
||||||
|
if (p > lastReported) {
|
||||||
|
onProgress(p)
|
||||||
|
lastReported = p
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
onProgress(1f)
|
||||||
|
}
|
||||||
|
404 -> throw ProviderError.NotFound(ref.key)
|
||||||
|
else -> throw ProviderError.Transient("Download failed (HTTP ${resp.code})")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
override suspend fun delete(ref: RemoteRef) {
|
||||||
|
withContext(Dispatchers.IO) {
|
||||||
|
ensureConnected()
|
||||||
|
val base = currentOrigin()
|
||||||
|
executeAuthed(base) { token ->
|
||||||
|
Request.Builder().url("$base/api/v1/recordings/${ref.key}?purge=true")
|
||||||
|
.delete().header("Authorization", "Bearer $token").build()
|
||||||
|
}.use { resp ->
|
||||||
|
when (resp.code) {
|
||||||
|
204, 200 -> {
|
||||||
|
sidecarDir(ref).deleteRecursively()
|
||||||
|
}
|
||||||
|
404 -> throw ProviderError.NotFound(ref.key)
|
||||||
|
else -> throw ProviderError.Transient("Delete failed (HTTP ${resp.code})")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
override suspend fun list(cursor: String?): Page<RemoteRecording> = withContext(Dispatchers.IO) {
|
||||||
|
ensureConnected()
|
||||||
|
val base = currentOrigin()
|
||||||
|
val offset = cursor?.toIntOrNull()?.coerceAtLeast(0) ?: 0
|
||||||
|
val limit = 200
|
||||||
|
executeAuthed(base) { token ->
|
||||||
|
get(base, "/api/v1/recordings?limit=$limit&offset=$offset&sort=recorded_at&order=desc", token)
|
||||||
|
}.use { resp ->
|
||||||
|
if (resp.code != 200) throw ProviderError.Transient("Listing failed (HTTP ${resp.code})")
|
||||||
|
val root = org.json.JSONObject(resp.body?.string().orEmpty())
|
||||||
|
val total = root.optInt("total", 0)
|
||||||
|
val items = root.optJSONArray("items") ?: org.json.JSONArray()
|
||||||
|
val out = mutableListOf<RemoteRecording>()
|
||||||
|
for (i in 0 until items.length()) {
|
||||||
|
parseRecording(items.getJSONObject(i))?.let { out += it }
|
||||||
|
}
|
||||||
|
val next = if (offset + limit < total) (offset + limit).toString() else null
|
||||||
|
Page(out, next)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// ---- AI content (M7 endpoints, consumed by the M8 details screen) ----
|
||||||
|
//
|
||||||
|
// The recording id here is the server-side UUID — the local row's
|
||||||
|
// [com.shonar.recording.RecordingEntity.remoteKey] after upload.
|
||||||
|
// 404 surfaces as [ProviderError.NotFound] ("no transcript yet" is a
|
||||||
|
// normal state, not a failure). Raw JSON comes back so parsing stays in
|
||||||
|
// the testable `recording.AiContent` helpers, not in this IO class.
|
||||||
|
|
||||||
|
suspend fun fetchTranscript(recordingId: String): String = getAiJson(
|
||||||
|
"/api/v1/recordings/$recordingId/transcript", recordingId
|
||||||
|
)
|
||||||
|
|
||||||
|
suspend fun fetchSummary(recordingId: String): String = getAiJson(
|
||||||
|
"/api/v1/recordings/$recordingId/summary", recordingId
|
||||||
|
)
|
||||||
|
|
||||||
|
suspend fun fetchJobs(recordingId: String): String = getAiJson(
|
||||||
|
"/api/v1/recordings/$recordingId/jobs", recordingId
|
||||||
|
)
|
||||||
|
|
||||||
|
/** Force one pipeline stage to re-run on the server (no re-upload).
|
||||||
|
* [model] only applies to job="transcribe" (switches the saved
|
||||||
|
* per-recording override). 409 = stage already running or nothing to
|
||||||
|
* summarize; surfaced as ProviderError.Transient with the server's
|
||||||
|
* detail message. */
|
||||||
|
suspend fun reprocess(recordingId: String, job: String, model: String? = null): String =
|
||||||
|
withContext(Dispatchers.IO) {
|
||||||
|
ensureConnected()
|
||||||
|
val base = currentOrigin()
|
||||||
|
val q = if (model != null) "&model=$model" else ""
|
||||||
|
executeAuthed(base) { token ->
|
||||||
|
post(base, "/api/v1/recordings/$recordingId/reprocess?job=$job$q", token, "{}")
|
||||||
|
}.use { resp ->
|
||||||
|
when (resp.code) {
|
||||||
|
200, 201 -> resp.body?.string().orEmpty()
|
||||||
|
404 -> throw ProviderError.NotFound(recordingId)
|
||||||
|
else -> {
|
||||||
|
val body = runCatching {
|
||||||
|
org.json.JSONObject(resp.body?.string().orEmpty())
|
||||||
|
.optString("detail")
|
||||||
|
}.getOrNull()
|
||||||
|
throw ProviderError.Transient(
|
||||||
|
body?.takeIf { it.isNotBlank() }
|
||||||
|
?: "Reprocess failed (HTTP ${resp.code})")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
/** User edit → new server version with edited_by_user=true (pipeline won't clobber). */
|
||||||
|
suspend fun updateTranscript(
|
||||||
|
recordingId: String,
|
||||||
|
payloadJson: String,
|
||||||
|
): String = withContext(Dispatchers.IO) {
|
||||||
|
ensureConnected()
|
||||||
|
val base = currentOrigin()
|
||||||
|
executeAuthed(base) { token -> put(base, "/api/v1/recordings/$recordingId/transcript", token, payloadJson) }
|
||||||
|
.use { resp ->
|
||||||
|
when (resp.code) {
|
||||||
|
200 -> resp.body?.string().orEmpty()
|
||||||
|
404 -> throw ProviderError.NotFound(recordingId)
|
||||||
|
else -> throw ProviderError.Transient("Transcript update failed (HTTP ${resp.code})")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
suspend fun updateSummary(recordingId: String, payloadJson: String): String =
|
||||||
|
withContext(Dispatchers.IO) {
|
||||||
|
ensureConnected()
|
||||||
|
val base = currentOrigin()
|
||||||
|
executeAuthed(base) { token -> put(base, "/api/v1/recordings/$recordingId/summary", token, payloadJson) }
|
||||||
|
.use { resp ->
|
||||||
|
when (resp.code) {
|
||||||
|
200 -> resp.body?.string().orEmpty()
|
||||||
|
404 -> throw ProviderError.NotFound(recordingId)
|
||||||
|
else -> throw ProviderError.Transient("Summary update failed (HTTP ${resp.code})")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
/** Title/notes rename from the details screen (server + Room updated by caller). */
|
||||||
|
suspend fun patchRecording(recordingId: String, payloadJson: String): String =
|
||||||
|
withContext(Dispatchers.IO) {
|
||||||
|
ensureConnected()
|
||||||
|
val base = currentOrigin()
|
||||||
|
executeAuthed(base) { token -> patch(base, "/api/v1/recordings/$recordingId", token, payloadJson) }
|
||||||
|
.use { resp ->
|
||||||
|
when (resp.code) {
|
||||||
|
200 -> resp.body?.string().orEmpty()
|
||||||
|
404 -> throw ProviderError.NotFound(recordingId)
|
||||||
|
else -> throw ProviderError.Transient("Recording update failed (HTTP ${resp.code})")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
private suspend fun getAiJson(path: String, recordingId: String): String =
|
||||||
|
withContext(Dispatchers.IO) {
|
||||||
|
ensureConnected()
|
||||||
|
val base = currentOrigin()
|
||||||
|
executeAuthed(base) { token -> get(base, path, token) }.use { resp ->
|
||||||
|
when (resp.code) {
|
||||||
|
200 -> resp.body?.string().orEmpty()
|
||||||
|
404 -> throw ProviderError.NotFound(recordingId)
|
||||||
|
else -> throw ProviderError.Transient("Request failed (HTTP ${resp.code})")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// ---- transcription models (Stage 1 endpoints, desktop Settings) --------
|
||||||
|
// Raw JSON throughout; parsing lives in the testable AiContent helpers.
|
||||||
|
|
||||||
|
suspend fun fetchModels(): String = withContext(Dispatchers.IO) {
|
||||||
|
ensureConnected()
|
||||||
|
val base = currentOrigin()
|
||||||
|
executeAuthed(base) { token -> get(base, "/api/v1/models", token) }.use { resp ->
|
||||||
|
if (resp.code != 200) throw ProviderError.Transient("Models request failed (HTTP ${resp.code})")
|
||||||
|
resp.body?.string().orEmpty()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
suspend fun setDefaultModel(model: String): String = withContext(Dispatchers.IO) {
|
||||||
|
ensureConnected()
|
||||||
|
val base = currentOrigin()
|
||||||
|
val payload = """{"model":${jsonStr(model)}}"""
|
||||||
|
executeAuthed(base) { token -> put(base, "/api/v1/models/default", token, payload) }
|
||||||
|
.use { resp ->
|
||||||
|
when (resp.code) {
|
||||||
|
200 -> resp.body?.string().orEmpty()
|
||||||
|
422 -> throw ProviderError.InvalidUrl("Unsupported model: $model")
|
||||||
|
else -> throw ProviderError.Transient("Default model update failed (HTTP ${resp.code})")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
/** Fetch a model into the server cache (needs internet once; may take minutes). */
|
||||||
|
suspend fun downloadModel(model: String): String = withContext(Dispatchers.IO) {
|
||||||
|
ensureConnected()
|
||||||
|
val base = currentOrigin()
|
||||||
|
executeAuthed(base) { token ->
|
||||||
|
post(base, "/api/v1/models/$model/download", token, "{}")
|
||||||
|
}.use { resp ->
|
||||||
|
when (resp.code) {
|
||||||
|
200 -> resp.body?.string().orEmpty()
|
||||||
|
422 -> throw ProviderError.InvalidUrl("Unsupported model: $model")
|
||||||
|
501 -> throw ProviderError.Transient("faster-whisper is not installed on the server")
|
||||||
|
else -> throw ProviderError.Transient("Model download failed (HTTP ${resp.code})")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// ---- sidecars (local cache until the backend gains endpoints in M7) ----
|
||||||
|
|
||||||
|
private fun sidecarDir(ref: RemoteRef) = File(sidecarRoot, ref.key)
|
||||||
|
|
||||||
|
override suspend fun putSidecar(ref: RemoteRef, kind: SidecarKind, bytes: ByteArray) =
|
||||||
|
withContext(Dispatchers.IO) {
|
||||||
|
val f = File(sidecarDir(ref), kind.fileName)
|
||||||
|
f.parentFile?.mkdirs()
|
||||||
|
f.writeBytes(bytes)
|
||||||
|
}
|
||||||
|
|
||||||
|
override suspend fun getSidecar(ref: RemoteRef, kind: SidecarKind): ByteArray? =
|
||||||
|
withContext(Dispatchers.IO) {
|
||||||
|
val f = File(sidecarDir(ref), kind.fileName)
|
||||||
|
if (f.exists()) f.readBytes() else null
|
||||||
|
}
|
||||||
|
|
||||||
|
// ---- status ------------------------------------------------------------
|
||||||
|
|
||||||
|
override suspend fun storageLocationSummary(): StorageLocation =
|
||||||
|
withContext(Dispatchers.IO) {
|
||||||
|
ensureConnected()
|
||||||
|
val base = currentOrigin()
|
||||||
|
val host = runCatching { java.net.URI(base).host }.getOrNull() ?: base
|
||||||
|
val email = auth.load()?.email.orEmpty()
|
||||||
|
var count = 0
|
||||||
|
var bytes = 0L
|
||||||
|
var cursor: String? = null
|
||||||
|
do {
|
||||||
|
val page = list(cursor)
|
||||||
|
count += page.items.size
|
||||||
|
bytes += page.items.sumOf { it.ref.sizeBytes }
|
||||||
|
cursor = page.nextCursor
|
||||||
|
} while (cursor != null)
|
||||||
|
StorageLocation(
|
||||||
|
headline = "Custom SHONAR server at $host",
|
||||||
|
detail = if (email.isBlank()) "$count recordings synced."
|
||||||
|
else "Signed in as $email · $count recordings synced.",
|
||||||
|
syncedCount = count,
|
||||||
|
localOnlyCount = 0,
|
||||||
|
bytesUsed = bytes,
|
||||||
|
)
|
||||||
|
}
|
||||||
|
|
||||||
|
// ---- internals ---------------------------------------------------------
|
||||||
|
|
||||||
|
private fun ensureConnected() {
|
||||||
|
if (_authState.value != AuthState.CONNECTED) throw ProviderError.NotConnected()
|
||||||
|
}
|
||||||
|
|
||||||
|
private suspend fun currentOrigin(): String =
|
||||||
|
origin ?: auth.load()?.baseUrl ?: throw ProviderError.NotConnected()
|
||||||
|
|
||||||
|
private suspend fun login(base: String, email: String, password: String) {
|
||||||
|
val body = """{"email":${jsonStr(email)},"password":${jsonStr(password)},""" +
|
||||||
|
""""device_name":${jsonStr(deviceName)},"platform":${jsonStr(platform)}}"""
|
||||||
|
postUnauthed(base, "/api/v1/auth/login", body).use { resp ->
|
||||||
|
when (resp.code) {
|
||||||
|
200 -> {
|
||||||
|
persistSession(base, email, org.json.JSONObject(resp.body?.string().orEmpty()))
|
||||||
|
origin = base
|
||||||
|
_authState.value = AuthState.CONNECTED
|
||||||
|
}
|
||||||
|
401 -> throw ProviderError.Transient("Invalid email or password")
|
||||||
|
else -> throw ProviderError.Transient("Login failed (HTTP ${resp.code})")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
private suspend fun resumeWithTokens(credential: ProviderCredential.OAuthTokens) {
|
||||||
|
val base = origin ?: auth.load()?.baseUrl
|
||||||
|
?: throw ProviderError.InvalidUrl("No SHONAR server configured yet")
|
||||||
|
origin = base
|
||||||
|
auth.save(
|
||||||
|
ShonarSession(
|
||||||
|
baseUrl = base,
|
||||||
|
email = credential.accountLabel,
|
||||||
|
accessToken = credential.accessToken,
|
||||||
|
refreshToken = credential.refreshToken.orEmpty(),
|
||||||
|
expiresAtEpochSec = credential.expiresAtEpochSec,
|
||||||
|
deviceId = null,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
executeAuthed(base) { token -> get(base, "/api/v1/auth/me", token) }.use { resp ->
|
||||||
|
if (resp.code != 200) throw ProviderError.AuthExpired()
|
||||||
|
}
|
||||||
|
_authState.value = AuthState.CONNECTED
|
||||||
|
}
|
||||||
|
|
||||||
|
private suspend fun persistSession(base: String, email: String, json: org.json.JSONObject) {
|
||||||
|
val nowSec = Instant.now().epochSecond
|
||||||
|
auth.save(
|
||||||
|
ShonarSession(
|
||||||
|
baseUrl = base,
|
||||||
|
email = email,
|
||||||
|
accessToken = json.getString("access_token"),
|
||||||
|
refreshToken = json.getString("refresh_token"),
|
||||||
|
expiresAtEpochSec = nowSec + json.optInt("expires_in", 900),
|
||||||
|
deviceId = json.optString("device_id", null).takeUnless { it.isNullOrBlank() },
|
||||||
|
)
|
||||||
|
)
|
||||||
|
}
|
||||||
|
|
||||||
|
/** Single-flight refresh; concurrent 401s must not double-refresh. */
|
||||||
|
private suspend fun refreshLocked(failedAccess: String): String = refreshMutex.withLock {
|
||||||
|
val current = auth.load() ?: throw ProviderError.NotConnected()
|
||||||
|
// A peer already refreshed while we queued — reuse its tokens.
|
||||||
|
if (current.accessToken != failedAccess && current.accessToken.isNotBlank()) {
|
||||||
|
return@withLock current.accessToken
|
||||||
|
}
|
||||||
|
if (current.refreshToken.isBlank()) {
|
||||||
|
_authState.value = AuthState.EXPIRED
|
||||||
|
throw ProviderError.AuthExpired()
|
||||||
|
}
|
||||||
|
val body = """{"refresh_token":${jsonStr(current.refreshToken)}}"""
|
||||||
|
try {
|
||||||
|
postUnauthed(current.baseUrl, "/api/v1/auth/refresh", body).use { resp ->
|
||||||
|
if (resp.code != 200) {
|
||||||
|
// Reuse detected or revoked family: stored tokens are dead.
|
||||||
|
auth.clearTokens()
|
||||||
|
_authState.value = AuthState.EXPIRED
|
||||||
|
throw ProviderError.AuthExpired()
|
||||||
|
}
|
||||||
|
persistSession(current.baseUrl, current.email, org.json.JSONObject(resp.body?.string().orEmpty()))
|
||||||
|
_authState.value = AuthState.CONNECTED
|
||||||
|
return@withLock auth.load()?.accessToken ?: throw ProviderError.AuthExpired()
|
||||||
|
}
|
||||||
|
} catch (e: ProviderError) {
|
||||||
|
throw e
|
||||||
|
} catch (e: Exception) {
|
||||||
|
throw ProviderError.Transient("Token refresh failed (${e.javaClass.simpleName})")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
/** Execute an authenticated request; one transparent refresh+retry on 401. */
|
||||||
|
private suspend fun executeAuthed(
|
||||||
|
base: String,
|
||||||
|
build: (access: String) -> Request,
|
||||||
|
): Response = withContext(Dispatchers.IO) {
|
||||||
|
val session = auth.load() ?: throw ProviderError.NotConnected()
|
||||||
|
val first = client.newCall(build(session.accessToken)).execute()
|
||||||
|
if (first.code != 401) return@withContext first
|
||||||
|
first.close()
|
||||||
|
val fresh = refreshLocked(session.accessToken)
|
||||||
|
val retry = client.newCall(build(fresh)).execute()
|
||||||
|
if (retry.code == 401) {
|
||||||
|
retry.close()
|
||||||
|
_authState.value = AuthState.EXPIRED
|
||||||
|
throw ProviderError.AuthExpired()
|
||||||
|
}
|
||||||
|
retry
|
||||||
|
}
|
||||||
|
|
||||||
|
private fun get(base: String, path: String, access: String): Request =
|
||||||
|
Request.Builder().url(base + path).get()
|
||||||
|
.header("Authorization", "Bearer $access").build()
|
||||||
|
|
||||||
|
private fun post(base: String, path: String, access: String, json: String): Request =
|
||||||
|
Request.Builder().url(base + path)
|
||||||
|
.post(json.toRequestBody("application/json; charset=utf-8".toMediaType()))
|
||||||
|
.header("Authorization", "Bearer $access").build()
|
||||||
|
|
||||||
|
private fun put(base: String, path: String, access: String, json: String): Request =
|
||||||
|
Request.Builder().url(base + path)
|
||||||
|
.put(json.toRequestBody("application/json; charset=utf-8".toMediaType()))
|
||||||
|
.header("Authorization", "Bearer $access").build()
|
||||||
|
|
||||||
|
private fun patch(base: String, path: String, access: String, json: String): Request =
|
||||||
|
Request.Builder().url(base + path)
|
||||||
|
.patch(json.toRequestBody("application/json; charset=utf-8".toMediaType()))
|
||||||
|
.header("Authorization", "Bearer $access").build()
|
||||||
|
|
||||||
|
private fun postUnauthed(base: String, path: String, json: String): Response {
|
||||||
|
val req = Request.Builder().url(base + path)
|
||||||
|
.post(json.toRequestBody("application/json; charset=utf-8".toMediaType())).build()
|
||||||
|
return client.newCall(req).execute()
|
||||||
|
}
|
||||||
|
|
||||||
|
// Upload-session flow returns Triple(sessionId, chunkSize, declaredMime echo not needed).
|
||||||
|
private suspend fun createSession(base: String, draft: RecordingDraft): Pair<String, Int> {
|
||||||
|
val body = buildString {
|
||||||
|
append("""{"declared_mime_type":${jsonStr(draft.mime)},""")
|
||||||
|
append(""""declared_size_bytes":${draft.sizeBytes},""")
|
||||||
|
append(""""title":${jsonStr(draft.title)},""")
|
||||||
|
append(""""client_recording_id":${jsonStr(draft.id)}""")
|
||||||
|
if (draft.transcriptionModel != null) {
|
||||||
|
append(""","transcription_model":${jsonStr(draft.transcriptionModel)}""")
|
||||||
|
}
|
||||||
|
append("}")
|
||||||
|
}
|
||||||
|
executeAuthed(base) { token -> post(base, "/api/v1/uploads", token, body) }.use { resp ->
|
||||||
|
when (resp.code) {
|
||||||
|
201 -> {
|
||||||
|
val json = org.json.JSONObject(resp.body?.string().orEmpty())
|
||||||
|
return json.getString("id") to json.optInt("chunk_size_bytes", 16 * 1024 * 1024)
|
||||||
|
}
|
||||||
|
413 -> throw ProviderError.Transient("Recording exceeds the server size limit")
|
||||||
|
415 -> throw ProviderError.Transient("Audio type not accepted by the server")
|
||||||
|
else -> throw ProviderError.Transient("Upload rejected (HTTP ${resp.code})")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
private suspend fun uploadStatus(base: String, sessionId: String): Set<Int> {
|
||||||
|
executeAuthed(base) { token -> get(base, "/api/v1/uploads/$sessionId", token) }.use { resp ->
|
||||||
|
if (resp.code != 200) throw ProviderError.Transient("Upload status failed (HTTP ${resp.code})")
|
||||||
|
val arr = org.json.JSONObject(resp.body?.string().orEmpty())
|
||||||
|
.optJSONArray("received_chunk_indexes") ?: return emptySet()
|
||||||
|
return (0 until arr.length()).map { arr.getInt(it) }.toSet()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
private suspend fun putChunk(base: String, sessionId: String, index: Int, bytes: ByteArray) {
|
||||||
|
val digest = MessageDigest.getInstance("SHA-256").digest(bytes).toHex()
|
||||||
|
val req = { token: String ->
|
||||||
|
Request.Builder().url("$base/api/v1/uploads/$sessionId/chunks/$index")
|
||||||
|
.put(bytes.toRequestBody("application/octet-stream".toMediaType()))
|
||||||
|
.header("Authorization", "Bearer $token")
|
||||||
|
.header("X-Chunk-Sha256", digest).build()
|
||||||
|
}
|
||||||
|
executeAuthed(base, req).use { resp ->
|
||||||
|
if (resp.code != 201) throw ProviderError.Transient("Chunk $index rejected (HTTP ${resp.code})")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
private suspend fun finalize(base: String, sessionId: String, draft: RecordingDraft): String {
|
||||||
|
val recordedAt = Instant.ofEpochMilli(draft.createdAtEpochMs).toString()
|
||||||
|
val body = buildString {
|
||||||
|
append("""{"recorded_at":${jsonStr(recordedAt)},""")
|
||||||
|
append(""""duration_seconds":${draft.durationMs / 1000.0}""")
|
||||||
|
if (draft.transcriptionModel != null) {
|
||||||
|
append(""","transcription_model":${jsonStr(draft.transcriptionModel)}""")
|
||||||
|
}
|
||||||
|
append("}")
|
||||||
|
}
|
||||||
|
executeAuthed(base) { token -> post(base, "/api/v1/uploads/$sessionId/finalize", token, body) }
|
||||||
|
.use { resp ->
|
||||||
|
if (resp.code != 201) {
|
||||||
|
throw ProviderError.Transient("Upload finalize failed (HTTP ${resp.code})")
|
||||||
|
}
|
||||||
|
return org.json.JSONObject(resp.body?.string().orEmpty()).getString("id")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
private fun parseRecording(json: org.json.JSONObject): RemoteRecording? {
|
||||||
|
if (!json.optBoolean("has_audio", false)) return null
|
||||||
|
val id = json.optString("id", "")
|
||||||
|
if (id.isBlank()) return null
|
||||||
|
val createdAt = runCatching {
|
||||||
|
Instant.parse(json.getString("recorded_at")).toEpochMilli()
|
||||||
|
}.getOrDefault(0L)
|
||||||
|
return RemoteRecording(
|
||||||
|
ref = RemoteRef(
|
||||||
|
providerId = ProviderRegistry.CUSTOM_SHONAR_ID,
|
||||||
|
key = id,
|
||||||
|
etag = null,
|
||||||
|
sizeBytes = json.optLong("size_bytes", 0),
|
||||||
|
),
|
||||||
|
title = json.optString("title", id),
|
||||||
|
createdAtEpochMs = createdAt,
|
||||||
|
durationMs = (json.optDouble("duration_seconds", 0.0) * 1000).toLong(),
|
||||||
|
mime = json.optString("mime_type", null) ?: "application/octet-stream",
|
||||||
|
)
|
||||||
|
}
|
||||||
|
|
||||||
|
companion object {
|
||||||
|
fun defaultClient(): OkHttpClient = OkHttpClient.Builder()
|
||||||
|
// No logging interceptor by policy (docs/server-providers.md §4):
|
||||||
|
// bodies would carry audio bytes and transcripts.
|
||||||
|
.connectTimeout(15, TimeUnit.SECONDS)
|
||||||
|
.readTimeout(60, TimeUnit.SECONDS)
|
||||||
|
.writeTimeout(60, TimeUnit.SECONDS)
|
||||||
|
.build()
|
||||||
|
|
||||||
|
/** JSON string literal with escaping; never pass secrets to message strings. */
|
||||||
|
internal fun jsonStr(raw: String): String = buildString {
|
||||||
|
append('"')
|
||||||
|
for (c in raw) when (c) {
|
||||||
|
'"' -> append("\\\"")
|
||||||
|
'\\' -> append("\\\\")
|
||||||
|
'\n' -> append("\\n")
|
||||||
|
'\r' -> append("\\r")
|
||||||
|
'\t' -> append("\\t")
|
||||||
|
else -> if (c < ' ') append("\\u%04x".format(c.code)) else append(c)
|
||||||
|
}
|
||||||
|
append('"')
|
||||||
|
}
|
||||||
|
|
||||||
|
private fun ByteArray.toHex(): String = joinToString("") { "%02x".format(it) }
|
||||||
|
|
||||||
|
private fun readSlice(file: File, offset: Long, max: Int): ByteArray {
|
||||||
|
file.inputStream().use { input ->
|
||||||
|
var skipped = 0L
|
||||||
|
while (skipped < offset) {
|
||||||
|
val n = input.skip(offset - skipped)
|
||||||
|
if (n <= 0) break
|
||||||
|
skipped += n
|
||||||
|
}
|
||||||
|
val buf = ByteArray(max)
|
||||||
|
var read = 0
|
||||||
|
while (read < max) {
|
||||||
|
val n = input.read(buf, read, max - read)
|
||||||
|
if (n < 0) break
|
||||||
|
read += n
|
||||||
|
}
|
||||||
|
return if (read == max) buf else buf.copyOf(read)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
194
shared/com/shonar/provider/FolderSyncProvider.kt
Normal file
194
shared/com/shonar/provider/FolderSyncProvider.kt
Normal file
|
|
@ -0,0 +1,194 @@
|
||||||
|
package com.shonar.provider
|
||||||
|
|
||||||
|
import java.io.File
|
||||||
|
import kotlinx.coroutines.Dispatchers
|
||||||
|
import kotlinx.coroutines.flow.MutableStateFlow
|
||||||
|
import kotlinx.coroutines.flow.StateFlow
|
||||||
|
import kotlinx.coroutines.withContext
|
||||||
|
|
||||||
|
/**
|
||||||
|
* P4: sync-folder provider — "bring your own sync". The app reads and
|
||||||
|
* writes plain files under a user-chosen directory; an external tool
|
||||||
|
* (Syncthing, the Nextcloud desktop client, rsync, …) moves those bytes
|
||||||
|
* between devices. The app never talks to a network for this provider —
|
||||||
|
* provable by construction (no HTTP imports in this file).
|
||||||
|
*
|
||||||
|
* The on-disk layout is identical to [LocalOnlyProvider] (`audio/{id}`,
|
||||||
|
* `sidecars/…`), so switching between local-only and a sync folder is a
|
||||||
|
* copy, not a migration, and anything already syncing the folder picks
|
||||||
|
* the recordings up with no special handling.
|
||||||
|
*
|
||||||
|
* Two deliberate differences from local-only:
|
||||||
|
* - [connect] validates the directory (must exist, be a directory, be
|
||||||
|
* readable AND writable) and refuses anything else. A typo must be an
|
||||||
|
* error, never a silently created folder somewhere surprising. Paths
|
||||||
|
* escaping via `..` are rejected for the same reason.
|
||||||
|
* - [deleteAccountAndData] NEVER deletes the folder's contents. That
|
||||||
|
* directory belongs to the user and their sync tool, not to the app —
|
||||||
|
* forgetting the path is the whole operation.
|
||||||
|
*/
|
||||||
|
class FolderSyncProvider(
|
||||||
|
private val pathStore: com.shonar.settings.SettingsStore =
|
||||||
|
com.shonar.settings.InMemorySettingsStore(),
|
||||||
|
) : ShonarProvider {
|
||||||
|
|
||||||
|
override val descriptor = ProviderDescriptor(
|
||||||
|
id = ID,
|
||||||
|
displayName = "Sync folder",
|
||||||
|
capabilities = setOf(), // the sync tool owns the protocol, not us
|
||||||
|
)
|
||||||
|
|
||||||
|
private val _authState = MutableStateFlow(AuthState.DISCONNECTED)
|
||||||
|
override val authState: StateFlow<AuthState> = _authState
|
||||||
|
|
||||||
|
private suspend fun storedRoot(): File? {
|
||||||
|
val raw = pathStore.getString(KEY_ROOT) ?: return null
|
||||||
|
return File(raw)
|
||||||
|
}
|
||||||
|
|
||||||
|
private fun checkDir(dir: File): File {
|
||||||
|
if (".." in dir.path.split(File.separatorChar)) {
|
||||||
|
throw ProviderError.InvalidUrl("Folder path must not contain '..'")
|
||||||
|
}
|
||||||
|
if (!dir.exists()) throw ProviderError.InvalidUrl(
|
||||||
|
"Folder does not exist: ${dir.path}. Create it (or let Syncthing create it) first."
|
||||||
|
)
|
||||||
|
if (!dir.isDirectory) throw ProviderError.InvalidUrl("Not a folder: ${dir.path}")
|
||||||
|
if (!dir.canRead() || !dir.canWrite()) throw ProviderError.InvalidUrl(
|
||||||
|
"Folder is not readable and writable: ${dir.path}"
|
||||||
|
)
|
||||||
|
return dir
|
||||||
|
}
|
||||||
|
|
||||||
|
override suspend fun probe(baseUrl: ServerUrl): ProbeResult =
|
||||||
|
ProbeResult.Incompatible // no server involved — nothing to probe
|
||||||
|
|
||||||
|
override suspend fun connect(credential: ProviderCredential) = withContext(Dispatchers.IO) {
|
||||||
|
val folder = credential as? ProviderCredential.FolderPath
|
||||||
|
?: throw ProviderError.InvalidUrl(
|
||||||
|
"Sync folder needs a folder path to sync through"
|
||||||
|
)
|
||||||
|
val dir = checkDir(File(folder.path))
|
||||||
|
pathStore.putString(KEY_ROOT, dir.canonicalPath)
|
||||||
|
_authState.value = AuthState.CONNECTED
|
||||||
|
}
|
||||||
|
|
||||||
|
/**
|
||||||
|
* Create the folder, then connect. Safety rules (a typo must never
|
||||||
|
* spray directories across storage):
|
||||||
|
* - no `..` segments, never a blank path;
|
||||||
|
* - the direct parent must already exist, be a directory, and be
|
||||||
|
* writable — only the final segment is ever created;
|
||||||
|
* - when the folder already exists this is just [connect].
|
||||||
|
*/
|
||||||
|
suspend fun createRoot(rawPath: String) = withContext(Dispatchers.IO) {
|
||||||
|
val trimmed = rawPath.trim()
|
||||||
|
if (trimmed.isEmpty()) throw ProviderError.InvalidUrl("Folder path is empty")
|
||||||
|
val dir = File(trimmed)
|
||||||
|
if (".." in dir.path.split(File.separatorChar)) {
|
||||||
|
throw ProviderError.InvalidUrl("Folder path must not contain '..'")
|
||||||
|
}
|
||||||
|
if (!dir.exists()) {
|
||||||
|
val parent = dir.absoluteFile.parentFile
|
||||||
|
?: throw ProviderError.InvalidUrl("Cannot create folder here: $trimmed")
|
||||||
|
if (!parent.isDirectory) throw ProviderError.InvalidUrl(
|
||||||
|
"Parent is not a folder: ${parent.path}"
|
||||||
|
)
|
||||||
|
if (!parent.canWrite()) throw ProviderError.InvalidUrl(
|
||||||
|
"Parent folder is not writable: ${parent.path}"
|
||||||
|
)
|
||||||
|
if (!dir.mkdirs() && !dir.isDirectory) throw ProviderError.InvalidUrl(
|
||||||
|
"Could not create folder: ${dir.path}"
|
||||||
|
)
|
||||||
|
}
|
||||||
|
val checked = checkDir(dir)
|
||||||
|
pathStore.putString(KEY_ROOT, checked.canonicalPath)
|
||||||
|
_authState.value = AuthState.CONNECTED
|
||||||
|
}
|
||||||
|
|
||||||
|
override suspend fun reconnect(): AuthState = withContext(Dispatchers.IO) {
|
||||||
|
val root = storedRoot()
|
||||||
|
if (root == null) {
|
||||||
|
_authState.value = AuthState.DISCONNECTED
|
||||||
|
return@withContext _authState.value
|
||||||
|
}
|
||||||
|
runCatching { checkDir(root) }
|
||||||
|
.onSuccess { _authState.value = AuthState.CONNECTED }
|
||||||
|
.onFailure { _authState.value = AuthState.DISCONNECTED }
|
||||||
|
_authState.value
|
||||||
|
}
|
||||||
|
|
||||||
|
override suspend fun disconnect(revokeOnServer: Boolean) {
|
||||||
|
// Nothing remote to revoke; the path is kept so reconnect is one tap.
|
||||||
|
_authState.value = AuthState.DISCONNECTED
|
||||||
|
}
|
||||||
|
|
||||||
|
override suspend fun deleteAccountAndData() = withContext(Dispatchers.IO) {
|
||||||
|
// Forget the folder. The files stay — they belong to the user and
|
||||||
|
// their sync tool, and deleting someone's Syncthing folder because
|
||||||
|
// they tapped "disconnect" would be unforgivable.
|
||||||
|
pathStore.remove(KEY_ROOT)
|
||||||
|
_authState.value = AuthState.DISCONNECTED
|
||||||
|
}
|
||||||
|
|
||||||
|
// ---- storage: delegate with rewritten identity --------------------------
|
||||||
|
|
||||||
|
private suspend fun root(): File {
|
||||||
|
if (_authState.value != AuthState.CONNECTED) throw ProviderError.NotConnected()
|
||||||
|
return storedRoot()?.let { checkDir(it) } ?: throw ProviderError.NotConnected()
|
||||||
|
}
|
||||||
|
|
||||||
|
override suspend fun upload(draft: RecordingDraft, onProgress: (Float) -> Unit): RemoteRef =
|
||||||
|
withContext(Dispatchers.IO) {
|
||||||
|
val ref = LocalOnlyProvider(root()).upload(draft, onProgress)
|
||||||
|
ref.copy(providerId = ID)
|
||||||
|
}
|
||||||
|
|
||||||
|
override suspend fun download(ref: RemoteRef, dest: File, onProgress: (Float) -> Unit) =
|
||||||
|
withContext(Dispatchers.IO) {
|
||||||
|
LocalOnlyProvider(root()).download(ref.copy(providerId = LocalOnlyProvider.ID), dest, onProgress)
|
||||||
|
}
|
||||||
|
|
||||||
|
override suspend fun delete(ref: RemoteRef) = withContext(Dispatchers.IO) {
|
||||||
|
LocalOnlyProvider(root()).delete(ref.copy(providerId = LocalOnlyProvider.ID))
|
||||||
|
}
|
||||||
|
|
||||||
|
override suspend fun list(cursor: String?): Page<RemoteRecording> = withContext(Dispatchers.IO) {
|
||||||
|
val page = LocalOnlyProvider(root()).list(cursor)
|
||||||
|
Page(
|
||||||
|
page.items.map { it.copy(ref = it.ref.copy(providerId = ID)) },
|
||||||
|
page.nextCursor,
|
||||||
|
)
|
||||||
|
}
|
||||||
|
|
||||||
|
override suspend fun putSidecar(ref: RemoteRef, kind: SidecarKind, bytes: ByteArray) =
|
||||||
|
withContext(Dispatchers.IO) {
|
||||||
|
LocalOnlyProvider(root())
|
||||||
|
.putSidecar(ref.copy(providerId = LocalOnlyProvider.ID), kind, bytes)
|
||||||
|
}
|
||||||
|
|
||||||
|
override suspend fun getSidecar(ref: RemoteRef, kind: SidecarKind): ByteArray? =
|
||||||
|
withContext(Dispatchers.IO) {
|
||||||
|
LocalOnlyProvider(root())
|
||||||
|
.getSidecar(ref.copy(providerId = LocalOnlyProvider.ID), kind)
|
||||||
|
}
|
||||||
|
|
||||||
|
override suspend fun storageLocationSummary(): StorageLocation =
|
||||||
|
withContext(Dispatchers.IO) {
|
||||||
|
val r = root()
|
||||||
|
val audio = File(r, "audio").listFiles()?.filter { it.isFile } ?: emptyList()
|
||||||
|
StorageLocation(
|
||||||
|
headline = "Sync folder",
|
||||||
|
detail = "${audio.size} recordings — synced by your sync tool, not by this app.",
|
||||||
|
syncedCount = 0,
|
||||||
|
localOnlyCount = audio.size,
|
||||||
|
bytesUsed = audio.sumOf { it.length() },
|
||||||
|
path = r.path,
|
||||||
|
)
|
||||||
|
}
|
||||||
|
|
||||||
|
companion object {
|
||||||
|
const val ID = "sync-folder"
|
||||||
|
const val KEY_ROOT = "provider.sync-folder.root"
|
||||||
|
}
|
||||||
|
}
|
||||||
249
shared/com/shonar/provider/LocalOnlyProvider.kt
Normal file
249
shared/com/shonar/provider/LocalOnlyProvider.kt
Normal file
|
|
@ -0,0 +1,249 @@
|
||||||
|
package com.shonar.provider
|
||||||
|
|
||||||
|
import kotlinx.coroutines.flow.MutableStateFlow
|
||||||
|
import kotlinx.coroutines.flow.StateFlow
|
||||||
|
import java.io.File
|
||||||
|
|
||||||
|
/**
|
||||||
|
* Local-only provider: everything stays in app-private storage. First-class
|
||||||
|
* so the app has no "no server" special case. No networking code paths at
|
||||||
|
* all — provable by construction (this file imports no HTTP library).
|
||||||
|
*
|
||||||
|
* The storage root is the app-private [defaultRoot], unless the user picked
|
||||||
|
* their own folder ([connect]/[createRoot] with a [ProviderCredential.FolderPath]):
|
||||||
|
* then files live there and [deleteAccountAndData] only forgets the path —
|
||||||
|
* a user folder is never deleted by this app.
|
||||||
|
*/
|
||||||
|
class LocalOnlyProvider(
|
||||||
|
private val defaultRoot: File,
|
||||||
|
private val pathStore: com.shonar.settings.SettingsStore =
|
||||||
|
com.shonar.settings.InMemorySettingsStore(),
|
||||||
|
private val rootKey: String = KEY_ROOT,
|
||||||
|
) : ShonarProvider {
|
||||||
|
|
||||||
|
override val descriptor = ProviderDescriptor(
|
||||||
|
id = ID,
|
||||||
|
displayName = "Local-only storage",
|
||||||
|
capabilities = setOf(), // no chunked protocol needed; no server AI
|
||||||
|
)
|
||||||
|
|
||||||
|
private val _authState = MutableStateFlow(AuthState.CONNECTED)
|
||||||
|
override val authState: StateFlow<AuthState> = _authState
|
||||||
|
|
||||||
|
private suspend fun storedRoot(): File? {
|
||||||
|
val raw = pathStore.getString(rootKey) ?: return null
|
||||||
|
return File(raw).takeIf { it.exists() }
|
||||||
|
}
|
||||||
|
|
||||||
|
/** App-private dir by default; the user's own folder once picked. */
|
||||||
|
private suspend fun effectiveRoot(): File = storedRoot() ?: defaultRoot
|
||||||
|
|
||||||
|
/** Where new recordings belong right now. Read by the recording repository. */
|
||||||
|
suspend fun currentRoot(): File = effectiveRoot()
|
||||||
|
|
||||||
|
private fun checkDir(dir: File): File {
|
||||||
|
if (".." in dir.path.split(File.separatorChar)) {
|
||||||
|
throw ProviderError.InvalidUrl("Folder path must not contain '..'")
|
||||||
|
}
|
||||||
|
if (!dir.exists()) throw ProviderError.InvalidUrl(
|
||||||
|
"Folder does not exist: ${dir.path}."
|
||||||
|
)
|
||||||
|
if (!dir.isDirectory) throw ProviderError.InvalidUrl("Not a folder: ${dir.path}")
|
||||||
|
if (!dir.canRead() || !dir.canWrite()) throw ProviderError.InvalidUrl(
|
||||||
|
"Folder is not readable and writable: ${dir.path}"
|
||||||
|
)
|
||||||
|
return dir
|
||||||
|
}
|
||||||
|
|
||||||
|
override suspend fun probe(baseUrl: ServerUrl): ProbeResult =
|
||||||
|
ProbeResult.Incompatible // local-only never talks to servers
|
||||||
|
|
||||||
|
override suspend fun connect(credential: ProviderCredential) {
|
||||||
|
when (credential) {
|
||||||
|
ProviderCredential.None -> _authState.value = AuthState.CONNECTED
|
||||||
|
is ProviderCredential.FolderPath -> {
|
||||||
|
val dir = checkDir(File(credential.path))
|
||||||
|
pathStore.putString(rootKey, dir.canonicalPath)
|
||||||
|
_authState.value = AuthState.CONNECTED
|
||||||
|
}
|
||||||
|
else -> throw ProviderError.InvalidUrl(
|
||||||
|
"Local-only storage takes no credentials"
|
||||||
|
)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
/**
|
||||||
|
* Use the app-private folder again (forget a previously picked folder;
|
||||||
|
* its files stay where they are).
|
||||||
|
*/
|
||||||
|
suspend fun useDefaultRoot() {
|
||||||
|
pathStore.remove(rootKey)
|
||||||
|
_authState.value = AuthState.CONNECTED
|
||||||
|
}
|
||||||
|
|
||||||
|
/**
|
||||||
|
* Create the folder, then use it. Same safety rules as the sync
|
||||||
|
* folder: only the final segment is ever created, the parent must
|
||||||
|
* already exist, no `..`. When it exists this is just [connect].
|
||||||
|
*/
|
||||||
|
suspend fun createRoot(rawPath: String) {
|
||||||
|
val trimmed = rawPath.trim()
|
||||||
|
if (trimmed.isEmpty()) throw ProviderError.InvalidUrl("Folder path is empty")
|
||||||
|
val dir = File(trimmed)
|
||||||
|
if (".." in dir.path.split(File.separatorChar)) {
|
||||||
|
throw ProviderError.InvalidUrl("Folder path must not contain '..'")
|
||||||
|
}
|
||||||
|
if (!dir.exists()) {
|
||||||
|
val parent = dir.absoluteFile.parentFile
|
||||||
|
?: throw ProviderError.InvalidUrl("Cannot create folder here: $trimmed")
|
||||||
|
if (!parent.isDirectory) throw ProviderError.InvalidUrl(
|
||||||
|
"Parent is not a folder: ${parent.path}"
|
||||||
|
)
|
||||||
|
if (!parent.canWrite()) throw ProviderError.InvalidUrl(
|
||||||
|
"Parent folder is not writable: ${parent.path}"
|
||||||
|
)
|
||||||
|
if (!dir.mkdirs() && !dir.isDirectory) throw ProviderError.InvalidUrl(
|
||||||
|
"Could not create folder: ${dir.path}"
|
||||||
|
)
|
||||||
|
}
|
||||||
|
val checked = checkDir(dir)
|
||||||
|
pathStore.putString(rootKey, checked.canonicalPath)
|
||||||
|
_authState.value = AuthState.CONNECTED
|
||||||
|
}
|
||||||
|
|
||||||
|
override suspend fun reconnect(): AuthState {
|
||||||
|
// A picked folder may have been deleted out from under us; fall
|
||||||
|
// back to the app-private dir rather than stranding recordings.
|
||||||
|
_authState.value = AuthState.CONNECTED
|
||||||
|
return _authState.value
|
||||||
|
}
|
||||||
|
|
||||||
|
override suspend fun disconnect(revokeOnServer: Boolean) {
|
||||||
|
// Nothing to revoke; DISCONNECTED means "user left local mode" — the
|
||||||
|
// sync layer treats it as paused, files remain on disk.
|
||||||
|
_authState.value = AuthState.DISCONNECTED
|
||||||
|
}
|
||||||
|
|
||||||
|
override suspend fun deleteAccountAndData() {
|
||||||
|
if (storedRoot() != null) {
|
||||||
|
// A user-picked folder belongs to the user — forget it, never
|
||||||
|
// delete it.
|
||||||
|
pathStore.remove(rootKey)
|
||||||
|
} else if (defaultRoot.exists()) {
|
||||||
|
defaultRoot.deleteRecursively()
|
||||||
|
}
|
||||||
|
_authState.value = AuthState.DISCONNECTED
|
||||||
|
}
|
||||||
|
|
||||||
|
// ---- storage -----------------------------------------------------------
|
||||||
|
|
||||||
|
private suspend fun audioFile(ref: RemoteRef) = File(effectiveRoot(), ref.key)
|
||||||
|
private suspend fun sidecarFile(ref: RemoteRef, kind: SidecarKind) =
|
||||||
|
File(effectiveRoot(), "sidecars/${ref.key}/${kind.fileName}")
|
||||||
|
|
||||||
|
override suspend fun upload(draft: RecordingDraft, onProgress: (Float) -> Unit): RemoteRef {
|
||||||
|
val key = "audio/${draft.id}"
|
||||||
|
val dest = File(effectiveRoot(), key)
|
||||||
|
dest.parentFile?.mkdirs()
|
||||||
|
// "Upload" locally = copy; report progress in slices so UI behaves uniformly
|
||||||
|
draft.sourceFile.inputStream().use { input ->
|
||||||
|
dest.outputStream().use { output ->
|
||||||
|
val buf = ByteArray(64 * 1024)
|
||||||
|
val total = draft.sizeBytes.coerceAtLeast(1)
|
||||||
|
var written = 0L
|
||||||
|
while (true) {
|
||||||
|
val n = input.read(buf)
|
||||||
|
if (n < 0) break
|
||||||
|
output.write(buf, 0, n)
|
||||||
|
written += n
|
||||||
|
onProgress((written.toFloat() / total).coerceIn(0f, 1f))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return RemoteRef(ID, key, etag = "sha256:" + dest.sha256Hex(), sizeBytes = dest.length())
|
||||||
|
}
|
||||||
|
|
||||||
|
override suspend fun download(ref: RemoteRef, dest: File, onProgress: (Float) -> Unit) {
|
||||||
|
val src = audioFile(ref)
|
||||||
|
if (!src.exists()) throw ProviderError.NotFound(ref.key)
|
||||||
|
src.copyTo(dest, overwrite = true)
|
||||||
|
onProgress(1f)
|
||||||
|
}
|
||||||
|
|
||||||
|
override suspend fun delete(ref: RemoteRef) {
|
||||||
|
audioFile(ref).delete()
|
||||||
|
File(effectiveRoot(), "sidecars/${ref.key}").deleteRecursively()
|
||||||
|
}
|
||||||
|
|
||||||
|
override suspend fun list(cursor: String?): Page<RemoteRecording> {
|
||||||
|
val audioRoot = File(effectiveRoot(), "audio")
|
||||||
|
val files = audioRoot.listFiles()?.filter { it.isFile } ?: emptyList()
|
||||||
|
// single page; no paging for local storage
|
||||||
|
val items = files.map { f ->
|
||||||
|
RemoteRecording(
|
||||||
|
ref = RemoteRef(ID, "audio/${f.name}", etag = "sha256:" + f.sha256Hex(), sizeBytes = f.length()),
|
||||||
|
title = f.nameWithoutExtension,
|
||||||
|
createdAtEpochMs = f.lastModified(),
|
||||||
|
durationMs = 0, // duration is tracked in Room, not on disk
|
||||||
|
mime = "application/octet-stream",
|
||||||
|
)
|
||||||
|
}
|
||||||
|
return Page(items, nextCursor = null)
|
||||||
|
}
|
||||||
|
|
||||||
|
override suspend fun putSidecar(ref: RemoteRef, kind: SidecarKind, bytes: ByteArray) {
|
||||||
|
val f = sidecarFile(ref, kind)
|
||||||
|
f.parentFile?.mkdirs()
|
||||||
|
f.writeBytes(bytes)
|
||||||
|
}
|
||||||
|
|
||||||
|
override suspend fun getSidecar(ref: RemoteRef, kind: SidecarKind): ByteArray? {
|
||||||
|
val f = sidecarFile(ref, kind)
|
||||||
|
return if (f.exists()) f.readBytes() else null
|
||||||
|
}
|
||||||
|
|
||||||
|
override suspend fun storageLocationSummary(): StorageLocation {
|
||||||
|
val raw = pathStore.getString(rootKey)
|
||||||
|
val custom = storedRoot()
|
||||||
|
if (raw != null && custom == null) {
|
||||||
|
return StorageLocation(
|
||||||
|
headline = "Folder unavailable",
|
||||||
|
detail = "Your picked folder is gone or unreadable. " +
|
||||||
|
"Reconnect storage or pick again — new recordings use the app folder meanwhile.",
|
||||||
|
syncedCount = 0,
|
||||||
|
localOnlyCount = 0,
|
||||||
|
bytesUsed = 0,
|
||||||
|
path = raw,
|
||||||
|
)
|
||||||
|
}
|
||||||
|
val r = custom ?: defaultRoot
|
||||||
|
val audio = File(r, "audio").listFiles()?.filter { it.isFile } ?: emptyList()
|
||||||
|
return StorageLocation(
|
||||||
|
headline = if (custom != null) "Your folder" else "On this device only",
|
||||||
|
detail = if (custom != null) "Recordings stay in your folder and never leave your phone."
|
||||||
|
else "Recordings and transcripts never leave your phone.",
|
||||||
|
syncedCount = 0,
|
||||||
|
localOnlyCount = audio.size,
|
||||||
|
bytesUsed = audio.sumOf { it.length() },
|
||||||
|
path = r.absolutePath,
|
||||||
|
)
|
||||||
|
}
|
||||||
|
|
||||||
|
companion object {
|
||||||
|
const val ID = "local-only"
|
||||||
|
const val KEY_ROOT = "provider.local-only.root"
|
||||||
|
|
||||||
|
private fun File.sha256Hex(): String {
|
||||||
|
val md = java.security.MessageDigest.getInstance("SHA-256")
|
||||||
|
inputStream().use { inn ->
|
||||||
|
val buf = ByteArray(64 * 1024)
|
||||||
|
while (true) {
|
||||||
|
val n = inn.read(buf)
|
||||||
|
if (n < 0) break
|
||||||
|
md.update(buf, 0, n)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return md.digest().joinToString("") { "%02x".format(it) }
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
157
shared/com/shonar/provider/NextcloudAuth.kt
Normal file
157
shared/com/shonar/provider/NextcloudAuth.kt
Normal file
|
|
@ -0,0 +1,157 @@
|
||||||
|
package com.shonar.provider
|
||||||
|
|
||||||
|
import java.util.concurrent.TimeUnit
|
||||||
|
import kotlinx.coroutines.Dispatchers
|
||||||
|
import kotlinx.coroutines.withContext
|
||||||
|
import okhttp3.MediaType.Companion.toMediaType
|
||||||
|
import okhttp3.OkHttpClient
|
||||||
|
import okhttp3.Request
|
||||||
|
import okhttp3.RequestBody.Companion.toRequestBody
|
||||||
|
import org.json.JSONObject
|
||||||
|
|
||||||
|
/**
|
||||||
|
* Nextcloud server identification + login flow v2 (docs/server-providers.md
|
||||||
|
* §3, §6). No credentials are ever sent here: [startLogin] is anonymous,
|
||||||
|
* [poll] carries only the one-time poll token.
|
||||||
|
*
|
||||||
|
* Login flow v2 (official, supported):
|
||||||
|
* 1. POST {base}/index.php/login/v2 -> {poll:{token,endpoint}, login:url}
|
||||||
|
* 2. user approves {login} in a browser (server-owned consent screen)
|
||||||
|
* 3. POST {poll.endpoint} {"token":...} -> 404 while pending,
|
||||||
|
* 200 {server, loginName, appPassword} once approved
|
||||||
|
*/
|
||||||
|
class NextcloudAuth(
|
||||||
|
private val client: OkHttpClient = defaultClient(),
|
||||||
|
private val tofu: TofuManager? = null,
|
||||||
|
) {
|
||||||
|
|
||||||
|
/** Is there a Nextcloud at [url]? Read-only, no credentials. */
|
||||||
|
suspend fun probe(url: ServerUrl): ProbeResult = withContext(Dispatchers.IO) {
|
||||||
|
val req = Request.Builder().url(url.origin + "/status.php").get().build()
|
||||||
|
try {
|
||||||
|
client.newCall(req).execute().use { resp ->
|
||||||
|
if (resp.code != 200) return@withContext ProbeResult.Incompatible
|
||||||
|
val body = JSONObject(resp.body?.string().orEmpty())
|
||||||
|
// status.php: {productname, versionstring, ...}. Some forks
|
||||||
|
// report "Nextcloud" with different casing — accept any.
|
||||||
|
val product = body.optString("productname", "")
|
||||||
|
if (!product.equals("nextcloud", ignoreCase = true)) {
|
||||||
|
return@withContext ProbeResult.Incompatible
|
||||||
|
}
|
||||||
|
ProbeResult.Compatible(
|
||||||
|
descriptor = ProviderDescriptor(
|
||||||
|
id = ProviderRegistry.NEXTCLOUD_ID,
|
||||||
|
displayName = "Nextcloud",
|
||||||
|
isDefault = true,
|
||||||
|
capabilities = setOf(
|
||||||
|
ProviderDescriptor.Capability.CHUNKED_UPLOAD,
|
||||||
|
ProviderDescriptor.Capability.QUOTA_INFO,
|
||||||
|
),
|
||||||
|
),
|
||||||
|
serverName = product,
|
||||||
|
version = body.optString("versionstring", "unknown"),
|
||||||
|
)
|
||||||
|
}
|
||||||
|
} catch (e: javax.net.ssl.SSLHandshakeException) {
|
||||||
|
ProbeResult.TlsFailure(tofu?.failureFor(url.host)?.spkiHex ?: "unknown")
|
||||||
|
} catch (e: javax.net.ssl.SSLPeerUnverifiedException) {
|
||||||
|
ProbeResult.TlsFailure(tofu?.failureFor(url.host)?.spkiHex ?: "unknown")
|
||||||
|
} catch (e: Exception) {
|
||||||
|
ProbeResult.NetworkError(e.javaClass.simpleName)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
/** Begins login flow v2. Returns the browser URL + poll handle. */
|
||||||
|
suspend fun startLogin(url: ServerUrl): LoginFlowSession = withContext(Dispatchers.IO) {
|
||||||
|
val req = Request.Builder().url(url.origin + "/index.php/login/v2")
|
||||||
|
.post(ByteArray(0).toRequestBody(null)).build()
|
||||||
|
try {
|
||||||
|
client.newCall(req).execute().use { resp ->
|
||||||
|
if (resp.code != 200) {
|
||||||
|
throw ProviderError.Transient(
|
||||||
|
"Server refused the login request (HTTP ${resp.code}). " +
|
||||||
|
"Is this a Nextcloud?"
|
||||||
|
)
|
||||||
|
}
|
||||||
|
val body = JSONObject(resp.body?.string().orEmpty())
|
||||||
|
val poll = body.optJSONObject("poll")
|
||||||
|
val login = body.optString("login", "")
|
||||||
|
val token = poll?.optString("token", "").orEmpty()
|
||||||
|
val endpoint = poll?.optString("endpoint", "").orEmpty()
|
||||||
|
if (login.isBlank() || token.isBlank() || endpoint.isBlank()) {
|
||||||
|
throw ProviderError.Transient("Server gave an incomplete login response")
|
||||||
|
}
|
||||||
|
LoginFlowSession(
|
||||||
|
baseUrl = url.origin,
|
||||||
|
loginUrl = login,
|
||||||
|
pollToken = token,
|
||||||
|
pollEndpoint = endpoint,
|
||||||
|
)
|
||||||
|
}
|
||||||
|
} catch (e: ProviderError) {
|
||||||
|
throw e
|
||||||
|
} catch (e: Exception) {
|
||||||
|
throw ProviderError.Transient("Could not reach the server (${e.javaClass.simpleName})")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
/** One poll attempt. Call repeatedly until [PollResult.Approved]. */
|
||||||
|
suspend fun poll(flow: LoginFlowSession): PollResult = withContext(Dispatchers.IO) {
|
||||||
|
val json = """{"token":"${flow.pollToken}"}"""
|
||||||
|
val req = Request.Builder().url(flow.pollEndpoint)
|
||||||
|
.post(json.toRequestBody("application/json; charset=utf-8".toMediaType())).build()
|
||||||
|
try {
|
||||||
|
client.newCall(req).execute().use { resp ->
|
||||||
|
when (resp.code) {
|
||||||
|
200 -> {
|
||||||
|
val body = JSONObject(resp.body?.string().orEmpty())
|
||||||
|
val server = body.optString("server", flow.baseUrl)
|
||||||
|
val name = body.optString("loginName", "")
|
||||||
|
val pass = body.optString("appPassword", "")
|
||||||
|
if (name.isBlank() || pass.isBlank()) {
|
||||||
|
return@withContext PollResult.Failed("Server approved but sent no credentials")
|
||||||
|
}
|
||||||
|
PollResult.Approved(
|
||||||
|
ProviderCredential.AppPassword(
|
||||||
|
accountLabel = "$name@${baseHost(server)}",
|
||||||
|
loginUrl = server,
|
||||||
|
user = name,
|
||||||
|
password = pass,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
}
|
||||||
|
404 -> PollResult.Pending
|
||||||
|
else -> PollResult.Failed("Login poll failed (HTTP ${resp.code})")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
} catch (e: Exception) {
|
||||||
|
PollResult.Failed("Login poll failed (${e.javaClass.simpleName})")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
companion object {
|
||||||
|
fun defaultClient(): OkHttpClient = OkHttpClient.Builder()
|
||||||
|
.connectTimeout(10, TimeUnit.SECONDS)
|
||||||
|
.readTimeout(20, TimeUnit.SECONDS)
|
||||||
|
.followRedirects(false)
|
||||||
|
.build()
|
||||||
|
|
||||||
|
private fun baseHost(server: String): String =
|
||||||
|
runCatching { java.net.URI(server).host }.getOrNull() ?: server
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
/** Browser URL + one-time poll handle from [NextcloudAuth.startLogin]. */
|
||||||
|
data class LoginFlowSession(
|
||||||
|
val baseUrl: String,
|
||||||
|
val loginUrl: String,
|
||||||
|
val pollToken: String,
|
||||||
|
val pollEndpoint: String,
|
||||||
|
)
|
||||||
|
|
||||||
|
/** One poll attempt's outcome. The token itself never appears in messages. */
|
||||||
|
sealed class PollResult {
|
||||||
|
data object Pending : PollResult()
|
||||||
|
data class Approved(val credential: ProviderCredential.AppPassword) : PollResult()
|
||||||
|
data class Failed(val reason: String) : PollResult()
|
||||||
|
}
|
||||||
55
shared/com/shonar/provider/NextcloudAuthStore.kt
Normal file
55
shared/com/shonar/provider/NextcloudAuthStore.kt
Normal file
|
|
@ -0,0 +1,55 @@
|
||||||
|
package com.shonar.provider
|
||||||
|
|
||||||
|
import com.shonar.settings.SettingsStore
|
||||||
|
|
||||||
|
/**
|
||||||
|
* Secure persistence for the Nextcloud session: server origin, DAV user id,
|
||||||
|
* display name, and the app password from login flow v2. The user's *normal*
|
||||||
|
* account password is never requested, typed, or stored — only the
|
||||||
|
* server-issued app password lives here (docs/server-providers.md §3).
|
||||||
|
*/
|
||||||
|
class NextcloudAuthStore(private val secure: SettingsStore) {
|
||||||
|
|
||||||
|
suspend fun save(session: NcSession) {
|
||||||
|
secure.putString(KEY_BASE_URL, session.baseUrl)
|
||||||
|
secure.putString(KEY_USER_ID, session.userId)
|
||||||
|
secure.putString(KEY_USERNAME, session.username)
|
||||||
|
secure.putString(KEY_APP_PASSWORD, session.appPassword)
|
||||||
|
}
|
||||||
|
|
||||||
|
suspend fun load(): NcSession? {
|
||||||
|
val base = secure.getString(KEY_BASE_URL) ?: return null
|
||||||
|
val pass = secure.getString(KEY_APP_PASSWORD) ?: return null
|
||||||
|
return NcSession(
|
||||||
|
baseUrl = base,
|
||||||
|
userId = secure.getString(KEY_USER_ID).orEmpty(),
|
||||||
|
username = secure.getString(KEY_USERNAME).orEmpty(),
|
||||||
|
appPassword = pass,
|
||||||
|
)
|
||||||
|
}
|
||||||
|
|
||||||
|
suspend fun clear() {
|
||||||
|
secure.remove(KEY_BASE_URL)
|
||||||
|
secure.remove(KEY_USER_ID)
|
||||||
|
secure.remove(KEY_USERNAME)
|
||||||
|
secure.remove(KEY_APP_PASSWORD)
|
||||||
|
}
|
||||||
|
|
||||||
|
companion object {
|
||||||
|
private const val PREFIX = "provider.nextcloud."
|
||||||
|
const val KEY_BASE_URL = PREFIX + "base_url"
|
||||||
|
const val KEY_USER_ID = PREFIX + "user_id"
|
||||||
|
const val KEY_USERNAME = PREFIX + "username"
|
||||||
|
const val KEY_APP_PASSWORD = PREFIX + "app_password"
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
/** Authenticated Nextcloud session. Never logged (see toString). */
|
||||||
|
data class NcSession(
|
||||||
|
val baseUrl: String,
|
||||||
|
val userId: String, // DAV path user (OCS `id`, not the display name)
|
||||||
|
val username: String, // human hint only
|
||||||
|
val appPassword: String,
|
||||||
|
) {
|
||||||
|
override fun toString(): String = "NcSession(server=$baseUrl, user=$username, [redacted])"
|
||||||
|
}
|
||||||
644
shared/com/shonar/provider/NextcloudProvider.kt
Normal file
644
shared/com/shonar/provider/NextcloudProvider.kt
Normal file
|
|
@ -0,0 +1,644 @@
|
||||||
|
package com.shonar.provider
|
||||||
|
|
||||||
|
import java.io.File
|
||||||
|
import java.time.format.DateTimeFormatter
|
||||||
|
import java.util.UUID
|
||||||
|
import kotlinx.coroutines.Dispatchers
|
||||||
|
import kotlinx.coroutines.ensureActive
|
||||||
|
import kotlinx.coroutines.flow.MutableStateFlow
|
||||||
|
import kotlinx.coroutines.flow.StateFlow
|
||||||
|
import kotlinx.coroutines.withContext
|
||||||
|
import okhttp3.MediaType.Companion.toMediaType
|
||||||
|
import okhttp3.OkHttpClient
|
||||||
|
import okhttp3.Request
|
||||||
|
import okhttp3.RequestBody.Companion.toRequestBody
|
||||||
|
import okhttp3.Response
|
||||||
|
import org.json.JSONObject
|
||||||
|
|
||||||
|
/**
|
||||||
|
* P4: Nextcloud provider — the product default. Talks only official,
|
||||||
|
* supported endpoints (docs/server-providers.md §6):
|
||||||
|
*
|
||||||
|
* - GET {base}/status.php (probe)
|
||||||
|
* - login flow v2 (see [NextcloudAuth])
|
||||||
|
* - GET {base}/ocs/v2.php/cloud/user (identity + quota)
|
||||||
|
* - DELETE {base}/ocs/v2.php/core/apppassword (revoke own app password)
|
||||||
|
* - WebDAV {base}/remote.php/dav/files/{user}/… (PROPFIND/GET/PUT/MKCOL/MOVE/DELETE)
|
||||||
|
* - Chunked upload v2 {base}/remote.php/dav/uploads/{user}/{transfer}/
|
||||||
|
* (MKCOL, PUT chunks 00001..N, MOVE {transfer}/.file -> destination)
|
||||||
|
*
|
||||||
|
* Layout: `SHONAR/audio/{uuid}.m4a`, `SHONAR/sidecars/{uuid}/{kind}.json`.
|
||||||
|
* Originals are never overwritten by processing artifacts — uploads with
|
||||||
|
* the same draft id MOVE onto the same key (idempotent replace).
|
||||||
|
*
|
||||||
|
* Chunk naming follows the developer manual: chunks are numbered 1..10000
|
||||||
|
* and assembled in name order, so names are zero-padded to 5 digits
|
||||||
|
* ("00001".."10000") to keep lexical order == numeric order. Chunk size
|
||||||
|
* defaults to 16 MiB — the server requires 5 MiB..5 GiB per chunk (last
|
||||||
|
* chunk exempt), so never lower the default for production use; the
|
||||||
|
* constructor parameter exists for tests only.
|
||||||
|
*
|
||||||
|
* Transfer ids are deterministic per recording (`shonar-{draft.id}`), so a
|
||||||
|
* killed upload resumes by PROPFIND-ing the transfer folder and skipping
|
||||||
|
* present chunks — including across process restarts.
|
||||||
|
*/
|
||||||
|
class NextcloudProvider(
|
||||||
|
private val auth: NextcloudAuthStore,
|
||||||
|
private val client: OkHttpClient = NextcloudAuth.defaultClient(),
|
||||||
|
private val loginFlow: NextcloudAuth = NextcloudAuth(client),
|
||||||
|
private val chunkSizeBytes: Long = 16 * 1024 * 1024,
|
||||||
|
) : ShonarProvider {
|
||||||
|
|
||||||
|
override val descriptor = ProviderDescriptor(
|
||||||
|
id = ProviderRegistry.NEXTCLOUD_ID,
|
||||||
|
displayName = "Nextcloud",
|
||||||
|
isDefault = true,
|
||||||
|
capabilities = setOf(
|
||||||
|
ProviderDescriptor.Capability.CHUNKED_UPLOAD,
|
||||||
|
ProviderDescriptor.Capability.QUOTA_INFO,
|
||||||
|
),
|
||||||
|
)
|
||||||
|
|
||||||
|
private val _authState = MutableStateFlow(AuthState.DISCONNECTED)
|
||||||
|
override val authState: StateFlow<AuthState> = _authState
|
||||||
|
|
||||||
|
// ---- lifecycle ---------------------------------------------------------
|
||||||
|
|
||||||
|
override suspend fun probe(baseUrl: ServerUrl): ProbeResult = loginFlow.probe(baseUrl)
|
||||||
|
|
||||||
|
override suspend fun connect(credential: ProviderCredential) = withContext(Dispatchers.IO) {
|
||||||
|
val app = credential as? ProviderCredential.AppPassword
|
||||||
|
?: throw ProviderError.InvalidUrl(
|
||||||
|
"Nextcloud connects with an app password from the browser login flow"
|
||||||
|
)
|
||||||
|
val origin = ServerUrl.parse(app.loginUrl).getOrNull()?.origin
|
||||||
|
?: throw ProviderError.InvalidUrl("Not a valid server URL")
|
||||||
|
// Validate before persisting: a wrong password must not overwrite a
|
||||||
|
// working session.
|
||||||
|
val userId = try {
|
||||||
|
ocsUserId(origin, app.user, app.password)
|
||||||
|
} catch (e: ProviderError.AuthExpired) {
|
||||||
|
throw ProviderError.Transient(
|
||||||
|
"Nextcloud rejected the login — approve it in the browser again"
|
||||||
|
)
|
||||||
|
}
|
||||||
|
auth.save(
|
||||||
|
NcSession(
|
||||||
|
baseUrl = origin,
|
||||||
|
userId = userId,
|
||||||
|
username = app.user,
|
||||||
|
appPassword = app.password,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
_authState.value = AuthState.CONNECTED
|
||||||
|
}
|
||||||
|
|
||||||
|
override suspend fun reconnect(): AuthState = withContext(Dispatchers.IO) {
|
||||||
|
val saved = auth.load()
|
||||||
|
if (saved == null || saved.appPassword.isBlank() || saved.userId.isBlank()) {
|
||||||
|
_authState.value = AuthState.DISCONNECTED
|
||||||
|
return@withContext _authState.value
|
||||||
|
}
|
||||||
|
try {
|
||||||
|
ocsUserId(saved.baseUrl, saved.userId, saved.appPassword)
|
||||||
|
_authState.value = AuthState.CONNECTED
|
||||||
|
} catch (e: ProviderError.AuthExpired) {
|
||||||
|
_authState.value = AuthState.EXPIRED
|
||||||
|
} catch (e: ProviderError) {
|
||||||
|
_authState.value = AuthState.OFFLINE
|
||||||
|
}
|
||||||
|
_authState.value
|
||||||
|
}
|
||||||
|
|
||||||
|
override suspend fun disconnect(revokeOnServer: Boolean) = withContext(Dispatchers.IO) {
|
||||||
|
if (revokeOnServer) {
|
||||||
|
// Best effort: local state is cleared even if revoke fails.
|
||||||
|
runCatching {
|
||||||
|
val saved = auth.load()
|
||||||
|
if (saved != null && saved.appPassword.isNotBlank()) {
|
||||||
|
dav(saved, "DELETE", ocsPath("/core/apppassword"), null).close()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
auth.clear()
|
||||||
|
_authState.value = AuthState.DISCONNECTED
|
||||||
|
}
|
||||||
|
|
||||||
|
/**
|
||||||
|
* Revokes the app password and forgets the session. There is no API for
|
||||||
|
* deleting the whole Nextcloud account — that stays a manual step on the
|
||||||
|
* server, and the message says so.
|
||||||
|
*/
|
||||||
|
override suspend fun deleteAccountAndData() = withContext(Dispatchers.IO) {
|
||||||
|
disconnect(revokeOnServer = true)
|
||||||
|
}
|
||||||
|
|
||||||
|
// ---- storage -----------------------------------------------------------
|
||||||
|
|
||||||
|
override suspend fun upload(draft: RecordingDraft, onProgress: (Float) -> Unit): RemoteRef =
|
||||||
|
withContext(Dispatchers.IO) {
|
||||||
|
val session = connectedSession()
|
||||||
|
val key = "SHONAR/audio/${draft.id}${extFor(draft.mime)}"
|
||||||
|
val total = draft.sizeBytes.coerceAtLeast(1)
|
||||||
|
ensureDir(session, "SHONAR")
|
||||||
|
ensureDir(session, "SHONAR/audio")
|
||||||
|
|
||||||
|
val transfer = "shonar-${draft.id}"
|
||||||
|
mkcol(session, transfer)
|
||||||
|
val received = transferChunks(session, transfer)
|
||||||
|
val chunkCount = ((draft.sizeBytes + chunkSizeBytes - 1) / chunkSizeBytes)
|
||||||
|
.toInt().coerceAtLeast(1)
|
||||||
|
var sent = 0L
|
||||||
|
for (i in received) {
|
||||||
|
sent += if (i < chunkCount) chunkSizeBytes else 0L
|
||||||
|
}
|
||||||
|
sent = sent.coerceAtMost(draft.sizeBytes)
|
||||||
|
if (sent > 0) onProgress((sent.toFloat() / total).coerceIn(0f, 1f))
|
||||||
|
|
||||||
|
var idx = 1
|
||||||
|
while (idx <= chunkCount) {
|
||||||
|
ensureActive()
|
||||||
|
if (idx !in received) {
|
||||||
|
val slice = readSlice(draft.sourceFile, (idx - 1) * chunkSizeBytes, chunkSizeBytes)
|
||||||
|
putChunk(session, transfer, idx, slice, draft.sizeBytes, key)
|
||||||
|
sent += slice.size
|
||||||
|
onProgress((sent.toFloat() / total).coerceIn(0f, 1f))
|
||||||
|
}
|
||||||
|
idx++
|
||||||
|
}
|
||||||
|
assemble(session, transfer, key, draft)
|
||||||
|
// Best-effort cleanup of the transfer folder; the server also
|
||||||
|
// expires stale upload dirs on its own.
|
||||||
|
runCatching {
|
||||||
|
dav(session, "DELETE", uploadsPath(session, transfer) + "/", null).close()
|
||||||
|
}
|
||||||
|
onProgress(1f)
|
||||||
|
RemoteRef(ProviderRegistry.NEXTCLOUD_ID, key, etag = null, sizeBytes = draft.sizeBytes)
|
||||||
|
}
|
||||||
|
|
||||||
|
override suspend fun download(ref: RemoteRef, dest: File, onProgress: (Float) -> Unit) =
|
||||||
|
withContext(Dispatchers.IO) {
|
||||||
|
val session = connectedSession()
|
||||||
|
dav(session, "GET", filesPath(session, ref.key), null).use { resp ->
|
||||||
|
when (resp.code) {
|
||||||
|
200 -> {
|
||||||
|
val body = resp.body ?: throw ProviderError.Transient("Empty download response")
|
||||||
|
val total = body.contentLength().takeIf { it > 0 } ?: -1
|
||||||
|
dest.parentFile?.mkdirs()
|
||||||
|
body.byteStream().use { input ->
|
||||||
|
dest.outputStream().use { output ->
|
||||||
|
val buf = ByteArray(64 * 1024)
|
||||||
|
var written = 0L
|
||||||
|
var last = -1f
|
||||||
|
while (true) {
|
||||||
|
val n = input.read(buf)
|
||||||
|
if (n < 0) break
|
||||||
|
output.write(buf, 0, n)
|
||||||
|
written += n
|
||||||
|
if (total > 0) {
|
||||||
|
val p = (written.toFloat() / total).coerceIn(0f, 1f)
|
||||||
|
if (p > last) {
|
||||||
|
onProgress(p)
|
||||||
|
last = p
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
onProgress(1f)
|
||||||
|
}
|
||||||
|
401 -> throw ProviderError.AuthExpired()
|
||||||
|
404 -> throw ProviderError.NotFound(ref.key)
|
||||||
|
else -> throw ProviderError.Transient("Download failed (HTTP ${resp.code})")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
override suspend fun delete(ref: RemoteRef) {
|
||||||
|
withContext(Dispatchers.IO) {
|
||||||
|
val session = connectedSession()
|
||||||
|
dav(session, "DELETE", filesPath(session, ref.key), null).use { resp ->
|
||||||
|
when (resp.code) {
|
||||||
|
200, 201, 204 -> Unit
|
||||||
|
401 -> throw ProviderError.AuthExpired()
|
||||||
|
404 -> throw ProviderError.NotFound(ref.key)
|
||||||
|
else -> throw ProviderError.Transient("Delete failed (HTTP ${resp.code})")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
// Sidecars go with the recording; ignore failures (may not exist).
|
||||||
|
runCatching {
|
||||||
|
dav(session, "DELETE", filesPath(session, "SHONAR/sidecars/${uuidForKey(ref.key)}/"), null)
|
||||||
|
.close()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
override suspend fun list(cursor: String?): Page<RemoteRecording> = withContext(Dispatchers.IO) {
|
||||||
|
val session = connectedSession()
|
||||||
|
// Single page: a PROPFIND Depth:1 returns the whole folder. Paging
|
||||||
|
// stays null until a library outgrows one response.
|
||||||
|
if (cursor != null) return@withContext Page(emptyList(), null)
|
||||||
|
val items = propfind(session, filesPath(session, "SHONAR/audio/"), depth = "1")
|
||||||
|
.filter { it.isFile && it.relativePath != "SHONAR/audio/" && !it.relativePath.removePrefix("SHONAR/audio/").contains('/') }
|
||||||
|
.map { e ->
|
||||||
|
val name = e.relativePath.removePrefix("SHONAR/audio/")
|
||||||
|
RemoteRecording(
|
||||||
|
ref = RemoteRef(ProviderRegistry.NEXTCLOUD_ID, e.relativePath, e.etag, e.size),
|
||||||
|
title = name.substringBeforeLast('.'),
|
||||||
|
createdAtEpochMs = e.lastModified,
|
||||||
|
durationMs = 0, // duration is tracked locally, not over DAV
|
||||||
|
mime = e.contentType ?: "application/octet-stream",
|
||||||
|
)
|
||||||
|
}
|
||||||
|
Page(items, nextCursor = null)
|
||||||
|
}
|
||||||
|
|
||||||
|
// ---- sidecars: real remote files under SHONAR/sidecars/{uuid}/ --------
|
||||||
|
|
||||||
|
override suspend fun putSidecar(ref: RemoteRef, kind: SidecarKind, bytes: ByteArray) =
|
||||||
|
withContext(Dispatchers.IO) {
|
||||||
|
val session = connectedSession()
|
||||||
|
val dir = "SHONAR/sidecars/${uuidForKey(ref.key)}"
|
||||||
|
ensureDir(session, "SHONAR/sidecars")
|
||||||
|
ensureDir(session, dir)
|
||||||
|
dav(
|
||||||
|
session, "PUT", filesPath(session, "$dir/${kind.fileName}"),
|
||||||
|
bytes.toRequestBody("application/json; charset=utf-8".toMediaType()),
|
||||||
|
sensitiveBody = true, // transcripts are never logged, in any mode
|
||||||
|
).use { resp ->
|
||||||
|
if (resp.code !in 200..201 && resp.code != 204) {
|
||||||
|
throw mapError(resp.code, "Sidecar upload failed")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
override suspend fun getSidecar(ref: RemoteRef, kind: SidecarKind): ByteArray? =
|
||||||
|
withContext(Dispatchers.IO) {
|
||||||
|
val session = connectedSession()
|
||||||
|
dav(
|
||||||
|
session, "GET",
|
||||||
|
filesPath(session, "SHONAR/sidecars/${uuidForKey(ref.key)}/${kind.fileName}"), null,
|
||||||
|
sensitiveBody = true, // transcripts are never logged, in any mode
|
||||||
|
).use { resp ->
|
||||||
|
when (resp.code) {
|
||||||
|
200 -> resp.body?.bytes()
|
||||||
|
401 -> throw ProviderError.AuthExpired()
|
||||||
|
404 -> null
|
||||||
|
else -> throw mapError(resp.code, "Sidecar download failed")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// ---- status ------------------------------------------------------------
|
||||||
|
|
||||||
|
override suspend fun storageLocationSummary(): StorageLocation =
|
||||||
|
withContext(Dispatchers.IO) {
|
||||||
|
val session = connectedSession()
|
||||||
|
val quota = ocsQuota(session)
|
||||||
|
var count = 0
|
||||||
|
var bytes = 0L
|
||||||
|
// list() is single-page for now; keep the loop for when it pages.
|
||||||
|
var cursor: String? = null
|
||||||
|
do {
|
||||||
|
val page = list(cursor)
|
||||||
|
count += page.items.size
|
||||||
|
bytes += page.items.sumOf { it.ref.sizeBytes }
|
||||||
|
cursor = page.nextCursor
|
||||||
|
} while (cursor != null)
|
||||||
|
val host = runCatching { java.net.URI(session.baseUrl).host }.getOrNull()
|
||||||
|
?: session.baseUrl
|
||||||
|
StorageLocation(
|
||||||
|
headline = "Nextcloud at $host",
|
||||||
|
detail = "${session.username} · $count recordings synced" +
|
||||||
|
(quota?.let { " · ${formatBytes(it.free)} free of ${formatBytes(it.total)}" } ?: ""),
|
||||||
|
syncedCount = count,
|
||||||
|
localOnlyCount = 0,
|
||||||
|
bytesUsed = bytes,
|
||||||
|
)
|
||||||
|
}
|
||||||
|
|
||||||
|
// ---- HTTP + DAV plumbing -----------------------------------------------
|
||||||
|
|
||||||
|
private suspend fun connectedSession(): NcSession {
|
||||||
|
if (_authState.value != AuthState.CONNECTED) throw ProviderError.NotConnected()
|
||||||
|
return auth.load()?.takeIf { it.appPassword.isNotBlank() && it.userId.isNotBlank() }
|
||||||
|
?: throw ProviderError.NotConnected()
|
||||||
|
}
|
||||||
|
|
||||||
|
private fun basic(session: NcSession): String {
|
||||||
|
val raw = "${session.userId}:${session.appPassword}".toByteArray(Charsets.UTF_8)
|
||||||
|
return "Basic " + java.util.Base64.getEncoder().encodeToString(raw)
|
||||||
|
}
|
||||||
|
|
||||||
|
/** Raw DAV/OCS call. Caller closes the response. [path] starts with '/'. */
|
||||||
|
private fun dav(
|
||||||
|
session: NcSession,
|
||||||
|
method: String,
|
||||||
|
path: String,
|
||||||
|
body: okhttp3.RequestBody?,
|
||||||
|
sensitiveBody: Boolean = false,
|
||||||
|
): Response {
|
||||||
|
val builder = Request.Builder().url(session.baseUrl + path)
|
||||||
|
.header("Authorization", basic(session))
|
||||||
|
if (sensitiveBody) builder.header(RedactingLogger.SENSITIVE_BODY, "1")
|
||||||
|
if (method == "GET") builder.get()
|
||||||
|
else builder.method(method, body)
|
||||||
|
if (path.startsWith("/ocs/")) {
|
||||||
|
builder.header("OCS-APIRequest", "true")
|
||||||
|
builder.header("Accept", "application/json")
|
||||||
|
}
|
||||||
|
return client.newCall(builder.build()).execute()
|
||||||
|
}
|
||||||
|
|
||||||
|
private fun filesPath(session: NcSession, relative: String): String {
|
||||||
|
val segs = relative.split('/').filter { it.isNotEmpty() }.joinToString("/") { enc(it) }
|
||||||
|
return "/remote.php/dav/files/${enc(session.userId)}/$segs"
|
||||||
|
}
|
||||||
|
|
||||||
|
private fun uploadsPath(session: NcSession, transfer: String): String =
|
||||||
|
"/remote.php/dav/uploads/${enc(session.userId)}/${enc(transfer)}"
|
||||||
|
|
||||||
|
private fun ocsPath(suffix: String): String = "/ocs/v2.php$suffix"
|
||||||
|
|
||||||
|
/** OCS identity check; returns the DAV user id or throws. */
|
||||||
|
private suspend fun ocsUserId(origin: String, user: String, appPassword: String): String =
|
||||||
|
withContext(Dispatchers.IO) {
|
||||||
|
val raw = "$user:$appPassword".toByteArray(Charsets.UTF_8)
|
||||||
|
val req = Request.Builder().url(origin + ocsPath("/cloud/user")).get()
|
||||||
|
.header("Authorization", "Basic " + java.util.Base64.getEncoder().encodeToString(raw))
|
||||||
|
.header("OCS-APIRequest", "true")
|
||||||
|
.header("Accept", "application/json").build()
|
||||||
|
client.newCall(req).execute().use { resp ->
|
||||||
|
when (resp.code) {
|
||||||
|
200 -> {
|
||||||
|
val data = JSONObject(resp.body?.string().orEmpty())
|
||||||
|
.optJSONObject("ocs")?.optJSONObject("data")
|
||||||
|
val id = data?.optString("id", "").orEmpty()
|
||||||
|
if (id.isBlank()) throw ProviderError.Transient("Server identity reply was empty")
|
||||||
|
id
|
||||||
|
}
|
||||||
|
401 -> throw ProviderError.AuthExpired()
|
||||||
|
else -> throw ProviderError.Transient("Server identity check failed (HTTP ${resp.code})")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
private data class Quota(val free: Long, val total: Long)
|
||||||
|
|
||||||
|
private suspend fun ocsQuota(session: NcSession): Quota? = withContext(Dispatchers.IO) {
|
||||||
|
// Quota is informational; never fail the summary over it.
|
||||||
|
runCatching {
|
||||||
|
dav(session, "GET", ocsPath("/cloud/user"), null).use { resp ->
|
||||||
|
if (resp.code != 200) return@runCatching null
|
||||||
|
val q = JSONObject(resp.body?.string().orEmpty())
|
||||||
|
.optJSONObject("ocs")?.optJSONObject("data")?.optJSONObject("quota")
|
||||||
|
?: return@runCatching null
|
||||||
|
// Quota values may be numbers or numeric strings; total -3
|
||||||
|
// (or "unknown") means unlimited.
|
||||||
|
fun num(v: Any?): Long = when (v) {
|
||||||
|
is Number -> v.toLong()
|
||||||
|
is String -> v.toLongOrNull() ?: -3L
|
||||||
|
else -> -3L
|
||||||
|
}
|
||||||
|
val free = num(q.opt("free"))
|
||||||
|
val total = num(q.opt("total"))
|
||||||
|
if (total < 0) null else Quota(free.coerceAtLeast(0), total)
|
||||||
|
}
|
||||||
|
}.getOrNull()
|
||||||
|
}
|
||||||
|
|
||||||
|
/** MKCOL tolerant of "already exists". */
|
||||||
|
private suspend fun ensureDir(session: NcSession, relative: String) {
|
||||||
|
// Create level by level so a missing parent reads as progress, not 409.
|
||||||
|
val parts = relative.split('/').filter { it.isNotEmpty() }
|
||||||
|
var prefix = ""
|
||||||
|
for (part in parts) {
|
||||||
|
prefix = if (prefix.isEmpty()) part else "$prefix/$part"
|
||||||
|
dav(session, "MKCOL", filesPath(session, prefix), null).use { resp ->
|
||||||
|
if (resp.code == 401) throw ProviderError.AuthExpired()
|
||||||
|
if (resp.code != 201 && resp.code != 405) {
|
||||||
|
throw mapError(resp.code, "Could not create folder $prefix")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
private suspend fun mkcol(session: NcSession, transfer: String) {
|
||||||
|
dav(session, "MKCOL", uploadsPath(session, transfer) + "/", null).use { resp ->
|
||||||
|
// 405: transfer folder from a previous attempt — resume into it.
|
||||||
|
if (resp.code == 401) throw ProviderError.AuthExpired()
|
||||||
|
if (resp.code != 201 && resp.code != 405) {
|
||||||
|
throw mapError(resp.code, "Could not start the upload")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
/** Chunk names present in the transfer folder (resume). */
|
||||||
|
private suspend fun transferChunks(session: NcSession, transfer: String): Set<Int> {
|
||||||
|
val prefix = "/remote.php/dav/uploads/${session.userId}/"
|
||||||
|
val entries = propfind(session, uploadsPath(session, transfer) + "/", depth = "1", prefix = prefix)
|
||||||
|
return entries.mapNotNullTo(mutableSetOf()) { e ->
|
||||||
|
// Chunk names are "00001".. — compare by numeric value.
|
||||||
|
e.name.trimStart('0').ifEmpty { "0" }.toIntOrNull()
|
||||||
|
?.takeIf { it in 1..10000 }
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
private suspend fun putChunk(
|
||||||
|
session: NcSession,
|
||||||
|
transfer: String,
|
||||||
|
index: Int, // 1-based
|
||||||
|
bytes: ByteArray,
|
||||||
|
totalBytes: Long,
|
||||||
|
destKey: String,
|
||||||
|
) {
|
||||||
|
val name = "%05d".format(index)
|
||||||
|
val dest = session.baseUrl + filesPath(session, destKey)
|
||||||
|
val req = Request.Builder()
|
||||||
|
.url(session.baseUrl + uploadsPath(session, transfer) + "/" + name)
|
||||||
|
.put(bytes.toRequestBody("application/octet-stream".toMediaType()))
|
||||||
|
.header("Authorization", basic(session))
|
||||||
|
.header("OC-Total-Length", totalBytes.toString())
|
||||||
|
.header("Destination", dest).build()
|
||||||
|
client.newCall(req).execute().use { resp ->
|
||||||
|
when (resp.code) {
|
||||||
|
200, 201, 204 -> Unit
|
||||||
|
401 -> throw ProviderError.AuthExpired()
|
||||||
|
507 -> throw ProviderError.QuotaExceeded()
|
||||||
|
else -> throw mapError(resp.code, "Chunk $index rejected")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
private suspend fun assemble(
|
||||||
|
session: NcSession,
|
||||||
|
transfer: String,
|
||||||
|
destKey: String,
|
||||||
|
draft: RecordingDraft,
|
||||||
|
) {
|
||||||
|
val dest = session.baseUrl + filesPath(session, destKey)
|
||||||
|
val req = Request.Builder()
|
||||||
|
.url(session.baseUrl + uploadsPath(session, transfer) + "/.file")
|
||||||
|
.method("MOVE", null)
|
||||||
|
.header("Authorization", basic(session))
|
||||||
|
.header("Destination", dest)
|
||||||
|
.header("Overwrite", "T")
|
||||||
|
.header("OC-Total-Length", draft.sizeBytes.toString())
|
||||||
|
// Server mtime = recording time, so listings sort by when it
|
||||||
|
// was recorded, not when it finished uploading.
|
||||||
|
.header("X-OC-Mtime", (draft.createdAtEpochMs / 1000).toString()).build()
|
||||||
|
client.newCall(req).execute().use { resp ->
|
||||||
|
when (resp.code) {
|
||||||
|
200, 201, 204 -> Unit
|
||||||
|
401 -> throw ProviderError.AuthExpired()
|
||||||
|
404 -> throw ProviderError.Transient("Upload assembly failed — retry the upload")
|
||||||
|
507 -> throw ProviderError.QuotaExceeded()
|
||||||
|
else -> throw mapError(resp.code, "Upload assembly failed")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
internal data class DavEntry(
|
||||||
|
val relativePath: String, // relative to files/{user}/, decoded
|
||||||
|
val name: String,
|
||||||
|
val isFile: Boolean,
|
||||||
|
val size: Long,
|
||||||
|
val etag: String?,
|
||||||
|
val contentType: String?,
|
||||||
|
val lastModified: Long,
|
||||||
|
)
|
||||||
|
|
||||||
|
private suspend fun propfind(
|
||||||
|
session: NcSession,
|
||||||
|
path: String,
|
||||||
|
depth: String,
|
||||||
|
prefix: String = "/remote.php/dav/files/${session.userId}/",
|
||||||
|
): List<DavEntry> =
|
||||||
|
withContext(Dispatchers.IO) {
|
||||||
|
val body = """<?xml version="1.0"?>
|
||||||
|
<d:propfind xmlns:d="DAV:"><d:prop><d:getcontentlength/><d:getetag/><d:resourcetype/><d:getcontenttype/><d:getlastmodified/><d:displayname/></d:prop></d:propfind>"""
|
||||||
|
val req = Request.Builder().url(session.baseUrl + path)
|
||||||
|
.method("PROPFIND", body.toRequestBody("application/xml; charset=utf-8".toMediaType()))
|
||||||
|
.header("Authorization", basic(session))
|
||||||
|
.header("Depth", depth).build()
|
||||||
|
client.newCall(req).execute().use { resp ->
|
||||||
|
when (resp.code) {
|
||||||
|
200, 207 -> parseMultistatus(resp.body?.string().orEmpty(), prefix = prefix)
|
||||||
|
401 -> throw ProviderError.AuthExpired()
|
||||||
|
404 -> emptyList()
|
||||||
|
else -> throw mapError(resp.code, "Listing failed")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
private fun mapError(code: Int, fallback: String): ProviderError = when (code) {
|
||||||
|
401 -> ProviderError.AuthExpired()
|
||||||
|
507 -> ProviderError.QuotaExceeded()
|
||||||
|
else -> ProviderError.Transient("$fallback (HTTP $code)")
|
||||||
|
}
|
||||||
|
|
||||||
|
companion object {
|
||||||
|
/** One path segment, percent-encoded (spaces as %20, not '+'). */
|
||||||
|
internal fun enc(segment: String): String =
|
||||||
|
java.net.URLEncoder.encode(segment, "UTF-8").replace("+", "%20")
|
||||||
|
|
||||||
|
internal fun uuidForKey(key: String): String =
|
||||||
|
key.substringAfterLast('/').substringBeforeLast('.')
|
||||||
|
|
||||||
|
internal fun extFor(mime: String): String = when (mime.lowercase().substringBefore(';').trim()) {
|
||||||
|
"audio/mp4", "audio/m4a" -> ".m4a"
|
||||||
|
"audio/aac" -> ".aac"
|
||||||
|
"audio/wav", "audio/x-wav" -> ".wav"
|
||||||
|
"audio/ogg", "audio/opus" -> ".ogg"
|
||||||
|
"audio/webm" -> ".webm"
|
||||||
|
"audio/mpeg" -> ".mp3"
|
||||||
|
else -> ".m4a" // the app records m4a; unknown mimes keep a playable suffix
|
||||||
|
}
|
||||||
|
|
||||||
|
internal fun formatBytes(n: Long): String {
|
||||||
|
if (n < 1024) return "$n B"
|
||||||
|
val units = arrayOf("KB", "MB", "GB", "TB")
|
||||||
|
var v = n.toDouble() / 1024
|
||||||
|
var u = 0
|
||||||
|
while (v >= 1024 && u < units.size - 1) {
|
||||||
|
v /= 1024
|
||||||
|
u++
|
||||||
|
}
|
||||||
|
return "%s %s".format(if (v >= 100) "%.0f" else "%.1f".format(v), units[u])
|
||||||
|
}
|
||||||
|
|
||||||
|
internal fun parseMultistatus(xml: String, prefix: String): List<DavEntry> {
|
||||||
|
if (xml.isBlank()) return emptyList()
|
||||||
|
val out = mutableListOf<DavEntry>()
|
||||||
|
try {
|
||||||
|
val factory = javax.xml.parsers.DocumentBuilderFactory.newInstance()
|
||||||
|
factory.isNamespaceAware = true
|
||||||
|
// Harden against XXE: multistatus docs never need doctypes.
|
||||||
|
runCatching {
|
||||||
|
factory.setFeature("http://apache.org/xml/features/disallow-doctype-decl", true)
|
||||||
|
}
|
||||||
|
val doc = factory.newDocumentBuilder()
|
||||||
|
.parse(java.io.ByteArrayInputStream(xml.toByteArray(Charsets.UTF_8)))
|
||||||
|
val responses = doc.getElementsByTagNameNS("DAV:", "response")
|
||||||
|
for (i in 0 until responses.length) {
|
||||||
|
val el = responses.item(i) as? org.w3c.dom.Element ?: continue
|
||||||
|
fun text(tag: String): String? {
|
||||||
|
val nodes = el.getElementsByTagNameNS("DAV:", tag)
|
||||||
|
if (nodes.length == 0) return null
|
||||||
|
return nodes.item(0).textContent?.trim()?.takeIf { it.isNotEmpty() }
|
||||||
|
}
|
||||||
|
val href = text("href") ?: continue
|
||||||
|
val decoded = runCatching {
|
||||||
|
java.net.URLDecoder.decode(href, "UTF-8")
|
||||||
|
}.getOrNull() ?: href
|
||||||
|
// Strip scheme+host when the server returns absolute hrefs.
|
||||||
|
val pathOnly = runCatching { java.net.URI(decoded).path }.getOrNull() ?: decoded
|
||||||
|
val relative = pathOnly.removePrefix(prefix).trim('/')
|
||||||
|
val isCollection = runCatching {
|
||||||
|
val rt = el.getElementsByTagNameNS("DAV:", "resourcetype")
|
||||||
|
rt.length > 0 && (rt.item(0) as org.w3c.dom.Element)
|
||||||
|
.getElementsByTagNameNS("DAV:", "collection").length > 0
|
||||||
|
}.getOrDefault(false)
|
||||||
|
val lastMod = text("getlastmodified")?.let {
|
||||||
|
runCatching {
|
||||||
|
java.time.ZonedDateTime.parse(it, DateTimeFormatter.RFC_1123_DATE_TIME)
|
||||||
|
.toInstant().toEpochMilli()
|
||||||
|
}.getOrNull()
|
||||||
|
} ?: 0L
|
||||||
|
out += DavEntry(
|
||||||
|
relativePath = relative,
|
||||||
|
name = relative.substringAfterLast('/'),
|
||||||
|
isFile = !isCollection,
|
||||||
|
size = text("getcontentlength")?.toLongOrNull() ?: 0L,
|
||||||
|
etag = text("getetag")?.trim('"'),
|
||||||
|
contentType = text("getcontenttype"),
|
||||||
|
lastModified = lastMod,
|
||||||
|
)
|
||||||
|
}
|
||||||
|
} catch (e: Exception) {
|
||||||
|
throw ProviderError.Transient("Could not read the server listing")
|
||||||
|
}
|
||||||
|
return out
|
||||||
|
}
|
||||||
|
|
||||||
|
private fun readSlice(file: File, offset: Long, max: Long): ByteArray {
|
||||||
|
file.inputStream().use { input ->
|
||||||
|
var skipped = 0L
|
||||||
|
while (skipped < offset) {
|
||||||
|
val n = input.skip(offset - skipped)
|
||||||
|
if (n <= 0) break
|
||||||
|
skipped += n
|
||||||
|
}
|
||||||
|
val cap = max.coerceAtMost(Int.MAX_VALUE.toLong()).toInt()
|
||||||
|
val buf = ByteArray(cap)
|
||||||
|
var read = 0
|
||||||
|
while (read < cap) {
|
||||||
|
val n = input.read(buf, read, cap - read)
|
||||||
|
if (n < 0) break
|
||||||
|
read += n
|
||||||
|
}
|
||||||
|
return if (read == cap) buf else buf.copyOf(read)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
96
shared/com/shonar/provider/ProviderRegistry.kt
Normal file
96
shared/com/shonar/provider/ProviderRegistry.kt
Normal file
|
|
@ -0,0 +1,96 @@
|
||||||
|
package com.shonar.provider
|
||||||
|
|
||||||
|
import kotlinx.coroutines.flow.MutableStateFlow
|
||||||
|
import kotlinx.coroutines.flow.StateFlow
|
||||||
|
|
||||||
|
/**
|
||||||
|
* Maps provider ids to implementations. The active provider is whichever the
|
||||||
|
* user selected during setup; switching is just changing this id — the app
|
||||||
|
* never branches on concrete provider classes.
|
||||||
|
*/
|
||||||
|
class ProviderRegistry(
|
||||||
|
private val factories: Map<String, () -> ShonarProvider>,
|
||||||
|
/** Builds every provider HTTP client (TOFU trust + redacting logger). */
|
||||||
|
val tls: TlsPolicy = defaultTls(),
|
||||||
|
) {
|
||||||
|
|
||||||
|
private val _activeId = MutableStateFlow(LocalOnlyProvider.ID)
|
||||||
|
val activeId: StateFlow<String> = _activeId
|
||||||
|
|
||||||
|
/** Nextcloud is the product default once its factory registers (P4). */
|
||||||
|
val defaultProviderId: String
|
||||||
|
get() = factories.keys.sorted().let { ids ->
|
||||||
|
ids.firstOrNull { it == NEXTCLOUD_ID } ?: ids.firstOrNull() ?: LocalOnlyProvider.ID
|
||||||
|
}
|
||||||
|
|
||||||
|
fun provider(id: String): ShonarProvider =
|
||||||
|
factories[id]?.invoke() ?: throw ProviderError.InvalidUrl("Unknown provider: $id")
|
||||||
|
|
||||||
|
val active: ShonarProvider get() = provider(_activeId.value)
|
||||||
|
|
||||||
|
/** Selecting a provider does not touch existing local data. */
|
||||||
|
fun select(id: String): ShonarProvider {
|
||||||
|
if (id !in factories) throw ProviderError.InvalidUrl("Unknown provider: $id")
|
||||||
|
_activeId.value = id
|
||||||
|
return provider(id)
|
||||||
|
}
|
||||||
|
|
||||||
|
fun available(): List<ProviderDescriptor> =
|
||||||
|
factories.keys.map { provider(it).descriptor }
|
||||||
|
|
||||||
|
companion object {
|
||||||
|
const val NEXTCLOUD_ID = "nextcloud"
|
||||||
|
const val CUSTOM_SHONAR_ID = "custom-shonar"
|
||||||
|
const val SYNC_FOLDER_ID = FolderSyncProvider.ID
|
||||||
|
|
||||||
|
/** Behaviour-identical default: system trust, logging off, no pins. */
|
||||||
|
fun defaultTls(sink: (String) -> Unit = {}): TlsPolicy = TlsPolicy(
|
||||||
|
tofu = TofuManager(TofuStore(com.shonar.settings.InMemorySettingsStore())),
|
||||||
|
sink = sink,
|
||||||
|
)
|
||||||
|
|
||||||
|
/**
|
||||||
|
* Production factory map. P1 registered local-only, P3 the custom
|
||||||
|
* SHONAR server, P4 Nextcloud + the sync folder. Network providers
|
||||||
|
* are single shared instances (their sessions live in the secure
|
||||||
|
* store, not in the object) so connect/reconnect state survives
|
||||||
|
* `provider(id)` calls. P5: every HTTP client comes from [tlsPolicy].
|
||||||
|
*/
|
||||||
|
fun withDefaults(
|
||||||
|
appFilesDir: java.io.File,
|
||||||
|
secureStore: com.shonar.settings.SettingsStore =
|
||||||
|
com.shonar.settings.InMemorySettingsStore(),
|
||||||
|
plainStore: com.shonar.settings.SettingsStore =
|
||||||
|
com.shonar.settings.InMemorySettingsStore(),
|
||||||
|
tlsPolicy: TlsPolicy = defaultTls(),
|
||||||
|
): ProviderRegistry {
|
||||||
|
val custom = CustomShonarProvider(
|
||||||
|
auth = ShonarAuthStore(secureStore),
|
||||||
|
sidecarRoot = java.io.File(appFilesDir, "shonar-sidecars"),
|
||||||
|
client = tlsPolicy.apiClient(),
|
||||||
|
handshake = ShonarHandshake(tlsPolicy.probeClient(), tlsPolicy.tofu),
|
||||||
|
)
|
||||||
|
val nextcloud = NextcloudProvider(
|
||||||
|
auth = NextcloudAuthStore(secureStore),
|
||||||
|
client = tlsPolicy.nextcloudClient(),
|
||||||
|
loginFlow = NextcloudAuth(tlsPolicy.nextcloudClient(), tlsPolicy.tofu),
|
||||||
|
)
|
||||||
|
val folder = FolderSyncProvider(pathStore = plainStore)
|
||||||
|
return ProviderRegistry(
|
||||||
|
mapOf(
|
||||||
|
LocalOnlyProvider.ID to {
|
||||||
|
LocalOnlyProvider(
|
||||||
|
java.io.File(appFilesDir, "shonar-local"),
|
||||||
|
pathStore = plainStore,
|
||||||
|
)
|
||||||
|
},
|
||||||
|
CUSTOM_SHONAR_ID to { custom },
|
||||||
|
NEXTCLOUD_ID to { nextcloud },
|
||||||
|
SYNC_FOLDER_ID to { folder },
|
||||||
|
// P6: start9 / umbrel platform probes
|
||||||
|
),
|
||||||
|
tls = tlsPolicy,
|
||||||
|
)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
220
shared/com/shonar/provider/ProviderTypes.kt
Normal file
220
shared/com/shonar/provider/ProviderTypes.kt
Normal file
|
|
@ -0,0 +1,220 @@
|
||||||
|
package com.shonar.provider
|
||||||
|
|
||||||
|
import java.io.File
|
||||||
|
|
||||||
|
/**
|
||||||
|
* Shared vocabulary for server providers. Nothing here depends on any
|
||||||
|
* specific server product; the rest of the app talks to providers only
|
||||||
|
* through these types plus [ShonarProvider].
|
||||||
|
*/
|
||||||
|
|
||||||
|
/** What kind of provider this is and how the UI should present it. */
|
||||||
|
data class ProviderDescriptor(
|
||||||
|
val id: String, // "nextcloud" | "custom-shonar" | "local-only" | ...
|
||||||
|
val displayName: String, // "Nextcloud"
|
||||||
|
val isDefault: Boolean = false,
|
||||||
|
val capabilities: Set<Capability> = emptySet(),
|
||||||
|
) {
|
||||||
|
enum class Capability {
|
||||||
|
CHUNKED_UPLOAD, // resumable large-file upload protocol
|
||||||
|
SERVER_TRANSCRIPTION, // backend can transcribe (feature-gate AI)
|
||||||
|
SERVER_SUMMARY,
|
||||||
|
QUOTA_INFO, // storageLocationSummary can report quota
|
||||||
|
ACCOUNT_DELETION, // provider supports deleteAccountAndData
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
/**
|
||||||
|
* A validated server URL. Construction only via [parse]; guarantees:
|
||||||
|
* - scheme https, or http ONLY for private/LAN hosts (RFC1918, loopback, .local)
|
||||||
|
* - no userinfo (user:pass in URL is rejected)
|
||||||
|
* - no path traversal ("..")
|
||||||
|
* - non-blank host
|
||||||
|
*/
|
||||||
|
data class ServerUrl(val scheme: String, val host: String, val port: Int, val pathSegments: List<String>) {
|
||||||
|
|
||||||
|
val isCleartext: Boolean get() = scheme == "http"
|
||||||
|
|
||||||
|
/** Normalized base for API calls, e.g. https://cloud.example.com */
|
||||||
|
val origin: String
|
||||||
|
get() = buildString {
|
||||||
|
append(scheme).append("://").append(host)
|
||||||
|
val default = if (scheme == "https") 443 else 80
|
||||||
|
if (port != default) append(":").append(port)
|
||||||
|
}
|
||||||
|
|
||||||
|
companion object {
|
||||||
|
fun parse(raw: String): Result<ServerUrl> {
|
||||||
|
val trimmed = raw.trim().removeSuffix("/")
|
||||||
|
if (trimmed.isBlank()) return Result.failure(ProviderError.InvalidUrl("URL is empty"))
|
||||||
|
val uri = runCatching { java.net.URI(trimmed) }
|
||||||
|
.getOrElse { return Result.failure(ProviderError.InvalidUrl("Not a valid URL")) }
|
||||||
|
val scheme = (uri.scheme ?: "").lowercase()
|
||||||
|
if (scheme != "https" && scheme != "http")
|
||||||
|
return Result.failure(ProviderError.InvalidUrl("Only http(s) URLs are allowed"))
|
||||||
|
if (uri.userInfo != null)
|
||||||
|
return Result.failure(ProviderError.InvalidUrl("Credentials in URL are not allowed"))
|
||||||
|
val host = (uri.host ?: "").lowercase()
|
||||||
|
if (host.isBlank())
|
||||||
|
return Result.failure(ProviderError.InvalidUrl("Missing host"))
|
||||||
|
val segs = uri.path.split('/').filter { it.isNotBlank() }
|
||||||
|
if (segs.any { it == ".." })
|
||||||
|
return Result.failure(ProviderError.InvalidUrl("Path traversal is not allowed"))
|
||||||
|
if (scheme == "http" && !isPrivateHost(host))
|
||||||
|
return Result.failure(
|
||||||
|
ProviderError.InvalidUrl("Cleartext http:// is only allowed for local/LAN servers; use https:// for $host")
|
||||||
|
)
|
||||||
|
val port = if (uri.port > 0) uri.port else if (scheme == "https") 443 else 80
|
||||||
|
return Result.success(ServerUrl(scheme, host, port, segs))
|
||||||
|
}
|
||||||
|
|
||||||
|
fun isPrivateHost(host: String): Boolean {
|
||||||
|
if (host == "localhost" || host.endsWith(".local")) return true
|
||||||
|
val octets = host.split('.').takeIf { it.size == 4 }?.map { it.toIntOrNull() ?: -1 } ?: return false
|
||||||
|
if (octets.any { it !in 0..255 }) return false
|
||||||
|
return when {
|
||||||
|
octets[0] == 10 -> true // 10/8
|
||||||
|
octets[0] == 127 -> true // loopback
|
||||||
|
octets[0] == 192 && octets[1] == 168 -> true // 192.168/16
|
||||||
|
octets[0] == 172 && octets[1] in 16..31 -> true // 172.16/12
|
||||||
|
octets[0] == 169 && octets[1] == 254 -> true // link-local
|
||||||
|
else -> false
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
/** Credential material for a provider. NEVER put this in logs or Room. */
|
||||||
|
sealed class ProviderCredential {
|
||||||
|
abstract val accountLabel: String // human hint only (e.g. "user@cloud")
|
||||||
|
|
||||||
|
/** OAuth2/OIDC access token + refresh token (PKCE flow). */
|
||||||
|
data class OAuthTokens(
|
||||||
|
override val accountLabel: String,
|
||||||
|
val accessToken: String,
|
||||||
|
val refreshToken: String?,
|
||||||
|
val expiresAtEpochSec: Long?,
|
||||||
|
) : ProviderCredential() {
|
||||||
|
override fun toString(): String = "OAuthTokens(account=$accountLabel, [redacted])"
|
||||||
|
}
|
||||||
|
|
||||||
|
/** Nextcloud login-flow-v2 / manually created app password. */
|
||||||
|
data class AppPassword(
|
||||||
|
override val accountLabel: String,
|
||||||
|
val loginUrl: String,
|
||||||
|
val user: String,
|
||||||
|
val password: String,
|
||||||
|
) : ProviderCredential() {
|
||||||
|
override fun toString(): String = "AppPassword(account=$accountLabel, [redacted])"
|
||||||
|
}
|
||||||
|
|
||||||
|
/**
|
||||||
|
* Custom SHONAR server login: email + password exchanged for a rotating
|
||||||
|
* token pair (backend M1). The password is held in memory only for the
|
||||||
|
* login call and never persisted — only the resulting tokens are stored
|
||||||
|
* (see ShonarAuthStore).
|
||||||
|
*/
|
||||||
|
data class ShonarLogin(
|
||||||
|
override val accountLabel: String,
|
||||||
|
val serverUrl: ServerUrl,
|
||||||
|
val email: String,
|
||||||
|
val password: String,
|
||||||
|
) : ProviderCredential() {
|
||||||
|
override fun toString(): String =
|
||||||
|
"ShonarLogin(account=$accountLabel, server=${serverUrl.origin}, [redacted])"
|
||||||
|
}
|
||||||
|
|
||||||
|
/** No credential needed (local-only provider). */
|
||||||
|
data object None : ProviderCredential() {
|
||||||
|
override val accountLabel: String get() = "local"
|
||||||
|
override fun toString(): String = "None"
|
||||||
|
}
|
||||||
|
|
||||||
|
/**
|
||||||
|
* Sync-folder provider: an absolute directory path owned by an external
|
||||||
|
* sync tool (Syncthing, the Nextcloud desktop client, rsync…). The path
|
||||||
|
* is not secret, but it is validated — connect refuses a missing or
|
||||||
|
* non-writable directory rather than creating whatever was typed.
|
||||||
|
*/
|
||||||
|
data class FolderPath(
|
||||||
|
override val accountLabel: String,
|
||||||
|
val path: String,
|
||||||
|
) : ProviderCredential()
|
||||||
|
}
|
||||||
|
|
||||||
|
/** Result of probing a base URL for a compatible service. */
|
||||||
|
sealed class ProbeResult {
|
||||||
|
data class Compatible(val descriptor: ProviderDescriptor, val serverName: String, val version: String) : ProbeResult()
|
||||||
|
/** Reachable, SHONAR-compatible services were found; user must pick one (Start9/Umbrel platforms). */
|
||||||
|
data class ServicesFound(val services: List<DiscoveredService>) : ProbeResult()
|
||||||
|
data object Incompatible : ProbeResult()
|
||||||
|
data class NetworkError(val reason: String) : ProbeResult() // reason must be secret-free
|
||||||
|
data class TlsFailure(val fingerprintSha256: String) : ProbeResult() // triggers TOFU approval UI
|
||||||
|
}
|
||||||
|
|
||||||
|
/** A SHONAR-compatible service discovered on a platform (Start9/Umbrel). */
|
||||||
|
data class DiscoveredService(
|
||||||
|
val platformId: String, // "start9" | "umbrel"
|
||||||
|
val serviceName: String, // app/service display name
|
||||||
|
val baseUrl: ServerUrl,
|
||||||
|
val providerKind: String, // provider to bind once confirmed
|
||||||
|
)
|
||||||
|
|
||||||
|
/** Auth lifecycle shown in UI. Transitions: see docs/server-providers.md §3. */
|
||||||
|
enum class AuthState { DISCONNECTED, CONNECTED, EXPIRED, REVOKED, OFFLINE }
|
||||||
|
|
||||||
|
/** Opaque pointer to a remote object. Providers map keys to their own layout. */
|
||||||
|
data class RemoteRef(
|
||||||
|
val providerId: String,
|
||||||
|
val key: String, // provider-relative key, never a local path
|
||||||
|
val etag: String?,
|
||||||
|
val sizeBytes: Long,
|
||||||
|
)
|
||||||
|
|
||||||
|
/** Input to [ShonarProvider.upload]. sourceFile is a local file owned by the app. */
|
||||||
|
data class RecordingDraft(
|
||||||
|
val id: String, // public UUID string
|
||||||
|
val title: String,
|
||||||
|
val createdAtEpochMs: Long,
|
||||||
|
val durationMs: Long,
|
||||||
|
val mime: String, // audio/mp4 | audio/wav | audio/ogg
|
||||||
|
val sourceFile: File,
|
||||||
|
val sizeBytes: Long,
|
||||||
|
/** Optional per-recording transcription model override (null = server default). */
|
||||||
|
val transcriptionModel: String? = null,
|
||||||
|
)
|
||||||
|
|
||||||
|
enum class SidecarKind(val fileName: String) {
|
||||||
|
TRANSCRIPT("transcript.json"),
|
||||||
|
SUMMARY("summary.json"),
|
||||||
|
ACTION_ITEMS("action-items.json"),
|
||||||
|
KEYWORDS("keywords.json"),
|
||||||
|
NOTES("notes.json"),
|
||||||
|
}
|
||||||
|
|
||||||
|
data class Page<T>(val items: List<T>, val nextCursor: String?)
|
||||||
|
data class RemoteRecording(val ref: RemoteRef, val title: String, val createdAtEpochMs: Long, val durationMs: Long, val mime: String)
|
||||||
|
|
||||||
|
/** "Where is my data?" — what the StorageLocation screen renders. */
|
||||||
|
data class StorageLocation(
|
||||||
|
val headline: String, // "On this device only" / "Nextcloud at cloud.example.com"
|
||||||
|
val detail: String, // human-readable path/prefix, quota, etc.
|
||||||
|
val syncedCount: Int,
|
||||||
|
val localOnlyCount: Int,
|
||||||
|
val bytesUsed: Long,
|
||||||
|
/** Concrete on-device path, when the files live in a real folder. */
|
||||||
|
val path: String? = null,
|
||||||
|
)
|
||||||
|
|
||||||
|
/** Provider-level errors. Messages must never contain credentials/tokens. */
|
||||||
|
sealed class ProviderError(message: String) : Exception(message) {
|
||||||
|
class InvalidUrl(message: String) : ProviderError(message)
|
||||||
|
class NotConnected : ProviderError("Provider is not connected")
|
||||||
|
class AuthExpired : ProviderError("Access token expired — re-authentication required")
|
||||||
|
class Revoked : ProviderError("Access was revoked on the server")
|
||||||
|
class TlsUntrusted(fingerprint: String) : ProviderError("Server certificate not trusted (SHA-256 $fingerprint)")
|
||||||
|
class NotFound(key: String) : ProviderError("Remote item not found: $key")
|
||||||
|
class Conflict(key: String) : ProviderError("Remote item changed since last read: $key")
|
||||||
|
class QuotaExceeded : ProviderError("Server storage quota exceeded")
|
||||||
|
class Transient(message: String) : ProviderError(message)
|
||||||
|
}
|
||||||
83
shared/com/shonar/provider/ShonarAuthStore.kt
Normal file
83
shared/com/shonar/provider/ShonarAuthStore.kt
Normal file
|
|
@ -0,0 +1,83 @@
|
||||||
|
package com.shonar.provider
|
||||||
|
|
||||||
|
import com.shonar.settings.SettingsStore
|
||||||
|
|
||||||
|
/**
|
||||||
|
* Secure persistence for the custom SHONAR server session.
|
||||||
|
*
|
||||||
|
* Stored (all in the Keystore-backed secure store, never in Room or logs):
|
||||||
|
* - base URL origin (e.g. https://shonar.example.com)
|
||||||
|
* - account email (hint only)
|
||||||
|
* - access token + expiry
|
||||||
|
* - refresh token + device id
|
||||||
|
*
|
||||||
|
* The user's password is NEVER stored here: it lives in memory for exactly
|
||||||
|
* one login call (see [ProviderCredential.ShonarLogin]).
|
||||||
|
*
|
||||||
|
* Rotation discipline (backend M1 has refresh reuse detection): every
|
||||||
|
* successful refresh overwrites the stored pair immediately, so a stored
|
||||||
|
* refresh token is always the newest one the server has issued.
|
||||||
|
*/
|
||||||
|
class ShonarAuthStore(private val secure: SettingsStore) {
|
||||||
|
|
||||||
|
suspend fun save(session: ShonarSession) {
|
||||||
|
secure.putString(KEY_BASE_URL, session.baseUrl)
|
||||||
|
secure.putString(KEY_EMAIL, session.email)
|
||||||
|
secure.putString(KEY_ACCESS, session.accessToken)
|
||||||
|
secure.putString(KEY_REFRESH, session.refreshToken)
|
||||||
|
secure.putString(KEY_EXPIRES_AT, session.expiresAtEpochSec.toString())
|
||||||
|
if (session.deviceId != null) secure.putString(KEY_DEVICE_ID, session.deviceId)
|
||||||
|
else secure.remove(KEY_DEVICE_ID)
|
||||||
|
}
|
||||||
|
|
||||||
|
suspend fun load(): ShonarSession? {
|
||||||
|
val baseUrl = secure.getString(KEY_BASE_URL) ?: return null
|
||||||
|
val access = secure.getString(KEY_ACCESS) ?: return null
|
||||||
|
val refresh = secure.getString(KEY_REFRESH) ?: return null
|
||||||
|
return ShonarSession(
|
||||||
|
baseUrl = baseUrl,
|
||||||
|
email = secure.getString(KEY_EMAIL).orEmpty(),
|
||||||
|
accessToken = access,
|
||||||
|
refreshToken = refresh,
|
||||||
|
expiresAtEpochSec = secure.getString(KEY_EXPIRES_AT)?.toLongOrNull(),
|
||||||
|
deviceId = secure.getString(KEY_DEVICE_ID),
|
||||||
|
)
|
||||||
|
}
|
||||||
|
|
||||||
|
/** Drop tokens but keep the URL + email so re-login is one step. */
|
||||||
|
suspend fun clearTokens() {
|
||||||
|
secure.remove(KEY_ACCESS)
|
||||||
|
secure.remove(KEY_REFRESH)
|
||||||
|
secure.remove(KEY_EXPIRES_AT)
|
||||||
|
secure.remove(KEY_DEVICE_ID)
|
||||||
|
}
|
||||||
|
|
||||||
|
/** Forget everything, including which server was configured. */
|
||||||
|
suspend fun clearAll() {
|
||||||
|
secure.remove(KEY_BASE_URL)
|
||||||
|
secure.remove(KEY_EMAIL)
|
||||||
|
clearTokens()
|
||||||
|
}
|
||||||
|
|
||||||
|
companion object {
|
||||||
|
private const val PREFIX = "provider.custom-shonar."
|
||||||
|
const val KEY_BASE_URL = PREFIX + "base_url"
|
||||||
|
const val KEY_EMAIL = PREFIX + "email"
|
||||||
|
const val KEY_ACCESS = PREFIX + "access_token"
|
||||||
|
const val KEY_REFRESH = PREFIX + "refresh_token"
|
||||||
|
const val KEY_EXPIRES_AT = PREFIX + "expires_at"
|
||||||
|
const val KEY_DEVICE_ID = PREFIX + "device_id"
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
/** In-memory session handed between login/refresh calls and the store. */
|
||||||
|
data class ShonarSession(
|
||||||
|
val baseUrl: String,
|
||||||
|
val email: String,
|
||||||
|
val accessToken: String,
|
||||||
|
val refreshToken: String,
|
||||||
|
val expiresAtEpochSec: Long?,
|
||||||
|
val deviceId: String?,
|
||||||
|
) {
|
||||||
|
override fun toString(): String = "ShonarSession(server=$baseUrl, account=$email, [redacted])"
|
||||||
|
}
|
||||||
91
shared/com/shonar/provider/ShonarHandshake.kt
Normal file
91
shared/com/shonar/provider/ShonarHandshake.kt
Normal file
|
|
@ -0,0 +1,91 @@
|
||||||
|
package com.shonar.provider
|
||||||
|
|
||||||
|
import kotlinx.coroutines.Dispatchers
|
||||||
|
import kotlinx.coroutines.withContext
|
||||||
|
import okhttp3.OkHttpClient
|
||||||
|
import okhttp3.Request
|
||||||
|
import java.util.concurrent.TimeUnit
|
||||||
|
import javax.net.ssl.SSLHandshakeException
|
||||||
|
|
||||||
|
/**
|
||||||
|
* Unauthenticated handshake against the SHONAR backend's
|
||||||
|
* GET /api/v1/provider-info. Used by the provider-selection screen and by
|
||||||
|
* Start9/Umbrel platform probes to identify SHONAR-compatible services.
|
||||||
|
*
|
||||||
|
* TLS policy (docs/server-providers.md §4):
|
||||||
|
* - full system verification by default, no bypass, ever.
|
||||||
|
* - a handshake failure yields ProbeResult.TlsFailure with the peer cert's
|
||||||
|
* SPKI SHA-256 (recorded by the TOFU trust manager) so the UI can run
|
||||||
|
* explicit trust-on-first-use approval; this client NEVER retries with
|
||||||
|
* verification disabled.
|
||||||
|
*/
|
||||||
|
class ShonarHandshake(
|
||||||
|
private val client: OkHttpClient = defaultClient(),
|
||||||
|
private val tofu: TofuManager? = null,
|
||||||
|
) {
|
||||||
|
|
||||||
|
suspend fun probe(url: ServerUrl): ProbeResult = withContext(Dispatchers.IO) {
|
||||||
|
val endpoint = url.origin + "/api/v1/provider-info"
|
||||||
|
val request = Request.Builder().url(endpoint).get().build()
|
||||||
|
try {
|
||||||
|
client.newCall(request).execute().use { resp ->
|
||||||
|
if (resp.code == 200) {
|
||||||
|
parseBody(resp.body?.string().orEmpty())
|
||||||
|
} else {
|
||||||
|
ProbeResult.Incompatible
|
||||||
|
}
|
||||||
|
}
|
||||||
|
} catch (e: SSLHandshakeException) {
|
||||||
|
ProbeResult.TlsFailure(fingerprintFor(url.host))
|
||||||
|
} catch (e: javax.net.ssl.SSLPeerUnverifiedException) {
|
||||||
|
// Hostname mismatch: strict verifier stays strict, but the UI
|
||||||
|
// should still say *why* instead of a generic network error.
|
||||||
|
ProbeResult.TlsFailure(fingerprintFor(url.host))
|
||||||
|
} catch (e: Exception) {
|
||||||
|
// message must stay generic: exception text can contain URLs but
|
||||||
|
// never credentials (this call sends no credentials at all)
|
||||||
|
ProbeResult.NetworkError(e.javaClass.simpleName)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
private fun parseBody(json: String): ProbeResult = try {
|
||||||
|
val root = org.json.JSONObject(json)
|
||||||
|
if (root.optString("kind") != "shonar") {
|
||||||
|
ProbeResult.Incompatible
|
||||||
|
} else {
|
||||||
|
val caps = root.optJSONObject("capabilities") ?: org.json.JSONObject()
|
||||||
|
val set = mutableSetOf<ProviderDescriptor.Capability>()
|
||||||
|
if (caps.optBoolean("chunked_upload")) set += ProviderDescriptor.Capability.CHUNKED_UPLOAD
|
||||||
|
if (caps.optBoolean("server_transcription")) set += ProviderDescriptor.Capability.SERVER_TRANSCRIPTION
|
||||||
|
if (caps.optBoolean("server_summary")) set += ProviderDescriptor.Capability.SERVER_SUMMARY
|
||||||
|
if (caps.optBoolean("account_deletion")) set += ProviderDescriptor.Capability.ACCOUNT_DELETION
|
||||||
|
ProbeResult.Compatible(
|
||||||
|
descriptor = ProviderDescriptor(
|
||||||
|
id = ProviderRegistry.CUSTOM_SHONAR_ID,
|
||||||
|
displayName = "Custom SHONAR server",
|
||||||
|
capabilities = set,
|
||||||
|
),
|
||||||
|
serverName = root.optString("storage_backend", "server"),
|
||||||
|
version = root.optString("version", "unknown"),
|
||||||
|
)
|
||||||
|
}
|
||||||
|
} catch (e: Exception) {
|
||||||
|
ProbeResult.Incompatible
|
||||||
|
}
|
||||||
|
|
||||||
|
/**
|
||||||
|
* SPKI SHA-256 recorded by the TOFU trust manager during the failed
|
||||||
|
* handshake, or "unknown" when nothing was captured (plain HTTP,
|
||||||
|
* pre-handshake failure, or no TOFU manager wired).
|
||||||
|
*/
|
||||||
|
private fun fingerprintFor(host: String): String =
|
||||||
|
tofu?.failureFor(host)?.spkiHex ?: "unknown"
|
||||||
|
|
||||||
|
companion object {
|
||||||
|
fun defaultClient(): OkHttpClient = OkHttpClient.Builder()
|
||||||
|
.connectTimeout(10, TimeUnit.SECONDS)
|
||||||
|
.readTimeout(15, TimeUnit.SECONDS)
|
||||||
|
.followRedirects(false) // do not silently follow redirects to other hosts
|
||||||
|
.build()
|
||||||
|
}
|
||||||
|
}
|
||||||
77
shared/com/shonar/provider/ShonarProvider.kt
Normal file
77
shared/com/shonar/provider/ShonarProvider.kt
Normal file
|
|
@ -0,0 +1,77 @@
|
||||||
|
package com.shonar.provider
|
||||||
|
|
||||||
|
import java.io.File
|
||||||
|
|
||||||
|
/**
|
||||||
|
* The single contract between SHONAR and any server (or the device itself).
|
||||||
|
*
|
||||||
|
* Everything the app needs from "the cloud" goes through this interface;
|
||||||
|
* UI, sync engine, and database reference providers by descriptor id only.
|
||||||
|
* LocalOnlyProvider is a first-class implementation, so a missing server is
|
||||||
|
* never a special case.
|
||||||
|
*
|
||||||
|
* Contract rules (enforced by the shared provider contract test suite):
|
||||||
|
* - all suspends are safe to cancel; partial uploads leave no visible object
|
||||||
|
* - upload() is idempotent per RecordingDraft.id (re-upload replaces the
|
||||||
|
* same key, never duplicates)
|
||||||
|
* - originals are immutable after successful upload until delete()
|
||||||
|
* - no method ever logs credentials, tokens, audio bytes, or transcripts
|
||||||
|
* - errors are ProviderError subtypes with secret-free messages
|
||||||
|
*/
|
||||||
|
interface ShonarProvider {
|
||||||
|
|
||||||
|
val descriptor: ProviderDescriptor
|
||||||
|
|
||||||
|
/** Current auth lifecycle; providers push updates here. */
|
||||||
|
val authState: kotlinx.coroutines.flow.StateFlow<AuthState>
|
||||||
|
|
||||||
|
/**
|
||||||
|
* Is there a SHONAR-compatible service at [baseUrl]? Purely read-only;
|
||||||
|
* must not require or request credentials. TlsFailure carries the SPKI
|
||||||
|
* fingerprint so the UI can run explicit trust-on-first-use approval —
|
||||||
|
* never silently accept.
|
||||||
|
*/
|
||||||
|
suspend fun probe(baseUrl: ServerUrl): ProbeResult
|
||||||
|
|
||||||
|
/** Validate [credential] against the server, then hand it to the secure store. */
|
||||||
|
suspend fun connect(credential: ProviderCredential): Unit
|
||||||
|
|
||||||
|
/** Re-validate stored credential (app start / after network return). */
|
||||||
|
suspend fun reconnect(): AuthState
|
||||||
|
|
||||||
|
/**
|
||||||
|
* Disconnect locally; if [revokeOnServer] and the provider supports it,
|
||||||
|
* revoke the credential server-side first (best effort — local state is
|
||||||
|
* cleared even if revoke fails, with the failure surfaced).
|
||||||
|
*/
|
||||||
|
suspend fun disconnect(revokeOnServer: Boolean)
|
||||||
|
|
||||||
|
/** Provider-native account/data deletion, then wipe local state. */
|
||||||
|
suspend fun deleteAccountAndData()
|
||||||
|
|
||||||
|
// ---- storage -----------------------------------------------------------
|
||||||
|
|
||||||
|
/**
|
||||||
|
* Upload the original audio for [draft], reporting [onProgress] 0..1.
|
||||||
|
* Implementations use chunked/resumable transfers when the capability is
|
||||||
|
* advertised. Returns the ref for later download/delete/sidecars.
|
||||||
|
*/
|
||||||
|
suspend fun upload(draft: RecordingDraft, onProgress: (Float) -> Unit): RemoteRef
|
||||||
|
|
||||||
|
suspend fun download(ref: RemoteRef, dest: File, onProgress: (Float) -> Unit)
|
||||||
|
|
||||||
|
suspend fun delete(ref: RemoteRef)
|
||||||
|
|
||||||
|
suspend fun list(cursor: String?): Page<RemoteRecording>
|
||||||
|
|
||||||
|
// ---- sidecars (transcript/summary/... JSON, synced independently) ------
|
||||||
|
|
||||||
|
suspend fun putSidecar(ref: RemoteRef, kind: SidecarKind, bytes: ByteArray)
|
||||||
|
|
||||||
|
/** null when the sidecar does not exist remotely. */
|
||||||
|
suspend fun getSidecar(ref: RemoteRef, kind: SidecarKind): ByteArray?
|
||||||
|
|
||||||
|
// ---- status ------------------------------------------------------------
|
||||||
|
|
||||||
|
suspend fun storageLocationSummary(): StorageLocation
|
||||||
|
}
|
||||||
42
shared/com/shonar/provider/SyncState.kt
Normal file
42
shared/com/shonar/provider/SyncState.kt
Normal file
|
|
@ -0,0 +1,42 @@
|
||||||
|
package com.shonar.provider
|
||||||
|
|
||||||
|
/**
|
||||||
|
* Sync lifecycle for one local recording against the active provider
|
||||||
|
* (docs/server-providers.md §2). P3 defines the vocabulary and the legal
|
||||||
|
* transitions; the WorkManager driver that moves recordings through it is
|
||||||
|
* M5 — until then uploads go through [ShonarProvider.upload] directly and
|
||||||
|
* land in UPLOADED.
|
||||||
|
*
|
||||||
|
* Pausing/canceling is a state, not a job kill: QUEUED and ERROR are stable
|
||||||
|
* resting states a later run resumes from.
|
||||||
|
*/
|
||||||
|
enum class SyncState {
|
||||||
|
LOCAL_ONLY, // provider is local-only, or user never enabled sync
|
||||||
|
QUEUED, // waiting for constraints (network, Wi-Fi-only, charging-only)
|
||||||
|
UPLOADING, // bytes in flight (resumable via the provider's session)
|
||||||
|
UPLOADED, // audio on the server; sidecars may still be pending
|
||||||
|
SYNCED, // audio + all sidecars confirmed server-side
|
||||||
|
ERROR, // failed; [SyncStatus.reasonCode] says why, [SyncStatus.retryAtEpochMs] when
|
||||||
|
}
|
||||||
|
|
||||||
|
data class SyncStatus(
|
||||||
|
val state: SyncState,
|
||||||
|
val retryAtEpochMs: Long? = null,
|
||||||
|
val reasonCode: String? = null,
|
||||||
|
)
|
||||||
|
|
||||||
|
/**
|
||||||
|
* Legal transitions. Anything not listed here is a programming error, not
|
||||||
|
* a state the UI should ever render.
|
||||||
|
*/
|
||||||
|
fun SyncStatus.canTransitionTo(next: SyncState): Boolean = when (state) {
|
||||||
|
SyncState.LOCAL_ONLY -> next == SyncState.QUEUED || next == SyncState.LOCAL_ONLY
|
||||||
|
SyncState.QUEUED -> next == SyncState.UPLOADING || next == SyncState.LOCAL_ONLY ||
|
||||||
|
next == SyncState.ERROR
|
||||||
|
SyncState.UPLOADING -> next == SyncState.UPLOADED || next == SyncState.QUEUED ||
|
||||||
|
next == SyncState.ERROR || next == SyncState.LOCAL_ONLY
|
||||||
|
SyncState.UPLOADED -> next == SyncState.SYNCED || next == SyncState.UPLOADING ||
|
||||||
|
next == SyncState.ERROR || next == SyncState.LOCAL_ONLY
|
||||||
|
SyncState.SYNCED -> next == SyncState.UPLOADING || next == SyncState.LOCAL_ONLY
|
||||||
|
SyncState.ERROR -> next == SyncState.QUEUED || next == SyncState.LOCAL_ONLY
|
||||||
|
}
|
||||||
190
shared/com/shonar/provider/TlsPolicy.kt
Normal file
190
shared/com/shonar/provider/TlsPolicy.kt
Normal file
|
|
@ -0,0 +1,190 @@
|
||||||
|
package com.shonar.provider
|
||||||
|
|
||||||
|
import java.util.concurrent.TimeUnit
|
||||||
|
import okhttp3.Interceptor
|
||||||
|
import okhttp3.OkHttpClient
|
||||||
|
import okhttp3.Response
|
||||||
|
|
||||||
|
/**
|
||||||
|
* P5: builds every provider HTTP client with the TOFU trust manager and the
|
||||||
|
* redacting logger (docs/server-providers.md §4). One place, so no call
|
||||||
|
* site can accidentally build a client that skips either.
|
||||||
|
*
|
||||||
|
* The log [sink] is an explicit constructor argument (no platform default)
|
||||||
|
* so this file compiles on Android and JVM desktop alike; each platform
|
||||||
|
* passes its own sink (Logcat vs stdout).
|
||||||
|
*/
|
||||||
|
class TlsPolicy(
|
||||||
|
val tofu: TofuManager,
|
||||||
|
private val bodiesEnabled: () -> Boolean = { false },
|
||||||
|
private val sink: (String) -> Unit,
|
||||||
|
) {
|
||||||
|
fun newClient(
|
||||||
|
connectTimeoutS: Long,
|
||||||
|
readTimeoutS: Long,
|
||||||
|
writeTimeoutS: Long = readTimeoutS,
|
||||||
|
followRedirects: Boolean = true,
|
||||||
|
): OkHttpClient {
|
||||||
|
val tm = tofu.trustManager
|
||||||
|
val sslContext = javax.net.ssl.SSLContext.getInstance("TLS")
|
||||||
|
sslContext.init(null, arrayOf(tm), null)
|
||||||
|
return OkHttpClient.Builder()
|
||||||
|
.sslSocketFactory(sslContext.socketFactory, tm)
|
||||||
|
.connectTimeout(connectTimeoutS, TimeUnit.SECONDS)
|
||||||
|
.readTimeout(readTimeoutS, TimeUnit.SECONDS)
|
||||||
|
.writeTimeout(writeTimeoutS, TimeUnit.SECONDS)
|
||||||
|
.followRedirects(followRedirects)
|
||||||
|
.addInterceptor(RedactingLogger(bodiesEnabled, sink))
|
||||||
|
.build()
|
||||||
|
}
|
||||||
|
|
||||||
|
/** P3-era API shape (15/60/60s). */
|
||||||
|
fun apiClient(): OkHttpClient = newClient(15, 60, 60)
|
||||||
|
|
||||||
|
/** Short probing shape (10/15s, no redirects). */
|
||||||
|
fun probeClient(): OkHttpClient = newClient(10, 15, 15, followRedirects = false)
|
||||||
|
|
||||||
|
/** Nextcloud shape (10/20/60s, no redirects). */
|
||||||
|
fun nextcloudClient(): OkHttpClient = newClient(10, 20, 60, followRedirects = false)
|
||||||
|
}
|
||||||
|
|
||||||
|
/**
|
||||||
|
* Logging is OFF by default; when the `log_http_bodies` debug setting is
|
||||||
|
* on, requests log in redacted form. Invariants (leak-tested):
|
||||||
|
* - Authorization / Cookie / Set-Cookie headers are never logged.
|
||||||
|
* - bodies log only for JSON/XML/text under [MAX_BODY] bytes; audio and
|
||||||
|
* other binary bodies never log.
|
||||||
|
* - token-shaped JSON values (`…token…`, `password`, `appPassword`) are
|
||||||
|
* masked, and Basic credentials are masked, even inside bodies.
|
||||||
|
*/
|
||||||
|
class RedactingLogger(
|
||||||
|
private val bodiesEnabled: () -> Boolean,
|
||||||
|
private val sink: (String) -> Unit,
|
||||||
|
) : Interceptor {
|
||||||
|
|
||||||
|
override fun intercept(chain: Interceptor.Chain): Response {
|
||||||
|
val sensitive = chain.request().header(SENSITIVE_BODY) != null
|
||||||
|
val req = chain.request().newBuilder().removeHeader(SENSITIVE_BODY).build()
|
||||||
|
if (!bodiesEnabled()) return chain.proceed(req)
|
||||||
|
val t0 = System.currentTimeMillis()
|
||||||
|
val reqBody = req.body
|
||||||
|
val reqLen = reqBody?.contentLength()?.takeIf { it >= 0 }
|
||||||
|
sink("→ ${req.method} ${req.url.host}${req.url.encodedPath} body=${reqLen?.let { "$it B" } ?: "?"}")
|
||||||
|
for (i in 0 until req.headers.size) {
|
||||||
|
val name = req.headers.name(i)
|
||||||
|
sink(" $name: ${if (isSensitiveHeader(name)) "[redacted]" else req.headers.value(i)}")
|
||||||
|
}
|
||||||
|
if (sensitive) {
|
||||||
|
sink(" request-body: [sensitive, not logged]")
|
||||||
|
} else {
|
||||||
|
logRequestBody(reqBody?.contentType()?.toString(), reqBody)
|
||||||
|
}
|
||||||
|
try {
|
||||||
|
val resp = chain.proceed(req)
|
||||||
|
val ms = System.currentTimeMillis() - t0
|
||||||
|
sink("← ${resp.code} ${req.url.encodedPath} (${ms}ms)")
|
||||||
|
if (sensitive) {
|
||||||
|
sink(" response-body: [sensitive, not logged]")
|
||||||
|
return resp
|
||||||
|
}
|
||||||
|
val peek = resp.peekBody(MAX_BODY + 1)
|
||||||
|
logBody(
|
||||||
|
" response",
|
||||||
|
peek.contentType()?.toString(),
|
||||||
|
peek.bytes().toList(),
|
||||||
|
)
|
||||||
|
return resp
|
||||||
|
} catch (e: Exception) {
|
||||||
|
sink("✕ ${req.url.encodedPath} failed (${e.javaClass.simpleName})")
|
||||||
|
throw e
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
/**
|
||||||
|
* Request bodies need a size gate BEFORE reading: uploads are file
|
||||||
|
* bodies that must never be buffered just to log a prefix of them.
|
||||||
|
*/
|
||||||
|
private fun logRequestBody(contentType: String?, body: okhttp3.RequestBody?) {
|
||||||
|
if (body == null) {
|
||||||
|
sink(" request-body: none")
|
||||||
|
return
|
||||||
|
}
|
||||||
|
val len = try {
|
||||||
|
body.contentLength()
|
||||||
|
} catch (e: Exception) {
|
||||||
|
-1L
|
||||||
|
}
|
||||||
|
if (len < 0 || len > MAX_BODY) {
|
||||||
|
val what = if (len < 0) "streaming" else "$len bytes"
|
||||||
|
sink(" request-body: [$what, not logged]")
|
||||||
|
return
|
||||||
|
}
|
||||||
|
val bytes = try {
|
||||||
|
val buf = okio.Buffer()
|
||||||
|
body.writeTo(buf)
|
||||||
|
buf.readByteArray().toList()
|
||||||
|
} catch (e: Exception) {
|
||||||
|
null
|
||||||
|
}
|
||||||
|
logBody(" request", contentType, bytes)
|
||||||
|
}
|
||||||
|
|
||||||
|
private fun logBody(prefix: String, contentType: String?, bytes: List<Byte>?) {
|
||||||
|
if (bytes == null) {
|
||||||
|
sink("$prefix-body: none")
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if (!isLoggableType(contentType)) {
|
||||||
|
sink("$prefix-body: [${bytes.size} bytes, not logged]")
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if (bytes.size > MAX_BODY) {
|
||||||
|
sink("$prefix-body: [${bytes.size} bytes, over the $MAX_BODY B cap, not logged]")
|
||||||
|
return
|
||||||
|
}
|
||||||
|
sink("$prefix-body: ${redact(bytes.toByteArray().toString(Charsets.UTF_8))}")
|
||||||
|
}
|
||||||
|
|
||||||
|
companion object {
|
||||||
|
const val MAX_BODY: Long = 8192
|
||||||
|
|
||||||
|
/**
|
||||||
|
* Opt out of body logging per request (both directions). Providers
|
||||||
|
* set this on calls whose bodies are transcripts or other sensitive
|
||||||
|
* payloads; the logger strips it before sending so it never reaches
|
||||||
|
* the wire.
|
||||||
|
*/
|
||||||
|
const val SENSITIVE_BODY = "X-Shonar-Sensitive-Body"
|
||||||
|
|
||||||
|
internal fun isSensitiveHeader(name: String): Boolean = when (name.lowercase()) {
|
||||||
|
"authorization", "cookie", "set-cookie", "x-chunk-sha256" -> true
|
||||||
|
else -> false
|
||||||
|
}
|
||||||
|
|
||||||
|
internal fun isLoggableType(contentType: String?): Boolean {
|
||||||
|
if (contentType == null) return false
|
||||||
|
val t = contentType.lowercase().substringBefore(';').trim()
|
||||||
|
return t.startsWith("application/json") || t.endsWith("+json") ||
|
||||||
|
t.startsWith("text/") || t.endsWith("+xml") ||
|
||||||
|
t == "application/xml" || t == "application/x-www-form-urlencoded"
|
||||||
|
}
|
||||||
|
|
||||||
|
private val SECRET_JSON = Regex(
|
||||||
|
""""[^"]*(token|password|secret)[^"]*"\s*:\s*"[^"]*"""",
|
||||||
|
RegexOption.IGNORE_CASE,
|
||||||
|
)
|
||||||
|
private val BASIC = Regex("""Basic\s+[A-Za-z0-9+/=]{8,}""")
|
||||||
|
private val BEARER = Regex("""Bearer\s+[A-Za-z0-9\-_.~+/=]{8,}""")
|
||||||
|
|
||||||
|
/** Mask token-shaped values; safe to run on any text. */
|
||||||
|
internal fun redact(text: String): String {
|
||||||
|
var out = SECRET_JSON.replace(text) { m ->
|
||||||
|
val key = m.value.substringBefore(':')
|
||||||
|
"""$key:"***""""
|
||||||
|
}
|
||||||
|
out = BASIC.replace(out, "Basic [redacted]")
|
||||||
|
out = BEARER.replace(out, "Bearer [redacted]")
|
||||||
|
return out
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
287
shared/com/shonar/provider/TofuTrust.kt
Normal file
287
shared/com/shonar/provider/TofuTrust.kt
Normal file
|
|
@ -0,0 +1,287 @@
|
||||||
|
package com.shonar.provider
|
||||||
|
|
||||||
|
import com.shonar.settings.SettingsStore
|
||||||
|
import java.net.Socket
|
||||||
|
import java.security.MessageDigest
|
||||||
|
import java.security.cert.CertificateException
|
||||||
|
import java.security.cert.X509Certificate
|
||||||
|
import java.util.concurrent.ConcurrentHashMap
|
||||||
|
import javax.net.ssl.SSLEngine
|
||||||
|
import javax.net.ssl.SSLSession
|
||||||
|
import javax.net.ssl.X509ExtendedTrustManager
|
||||||
|
import javax.net.ssl.X509TrustManager
|
||||||
|
|
||||||
|
/**
|
||||||
|
* P5: trust-on-first-use (TOFU) for self-signed LAN servers
|
||||||
|
* (docs/server-providers.md §4).
|
||||||
|
*
|
||||||
|
* Policy, enforced by construction:
|
||||||
|
* - system CAs are always tried first; TOFU pins are a fallback, never a
|
||||||
|
* replacement. A host that later gets a real certificate just works.
|
||||||
|
* - pins are per-host: an approved cert for `nas.local` is trusted ONLY
|
||||||
|
* when `nas.local` presents it. Byte-equality on the leaf DER, so a
|
||||||
|
* rotation needs a fresh approval (correct TOFU semantics).
|
||||||
|
* - nothing is ever trusted silently: the first failure only RECORDS the
|
||||||
|
* chain, and trust requires an explicit [TofuManager.approve] call from
|
||||||
|
* a UI that showed the fingerprint.
|
||||||
|
* - the normal TLS hostname verifier stays strict; TOFU covers unknown
|
||||||
|
* CAs, not name mismatches.
|
||||||
|
*/
|
||||||
|
class TofuStore(private val secure: SettingsStore) {
|
||||||
|
|
||||||
|
suspend fun addPin(host: String, derBase64: String) {
|
||||||
|
val pins = pins(host).toMutableSet()
|
||||||
|
pins += derBase64
|
||||||
|
secure.putString(pinKey(host), org.json.JSONArray(pins.toList()).toString())
|
||||||
|
val hosts = hosts().toMutableSet()
|
||||||
|
if (hosts.add(host.lowercase())) {
|
||||||
|
secure.putString(KEY_HOSTS, org.json.JSONArray(hosts.toList()).toString())
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
suspend fun pins(host: String): Set<String> {
|
||||||
|
val raw = secure.getString(pinKey(host.lowercase())) ?: return emptySet()
|
||||||
|
return runCatching {
|
||||||
|
val arr = org.json.JSONArray(raw)
|
||||||
|
(0 until arr.length()).map { arr.getString(it) }.toSet()
|
||||||
|
}.getOrDefault(emptySet())
|
||||||
|
}
|
||||||
|
|
||||||
|
suspend fun hosts(): Set<String> {
|
||||||
|
val raw = secure.getString(KEY_HOSTS) ?: return emptySet()
|
||||||
|
return runCatching {
|
||||||
|
val arr = org.json.JSONArray(raw)
|
||||||
|
(0 until arr.length()).map { arr.getString(it) }.toSet()
|
||||||
|
}.getOrDefault(emptySet())
|
||||||
|
}
|
||||||
|
|
||||||
|
suspend fun removeHost(host: String) {
|
||||||
|
secure.remove(pinKey(host.lowercase()))
|
||||||
|
val hosts = hosts().toMutableSet()
|
||||||
|
if (hosts.remove(host.lowercase())) {
|
||||||
|
secure.putString(KEY_HOSTS, org.json.JSONArray(hosts.toList()).toString())
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
companion object {
|
||||||
|
private const val PREFIX = "tofu."
|
||||||
|
const val KEY_HOSTS = PREFIX + "hosts"
|
||||||
|
fun pinKey(host: String): String = PREFIX + "pins." + host.lowercase()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
/** A recorded untrusted chain, shown to the user for approval. */
|
||||||
|
data class TofuFailure(
|
||||||
|
val host: String,
|
||||||
|
/** SPKI SHA-256, browser-style `AB:CD:…` uppercase hex. */
|
||||||
|
val spkiHex: String,
|
||||||
|
val subject: String,
|
||||||
|
val issuer: String,
|
||||||
|
val validFrom: String,
|
||||||
|
val validUntil: String,
|
||||||
|
val leafDer: ByteArray,
|
||||||
|
val recordedAtMs: Long = System.currentTimeMillis(),
|
||||||
|
) {
|
||||||
|
override fun toString(): String =
|
||||||
|
"TofuFailure(host=$host, spki=$spkiHex, subject=$subject)"
|
||||||
|
|
||||||
|
override fun equals(other: Any?): Boolean {
|
||||||
|
if (this === other) return true
|
||||||
|
if (other !is TofuFailure) return false
|
||||||
|
return host == other.host && spkiHex == other.spkiHex &&
|
||||||
|
leafDer.contentEquals(other.leafDer)
|
||||||
|
}
|
||||||
|
|
||||||
|
override fun hashCode(): Int = 31 * host.hashCode() + spkiHex.hashCode()
|
||||||
|
}
|
||||||
|
|
||||||
|
/**
|
||||||
|
* Thrown (as a [CertificateException], so it propagates through the TLS
|
||||||
|
* stack untouched) when a chain is neither system-trusted nor pinned.
|
||||||
|
* Carries no secrets — DNs and fingerprints are public cert contents.
|
||||||
|
*/
|
||||||
|
class TofuUntrustedException(
|
||||||
|
val failure: TofuFailure,
|
||||||
|
message: String = "Untrusted certificate for ${failure.host} (SPKI ${failure.spkiHex})",
|
||||||
|
) : CertificateException(message)
|
||||||
|
|
||||||
|
/**
|
||||||
|
* System-first trust manager with per-host TOFU fallback. The pin cache is
|
||||||
|
* in-memory (handshakes run on TLS threads that cannot suspend); it is
|
||||||
|
* filled by [TofuManager.refresh] at startup and kept in sync by
|
||||||
|
* approve/forget.
|
||||||
|
*/
|
||||||
|
class TofuTrustManager(
|
||||||
|
private val system: X509TrustManager,
|
||||||
|
private val manager: TofuManager,
|
||||||
|
) : X509ExtendedTrustManager() {
|
||||||
|
|
||||||
|
override fun checkClientTrusted(chain: Array<X509Certificate>, authType: String) {
|
||||||
|
system.checkClientTrusted(chain, authType)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Client-auth paths: this app is always the TLS client, so these only
|
||||||
|
// need to exist (newer JDKs/Android declare them abstract). Delegate to
|
||||||
|
// the 2-arg system check.
|
||||||
|
override fun checkClientTrusted(
|
||||||
|
chain: Array<X509Certificate>, authType: String, socket: Socket,
|
||||||
|
) {
|
||||||
|
system.checkClientTrusted(chain, authType)
|
||||||
|
}
|
||||||
|
|
||||||
|
override fun checkClientTrusted(
|
||||||
|
chain: Array<X509Certificate>, authType: String, engine: SSLEngine,
|
||||||
|
) {
|
||||||
|
system.checkClientTrusted(chain, authType)
|
||||||
|
}
|
||||||
|
|
||||||
|
override fun checkServerTrusted(chain: Array<X509Certificate>, authType: String) {
|
||||||
|
checkServerTrusted(chain, authType, host = null)
|
||||||
|
}
|
||||||
|
|
||||||
|
override fun checkServerTrusted(
|
||||||
|
chain: Array<X509Certificate>, authType: String, socket: Socket,
|
||||||
|
) {
|
||||||
|
val host = runCatching {
|
||||||
|
(socket as? javax.net.ssl.SSLSocket)?.handshakeSession?.peerHost
|
||||||
|
}.getOrNull()
|
||||||
|
checkServerTrusted(chain, authType, host = host)
|
||||||
|
}
|
||||||
|
|
||||||
|
override fun checkServerTrusted(
|
||||||
|
chain: Array<X509Certificate>, authType: String, engine: SSLEngine,
|
||||||
|
) {
|
||||||
|
val host = runCatching {
|
||||||
|
(engine.session as? javax.net.ssl.ExtendedSSLSession)?.peerHost
|
||||||
|
}.getOrNull()
|
||||||
|
checkServerTrusted(chain, authType, host = host)
|
||||||
|
}
|
||||||
|
|
||||||
|
private fun checkServerTrusted(
|
||||||
|
chain: Array<X509Certificate>, authType: String, host: String?,
|
||||||
|
) {
|
||||||
|
// pins are checked first so an approved self-signed cert keeps
|
||||||
|
// working even where a system path would also exist; system trust
|
||||||
|
// is the fallback that makes later real certs frictionless.
|
||||||
|
if (chain.isNotEmpty() && host != null && manager.isPinned(host, chain[0])) {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
try {
|
||||||
|
system.checkServerTrusted(chain, authType)
|
||||||
|
} catch (e: CertificateException) {
|
||||||
|
val failure = TofuFailure(
|
||||||
|
host = (host ?: "").lowercase(),
|
||||||
|
spkiHex = spkiHex(chain[0]),
|
||||||
|
subject = chain[0].subjectX500Principal.name,
|
||||||
|
issuer = chain[0].issuerX500Principal.name,
|
||||||
|
validFrom = chain[0].notBefore.toString(),
|
||||||
|
validUntil = chain[0].notAfter.toString(),
|
||||||
|
leafDer = chain[0].encoded,
|
||||||
|
)
|
||||||
|
manager.record(failure)
|
||||||
|
throw TofuUntrustedException(failure)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
override fun getAcceptedIssuers(): Array<X509Certificate> = system.acceptedIssuers
|
||||||
|
|
||||||
|
companion object {
|
||||||
|
/** SPKI SHA-256 as browser-style colon-separated uppercase hex. */
|
||||||
|
internal fun spkiHex(cert: X509Certificate): String {
|
||||||
|
val digest = MessageDigest.getInstance("SHA-256")
|
||||||
|
.digest(cert.publicKey.encoded)
|
||||||
|
return digest.joinToString(":") { "%02X".format(it) }
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
/**
|
||||||
|
* Owns the pin store, the in-memory cache the trust manager reads, and the
|
||||||
|
* recorded failures the approval UI consumes.
|
||||||
|
*/
|
||||||
|
class TofuManager(
|
||||||
|
private val store: TofuStore,
|
||||||
|
system: X509TrustManager = defaultSystemTrustManager(),
|
||||||
|
) {
|
||||||
|
private val cache = ConcurrentHashMap<String, Set<String>>()
|
||||||
|
private val failures = ConcurrentHashMap<String, TofuFailure>()
|
||||||
|
|
||||||
|
val trustManager: TofuTrustManager by lazy { TofuTrustManager(system, this) }
|
||||||
|
|
||||||
|
/** Fill the cache from the store. Call at startup (and only there). */
|
||||||
|
suspend fun refresh() {
|
||||||
|
val fresh = mutableMapOf<String, Set<String>>()
|
||||||
|
for (host in store.hosts()) {
|
||||||
|
fresh[host.lowercase()] = store.pins(host)
|
||||||
|
}
|
||||||
|
cache.clear()
|
||||||
|
cache.putAll(fresh)
|
||||||
|
}
|
||||||
|
|
||||||
|
internal fun isPinned(host: String, leaf: X509Certificate): Boolean {
|
||||||
|
val pins = cache[host.lowercase()] ?: return false
|
||||||
|
val der = runCatching { leaf.encoded }.getOrNull() ?: return false
|
||||||
|
return pins.any { pin ->
|
||||||
|
runCatching {
|
||||||
|
MessageDigest.isEqual(
|
||||||
|
der,
|
||||||
|
java.util.Base64.getDecoder().decode(pin),
|
||||||
|
)
|
||||||
|
}.getOrDefault(false)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
internal fun record(failure: TofuFailure) {
|
||||||
|
failures[failure.host] = failure
|
||||||
|
}
|
||||||
|
|
||||||
|
/** Most recent failure for [host], else a fresh hostless one, else null. */
|
||||||
|
fun failureFor(host: String): TofuFailure? {
|
||||||
|
val key = host.lowercase()
|
||||||
|
failures[key]?.let { return it }
|
||||||
|
val fallback = failures[""] ?: return null
|
||||||
|
// A hostless record only belongs to this probe if it just happened
|
||||||
|
// (single onboarding flow, no concurrency to speak of).
|
||||||
|
if (System.currentTimeMillis() - fallback.recordedAtMs > FRESH_MS) return null
|
||||||
|
return fallback
|
||||||
|
}
|
||||||
|
|
||||||
|
/**
|
||||||
|
* Trust [host]'s recorded leaf from now on. Returns false when there is
|
||||||
|
* nothing recorded (never invent trust).
|
||||||
|
*/
|
||||||
|
suspend fun approve(host: String): Boolean {
|
||||||
|
val key = host.lowercase()
|
||||||
|
val failure = failureFor(key) ?: return false
|
||||||
|
val der = java.util.Base64.getEncoder().encodeToString(failure.leafDer)
|
||||||
|
store.addPin(key, der)
|
||||||
|
cache[key] = store.pins(key)
|
||||||
|
failures.remove(key)
|
||||||
|
failures.remove("")
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
|
||||||
|
suspend fun decline(host: String) {
|
||||||
|
failures.remove(host.lowercase())
|
||||||
|
failures.remove("")
|
||||||
|
}
|
||||||
|
|
||||||
|
suspend fun forget(host: String) {
|
||||||
|
store.removeHost(host)
|
||||||
|
cache.remove(host.lowercase())
|
||||||
|
}
|
||||||
|
|
||||||
|
suspend fun pinnedHosts(): Set<String> = store.hosts()
|
||||||
|
|
||||||
|
companion object {
|
||||||
|
private const val FRESH_MS = 30_000L
|
||||||
|
|
||||||
|
fun defaultSystemTrustManager(): X509TrustManager {
|
||||||
|
val factory = javax.net.ssl.TrustManagerFactory.getInstance(
|
||||||
|
javax.net.ssl.TrustManagerFactory.getDefaultAlgorithm()
|
||||||
|
)
|
||||||
|
factory.init(null as java.security.KeyStore?)
|
||||||
|
return factory.trustManagers.filterIsInstance<X509TrustManager>().first()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
225
shared/com/shonar/recording/AiContent.kt
Normal file
225
shared/com/shonar/recording/AiContent.kt
Normal file
|
|
@ -0,0 +1,225 @@
|
||||||
|
package com.shonar.recording
|
||||||
|
|
||||||
|
import org.json.JSONArray
|
||||||
|
import org.json.JSONObject
|
||||||
|
|
||||||
|
/**
|
||||||
|
* M8: AI content models for the recording details screen.
|
||||||
|
*
|
||||||
|
* Pure Kotlin + org.json — no Android imports, so this file runs under
|
||||||
|
* plain JVM unit tests. The backend shapes mirror
|
||||||
|
* `backend/shonar/api/schemas_recordings.py`:
|
||||||
|
* - transcript: {version, language, provider, model, text,
|
||||||
|
* segments: [{start, end, text, speaker}], edited_by_user}
|
||||||
|
* - summary: {version, provider, model, content: {short, detailed,
|
||||||
|
* key_points, decisions, action_items, questions}, edited_by_user}
|
||||||
|
* - jobs: [{job_type, status, attempt, max_attempts, error, ...}]
|
||||||
|
*/
|
||||||
|
|
||||||
|
data class TranscriptSegment(
|
||||||
|
val startSec: Double,
|
||||||
|
val endSec: Double,
|
||||||
|
val text: String,
|
||||||
|
val speaker: String? = null,
|
||||||
|
)
|
||||||
|
|
||||||
|
data class TranscriptData(
|
||||||
|
val version: Int,
|
||||||
|
val language: String? = null,
|
||||||
|
val provider: String = "",
|
||||||
|
val model: String? = null,
|
||||||
|
val text: String = "",
|
||||||
|
val segments: List<TranscriptSegment> = emptyList(),
|
||||||
|
val editedByUser: Boolean = false,
|
||||||
|
)
|
||||||
|
|
||||||
|
data class SummaryData(
|
||||||
|
val version: Int,
|
||||||
|
val provider: String = "",
|
||||||
|
val model: String? = null,
|
||||||
|
/** Raw content dict; known keys read via [summaryString]/[summaryList]. */
|
||||||
|
val content: Map<String, Any?> = emptyMap(),
|
||||||
|
val editedByUser: Boolean = false,
|
||||||
|
) {
|
||||||
|
val short: String get() = summaryString(content, "short")
|
||||||
|
val detailed: String get() = summaryString(content, "detailed")
|
||||||
|
fun list(key: String): List<String> = summaryList(content, key)
|
||||||
|
}
|
||||||
|
|
||||||
|
data class JobInfo(
|
||||||
|
val jobType: String,
|
||||||
|
val status: String,
|
||||||
|
val attempt: Int = 0,
|
||||||
|
val maxAttempts: Int = 0,
|
||||||
|
val error: String? = null,
|
||||||
|
/** Display-only phase ("loading-model", "transcribing") + 0-100 progress. */
|
||||||
|
val stage: String? = null,
|
||||||
|
val progress: Int? = null,
|
||||||
|
)
|
||||||
|
|
||||||
|
/** Summary content keys, in display order. */
|
||||||
|
val SUMMARY_LIST_KEYS = listOf("key_points", "decisions", "action_items", "questions")
|
||||||
|
|
||||||
|
/**
|
||||||
|
* Index of the segment playing at [positionMs]: the last segment whose
|
||||||
|
* start is at or before the position. -1 when there are no segments or
|
||||||
|
* playback precedes the first one. End times are ignored on purpose —
|
||||||
|
* gaps between segments keep the last-spoken line highlighted instead of
|
||||||
|
* flickering to nothing.
|
||||||
|
*/
|
||||||
|
fun activeSegmentIndex(segments: List<TranscriptSegment>, positionMs: Long): Int {
|
||||||
|
val posSec = positionMs / 1000.0
|
||||||
|
var active = -1
|
||||||
|
for (i in segments.indices) {
|
||||||
|
if (segments[i].startSec <= posSec) active = i else break
|
||||||
|
}
|
||||||
|
return active
|
||||||
|
}
|
||||||
|
|
||||||
|
fun parseTranscript(raw: String): TranscriptData? = runCatching {
|
||||||
|
val json = JSONObject(raw)
|
||||||
|
val segments = mutableListOf<TranscriptSegment>()
|
||||||
|
json.optJSONArray("segments")?.let { arr ->
|
||||||
|
for (i in 0 until arr.length()) {
|
||||||
|
val s = arr.optJSONObject(i) ?: continue
|
||||||
|
segments += TranscriptSegment(
|
||||||
|
startSec = s.optDouble("start", 0.0),
|
||||||
|
endSec = s.optDouble("end", 0.0),
|
||||||
|
text = s.optString("text", ""),
|
||||||
|
speaker = s.optString("speaker", null).takeUnless { it.isNullOrBlank() },
|
||||||
|
)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
segments.sortBy { it.startSec }
|
||||||
|
TranscriptData(
|
||||||
|
version = json.optInt("version", 0),
|
||||||
|
language = json.optString("language", null).takeUnless { it.isNullOrBlank() },
|
||||||
|
provider = json.optString("provider", ""),
|
||||||
|
model = json.optString("model", null).takeUnless { it.isNullOrBlank() },
|
||||||
|
text = json.optString("text", ""),
|
||||||
|
segments = segments,
|
||||||
|
editedByUser = json.optBoolean("edited_by_user", false),
|
||||||
|
)
|
||||||
|
}.getOrNull()
|
||||||
|
|
||||||
|
fun parseSummary(raw: String): SummaryData? = runCatching {
|
||||||
|
val json = JSONObject(raw)
|
||||||
|
val content = mutableMapOf<String, Any?>()
|
||||||
|
json.optJSONObject("content")?.let { obj ->
|
||||||
|
for (key in obj.keys()) content[key] = obj.get(key)
|
||||||
|
}
|
||||||
|
SummaryData(
|
||||||
|
version = json.optInt("version", 0),
|
||||||
|
provider = json.optString("provider", ""),
|
||||||
|
model = json.optString("model", null).takeUnless { it.isNullOrBlank() },
|
||||||
|
content = content,
|
||||||
|
editedByUser = json.optBoolean("edited_by_user", false),
|
||||||
|
)
|
||||||
|
}.getOrNull()
|
||||||
|
|
||||||
|
fun parseJobs(raw: String): List<JobInfo> = runCatching {
|
||||||
|
val arr = JSONArray(raw)
|
||||||
|
(0 until arr.length()).mapNotNull { i ->
|
||||||
|
val o = arr.optJSONObject(i) ?: return@mapNotNull null
|
||||||
|
JobInfo(
|
||||||
|
jobType = o.optString("job_type", ""),
|
||||||
|
status = o.optString("status", ""),
|
||||||
|
attempt = o.optInt("attempt", 0),
|
||||||
|
maxAttempts = o.optInt("max_attempts", 0),
|
||||||
|
error = o.optString("error", null).takeUnless { it.isNullOrBlank() },
|
||||||
|
stage = o.optString("stage", null).takeUnless { it.isNullOrBlank() },
|
||||||
|
progress = if (o.has("progress") && !o.isNull("progress")) o.optInt("progress") else null,
|
||||||
|
)
|
||||||
|
}
|
||||||
|
}.getOrDefault(emptyList())
|
||||||
|
|
||||||
|
/** Payload for PUT /recordings/{id}/transcript. Null segments = text-only correction. */
|
||||||
|
fun transcriptUpdatePayload(text: String, segments: List<TranscriptSegment>? = null): String {
|
||||||
|
val obj = JSONObject().put("text", text)
|
||||||
|
if (segments != null) {
|
||||||
|
obj.put(
|
||||||
|
"segments",
|
||||||
|
JSONArray().also { arr ->
|
||||||
|
segments.forEach { s ->
|
||||||
|
arr.put(
|
||||||
|
JSONObject()
|
||||||
|
.put("start", s.startSec)
|
||||||
|
.put("end", s.endSec)
|
||||||
|
.put("text", s.text)
|
||||||
|
.let { o -> if (s.speaker != null) o.put("speaker", s.speaker) else o },
|
||||||
|
)
|
||||||
|
}
|
||||||
|
},
|
||||||
|
)
|
||||||
|
}
|
||||||
|
return obj.toString()
|
||||||
|
}
|
||||||
|
|
||||||
|
/** Payload for PUT /recordings/{id}/summary. */
|
||||||
|
fun summaryUpdatePayload(content: Map<String, Any?>): String =
|
||||||
|
JSONObject().put("content", JSONObject(content)).toString()
|
||||||
|
|
||||||
|
/** Payload for PATCH /recordings/{id} (title/notes rename from details). */
|
||||||
|
fun recordingPatchPayload(title: String? = null, notes: String? = null): String {
|
||||||
|
val obj = JSONObject()
|
||||||
|
if (title != null) obj.put("title", title)
|
||||||
|
if (notes != null) obj.put("notes", notes)
|
||||||
|
return obj.toString()
|
||||||
|
}
|
||||||
|
|
||||||
|
// ---- transcription models (Stage 1: GET /api/v1/models) ------------------------
|
||||||
|
|
||||||
|
data class TranscriptionModelInfo(
|
||||||
|
val name: String,
|
||||||
|
val displayName: String,
|
||||||
|
val description: String,
|
||||||
|
val params: String = "",
|
||||||
|
val approxMemory: String = "",
|
||||||
|
val relativeSpeed: String = "",
|
||||||
|
val isDefault: Boolean = false,
|
||||||
|
val downloaded: Boolean = false,
|
||||||
|
val available: Boolean = false,
|
||||||
|
)
|
||||||
|
|
||||||
|
data class ModelsData(
|
||||||
|
val defaultModel: String,
|
||||||
|
val fasterWhisperInstalled: Boolean,
|
||||||
|
val models: List<TranscriptionModelInfo> = emptyList(),
|
||||||
|
) {
|
||||||
|
fun byName(name: String): TranscriptionModelInfo? = models.firstOrNull { it.name == name }
|
||||||
|
}
|
||||||
|
|
||||||
|
fun parseModels(raw: String): ModelsData? = runCatching {
|
||||||
|
val json = JSONObject(raw)
|
||||||
|
val models = mutableListOf<TranscriptionModelInfo>()
|
||||||
|
json.optJSONArray("models")?.let { arr ->
|
||||||
|
for (i in 0 until arr.length()) {
|
||||||
|
val o = arr.optJSONObject(i) ?: continue
|
||||||
|
models += TranscriptionModelInfo(
|
||||||
|
name = o.optString("name", ""),
|
||||||
|
displayName = o.optString("display_name", o.optString("name", "")),
|
||||||
|
description = o.optString("description", ""),
|
||||||
|
params = o.optString("params", ""),
|
||||||
|
approxMemory = o.optString("approx_memory", ""),
|
||||||
|
relativeSpeed = o.optString("relative_speed", ""),
|
||||||
|
isDefault = o.optBoolean("is_default", false),
|
||||||
|
downloaded = o.optBoolean("downloaded", false),
|
||||||
|
available = o.optBoolean("available", false),
|
||||||
|
)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
ModelsData(
|
||||||
|
defaultModel = json.optString("default_model", "base"),
|
||||||
|
fasterWhisperInstalled = json.optBoolean("faster_whisper_installed", false),
|
||||||
|
models = models.filter { it.name.isNotBlank() },
|
||||||
|
)
|
||||||
|
}.getOrNull()
|
||||||
|
|
||||||
|
private fun summaryString(content: Map<String, Any?>, key: String): String =
|
||||||
|
(content[key] as? String).orEmpty()
|
||||||
|
|
||||||
|
private fun summaryList(content: Map<String, Any?>, key: String): List<String> = when (val v = content[key]) {
|
||||||
|
is JSONArray -> (0 until v.length()).mapNotNull { v.optString(it, null)?.takeIf { it.isNotBlank() } }
|
||||||
|
is List<*> -> v.mapNotNull { (it as? String)?.takeIf { it.isNotBlank() } }
|
||||||
|
else -> emptyList()
|
||||||
|
}
|
||||||
Some files were not shown because too many files have changed in this diff Show more
Loading…
Add table
Add a link
Reference in a new issue