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

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