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:
avi 2026-09-14 17:14:54 -05:00
commit 76c867fca4
136 changed files with 21099 additions and 0 deletions

9
.gitignore vendored Normal file
View 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
View 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()
}

File diff suppressed because it is too large Load diff

View 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
}
}

View 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)
}
}
}
}
}
}
}

View 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))
}

View 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)
}

View 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) }
}
}

View 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)
}

View 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()
}
}

View 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)
}
}

View 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 })
}
}

View 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))
}
}

View 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)
}
}

View 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
View 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
View 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
View 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()

View 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"}

View 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 ###

View 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;")

View 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
View 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

View file

@ -0,0 +1,3 @@
"""S.H.O.N.A.R. — Self-hosted Oral Notes and Audio Recorder."""
__version__ = "0.1.0"

View file

View 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)]

View 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

View 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)

View 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

View 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)

View 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}

View 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(),
}

View 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),
)

View 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,
}

View 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

View 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",
},
)

View 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

View file

View 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()

View 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)

View 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)

View 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
View 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
)

View 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
View 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)

View 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
View 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"}

View file

View 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}")

View 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)

View 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)

View 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

View 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)

View 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)

View 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")

View 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).

View 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()

View 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()

View 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")

View 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

View 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

View 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)

View 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

View 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
View 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

View file

113
backend/tests/conftest.py Normal file
View 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)

View 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")

View 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
View 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

View 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

View 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
View 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)

View 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)

View 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)

View 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"]

View 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
View 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

Binary file not shown.

View 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
View 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
View 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
View file

@ -0,0 +1,17 @@
pluginManagement {
repositories {
mavenCentral()
gradlePluginPortal()
google()
}
}
dependencyResolutionManagement {
repositories {
mavenCentral()
google()
}
}
rootProject.name = "shonar-desktop"
include(":app")

View 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
}
}

View 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)
)
}
}

View 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)
}
}
}
}

View 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"
}
}

View 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) }
}
}
}

View 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()
}

View 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])"
}

View 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)
}
}
}
}

View 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,
)
}
}
}

View 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)
}

View 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])"
}

View 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()
}
}

View 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
}

View 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
}

View 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
}
}
}

View 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()
}
}
}

View 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