Standalone Shonar Desktop: vendor portable sources + local engine; decouple from ~/Projects/Shonar
- shared/ = portable Android-origin sources vendored from deferred/desktop-server (app/build.gradle.kts srcDir repointed; PlaybackController.kt excluded as Android-only) - backend/ = bundled-lite engine (SQLite + inline queue); .venv symlinked from the old checkout, PYTHONPATH pins THIS backend's code over any editable install - repoRoot() resolves this project dir (env SHONAR_REPO still wins); desktop-dev.sh watches shared/ + backend/ - Verified: :app:compileKotlin + :app:test green (23 tests); engine boots on :8010, self-migrates, /healthz ok
This commit is contained in:
commit
76c867fca4
136 changed files with 21099 additions and 0 deletions
0
backend/shonar/services/__init__.py
Normal file
0
backend/shonar/services/__init__.py
Normal file
160
backend/shonar/services/ai/__init__.py
Normal file
160
backend/shonar/services/ai/__init__.py
Normal file
|
|
@ -0,0 +1,160 @@
|
|||
"""AI provider interfaces (M7).
|
||||
|
||||
Two independent axes, both optional and both configured only through
|
||||
environment variables (never hard-coded keys):
|
||||
|
||||
- transcription: none | whisper_http | faster_whisper
|
||||
- LLM (summary/action items): none | openai_compat | ollama
|
||||
|
||||
"none" is a first-class choice: recording, sync, playback, and manual
|
||||
transcripts work with no AI configured at all. The pipeline treats a
|
||||
missing provider as "skip this stage", never as an error.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Protocol
|
||||
|
||||
from shonar.core.config import Settings
|
||||
|
||||
|
||||
class AIError(Exception):
|
||||
"""Base for AI failures. Messages must be user-safe: they surface in
|
||||
``processing_error`` and therefore on screens."""
|
||||
|
||||
|
||||
class ProviderConfigError(AIError):
|
||||
"""Persistent misconfiguration (bad credentials, unknown model, missing
|
||||
dependency). Fails the job immediately — retrying cannot help."""
|
||||
|
||||
|
||||
class ProviderTransientError(AIError):
|
||||
"""May succeed on retry (timeouts, 429/5xx). The worker requeues these
|
||||
up to the job's max attempts."""
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class Segment:
|
||||
start: float
|
||||
end: float
|
||||
text: str
|
||||
speaker: str | None = None
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class TranscriptResult:
|
||||
text: str
|
||||
language: str | None
|
||||
segments: list[Segment] = field(default_factory=list)
|
||||
model: str = ""
|
||||
|
||||
|
||||
class TranscriptionProvider(Protocol):
|
||||
name: str
|
||||
|
||||
async def transcribe(
|
||||
self,
|
||||
audio: bytes,
|
||||
mime: str,
|
||||
*,
|
||||
language_hint: str | None = None,
|
||||
on_progress=None, # optional Callable[[int], None], 0..99
|
||||
) -> TranscriptResult: ...
|
||||
|
||||
|
||||
SUMMARY_KEYS = ("short", "detailed", "key_points", "decisions", "action_items", "questions")
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class SummaryResult:
|
||||
short: str = ""
|
||||
detailed: str = ""
|
||||
key_points: tuple[str, ...] = ()
|
||||
decisions: tuple[str, ...] = ()
|
||||
action_items: tuple[str, ...] = ()
|
||||
questions: tuple[str, ...] = ()
|
||||
model: str = ""
|
||||
|
||||
def to_dict(self) -> dict:
|
||||
return {
|
||||
"short": self.short,
|
||||
"detailed": self.detailed,
|
||||
"key_points": list(self.key_points),
|
||||
"decisions": list(self.decisions),
|
||||
"action_items": list(self.action_items),
|
||||
"questions": list(self.questions),
|
||||
}
|
||||
|
||||
@classmethod
|
||||
def from_dict(cls, raw: dict, model: str = "") -> SummaryResult:
|
||||
def text(key: str) -> str:
|
||||
v = raw.get(key)
|
||||
return v if isinstance(v, str) else ""
|
||||
|
||||
def strs(key: str) -> tuple[str, ...]:
|
||||
v = raw.get(key)
|
||||
if not isinstance(v, list):
|
||||
return ()
|
||||
return tuple(s for s in v if isinstance(s, str) and s.strip())
|
||||
|
||||
return cls(
|
||||
short=text("short"),
|
||||
detailed=text("detailed"),
|
||||
key_points=strs("key_points"),
|
||||
decisions=strs("decisions"),
|
||||
action_items=strs("action_items"),
|
||||
questions=strs("questions"),
|
||||
model=model,
|
||||
)
|
||||
|
||||
|
||||
class LlmProvider(Protocol):
|
||||
name: str
|
||||
|
||||
async def summarize(self, transcript: str, *, title: str | None = None) -> SummaryResult: ...
|
||||
|
||||
|
||||
def get_transcription_provider(settings: Settings) -> TranscriptionProvider | None:
|
||||
"""None means "transcription stage skipped", never an error."""
|
||||
kind = settings.transcription_provider.strip().lower()
|
||||
if kind in ("", "none"):
|
||||
return None
|
||||
if kind == "whisper_http":
|
||||
from shonar.services.ai.whisper_http import WhisperHttpProvider
|
||||
|
||||
return WhisperHttpProvider(
|
||||
base_url=settings.transcription_base_url,
|
||||
model=settings.transcription_model,
|
||||
api_key=settings.transcription_api_key,
|
||||
)
|
||||
if kind == "faster_whisper":
|
||||
from shonar.services.ai.faster_whisper import FasterWhisperProvider
|
||||
|
||||
return FasterWhisperProvider(model=settings.transcription_model)
|
||||
raise ProviderConfigError(
|
||||
f"Unknown transcription provider: {settings.transcription_provider!r}"
|
||||
)
|
||||
|
||||
|
||||
def get_llm_provider(settings: Settings) -> LlmProvider | None:
|
||||
"""None means "summary stage skipped", never an error."""
|
||||
kind = settings.llm_provider.strip().lower()
|
||||
if kind in ("", "none"):
|
||||
return None
|
||||
if kind == "openai_compat":
|
||||
from shonar.services.ai.openai_compat import OpenAICompatProvider
|
||||
|
||||
return OpenAICompatProvider(
|
||||
base_url=settings.llm_base_url,
|
||||
model=settings.llm_model,
|
||||
api_key=settings.llm_api_key,
|
||||
)
|
||||
if kind == "ollama":
|
||||
from shonar.services.ai.ollama import OllamaProvider
|
||||
|
||||
return OllamaProvider(
|
||||
base_url=settings.llm_base_url,
|
||||
model=settings.llm_model,
|
||||
)
|
||||
raise ProviderConfigError(f"Unknown LLM provider: {settings.llm_provider!r}")
|
||||
35
backend/shonar/services/ai/_llm.py
Normal file
35
backend/shonar/services/ai/_llm.py
Normal file
|
|
@ -0,0 +1,35 @@
|
|||
"""Shared summary contract: one system prompt, one JSON shape.
|
||||
|
||||
The model replies with JSON only; partial replies are accepted and missing
|
||||
keys default to empty (a terse-but-valid summary beats a failed job).
|
||||
Transcript input is truncated to bound context — a way in, not a report.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from shonar.services.ai import SummaryResult
|
||||
|
||||
SYSTEM_PROMPT = (
|
||||
"You summarize voice recordings for the speaker's own later reference. "
|
||||
"Reply with JSON only, exactly these keys: "
|
||||
'{"short": "1-2 sentences", "detailed": "a faithful paragraph", '
|
||||
'"key_points": [], "decisions": [], "action_items": [], "questions": []}. '
|
||||
"Empty arrays when absent. Never invent names, dates, or commitments "
|
||||
"not stated in the transcript."
|
||||
)
|
||||
|
||||
MAX_TRANSCRIPT_CHARS = 12_000
|
||||
|
||||
|
||||
def build_user_message(transcript: str, title: str | None) -> str:
|
||||
text = transcript[:MAX_TRANSCRIPT_CHARS]
|
||||
if len(transcript) > MAX_TRANSCRIPT_CHARS:
|
||||
text += f"\n\n[truncated from {len(transcript)} chars]"
|
||||
head = f'Title: "{title}"\n\n' if title else ""
|
||||
return head + "Transcript:\n" + text
|
||||
|
||||
|
||||
def parse_summary(data: object, model: str) -> SummaryResult:
|
||||
if not isinstance(data, dict):
|
||||
return SummaryResult(model=model)
|
||||
return SummaryResult.from_dict(data, model=model)
|
||||
108
backend/shonar/services/ai/faster_whisper.py
Normal file
108
backend/shonar/services/ai/faster_whisper.py
Normal file
|
|
@ -0,0 +1,108 @@
|
|||
"""Local transcription via faster-whisper (optional dependency).
|
||||
|
||||
Runs fully on this machine: audio never leaves the server for this stage.
|
||||
The import is lazy so the base install (and every test run) works without
|
||||
the heavyweight dependency.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import tempfile
|
||||
import threading
|
||||
from pathlib import Path
|
||||
|
||||
from shonar.services.ai import ProviderConfigError, Segment, TranscriptResult
|
||||
from shonar.services.ai.model_registry import (
|
||||
download_instructions,
|
||||
is_model_downloaded,
|
||||
validate_model_name,
|
||||
)
|
||||
|
||||
# One loaded model per name, shared across jobs in this worker process.
|
||||
# ctranslate2 inference is thread-safe; a lock serializes first-load only.
|
||||
_MODEL_CACHE: dict[str, object] = {}
|
||||
_MODEL_CACHE_LOCK = threading.Lock()
|
||||
|
||||
|
||||
def _load_model(name: str) -> object:
|
||||
from faster_whisper import WhisperModel
|
||||
|
||||
with _MODEL_CACHE_LOCK:
|
||||
model = _MODEL_CACHE.get(name)
|
||||
if model is None:
|
||||
model = WhisperModel(name, device="auto")
|
||||
_MODEL_CACHE[name] = model
|
||||
return model
|
||||
|
||||
|
||||
def ensure_model_available(name: str) -> None:
|
||||
"""Fail fast with download instructions instead of triggering a surprise
|
||||
multi-GB download inside a transcription job."""
|
||||
if not is_model_downloaded(name):
|
||||
raise ProviderConfigError(download_instructions(name))
|
||||
|
||||
|
||||
class FasterWhisperProvider:
|
||||
name = "faster_whisper"
|
||||
|
||||
def __init__(self, model: str = "base") -> None:
|
||||
try:
|
||||
import faster_whisper # noqa: F401
|
||||
except ImportError as e:
|
||||
raise ProviderConfigError(
|
||||
"faster-whisper is not installed (pip install shonar-backend[faster-whisper])."
|
||||
) from e
|
||||
self.model = validate_model_name(model or "base")
|
||||
|
||||
async def transcribe(
|
||||
self,
|
||||
audio: bytes,
|
||||
mime: str,
|
||||
*,
|
||||
language_hint: str | None = None,
|
||||
on_progress=None, # Callable[[int], None] | None — 0..99 percent
|
||||
) -> TranscriptResult:
|
||||
# faster-whisper is blocking CPU work: keep it off the event loop.
|
||||
return await asyncio.to_thread(self._run, audio, language_hint, on_progress)
|
||||
|
||||
def _run(
|
||||
self,
|
||||
audio: bytes,
|
||||
language_hint: str | None,
|
||||
on_progress=None,
|
||||
) -> TranscriptResult:
|
||||
path: Path | None = None
|
||||
try:
|
||||
with tempfile.NamedTemporaryFile(suffix=".m4a", delete=False) as f:
|
||||
f.write(audio)
|
||||
path = Path(f.name)
|
||||
ensure_model_available(self.model)
|
||||
model = _load_model(self.model)
|
||||
segments_iter, info = model.transcribe( # type: ignore[union-attr]
|
||||
str(path),
|
||||
beam_size=5,
|
||||
language=language_hint,
|
||||
)
|
||||
duration = float(getattr(info, "duration", 0.0) or 0.0)
|
||||
segments = []
|
||||
last_pct = -1
|
||||
for s in segments_iter:
|
||||
segments.append(Segment(start=s.start, end=s.end, text=s.text.strip()))
|
||||
if on_progress is not None and duration > 0:
|
||||
# Throttle: report only on whole-percent gains. Capped
|
||||
# at 99 — the caller commits 100 when the row finishes.
|
||||
pct = min(99, int(s.end / duration * 100))
|
||||
if pct > last_pct:
|
||||
last_pct = pct
|
||||
on_progress(pct)
|
||||
text = " ".join(s.text for s in segments).strip()
|
||||
return TranscriptResult(
|
||||
text=text,
|
||||
language=getattr(info, "language", None),
|
||||
segments=segments,
|
||||
model=self.model,
|
||||
)
|
||||
finally:
|
||||
if path is not None:
|
||||
path.unlink(missing_ok=True)
|
||||
165
backend/shonar/services/ai/model_registry.py
Normal file
165
backend/shonar/services/ai/model_registry.py
Normal file
|
|
@ -0,0 +1,165 @@
|
|||
"""Transcription model registry (Stage 1).
|
||||
|
||||
The supported faster-whisper sizes, their display metadata, validation, and
|
||||
local availability checks. Nothing here downloads anything: faster-whisper
|
||||
fetches from HuggingFace on first use, so "downloaded" is answered by
|
||||
inspecting the HF hub cache, and "available" additionally requires the
|
||||
faster-whisper package itself.
|
||||
|
||||
No silent substitution anywhere: unknown names are rejected, missing
|
||||
downloads fail fast with instructions.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
from dataclasses import dataclass
|
||||
from pathlib import Path
|
||||
|
||||
from sqlalchemy import select
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from shonar.db.models import AppSetting
|
||||
from shonar.services.ai import ProviderConfigError
|
||||
|
||||
DEFAULT_MODEL = "base"
|
||||
DEFAULT_MODEL_KEY = "transcription.default_model"
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class TranscriptionModelInfo:
|
||||
name: str
|
||||
display_name: str
|
||||
description: str
|
||||
params: str
|
||||
approx_memory: str
|
||||
relative_speed: str
|
||||
|
||||
|
||||
SUPPORTED_TRANSCRIPTION_MODELS: dict[str, TranscriptionModelInfo] = {
|
||||
"tiny": TranscriptionModelInfo(
|
||||
name="tiny",
|
||||
display_name="Tiny",
|
||||
description="Fastest and lightest. Good for quick drafts and slow machines.",
|
||||
params="~39M",
|
||||
approx_memory="~1 GB RAM",
|
||||
relative_speed="~10x real-time (CPU)",
|
||||
),
|
||||
"base": TranscriptionModelInfo(
|
||||
name="base",
|
||||
display_name="Base (default)",
|
||||
description="Balanced default. Works reasonably well on ordinary computers.",
|
||||
params="~74M",
|
||||
approx_memory="~1 GB RAM",
|
||||
relative_speed="~7x real-time (CPU)",
|
||||
),
|
||||
"small": TranscriptionModelInfo(
|
||||
name="small",
|
||||
display_name="Small",
|
||||
description="Better accuracy with higher resource usage.",
|
||||
params="~244M",
|
||||
approx_memory="~2 GB RAM",
|
||||
relative_speed="~4x real-time (CPU)",
|
||||
),
|
||||
"medium": TranscriptionModelInfo(
|
||||
name="medium",
|
||||
display_name="Medium",
|
||||
description="Higher accuracy and slower performance.",
|
||||
params="~769M",
|
||||
approx_memory="~5 GB RAM",
|
||||
relative_speed="~2x real-time (CPU)",
|
||||
),
|
||||
"large-v3": TranscriptionModelInfo(
|
||||
name="large-v3",
|
||||
display_name="Large v3",
|
||||
description="Highest accuracy and greatest resource requirements.",
|
||||
params="~1.5B",
|
||||
approx_memory="~10 GB RAM",
|
||||
relative_speed="~1x real-time (CPU)",
|
||||
),
|
||||
}
|
||||
|
||||
|
||||
def normalize_model_name(raw: str | None) -> str:
|
||||
"""Case/whitespace-tolerant normalization. Never maps one model to another."""
|
||||
return (raw or "").strip().lower()
|
||||
|
||||
|
||||
def validate_model_name(raw: str | None) -> str:
|
||||
"""Return the normalized name, or raise with a helpful message."""
|
||||
name = normalize_model_name(raw)
|
||||
if name in SUPPORTED_TRANSCRIPTION_MODELS:
|
||||
return name
|
||||
supported = ", ".join(sorted(SUPPORTED_TRANSCRIPTION_MODELS))
|
||||
raise ProviderConfigError(
|
||||
f"Unsupported transcription model {raw!r}. Supported models: {supported}. "
|
||||
"Check the spelling — a different model is never substituted silently."
|
||||
)
|
||||
|
||||
|
||||
def _hf_hub_cache() -> Path:
|
||||
try:
|
||||
from huggingface_hub.constants import HF_HUB_CACHE
|
||||
|
||||
return Path(HF_HUB_CACHE)
|
||||
except ImportError:
|
||||
return Path(
|
||||
os.environ.get("HF_HUB_CACHE", str(Path.home() / ".cache" / "huggingface" / "hub"))
|
||||
)
|
||||
|
||||
|
||||
def is_model_downloaded(name: str) -> bool:
|
||||
"""True when a non-empty faster-whisper snapshot for `name` sits in the
|
||||
HuggingFace hub cache (repo Systran/faster-whisper-<name>)."""
|
||||
repo_dir = _hf_hub_cache() / f"models--Systran--faster-whisper-{name}"
|
||||
snapshots = repo_dir / "snapshots"
|
||||
if not snapshots.is_dir():
|
||||
return False
|
||||
return any(s.is_dir() and any(s.iterdir()) for s in snapshots.iterdir())
|
||||
|
||||
|
||||
def is_faster_whisper_installed() -> bool:
|
||||
try:
|
||||
import faster_whisper # noqa: F401
|
||||
|
||||
return True
|
||||
except ImportError:
|
||||
return False
|
||||
|
||||
|
||||
def download_instructions(name: str) -> str:
|
||||
return (
|
||||
f'Model "{name}" is not downloaded. Download it with: '
|
||||
f"POST /api/v1/models/{name}/download "
|
||||
"(needs internet once), or run any transcription with that model selected — "
|
||||
"faster-whisper fetches it from HuggingFace automatically."
|
||||
)
|
||||
|
||||
|
||||
async def get_global_default_model(session: AsyncSession) -> str:
|
||||
row = await session.scalar(select(AppSetting).where(AppSetting.key == DEFAULT_MODEL_KEY))
|
||||
if row is None:
|
||||
return DEFAULT_MODEL
|
||||
model = row.value.get("model") if isinstance(row.value, dict) else None
|
||||
return model if model in SUPPORTED_TRANSCRIPTION_MODELS else DEFAULT_MODEL
|
||||
|
||||
|
||||
async def set_global_default_model(session: AsyncSession, raw: str) -> str:
|
||||
"""Validate + persist the global default. Affects future recordings only —
|
||||
existing rows keep their saved model."""
|
||||
name = validate_model_name(raw)
|
||||
row = await session.scalar(select(AppSetting).where(AppSetting.key == DEFAULT_MODEL_KEY))
|
||||
if row is None:
|
||||
row = AppSetting(key=DEFAULT_MODEL_KEY, value={"model": name})
|
||||
session.add(row)
|
||||
else:
|
||||
row.value = {"model": name}
|
||||
await session.flush()
|
||||
return name
|
||||
|
||||
|
||||
def effective_model(recording_model: str | None, global_default: str) -> str:
|
||||
"""Per-recording override wins; otherwise the global default."""
|
||||
if recording_model and recording_model in SUPPORTED_TRANSCRIPTION_MODELS:
|
||||
return recording_model
|
||||
return global_default
|
||||
78
backend/shonar/services/ai/ollama.py
Normal file
78
backend/shonar/services/ai/ollama.py
Normal file
|
|
@ -0,0 +1,78 @@
|
|||
"""Summaries via a local Ollama server (`/api/chat`, JSON mode)."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import AsyncIterator
|
||||
from contextlib import asynccontextmanager
|
||||
|
||||
import httpx
|
||||
|
||||
from shonar.services.ai import ProviderConfigError, ProviderTransientError, SummaryResult
|
||||
from shonar.services.ai._llm import SYSTEM_PROMPT, build_user_message, parse_summary
|
||||
|
||||
|
||||
class OllamaProvider:
|
||||
name = "ollama"
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
base_url: str,
|
||||
model: str = "",
|
||||
timeout_s: float = 900.0,
|
||||
http_client: httpx.AsyncClient | None = None,
|
||||
) -> None:
|
||||
if not base_url.strip():
|
||||
raise ProviderConfigError("ollama needs SHONAR_LLM_BASE_URL.")
|
||||
if not model.strip():
|
||||
raise ProviderConfigError("ollama needs SHONAR_LLM_MODEL.")
|
||||
self.base_url = base_url.rstrip("/")
|
||||
self.model = model
|
||||
self.timeout_s = timeout_s
|
||||
self.http_client = http_client
|
||||
|
||||
@asynccontextmanager
|
||||
async def _client(self) -> AsyncIterator[httpx.AsyncClient]:
|
||||
if self.http_client is not None:
|
||||
yield self.http_client
|
||||
else:
|
||||
async with httpx.AsyncClient(timeout=self.timeout_s) as client:
|
||||
yield client
|
||||
|
||||
async def summarize(self, transcript: str, *, title: str | None = None) -> SummaryResult:
|
||||
payload = {
|
||||
"model": self.model,
|
||||
"stream": False,
|
||||
"format": "json",
|
||||
# qwen3-family models "think" by default: a long reasoning
|
||||
# chain before the JSON answer, brutally slow on CPU and it
|
||||
# does not improve the summary. Ask for the answer directly
|
||||
# (ignored by non-thinking models).
|
||||
"think": False,
|
||||
"options": {"num_ctx": 8192},
|
||||
"messages": [
|
||||
{"role": "system", "content": SYSTEM_PROMPT},
|
||||
{"role": "user", "content": build_user_message(transcript, title)},
|
||||
],
|
||||
}
|
||||
try:
|
||||
async with self._client() as client:
|
||||
resp = await client.post(f"{self.base_url}/api/chat", json=payload)
|
||||
except (httpx.TimeoutException, httpx.TransportError) as e:
|
||||
raise ProviderTransientError(f"Ollama unreachable: {type(e).__name__}") from e
|
||||
if resp.status_code == 404:
|
||||
# Missing model and missing route both 404 here; both are
|
||||
# configuration, not weather.
|
||||
raise ProviderConfigError("Ollama has no such model or route (HTTP 404).")
|
||||
if resp.status_code != 200:
|
||||
raise ProviderTransientError(f"Summarization failed (HTTP {resp.status_code}).")
|
||||
try:
|
||||
content = resp.json()["message"]["content"]
|
||||
except (ValueError, KeyError, TypeError) as e:
|
||||
raise ProviderTransientError("Ollama sent an unreadable reply.") from e
|
||||
import json as _json
|
||||
|
||||
try:
|
||||
data = _json.loads(content)
|
||||
except ValueError as e:
|
||||
raise ProviderTransientError("Ollama reply was not JSON.") from e
|
||||
return parse_summary(data, self.model)
|
||||
82
backend/shonar/services/ai/openai_compat.py
Normal file
82
backend/shonar/services/ai/openai_compat.py
Normal file
|
|
@ -0,0 +1,82 @@
|
|||
"""Summaries via any OpenAI-compatible chat endpoint (self-hosted
|
||||
vLLM/llama.cpp server, commercial API, …) with JSON mode.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import AsyncIterator
|
||||
from contextlib import asynccontextmanager
|
||||
|
||||
import httpx
|
||||
|
||||
from shonar.services.ai import ProviderConfigError, ProviderTransientError, SummaryResult
|
||||
from shonar.services.ai._llm import SYSTEM_PROMPT, build_user_message, parse_summary
|
||||
|
||||
|
||||
class OpenAICompatProvider:
|
||||
name = "openai_compat"
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
base_url: str,
|
||||
model: str = "",
|
||||
api_key: str = "",
|
||||
timeout_s: float = 180.0,
|
||||
http_client: httpx.AsyncClient | None = None,
|
||||
) -> None:
|
||||
if not base_url.strip():
|
||||
raise ProviderConfigError("openai_compat needs SHONAR_LLM_BASE_URL.")
|
||||
if not model.strip():
|
||||
raise ProviderConfigError("openai_compat needs SHONAR_LLM_MODEL.")
|
||||
self.base_url = base_url.rstrip("/")
|
||||
self.model = model
|
||||
self.api_key = api_key
|
||||
self.timeout_s = timeout_s
|
||||
self.http_client = http_client
|
||||
|
||||
@asynccontextmanager
|
||||
async def _client(self) -> AsyncIterator[httpx.AsyncClient]:
|
||||
if self.http_client is not None:
|
||||
yield self.http_client
|
||||
else:
|
||||
async with httpx.AsyncClient(timeout=self.timeout_s) as client:
|
||||
yield client
|
||||
|
||||
async def summarize(self, transcript: str, *, title: str | None = None) -> SummaryResult:
|
||||
headers = (
|
||||
{"Authorization": f"Bearer {self.api_key}"} if self.api_key else {}
|
||||
)
|
||||
payload = {
|
||||
"model": self.model,
|
||||
"temperature": 0.2,
|
||||
"response_format": {"type": "json_object"},
|
||||
"messages": [
|
||||
{"role": "system", "content": SYSTEM_PROMPT},
|
||||
{"role": "user", "content": build_user_message(transcript, title)},
|
||||
],
|
||||
}
|
||||
try:
|
||||
async with self._client() as client:
|
||||
resp = await client.post(
|
||||
f"{self.base_url}/v1/chat/completions", headers=headers, json=payload
|
||||
)
|
||||
except (httpx.TimeoutException, httpx.TransportError) as e:
|
||||
raise ProviderTransientError(f"LLM unreachable: {type(e).__name__}") from e
|
||||
if resp.status_code in (401, 403, 404):
|
||||
raise ProviderConfigError(f"LLM refused the request (HTTP {resp.status_code}).")
|
||||
if resp.status_code == 429 or resp.status_code >= 500:
|
||||
raise ProviderTransientError(f"LLM busy (HTTP {resp.status_code}).")
|
||||
if resp.status_code != 200:
|
||||
raise ProviderTransientError(f"Summarization failed (HTTP {resp.status_code}).")
|
||||
try:
|
||||
body = resp.json()
|
||||
content = body["choices"][0]["message"]["content"]
|
||||
except (ValueError, KeyError, IndexError, TypeError) as e:
|
||||
raise ProviderTransientError("LLM sent an unreadable reply.") from e
|
||||
import json as _json
|
||||
|
||||
try:
|
||||
data = _json.loads(content)
|
||||
except ValueError as e:
|
||||
raise ProviderTransientError("LLM reply was not JSON.") from e
|
||||
return parse_summary(data, self.model)
|
||||
128
backend/shonar/services/ai/whisper_http.py
Normal file
128
backend/shonar/services/ai/whisper_http.py
Normal file
|
|
@ -0,0 +1,128 @@
|
|||
"""Transcription via any OpenAI-compatible `/v1/audio/transcriptions`
|
||||
endpoint (self-hosted whisper.cpp server, commercial Whisper API, …).
|
||||
Sends `verbose_json` so segment timings come back with the text.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import AsyncIterator
|
||||
from contextlib import asynccontextmanager
|
||||
|
||||
import httpx
|
||||
|
||||
from shonar.services.ai import (
|
||||
ProviderConfigError,
|
||||
ProviderTransientError,
|
||||
Segment,
|
||||
TranscriptResult,
|
||||
)
|
||||
|
||||
|
||||
class WhisperHttpProvider:
|
||||
name = "whisper_http"
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
base_url: str,
|
||||
model: str = "base",
|
||||
api_key: str = "",
|
||||
timeout_s: float = 300.0,
|
||||
http_client: httpx.AsyncClient | None = None,
|
||||
) -> None:
|
||||
if not base_url.strip():
|
||||
raise ProviderConfigError(
|
||||
"whisper_http needs SHONAR_TRANSCRIPTION_BASE_URL."
|
||||
)
|
||||
self.base_url = base_url.rstrip("/")
|
||||
self.model = model
|
||||
self.api_key = api_key
|
||||
self.timeout_s = timeout_s
|
||||
self.http_client = http_client
|
||||
|
||||
@asynccontextmanager
|
||||
async def _client(self) -> AsyncIterator[httpx.AsyncClient]:
|
||||
if self.http_client is not None:
|
||||
yield self.http_client
|
||||
else:
|
||||
async with httpx.AsyncClient(timeout=self.timeout_s) as client:
|
||||
yield client
|
||||
|
||||
async def transcribe(
|
||||
self,
|
||||
audio: bytes,
|
||||
mime: str,
|
||||
*,
|
||||
language_hint: str | None = None,
|
||||
on_progress=None, # accepted for protocol parity; not reported
|
||||
) -> TranscriptResult:
|
||||
headers = (
|
||||
{"Authorization": f"Bearer {self.api_key}"} if self.api_key else {}
|
||||
)
|
||||
data: dict[str, str] = {"model": self.model, "response_format": "verbose_json"}
|
||||
if language_hint:
|
||||
data["language"] = language_hint
|
||||
files = {"file": (f"audio.{_ext(mime)}", audio, mime or "application/octet-stream")}
|
||||
try:
|
||||
async with self._client() as client:
|
||||
resp = await client.post(
|
||||
f"{self.base_url}/v1/audio/transcriptions",
|
||||
headers=headers,
|
||||
data=data,
|
||||
files=files,
|
||||
)
|
||||
except (httpx.TimeoutException, httpx.TransportError) as e:
|
||||
raise ProviderTransientError(
|
||||
f"Transcription service unreachable: {type(e).__name__}"
|
||||
) from e
|
||||
if resp.status_code in (401, 403, 404):
|
||||
raise ProviderConfigError(
|
||||
f"Transcription service refused the request (HTTP {resp.status_code})."
|
||||
)
|
||||
if resp.status_code == 429 or resp.status_code >= 500:
|
||||
raise ProviderTransientError(
|
||||
f"Transcription service busy (HTTP {resp.status_code})."
|
||||
)
|
||||
if resp.status_code != 200:
|
||||
raise ProviderTransientError(
|
||||
f"Transcription failed (HTTP {resp.status_code})."
|
||||
)
|
||||
try:
|
||||
body = resp.json()
|
||||
except ValueError as e:
|
||||
raise ProviderTransientError("Transcription service sent no JSON.") from e
|
||||
segments = []
|
||||
raw_segs = body.get("segments")
|
||||
if isinstance(raw_segs, list):
|
||||
for s in raw_segs:
|
||||
if not isinstance(s, dict):
|
||||
continue
|
||||
try:
|
||||
segments.append(
|
||||
Segment(
|
||||
start=float(s.get("start", 0.0)),
|
||||
end=float(s.get("end", 0.0)),
|
||||
text=str(s.get("text", "")),
|
||||
)
|
||||
)
|
||||
except (TypeError, ValueError):
|
||||
continue
|
||||
text = body.get("text")
|
||||
return TranscriptResult(
|
||||
text=text if isinstance(text, str) else "",
|
||||
language=body.get("language") if isinstance(body.get("language"), str) else None,
|
||||
segments=segments,
|
||||
model=self.model,
|
||||
)
|
||||
|
||||
|
||||
def _ext(mime: str) -> str:
|
||||
return {
|
||||
"audio/mp4": "m4a",
|
||||
"audio/m4a": "m4a",
|
||||
"audio/wav": "wav",
|
||||
"audio/x-wav": "wav",
|
||||
"audio/ogg": "ogg",
|
||||
"audio/opus": "ogg",
|
||||
"audio/webm": "webm",
|
||||
"audio/mpeg": "mp3",
|
||||
}.get(mime.lower().split(";")[0].strip(), "bin")
|
||||
167
backend/shonar/services/auth.py
Normal file
167
backend/shonar/services/auth.py
Normal file
|
|
@ -0,0 +1,167 @@
|
|||
"""Authentication service: registration, login, rotating refresh tokens with
|
||||
reuse detection, logout, account deletion."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import uuid
|
||||
from datetime import UTC, datetime
|
||||
|
||||
from sqlalchemy import select, update
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from shonar.core.security import (
|
||||
generate_refresh_token,
|
||||
hash_password,
|
||||
hash_refresh_token,
|
||||
refresh_token_ttl,
|
||||
verify_password,
|
||||
)
|
||||
from shonar.db.models import Device, RefreshToken, User, utcnow
|
||||
|
||||
|
||||
class AuthError(Exception):
|
||||
"""Safe, client-displayable auth failure (never leaks which half failed
|
||||
beyond what the flow requires)."""
|
||||
|
||||
def __init__(self, message: str, status_code: int = 401):
|
||||
super().__init__(message)
|
||||
self.message = message
|
||||
self.status_code = status_code
|
||||
|
||||
|
||||
async def register_user(
|
||||
session: AsyncSession, email: str, password: str, display_name: str | None
|
||||
) -> User:
|
||||
email = email.strip().lower()
|
||||
existing = await session.scalar(select(User).where(User.email == email))
|
||||
if existing is not None:
|
||||
# Use a generic message; do not reveal whether the account exists in
|
||||
# flows where that matters. For self-hosted registration the UX cost
|
||||
# of "email already registered" is acceptable and helpful.
|
||||
raise AuthError("An account with this email already exists.", 409)
|
||||
user = User(
|
||||
email=email,
|
||||
password_hash=hash_password(password),
|
||||
display_name=display_name,
|
||||
)
|
||||
session.add(user)
|
||||
await session.flush()
|
||||
return user
|
||||
|
||||
|
||||
async def issue_refresh_token(
|
||||
session: AsyncSession,
|
||||
user_id: uuid.UUID,
|
||||
family: uuid.UUID | None,
|
||||
device_id: uuid.UUID | None,
|
||||
) -> tuple[str, RefreshToken]:
|
||||
token = generate_refresh_token()
|
||||
rt = RefreshToken(
|
||||
user_id=user_id,
|
||||
token_hash=hash_refresh_token(token),
|
||||
family=family or uuid.uuid4(),
|
||||
device_id=device_id,
|
||||
expires_at=datetime.now(UTC) + refresh_token_ttl(),
|
||||
)
|
||||
session.add(rt)
|
||||
await session.flush()
|
||||
return token, rt
|
||||
|
||||
|
||||
async def login(
|
||||
session: AsyncSession,
|
||||
email: str,
|
||||
password: str,
|
||||
device_name: str | None,
|
||||
platform: str,
|
||||
) -> tuple[User, str, Device]:
|
||||
"""Returns (user, refresh_token, device). Raises AuthError safely."""
|
||||
email = email.strip().lower()
|
||||
user = await session.scalar(select(User).where(User.email == email))
|
||||
if user is None or user.deleted_at is not None or not user.is_active:
|
||||
raise AuthError("Invalid email or password.")
|
||||
if not verify_password(user.password_hash, password):
|
||||
raise AuthError("Invalid email or password.")
|
||||
|
||||
device = Device(user_id=user.id, name=device_name or "Android device", platform=platform)
|
||||
session.add(device)
|
||||
await session.flush()
|
||||
|
||||
refresh_token, _ = await issue_refresh_token(session, user.id, None, device.id)
|
||||
return user, refresh_token, device
|
||||
|
||||
|
||||
async def rotate_refresh_token(
|
||||
session: AsyncSession, presented_token: str
|
||||
) -> tuple[User, str, uuid.UUID | None]:
|
||||
"""Consume a refresh token and issue a replacement in the same family.
|
||||
|
||||
Reuse detection: presenting an already-consumed/revoked token revokes the
|
||||
entire family (an attacker's stolen token dies along with the real one).
|
||||
"""
|
||||
token_hash = hash_refresh_token(presented_token)
|
||||
rt = await session.scalar(select(RefreshToken).where(RefreshToken.token_hash == token_hash))
|
||||
now = utcnow()
|
||||
|
||||
if rt is None:
|
||||
raise AuthError("Invalid refresh token.")
|
||||
|
||||
if rt.revoked_at is not None or rt.replaced_by is not None:
|
||||
# REUSE DETECTED — revoke the whole family. Commit BEFORE raising:
|
||||
# the request's transaction would otherwise roll back on the 401 and
|
||||
# silently undo the security-revocation.
|
||||
await session.execute(
|
||||
update(RefreshToken)
|
||||
.where(RefreshToken.family == rt.family, RefreshToken.revoked_at.is_(None))
|
||||
.values(revoked_at=now)
|
||||
)
|
||||
await session.commit()
|
||||
raise AuthError("Refresh token reuse detected. Please log in again.", 401)
|
||||
|
||||
if rt.expires_at < now:
|
||||
raise AuthError("Refresh token expired.", 401)
|
||||
|
||||
user = await session.get(User, rt.user_id)
|
||||
if user is None or user.deleted_at is not None or not user.is_active:
|
||||
raise AuthError("Account unavailable.", 401)
|
||||
|
||||
new_token, new_rt = await issue_refresh_token(session, user.id, rt.family, rt.device_id)
|
||||
rt.revoked_at = now
|
||||
rt.replaced_by = new_rt.id
|
||||
|
||||
if rt.device_id is not None:
|
||||
device = await session.get(Device, rt.device_id)
|
||||
if device is not None:
|
||||
device.last_seen_at = now
|
||||
await session.flush()
|
||||
return user, new_token, rt.device_id
|
||||
|
||||
|
||||
async def logout(session: AsyncSession, presented_token: str) -> None:
|
||||
"""Revoke the presented token's whole family (logs the device out)."""
|
||||
token_hash = hash_refresh_token(presented_token)
|
||||
rt = await session.scalar(select(RefreshToken).where(RefreshToken.token_hash == token_hash))
|
||||
if rt is None:
|
||||
return # idempotent
|
||||
await session.execute(
|
||||
update(RefreshToken)
|
||||
.where(RefreshToken.family == rt.family, RefreshToken.revoked_at.is_(None))
|
||||
.values(revoked_at=utcnow())
|
||||
)
|
||||
|
||||
|
||||
async def delete_account(session: AsyncSession, user: User, password: str) -> None:
|
||||
if not verify_password(user.password_hash, password):
|
||||
raise AuthError("Invalid password.", 403)
|
||||
now = utcnow()
|
||||
user.deleted_at = now
|
||||
user.is_active = False
|
||||
# Revoke every refresh token for the user.
|
||||
await session.execute(
|
||||
update(RefreshToken)
|
||||
.where(RefreshToken.user_id == user.id, RefreshToken.revoked_at.is_(None))
|
||||
.values(revoked_at=now)
|
||||
)
|
||||
# NOTE: hard deletion of rows/files is performed by the retention sweep
|
||||
# (services/retention.py) so an accidental deletion can be cancelled
|
||||
# within the grace window (see docs/security.md).
|
||||
251
backend/shonar/services/exports.py
Normal file
251
backend/shonar/services/exports.py
Normal file
|
|
@ -0,0 +1,251 @@
|
|||
"""Exports (M9): audio, transcript txt, notes markdown, bundle zip.
|
||||
|
||||
Synchronous generation — every artifact is small enough (text, or one
|
||||
audio file) that a background job adds failure modes, not speed. Each
|
||||
successful export records an ExportJob row and stores the produced bytes
|
||||
as an Asset(kind=export) so the audit trail exists; the response is the
|
||||
file itself (no separate download-asset round trip).
|
||||
|
||||
Formats:
|
||||
audio — the original upload, byte-identical, original mime/extension
|
||||
txt — current transcript text
|
||||
md — notes.md: title, metadata, notes, summary sections, transcript
|
||||
zip — bundle: original audio + transcript.txt + notes.md
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import io
|
||||
import uuid
|
||||
import zipfile
|
||||
from dataclasses import dataclass
|
||||
|
||||
from sqlalchemy import select
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from shonar.db.models import (
|
||||
Asset,
|
||||
AssetKind,
|
||||
ExportJob,
|
||||
JobStatus,
|
||||
Recording,
|
||||
Summary,
|
||||
Transcript,
|
||||
)
|
||||
|
||||
|
||||
class ExportError(Exception):
|
||||
def __init__(self, status_code: int, message: str):
|
||||
super().__init__(message)
|
||||
self.status_code = status_code
|
||||
self.message = message
|
||||
|
||||
|
||||
EXPORT_FORMATS = ("audio", "txt", "md", "zip")
|
||||
|
||||
|
||||
@dataclass
|
||||
class ExportResult:
|
||||
filename: str
|
||||
mime_type: str
|
||||
data: bytes
|
||||
|
||||
|
||||
def _safe_filename(rec: Recording) -> str:
|
||||
"""Slug the title; fall back to the recorded_at stamp. Never leaks ids."""
|
||||
base = "".join(
|
||||
c if (c.isalnum() or c in "-_ ") else " " for c in (rec.title or "")
|
||||
).strip()
|
||||
if not base:
|
||||
base = f"recording-{rec.recorded_at:%Y%m%d-%H%M%S}"
|
||||
return base[:120]
|
||||
|
||||
|
||||
async def _current_transcript(session: AsyncSession, rec_id: uuid.UUID) -> Transcript | None:
|
||||
return await session.scalar(
|
||||
select(Transcript)
|
||||
.where(Transcript.recording_id == rec_id, Transcript.superseded_at.is_(None))
|
||||
.order_by(Transcript.version.desc())
|
||||
)
|
||||
|
||||
|
||||
async def _current_summary(session: AsyncSession, rec_id: uuid.UUID) -> Summary | None:
|
||||
return await session.scalar(
|
||||
select(Summary)
|
||||
.where(Summary.recording_id == rec_id, Summary.superseded_at.is_(None))
|
||||
.order_by(Summary.version.desc())
|
||||
)
|
||||
|
||||
|
||||
async def _original_asset(session: AsyncSession, rec_id: uuid.UUID) -> Asset | None:
|
||||
return await session.scalar(
|
||||
select(Asset).where(Asset.recording_id == rec_id, Asset.kind == AssetKind.original)
|
||||
)
|
||||
|
||||
|
||||
def _summary_md(summary: Summary | None) -> str:
|
||||
if summary is None:
|
||||
return ""
|
||||
c = summary.content or {}
|
||||
out = ["## Summary\n"]
|
||||
if c.get("short"):
|
||||
out.append(f"{c['short']}\n")
|
||||
if c.get("detailed"):
|
||||
out.append(f"### Detailed\n\n{c['detailed']}\n")
|
||||
for key, header in (
|
||||
("key_points", "Key Points"),
|
||||
("decisions", "Decisions"),
|
||||
("action_items", "Action Items"),
|
||||
("questions", "Questions"),
|
||||
):
|
||||
items = c.get(key)
|
||||
if isinstance(items, list) and items:
|
||||
out.append(f"### {header}\n")
|
||||
out.extend(f"- {x}" for x in items)
|
||||
out.append("")
|
||||
return "\n".join(out)
|
||||
|
||||
|
||||
def _transcript_md(t: Transcript | None) -> str:
|
||||
if t is None:
|
||||
return ""
|
||||
lines = ["## Transcript\n"]
|
||||
segs = [s for s in (t.segments or []) if isinstance(s, dict)]
|
||||
if segs:
|
||||
for s in segs:
|
||||
start = float(s.get("start", 0.0))
|
||||
stamp = f"{int(start // 60):02d}:{start % 60:04.1f}"
|
||||
speaker = f"**{s['speaker']}**: " if s.get("speaker") else ""
|
||||
lines.append(f"- `[{stamp}]` {speaker}{s.get('text', '').strip()}")
|
||||
else:
|
||||
lines.append(t.text or "")
|
||||
return "\n".join(lines) + "\n"
|
||||
|
||||
|
||||
def _notes_md(
|
||||
rec: Recording, t: Transcript | None, summary: Summary | None,
|
||||
tag_names: list[str] | None = None,
|
||||
) -> str:
|
||||
parts = [
|
||||
f"# {rec.title}\n",
|
||||
f"- Recorded: {rec.recorded_at:%Y-%m-%d %H:%M} UTC",
|
||||
f"- Duration: {rec.duration_seconds:.1f}s",
|
||||
]
|
||||
if tag_names:
|
||||
parts.append("- Tags: " + ", ".join(f"`{x}`" for x in tag_names))
|
||||
parts.append("")
|
||||
if rec.notes:
|
||||
parts.append(f"## Notes\n\n{rec.notes}\n")
|
||||
s = _summary_md(summary)
|
||||
if s:
|
||||
parts.append(s)
|
||||
tr = _transcript_md(t)
|
||||
if tr:
|
||||
parts.append(tr)
|
||||
return "\n".join(parts)
|
||||
|
||||
|
||||
def _zip(files: list[tuple[str, bytes]]) -> bytes:
|
||||
buf = io.BytesIO()
|
||||
with zipfile.ZipFile(buf, "w", zipfile.ZIP_DEFLATED) as z:
|
||||
for name, data in files:
|
||||
z.writestr(name, data)
|
||||
return buf.getvalue()
|
||||
|
||||
|
||||
async def build_export(
|
||||
session: AsyncSession, rec: Recording, export_format: str
|
||||
) -> ExportResult:
|
||||
"""Build one export artifact for an owned, non-deleted recording."""
|
||||
from shonar.storage import get_storage
|
||||
|
||||
if export_format not in EXPORT_FORMATS:
|
||||
raise ExportError(422, f"Unknown export format. Use one of: {', '.join(EXPORT_FORMATS)}")
|
||||
|
||||
original = await _original_asset(session, rec.id)
|
||||
t = await _current_transcript(session, rec.id)
|
||||
summary = await _current_summary(session, rec.id)
|
||||
from shonar.services.search import tag_names
|
||||
|
||||
tags = await tag_names(session, rec.id)
|
||||
base = _safe_filename(rec)
|
||||
|
||||
if export_format == "audio":
|
||||
if original is None:
|
||||
raise ExportError(404, "No audio stored for this recording")
|
||||
data = await get_storage().get(original.storage_key)
|
||||
ext = original.storage_key[original.storage_key.rfind(".") :]
|
||||
return ExportResult(filename=f"{base}{ext}", mime_type=original.mime_type, data=data)
|
||||
|
||||
if export_format == "txt":
|
||||
if t is None:
|
||||
raise ExportError(404, "No transcript yet — transcribe first")
|
||||
return ExportResult(
|
||||
filename=f"{base}.txt", mime_type="text/plain; charset=utf-8",
|
||||
data=(t.text or "").encode("utf-8"),
|
||||
)
|
||||
|
||||
if export_format == "md":
|
||||
if t is None and summary is None and not rec.notes:
|
||||
raise ExportError(
|
||||
404, "Nothing to export — this recording has no notes, transcript, or summary"
|
||||
)
|
||||
return ExportResult(
|
||||
filename=f"{base}.md", mime_type="text/markdown; charset=utf-8",
|
||||
data=_notes_md(rec, t, summary, tags).encode("utf-8"),
|
||||
)
|
||||
|
||||
# zip bundle: whatever exists, always at least the audio when present.
|
||||
if original is None and t is None and summary is None and not rec.notes:
|
||||
raise ExportError(404, "Nothing to export for this recording")
|
||||
files: list[tuple[str, bytes]] = []
|
||||
if original is not None:
|
||||
audio = await get_storage().get(original.storage_key)
|
||||
ext = original.storage_key[original.storage_key.rfind(".") :]
|
||||
files.append((f"{base}{ext}", audio))
|
||||
if t is not None:
|
||||
files.append(("transcript.txt", (t.text or "").encode("utf-8")))
|
||||
files.append(("notes.md", _notes_md(rec, t, summary, tags).encode("utf-8")))
|
||||
return ExportResult(
|
||||
filename=f"{base}.zip", mime_type="application/zip", data=_zip(files)
|
||||
)
|
||||
|
||||
|
||||
async def record_export(
|
||||
session: AsyncSession, user_id: uuid.UUID, rec: Recording,
|
||||
export_format: str, result: ExportResult,
|
||||
) -> None:
|
||||
"""Persist the audit trail: ExportJob(succeeded) + Asset(kind=export).
|
||||
|
||||
Best-effort storage of the artifact bytes; a storage failure never
|
||||
fails the download the user already received.
|
||||
"""
|
||||
from shonar.storage import get_storage
|
||||
|
||||
job = ExportJob(
|
||||
user_id=user_id,
|
||||
recording_id=rec.id,
|
||||
export_type=export_format,
|
||||
status=JobStatus.succeeded,
|
||||
)
|
||||
session.add(job)
|
||||
try:
|
||||
key = f"exports/{user_id}/{rec.id}/{export_format}-{uuid.uuid4().hex}"
|
||||
await get_storage().put(key, result.data)
|
||||
import hashlib
|
||||
|
||||
asset = Asset(
|
||||
recording_id=rec.id,
|
||||
user_id=user_id,
|
||||
kind=AssetKind.export,
|
||||
storage_key=key,
|
||||
mime_type=result.mime_type,
|
||||
size_bytes=len(result.data),
|
||||
checksum_sha256=hashlib.sha256(result.data).hexdigest(),
|
||||
)
|
||||
session.add(asset)
|
||||
await session.flush()
|
||||
job.asset_id = asset.id
|
||||
except Exception: # noqa: BLE001 — audit copy is best-effort
|
||||
pass
|
||||
await session.flush()
|
||||
141
backend/shonar/services/inline_queue.py
Normal file
141
backend/shonar/services/inline_queue.py
Normal file
|
|
@ -0,0 +1,141 @@
|
|||
"""In-process job runner (desktop bundled-lite engine).
|
||||
|
||||
``queue_backend=inline`` replaces the arq/Redis transport with a single
|
||||
asyncio consumer inside the uvicorn process: DB rows stay the source of
|
||||
truth (ProcessingJob), this just runs the work. One job at a time —
|
||||
local faster-whisper is multi-GB per pass, mirroring the worker's
|
||||
``max_jobs=1`` rule.
|
||||
|
||||
Transient failures retry in-process with backoff up to MAX_TRIES (the
|
||||
same budget arq gives via ``max_tries``); the final failure is recorded
|
||||
by the task body itself (see processing._fail).
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import contextlib
|
||||
import logging
|
||||
|
||||
from shonar.db.models import JobType
|
||||
from shonar.services import processing
|
||||
|
||||
logger = logging.getLogger("shonar.inline_queue")
|
||||
|
||||
RETRY_DELAY_SECONDS = 5.0
|
||||
# Desktop engine: hard-delete expired soft-deletes once the app has been up
|
||||
# for a day, then daily. Delayed first run keeps startup snappy.
|
||||
RETENTION_INTERVAL_SECONDS = 24 * 3600.0
|
||||
|
||||
_queue: asyncio.Queue[tuple[str, str, int]] | None = None
|
||||
_consumer: asyncio.Task | None = None
|
||||
_retention: asyncio.Task | None = None
|
||||
|
||||
|
||||
def _get_queue() -> asyncio.Queue[tuple[str, str, int]]:
|
||||
global _queue
|
||||
if _queue is None:
|
||||
_queue = asyncio.Queue()
|
||||
return _queue
|
||||
|
||||
|
||||
async def start() -> None:
|
||||
"""Start the consumer and re-run anything the DB says is pending."""
|
||||
global _consumer, _retention
|
||||
_get_queue()
|
||||
# Inline engine = single process: any row still marked `running` at
|
||||
# startup is a corpse from the previous process (the worker died with
|
||||
# it). sweep_stale's 2h live-worker grace — correct for multi-worker
|
||||
# arq deployments — would starve these jobs, so reclaim them first.
|
||||
from sqlalchemy import update
|
||||
|
||||
from shonar.db.models import JobStatus, ProcessingJob
|
||||
from shonar.db.session import session_factory
|
||||
|
||||
async with session_factory()() as s:
|
||||
await s.execute(
|
||||
update(ProcessingJob)
|
||||
.where(ProcessingJob.status == JobStatus.running)
|
||||
.values(status=JobStatus.queued, started_at=None)
|
||||
)
|
||||
await s.commit()
|
||||
if _consumer is None or _consumer.done():
|
||||
_consumer = asyncio.create_task(_consume(), name="shonar-inline-queue")
|
||||
if _retention is None or _retention.done():
|
||||
_retention = asyncio.create_task(_retention_loop(), name="shonar-retention")
|
||||
# Crash recovery: queued rows (and orphaned running rows requeued by
|
||||
# the sweep) go back on the in-process queue.
|
||||
count = await processing.sweep_stale()
|
||||
if count:
|
||||
logger.info("inline queue startup sweep requeued %d jobs", count)
|
||||
|
||||
|
||||
async def _retention_loop() -> None:
|
||||
"""Daily hard-delete sweep for the desktop engine (no arq cron here).
|
||||
|
||||
First pass a few minutes after start (the app may only run for hours
|
||||
at a time, so a full-day initial sleep could starve the sweep), then
|
||||
once a day while running.
|
||||
"""
|
||||
from shonar.services import retention
|
||||
|
||||
await asyncio.sleep(120.0)
|
||||
while True:
|
||||
try:
|
||||
purged = await retention.sweep_deleted()
|
||||
if purged["recordings"] or purged["users"]:
|
||||
logger.info("inline retention sweep: %s", purged)
|
||||
except asyncio.CancelledError:
|
||||
raise
|
||||
except Exception: # noqa: BLE001 — the loop must survive any failure
|
||||
logger.exception("inline retention sweep failed")
|
||||
await asyncio.sleep(RETENTION_INTERVAL_SECONDS)
|
||||
|
||||
|
||||
async def stop() -> None:
|
||||
global _consumer, _retention
|
||||
if _consumer is not None:
|
||||
_consumer.cancel()
|
||||
with contextlib.suppress(BaseException): # noqa: BLE001 — shutdown is best-effort
|
||||
await _consumer
|
||||
_consumer = None
|
||||
if _retention is not None:
|
||||
_retention.cancel()
|
||||
with contextlib.suppress(BaseException): # noqa: BLE001
|
||||
await _retention
|
||||
_retention = None
|
||||
|
||||
|
||||
async def enqueue(job_type: JobType, recording_id: str) -> None:
|
||||
await _get_queue().put((job_type.value, str(recording_id), 1))
|
||||
|
||||
|
||||
async def _consume() -> None:
|
||||
q = _get_queue()
|
||||
while True:
|
||||
job_value, recording_id, attempt = await q.get()
|
||||
ctx = {"job_try": attempt}
|
||||
try:
|
||||
if job_value == JobType.transcribe.value:
|
||||
await processing.run_transcribe(ctx, recording_id)
|
||||
else:
|
||||
await processing.run_summarize(ctx, recording_id)
|
||||
except processing.ProviderTransientError as e:
|
||||
if attempt < processing.MAX_TRIES:
|
||||
logger.warning(
|
||||
"inline job %s:%s transient failure (%s); retry %d/%d",
|
||||
job_value, recording_id, e, attempt + 1, processing.MAX_TRIES,
|
||||
)
|
||||
await asyncio.sleep(RETRY_DELAY_SECONDS)
|
||||
await q.put((job_value, recording_id, attempt + 1))
|
||||
else:
|
||||
logger.error(
|
||||
"inline job %s:%s failed after %d tries: %s",
|
||||
job_value, recording_id, attempt, e,
|
||||
)
|
||||
except asyncio.CancelledError:
|
||||
raise
|
||||
except Exception: # noqa: BLE001 — consumer must survive any task crash
|
||||
logger.exception("inline job %s:%s crashed", job_value, recording_id)
|
||||
finally:
|
||||
q.task_done()
|
||||
74
backend/shonar/services/media.py
Normal file
74
backend/shonar/services/media.py
Normal file
|
|
@ -0,0 +1,74 @@
|
|||
"""Audio format validation: declared MIME type vs magic bytes.
|
||||
|
||||
Families decouple declared MIME aliases (audio/mp4 vs audio/m4a) from what
|
||||
the bytes actually are. The original file is stored exactly as uploaded —
|
||||
validation never rewrites it.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
FAMILY_BY_MIME = {
|
||||
"audio/mp4": "mp4",
|
||||
"audio/m4a": "mp4",
|
||||
"audio/aac": "aac",
|
||||
"audio/wav": "wav",
|
||||
"audio/x-wav": "wav",
|
||||
"audio/ogg": "ogg",
|
||||
"audio/opus": "ogg",
|
||||
"audio/webm": "webm",
|
||||
"audio/mpeg": "mpeg",
|
||||
}
|
||||
|
||||
EXT_BY_FAMILY = {
|
||||
"mp4": ".m4a",
|
||||
"aac": ".aac",
|
||||
"wav": ".wav",
|
||||
"ogg": ".ogg",
|
||||
"webm": ".webm",
|
||||
"mpeg": ".mp3",
|
||||
}
|
||||
|
||||
|
||||
def sniff_audio_family(data: bytes) -> str | None:
|
||||
"""Return the audio family from magic bytes, or None if unrecognized."""
|
||||
if len(data) >= 12:
|
||||
if data[4:8] == b"ftyp":
|
||||
return "mp4"
|
||||
if data[:4] == b"RIFF" and data[8:12] == b"WAVE":
|
||||
return "wav"
|
||||
if data[:4] == b"OggS":
|
||||
return "ogg"
|
||||
if data[:4] == b"\x1a\x45\xdf\xa3":
|
||||
return "webm"
|
||||
if data[:3] == b"ID3":
|
||||
return "mpeg"
|
||||
if len(data) >= 2 and data[0] == 0xFF and (data[1] & 0xF6) == 0xF0:
|
||||
# ADTS frame sync: AAC (also accepted as mpeg-family audio)
|
||||
return "aac"
|
||||
return None
|
||||
|
||||
|
||||
def declared_family(mime_type: str) -> str | None:
|
||||
return FAMILY_BY_MIME.get(mime_type.lower().split(";")[0].strip())
|
||||
|
||||
|
||||
def is_compatible(mime_type: str, data: bytes) -> bool:
|
||||
"""True when the declared MIME matches the sniffed magic bytes.
|
||||
|
||||
mpeg and aac are treated as one family: Android records AAC in ADTS or
|
||||
in MP4 containers and MIME reporting around these is inconsistent.
|
||||
"""
|
||||
declared = declared_family(mime_type)
|
||||
sniffed = sniff_audio_family(data)
|
||||
if declared is None or sniffed is None:
|
||||
return False
|
||||
if {declared, sniffed} == {"mpeg", "aac"}:
|
||||
return True
|
||||
return declared == sniffed
|
||||
|
||||
|
||||
def extension_for(mime_type: str, data: bytes) -> str:
|
||||
sniffed = sniff_audio_family(data)
|
||||
if sniffed is not None:
|
||||
return EXT_BY_FAMILY[sniffed]
|
||||
return EXT_BY_FAMILY.get(declared_family(mime_type) or "", ".bin")
|
||||
591
backend/shonar/services/processing.py
Normal file
591
backend/shonar/services/processing.py
Normal file
|
|
@ -0,0 +1,591 @@
|
|||
"""AI processing pipeline (M7): uploaded -> transcribed -> summarized.
|
||||
|
||||
State lives in the database (ProcessingJob rows); arq/redis is transport
|
||||
only. That ordering is deliberate: if redis is down, uploads still succeed
|
||||
(the jobs sit queued) and the worker sweep picks them up. The only race is
|
||||
a task running before the API transaction commits (recording invisible) —
|
||||
missing rows are transient failures, so arq retry absorbs it.
|
||||
|
||||
Entry points:
|
||||
- ``enqueue_for_recording``: called from upload finalize AND re-callable
|
||||
later (manual transcripts in M8 re-enter here). Idempotent.
|
||||
- ``run_transcribe`` / ``run_summarize``: arq task bodies. ``ctx`` is an
|
||||
arq context in production and a plain dict in tests; only
|
||||
``ctx.get("job_try", 1)`` is read.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import logging
|
||||
import uuid
|
||||
from datetime import timedelta
|
||||
|
||||
from sqlalchemy import func, select
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from shonar.core.config import get_settings
|
||||
from shonar.db.models import (
|
||||
Asset,
|
||||
AssetKind,
|
||||
JobStatus,
|
||||
JobType,
|
||||
ProcessingJob,
|
||||
ProcessingStatus,
|
||||
Recording,
|
||||
Summary,
|
||||
Transcript,
|
||||
utcnow,
|
||||
)
|
||||
from shonar.services.ai import (
|
||||
AIError,
|
||||
ProviderTransientError,
|
||||
get_llm_provider,
|
||||
get_transcription_provider,
|
||||
)
|
||||
from shonar.storage import get_storage
|
||||
|
||||
logger = logging.getLogger("shonar.processing")
|
||||
|
||||
|
||||
def _accepts_on_progress(provider) -> bool:
|
||||
"""True when the provider's transcribe() accepts on_progress=."""
|
||||
import inspect
|
||||
|
||||
try:
|
||||
return "on_progress" in inspect.signature(provider.transcribe).parameters
|
||||
except (TypeError, ValueError): # builtins / exotic callables
|
||||
return False
|
||||
|
||||
|
||||
def _thread_progress_reporter(job_id):
|
||||
"""Callback safe to invoke from a worker thread (faster-whisper runs
|
||||
via asyncio.to_thread): schedules a tiny session update on the loop.
|
||||
|
||||
Progress writes are best-effort display state — failures are logged
|
||||
and swallowed, never allowed to disturb the transcription itself."""
|
||||
loop = asyncio.get_running_loop()
|
||||
|
||||
def report(pct: int) -> None:
|
||||
async def _write() -> None:
|
||||
from sqlalchemy import update
|
||||
|
||||
from shonar.db.session import session_factory
|
||||
|
||||
try:
|
||||
async with session_factory()() as s:
|
||||
await s.execute(
|
||||
update(ProcessingJob)
|
||||
.where(ProcessingJob.id == job_id)
|
||||
.values(progress=int(pct))
|
||||
)
|
||||
await s.commit()
|
||||
except Exception: # pragma: no cover - display state only
|
||||
logger.debug("progress write failed for job %s", job_id, exc_info=True)
|
||||
|
||||
try:
|
||||
asyncio.run_coroutine_threadsafe(_write(), loop)
|
||||
except RuntimeError: # loop already gone (shutdown race)
|
||||
logger.debug("progress dropped for job %s (loop gone)", job_id)
|
||||
|
||||
return report
|
||||
|
||||
MAX_TRIES = 3
|
||||
|
||||
# A `running` job younger than this is treated as live work, not a crash
|
||||
# relic — long transcriptions legitimately outrun the 5-minute sweep, and
|
||||
# re-enqueuing them mid-flight caused a copy storm (one whisper per copy).
|
||||
STALE_RUNNING_AFTER = timedelta(hours=2)
|
||||
|
||||
|
||||
# --- entry -----------------------------------------------------------------
|
||||
|
||||
|
||||
async def enqueue_for_recording(session: AsyncSession, rec: Recording) -> list[JobType]:
|
||||
"""Queue whatever AI stages apply. Safe to call repeatedly: completed
|
||||
work is never redone, failed work is reset for another attempt."""
|
||||
if rec.deleted_at is not None:
|
||||
return []
|
||||
settings = get_settings()
|
||||
tprov = get_transcription_provider(settings)
|
||||
lprov = get_llm_provider(settings)
|
||||
queued: list[JobType] = []
|
||||
if tprov is not None and await reset_or_create(session, rec, JobType.transcribe):
|
||||
queued.append(JobType.transcribe)
|
||||
if (
|
||||
lprov is not None
|
||||
and await latest_transcript_text(session, rec.id) is not None
|
||||
and await reset_or_create(session, rec, JobType.summarize)
|
||||
):
|
||||
queued.append(JobType.summarize)
|
||||
if tprov is None and lprov is None:
|
||||
if rec.processing_status != ProcessingStatus.ai_disabled:
|
||||
rec.processing_status = ProcessingStatus.ai_disabled
|
||||
rec.processing_error = None
|
||||
elif queued and rec.processing_status not in (
|
||||
ProcessingStatus.processing,
|
||||
ProcessingStatus.completed,
|
||||
):
|
||||
rec.processing_status = ProcessingStatus.queued
|
||||
rec.processing_error = None
|
||||
await session.flush()
|
||||
for jt in queued:
|
||||
await transport_enqueue(jt, rec.id)
|
||||
return queued
|
||||
|
||||
|
||||
async def reset_or_create(
|
||||
session: AsyncSession, rec: Recording, job_type: JobType
|
||||
) -> bool:
|
||||
"""Ensure a queued job row. Returns True when (re)queued now: new rows,
|
||||
plus failed/skipped rows (a re-upload or re-entry deserves another
|
||||
attempt). Queued/running/succeeded rows are left alone."""
|
||||
existing = await session.scalar(
|
||||
select(ProcessingJob)
|
||||
.where(
|
||||
ProcessingJob.recording_id == rec.id,
|
||||
ProcessingJob.job_type == job_type,
|
||||
)
|
||||
.order_by(ProcessingJob.id.desc())
|
||||
)
|
||||
if existing is None:
|
||||
session.add(ProcessingJob(recording_id=rec.id, job_type=job_type))
|
||||
return True
|
||||
if existing.status in (JobStatus.queued, JobStatus.running, JobStatus.succeeded):
|
||||
return False
|
||||
existing.status = JobStatus.queued
|
||||
existing.attempt = 0
|
||||
existing.error = None
|
||||
existing.stage = None
|
||||
existing.progress = None
|
||||
existing.started_at = None
|
||||
existing.finished_at = None
|
||||
return True
|
||||
|
||||
|
||||
async def latest_transcript_text(
|
||||
session: AsyncSession, recording_id: uuid.UUID
|
||||
) -> str | None:
|
||||
"""Newest non-superseded transcript; user-edited rows win over newer
|
||||
machine rows (an edit is a verdict, not a draft)."""
|
||||
rows = (
|
||||
await session.scalars(
|
||||
select(Transcript)
|
||||
.where(
|
||||
Transcript.recording_id == recording_id,
|
||||
Transcript.superseded_at.is_(None),
|
||||
)
|
||||
.order_by(Transcript.version.desc())
|
||||
)
|
||||
).all()
|
||||
if not rows:
|
||||
return None
|
||||
for r in rows:
|
||||
if r.edited_by_user and r.text.strip():
|
||||
return r.text
|
||||
text = rows[0].text
|
||||
return text if text.strip() else None
|
||||
|
||||
|
||||
async def transport_enqueue(job_type: JobType, recording_id: uuid.UUID) -> None:
|
||||
"""Best-effort trigger for the configured backend ('arq' or 'inline').
|
||||
Failure only logs: the DB rows are the real queue and the sweep picks
|
||||
up anything the transport missed."""
|
||||
if get_settings().queue_backend.strip().lower() == "inline":
|
||||
from shonar.services import inline_queue
|
||||
|
||||
await inline_queue.enqueue(job_type, str(recording_id))
|
||||
return
|
||||
from arq import create_pool
|
||||
from arq.connections import RedisSettings
|
||||
from arq.constants import result_key_prefix
|
||||
|
||||
job_id = f"{job_type.value}:{recording_id}"
|
||||
try:
|
||||
pool = await create_pool(RedisSettings.from_dsn(get_settings().redis_url))
|
||||
try:
|
||||
# Deterministic _job_id: arq refuses a duplicate while a copy of
|
||||
# (job_type, recording) is queued or running, so sweep pokes and
|
||||
# retried enqueues can never pile up concurrent copies of the
|
||||
# same work.
|
||||
#
|
||||
# arq dedupes on job key OR result key, and failed runs leave a
|
||||
# result behind for keep_result days — that would silently swallow
|
||||
# deliberate re-runs of previously-finished work. Drop only the
|
||||
# stale RESULT key first (never the job key: that is exactly what
|
||||
# shields a live copy from duplicate enqueue).
|
||||
await pool.delete(result_key_prefix + job_id)
|
||||
await pool.enqueue_job(
|
||||
"run_" + job_type.value,
|
||||
str(recording_id),
|
||||
_job_id=job_id,
|
||||
)
|
||||
finally:
|
||||
await pool.aclose()
|
||||
except Exception as e: # noqa: BLE001 — transport must never break uploads
|
||||
logger.warning("arq enqueue failed (%s); worker sweep will pick it up", e)
|
||||
|
||||
|
||||
# --- tasks -----------------------------------------------------------------
|
||||
|
||||
|
||||
def _job_try(ctx: dict) -> int:
|
||||
try:
|
||||
return int(ctx.get("job_try", 1))
|
||||
except (TypeError, ValueError):
|
||||
return 1
|
||||
|
||||
|
||||
async def _load(session: AsyncSession, recording_id: str) -> Recording | None:
|
||||
try:
|
||||
rid = uuid.UUID(recording_id)
|
||||
except ValueError:
|
||||
return None
|
||||
return await session.get(Recording, rid)
|
||||
|
||||
|
||||
async def _job(
|
||||
session: AsyncSession, rec: Recording, job_type: JobType
|
||||
) -> ProcessingJob:
|
||||
job = await session.scalar(
|
||||
select(ProcessingJob)
|
||||
.where(
|
||||
ProcessingJob.recording_id == rec.id,
|
||||
ProcessingJob.job_type == job_type,
|
||||
)
|
||||
.order_by(ProcessingJob.id.desc())
|
||||
)
|
||||
if job is None:
|
||||
job = ProcessingJob(recording_id=rec.id, job_type=job_type)
|
||||
session.add(job)
|
||||
await session.flush()
|
||||
return job
|
||||
|
||||
|
||||
async def _fail(
|
||||
session: AsyncSession,
|
||||
rec: Recording,
|
||||
job: ProcessingJob,
|
||||
message: str,
|
||||
ctx: dict,
|
||||
exc: AIError | None = None,
|
||||
) -> None:
|
||||
"""Config errors fail now; transient errors fail only on the last try
|
||||
(returning normally), otherwise they raise for arq retry."""
|
||||
transient = exc is None or isinstance(exc, ProviderTransientError)
|
||||
if transient and _job_try(ctx) < MAX_TRIES:
|
||||
# Running state was committed before the long phase; put the row
|
||||
# back to queued for the retry and persist that (a raise no longer
|
||||
# rolls the pre-phase commit back).
|
||||
job.status = JobStatus.queued
|
||||
job.attempt = _job_try(ctx)
|
||||
job.stage = None
|
||||
await session.commit()
|
||||
raise ProviderTransientError(message)
|
||||
job.status = JobStatus.failed
|
||||
job.error = message
|
||||
job.stage = None
|
||||
job.finished_at = utcnow()
|
||||
if job.job_type == JobType.summarize and await latest_transcript_text(
|
||||
session, rec.id
|
||||
):
|
||||
# The transcript is usable; a summary timeout must not mark the
|
||||
# whole recording failed (the summary can be re-run separately).
|
||||
rec.processing_status = ProcessingStatus.completed
|
||||
rec.processing_error = f"Summary failed: {message}"
|
||||
else:
|
||||
rec.processing_status = ProcessingStatus.failed
|
||||
rec.processing_error = message
|
||||
await session.flush()
|
||||
|
||||
|
||||
async def run_transcribe(ctx: dict, recording_id: str) -> None:
|
||||
"""Transcribe the original audio; chain into summarization when an LLM
|
||||
is configured."""
|
||||
from shonar.db.session import session_factory
|
||||
|
||||
settings = get_settings()
|
||||
async with session_factory()() as session:
|
||||
rec = await _load(session, recording_id)
|
||||
if rec is None or rec.deleted_at is not None:
|
||||
# Finalize race: the API transaction may not have committed yet.
|
||||
raise ProviderTransientError("Recording not ready; retrying.")
|
||||
job = await _job(session, rec, JobType.transcribe)
|
||||
provider = get_transcription_provider(settings)
|
||||
if provider is None:
|
||||
job.status = JobStatus.skipped
|
||||
await session.flush()
|
||||
await session.commit()
|
||||
await _maybe_chain_summarize(session, rec)
|
||||
return
|
||||
# The recording's saved model is authoritative: a per-recording
|
||||
# override wins, otherwise the global default in force at finalize
|
||||
# time (changing the default never rewrites history).
|
||||
from shonar.services.ai.model_registry import (
|
||||
effective_model,
|
||||
get_global_default_model,
|
||||
)
|
||||
|
||||
model = effective_model(rec.transcription_model, await get_global_default_model(session))
|
||||
if provider.name == "faster_whisper":
|
||||
from shonar.services.ai.faster_whisper import FasterWhisperProvider
|
||||
|
||||
try:
|
||||
provider = FasterWhisperProvider(model=model)
|
||||
except AIError as e:
|
||||
await _fail(session, rec, job, str(e), ctx, e)
|
||||
await session.commit()
|
||||
return
|
||||
original = await session.scalar(
|
||||
select(Asset).where(
|
||||
Asset.recording_id == rec.id, Asset.kind == AssetKind.original
|
||||
)
|
||||
)
|
||||
if original is None:
|
||||
raise ProviderTransientError("Audio not ready; retrying.")
|
||||
job.status = JobStatus.running
|
||||
job.attempt = _job_try(ctx)
|
||||
job.started_at = utcnow()
|
||||
job.stage = "loading-model"
|
||||
job.progress = None
|
||||
rec.processing_status = ProcessingStatus.processing
|
||||
rec.processing_error = None
|
||||
# Commit before the long CPU phase: an open transaction is invisible
|
||||
# to other readers (and on SQLite it locks out the progress writer).
|
||||
await session.commit()
|
||||
try:
|
||||
audio = await get_storage().get(original.storage_key)
|
||||
job.stage = "transcribing"
|
||||
await session.commit()
|
||||
kwargs = {}
|
||||
if _accepts_on_progress(provider):
|
||||
kwargs["on_progress"] = _thread_progress_reporter(job.id)
|
||||
result = await provider.transcribe(audio, original.mime_type, **kwargs)
|
||||
except AIError as e:
|
||||
await _fail(session, rec, job, str(e), ctx, e)
|
||||
await session.commit()
|
||||
return
|
||||
await store_transcript(
|
||||
session, rec, result.text, result.segments, result.language,
|
||||
provider.name, getattr(result, "model", ""),
|
||||
)
|
||||
job.status = JobStatus.succeeded
|
||||
job.stage = None
|
||||
job.progress = 100
|
||||
job.finished_at = utcnow()
|
||||
await session.flush()
|
||||
await _maybe_chain_summarize(session, rec)
|
||||
await session.commit()
|
||||
|
||||
|
||||
async def _maybe_chain_summarize(session: AsyncSession, rec: Recording) -> None:
|
||||
"""After transcription (or a skip): summarize when possible, else finish."""
|
||||
settings = get_settings()
|
||||
if get_llm_provider(settings) is None:
|
||||
if rec.processing_status != ProcessingStatus.completed:
|
||||
rec.processing_status = ProcessingStatus.completed
|
||||
rec.processing_error = None
|
||||
await session.flush()
|
||||
return
|
||||
if await latest_transcript_text(session, rec.id) is None:
|
||||
# LLM configured but nothing to summarize (e.g. empty transcript).
|
||||
existing = await session.scalar(
|
||||
select(ProcessingJob).where(
|
||||
ProcessingJob.recording_id == rec.id,
|
||||
ProcessingJob.job_type == JobType.summarize,
|
||||
)
|
||||
)
|
||||
if existing is not None and existing.status == JobStatus.queued:
|
||||
existing.status = JobStatus.skipped
|
||||
if rec.processing_status != ProcessingStatus.completed:
|
||||
rec.processing_status = ProcessingStatus.completed
|
||||
await session.flush()
|
||||
return
|
||||
if await reset_or_create(session, rec, JobType.summarize):
|
||||
await transport_enqueue(JobType.summarize, rec.id)
|
||||
# Status stays `processing` until the summarize task lands.
|
||||
|
||||
|
||||
async def run_summarize(ctx: dict, recording_id: str) -> None:
|
||||
"""Summarize the latest transcript into the structured summary shape."""
|
||||
from shonar.db.session import session_factory
|
||||
|
||||
settings = get_settings()
|
||||
async with session_factory()() as session:
|
||||
rec = await _load(session, recording_id)
|
||||
if rec is None or rec.deleted_at is not None:
|
||||
raise ProviderTransientError("Recording not ready; retrying.")
|
||||
job = await _job(session, rec, JobType.summarize)
|
||||
provider = get_llm_provider(settings)
|
||||
if provider is None:
|
||||
job.status = JobStatus.skipped
|
||||
await session.flush()
|
||||
await session.commit()
|
||||
return
|
||||
text = await latest_transcript_text(session, rec.id)
|
||||
if text is None:
|
||||
job.status = JobStatus.skipped
|
||||
await session.flush()
|
||||
if rec.processing_status != ProcessingStatus.completed:
|
||||
rec.processing_status = ProcessingStatus.completed
|
||||
await session.commit()
|
||||
return
|
||||
job.status = JobStatus.running
|
||||
job.attempt = _job_try(ctx)
|
||||
job.started_at = utcnow()
|
||||
job.stage = "summarizing"
|
||||
rec.processing_status = ProcessingStatus.processing
|
||||
await session.commit() # visible before the long LLM call
|
||||
try:
|
||||
result = await provider.summarize(text, title=rec.title)
|
||||
except AIError as e:
|
||||
await _fail(session, rec, job, str(e), ctx, e)
|
||||
await session.commit()
|
||||
return
|
||||
await store_summary(session, rec, result.to_dict(), provider.name, result.model)
|
||||
job.status = JobStatus.succeeded
|
||||
job.stage = None
|
||||
job.progress = 100
|
||||
job.finished_at = utcnow()
|
||||
rec.processing_status = ProcessingStatus.completed
|
||||
rec.processing_error = None
|
||||
await session.flush()
|
||||
await session.commit()
|
||||
|
||||
|
||||
async def store_transcript(
|
||||
session: AsyncSession,
|
||||
rec: Recording,
|
||||
text: str,
|
||||
segments: list,
|
||||
language: str | None,
|
||||
provider_name: str,
|
||||
model: str,
|
||||
) -> None:
|
||||
"""Insert a new auto version; a newest user-edited row wins instead and
|
||||
nothing is inserted (edits are verdicts)."""
|
||||
existing = (
|
||||
await session.scalars(
|
||||
select(Transcript)
|
||||
.where(
|
||||
Transcript.recording_id == rec.id,
|
||||
Transcript.superseded_at.is_(None),
|
||||
)
|
||||
.order_by(Transcript.version.desc())
|
||||
)
|
||||
).all()
|
||||
if existing and existing[0].edited_by_user:
|
||||
return
|
||||
now = utcnow()
|
||||
max_version = await session.scalar(
|
||||
select(func.max(Transcript.version)).where(Transcript.recording_id == rec.id)
|
||||
)
|
||||
for row in existing:
|
||||
row.superseded_at = now
|
||||
session.add(
|
||||
Transcript(
|
||||
recording_id=rec.id,
|
||||
version=(max_version or 0) + 1,
|
||||
language=language,
|
||||
provider=provider_name,
|
||||
model=model,
|
||||
text=text,
|
||||
segments=[
|
||||
{"start": s.start, "end": s.end, "text": s.text, "speaker": s.speaker}
|
||||
for s in segments
|
||||
],
|
||||
edited_by_user=False,
|
||||
)
|
||||
)
|
||||
await session.flush()
|
||||
|
||||
|
||||
async def store_summary(
|
||||
session: AsyncSession,
|
||||
rec: Recording,
|
||||
content: dict,
|
||||
provider_name: str,
|
||||
model: str,
|
||||
) -> None:
|
||||
existing = (
|
||||
await session.scalars(
|
||||
select(Summary)
|
||||
.where(
|
||||
Summary.recording_id == rec.id,
|
||||
Summary.superseded_at.is_(None),
|
||||
)
|
||||
.order_by(Summary.version.desc())
|
||||
)
|
||||
).all()
|
||||
if existing and existing[0].edited_by_user:
|
||||
return
|
||||
now = utcnow()
|
||||
max_version = await session.scalar(
|
||||
select(func.max(Summary.version)).where(Summary.recording_id == rec.id)
|
||||
)
|
||||
for row in existing:
|
||||
row.superseded_at = now
|
||||
session.add(
|
||||
Summary(
|
||||
recording_id=rec.id,
|
||||
version=(max_version or 0) + 1,
|
||||
provider=provider_name,
|
||||
model=model,
|
||||
content=content,
|
||||
edited_by_user=False,
|
||||
)
|
||||
)
|
||||
await session.flush()
|
||||
|
||||
|
||||
async def sweep_stale(limit: int = 100) -> int:
|
||||
"""Crash recovery + transport-loss backstop: requeue jobs stuck running
|
||||
or sitting queued, newest catastrophe first. Returns jobs re-enqueued."""
|
||||
from shonar.db.session import session_factory
|
||||
|
||||
count = 0
|
||||
async with session_factory()() as session:
|
||||
rows = (
|
||||
await session.scalars(
|
||||
select(ProcessingJob)
|
||||
.where(ProcessingJob.status.in_([JobStatus.queued, JobStatus.running]))
|
||||
.order_by(ProcessingJob.id.desc())
|
||||
.limit(limit)
|
||||
)
|
||||
).all()
|
||||
for job in rows:
|
||||
rec = await session.get(Recording, job.recording_id)
|
||||
if rec is None or rec.deleted_at is not None:
|
||||
job.status = JobStatus.skipped
|
||||
continue
|
||||
if job.status == JobStatus.running:
|
||||
if job.attempt >= MAX_TRIES:
|
||||
job.status = JobStatus.failed
|
||||
job.error = "Worker died too many times."
|
||||
job.finished_at = utcnow()
|
||||
continue
|
||||
# Live work, not a crash relic: a running job that started
|
||||
# recently belongs to a worker still chewing it (long audio
|
||||
# outruns the sweep cadence). Only orphaned runs — no
|
||||
# started_at, or older than the grace window — get requeued.
|
||||
if job.started_at is not None and (
|
||||
utcnow() - job.started_at
|
||||
) < STALE_RUNNING_AFTER:
|
||||
continue
|
||||
job.status = JobStatus.queued
|
||||
job.error = None
|
||||
count += 1
|
||||
await session.commit()
|
||||
# Transport outside the transaction: rows are the queue, this just pokes.
|
||||
async with session_factory()() as session:
|
||||
rows = (
|
||||
await session.scalars(
|
||||
select(ProcessingJob)
|
||||
.where(ProcessingJob.status == JobStatus.queued)
|
||||
.order_by(ProcessingJob.id.desc())
|
||||
.limit(limit)
|
||||
)
|
||||
).all()
|
||||
for job in rows:
|
||||
await transport_enqueue(job.job_type, job.recording_id)
|
||||
return count
|
||||
101
backend/shonar/services/retention.py
Normal file
101
backend/shonar/services/retention.py
Normal file
|
|
@ -0,0 +1,101 @@
|
|||
"""Retention sweep (M9): hard-delete what passed its grace window.
|
||||
|
||||
Two policies, both driven by ``deleted_at``:
|
||||
|
||||
* **Recordings** soft-deleted more than ``retention_grace_days`` ago have
|
||||
their rows removed (cascade cleans transcripts/summaries/jobs/tags) and
|
||||
every stored asset file deleted best-effort.
|
||||
* **Accounts** deleted more than ``retention_grace_days`` ago are hard
|
||||
deleted (user cascade takes their recordings/assets/devices/tokens);
|
||||
their storage files are collected the same way.
|
||||
|
||||
Running inside the grace window is a no-op, so an accidental delete stays
|
||||
cancellable until the sweep actually fires. The sweep is idempotent and
|
||||
safe to run on any cadence (worker cron + inline-queue timer).
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import contextlib
|
||||
import logging
|
||||
from datetime import timedelta
|
||||
|
||||
from sqlalchemy import delete as sql_delete
|
||||
from sqlalchemy import select
|
||||
|
||||
from shonar.core.config import get_settings
|
||||
from shonar.db.models import Asset, Recording, User, utcnow
|
||||
from shonar.db.session import session_factory
|
||||
from shonar.storage import get_storage
|
||||
|
||||
logger = logging.getLogger("shonar.retention")
|
||||
|
||||
|
||||
async def sweep_deleted(limit: int = 200) -> dict[str, int]:
|
||||
"""Hard-purge expired recordings and accounts. Returns counts."""
|
||||
grace = timedelta(days=get_settings().retention_grace_days)
|
||||
cutoff = utcnow() - grace
|
||||
purged = {"recordings": 0, "users": 0, "files": 0}
|
||||
storage = get_storage()
|
||||
|
||||
async with session_factory()() as session:
|
||||
# --- recordings (skip rows whose account is also expiring: the
|
||||
# user cascade below collects their files in one pass) ---
|
||||
expiring_users = select(User.id).where(
|
||||
User.deleted_at.is_not(None), User.deleted_at < cutoff
|
||||
)
|
||||
recs = list(
|
||||
await session.scalars(
|
||||
select(Recording)
|
||||
.where(
|
||||
Recording.deleted_at.is_not(None),
|
||||
Recording.deleted_at < cutoff,
|
||||
Recording.user_id.notin_(expiring_users),
|
||||
)
|
||||
.limit(limit)
|
||||
)
|
||||
)
|
||||
for rec in recs:
|
||||
assets = list(
|
||||
await session.scalars(select(Asset).where(Asset.recording_id == rec.id))
|
||||
)
|
||||
await session.delete(rec)
|
||||
await session.flush()
|
||||
for a in assets:
|
||||
with contextlib.suppress(Exception): # best effort; DB row is gone
|
||||
await storage.delete(a.storage_key)
|
||||
purged["files"] += 1
|
||||
purged["recordings"] += 1
|
||||
|
||||
# --- accounts ---
|
||||
users = list(
|
||||
await session.scalars(
|
||||
select(User)
|
||||
.where(User.deleted_at.is_not(None), User.deleted_at < cutoff)
|
||||
.limit(limit)
|
||||
)
|
||||
)
|
||||
for user in users:
|
||||
assets = list(
|
||||
await session.scalars(select(Asset).where(Asset.user_id == user.id))
|
||||
)
|
||||
keys = [a.storage_key for a in assets]
|
||||
# DB-level delete: the ORM would null the NOT NULL FKs of the
|
||||
# user's recordings before the ON DELETE CASCADE could fire.
|
||||
await session.execute(sql_delete(User).where(User.id == user.id))
|
||||
await session.flush()
|
||||
for key in keys:
|
||||
with contextlib.suppress(Exception):
|
||||
await storage.delete(key)
|
||||
purged["files"] += 1
|
||||
purged["users"] += 1
|
||||
|
||||
await session.commit()
|
||||
|
||||
if purged["recordings"] or purged["users"]:
|
||||
logger.info(
|
||||
"retention sweep purged %d recordings, %d accounts, %d files (grace %dd)",
|
||||
purged["recordings"], purged["users"], purged["files"],
|
||||
get_settings().retention_grace_days,
|
||||
)
|
||||
return purged
|
||||
292
backend/shonar/services/search.py
Normal file
292
backend/shonar/services/search.py
Normal file
|
|
@ -0,0 +1,292 @@
|
|||
"""Full-text search across recordings (M9).
|
||||
|
||||
Postgres uses the tsvector columns from migration ``fts0000000001``
|
||||
(title/notes, transcript text, summary content, tag names). SQLite (the
|
||||
desktop bundled-lite engine) falls back to a substring scan; libraries
|
||||
there are single-user and small, and the endpoint contract is identical.
|
||||
|
||||
This module is the whole SearchBackend seam — a Meilisearch/OpenSearch
|
||||
implementation would replace it, not the callers.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import uuid
|
||||
from dataclasses import dataclass
|
||||
|
||||
from sqlalchemy import select, text
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from shonar.db.models import Recording, RecordingTag, Summary, Tag, Transcript
|
||||
|
||||
# ``scope`` values accepted by the endpoint.
|
||||
SCOPES = ("all", "title", "notes", "transcript", "summary", "tag")
|
||||
|
||||
# Hit ordering: a title hit outranks a transcript hit.
|
||||
_FIELD_RANK = {"title": 0, "tag": 1, "notes": 2, "summary": 3, "transcript": 4}
|
||||
|
||||
|
||||
@dataclass
|
||||
class SearchHit:
|
||||
recording: Recording
|
||||
field: str # where the best match landed
|
||||
snippet: str
|
||||
|
||||
|
||||
def _snippet_around(value: str, at: int, width: int = 160) -> str:
|
||||
"""~``width`` chars centred on ``at``, word-bounded, with ellipses."""
|
||||
half = width // 2
|
||||
start = max(0, at - half)
|
||||
end = min(len(value), at + half)
|
||||
if start > 0:
|
||||
start = value.rfind(" ", 0, start) + 1 or start
|
||||
if end < len(value):
|
||||
nxt = value.find(" ", end)
|
||||
end = nxt if nxt != -1 else end
|
||||
prefix = "…" if start > 0 else ""
|
||||
suffix = "…" if end < len(value) else ""
|
||||
return f"{prefix}{value[start:end].strip()}{suffix}"
|
||||
|
||||
|
||||
async def tag_names(session: AsyncSession, recording_id: uuid.UUID) -> list[str]:
|
||||
return list(
|
||||
await session.scalars(
|
||||
select(Tag.name)
|
||||
.join(RecordingTag, RecordingTag.tag_id == Tag.id)
|
||||
.where(RecordingTag.recording_id == recording_id)
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
async def _current_transcript_text(session: AsyncSession, recording_id: uuid.UUID) -> str:
|
||||
row = await session.scalar(
|
||||
select(Transcript.text)
|
||||
.where(Transcript.recording_id == recording_id, Transcript.superseded_at.is_(None))
|
||||
.order_by(Transcript.version.desc())
|
||||
)
|
||||
return row or ""
|
||||
|
||||
|
||||
def _summary_flat(content: dict | None) -> str:
|
||||
if not content:
|
||||
return ""
|
||||
parts = [str(content.get("short", "")), str(content.get("detailed", ""))]
|
||||
for key in ("key_points", "decisions", "action_items", "questions"):
|
||||
v = content.get(key)
|
||||
if isinstance(v, list):
|
||||
parts.extend(str(x) for x in v)
|
||||
return " ".join(p for p in parts if p)
|
||||
|
||||
|
||||
# --- SQLite fallback ----------------------------------------------------------
|
||||
|
||||
|
||||
async def _search_sqlite(
|
||||
session: AsyncSession, user_id: uuid.UUID, q: str, scope: str,
|
||||
limit: int, offset: int,
|
||||
) -> tuple[list[SearchHit], int]:
|
||||
needles = [t for t in q.lower().split() if t]
|
||||
if not needles:
|
||||
return [], 0
|
||||
recs = list(
|
||||
await session.scalars(
|
||||
select(Recording)
|
||||
.where(Recording.user_id == user_id, Recording.deleted_at.is_(None))
|
||||
.order_by(Recording.recorded_at.desc())
|
||||
)
|
||||
)
|
||||
hits: list[SearchHit] = []
|
||||
for rec in recs:
|
||||
fields: list[tuple[str, str]] = []
|
||||
if scope in ("all", "title"):
|
||||
fields.append(("title", rec.title or ""))
|
||||
if scope in ("all", "notes"):
|
||||
fields.append(("notes", rec.notes or ""))
|
||||
if scope in ("all", "tag"):
|
||||
fields.append(("tag", " ".join(await tag_names(session, rec.id))))
|
||||
if scope in ("all", "transcript"):
|
||||
fields.append(("transcript", await _current_transcript_text(session, rec.id)))
|
||||
if scope in ("all", "summary"):
|
||||
content = await session.scalar(
|
||||
select(Summary.content)
|
||||
.where(Summary.recording_id == rec.id, Summary.superseded_at.is_(None))
|
||||
.order_by(Summary.version.desc())
|
||||
)
|
||||
fields.append(("summary", _summary_flat(content)))
|
||||
# A field matches when EVERY needle appears in it (AND semantics,
|
||||
# matching plainto_tsquery on the Postgres side).
|
||||
best: SearchHit | None = None
|
||||
for field, value in fields:
|
||||
low = value.lower()
|
||||
if not all(n in low for n in needles):
|
||||
continue
|
||||
at = low.find(needles[0])
|
||||
cand = SearchHit(
|
||||
recording=rec, field=field, snippet=_snippet_around(value, max(at, 0))
|
||||
)
|
||||
if best is None or _FIELD_RANK[field] < _FIELD_RANK[best.field]:
|
||||
best = cand
|
||||
if best is not None:
|
||||
hits.append(best)
|
||||
hits.sort(key=lambda h: (_FIELD_RANK[h.field], h.recording.recorded_at), reverse=False)
|
||||
return hits[offset : offset + limit], len(hits)
|
||||
|
||||
|
||||
# --- PostgreSQL tsvector path -------------------------------------------------
|
||||
|
||||
# CTE ``q`` carries the parsed tsquery so it is computed once. Notes are not
|
||||
# in the recordings search_vector weights the way we want headlines, so notes
|
||||
# match by substring like the SQLite path (title/notes share the vector; the
|
||||
# field classifier prefers 'title' when the vector hits).
|
||||
_PG_MATCHES = """
|
||||
WITH q AS (SELECT plainto_tsquery('simple', :q) AS ts),
|
||||
lt AS (
|
||||
SELECT DISTINCT ON (recording_id) recording_id, search_vector, text
|
||||
FROM transcripts WHERE superseded_at IS NULL
|
||||
ORDER BY recording_id, version DESC
|
||||
),
|
||||
ls AS (
|
||||
SELECT DISTINCT ON (recording_id) recording_id, content, search_vector
|
||||
FROM summaries WHERE superseded_at IS NULL
|
||||
ORDER BY recording_id, version DESC
|
||||
),
|
||||
tm AS (
|
||||
SELECT DISTINCT rt.recording_id,
|
||||
ts_headline('simple', t.name, q.ts,
|
||||
'StartSel=,StopSel=,MaxFragments=0') AS snip
|
||||
FROM recording_tags rt
|
||||
JOIN tags t ON t.id = rt.tag_id AND t.user_id = :uid, q
|
||||
),
|
||||
matched AS (
|
||||
SELECT r.id AS id,
|
||||
CASE
|
||||
WHEN :want_title AND r.search_vector @@ q.ts
|
||||
AND coalesce(r.title, '') <> ''
|
||||
THEN 'title'
|
||||
WHEN :want_tag AND tm.recording_id IS NOT NULL THEN 'tag'
|
||||
WHEN :want_notes AND coalesce(r.notes, '') ILIKE '%' || :raw || '%' THEN 'notes'
|
||||
WHEN :want_summary AND ls.search_vector @@ q.ts THEN 'summary'
|
||||
WHEN :want_transcript AND lt.search_vector @@ q.ts THEN 'transcript'
|
||||
END AS field,
|
||||
CASE
|
||||
WHEN :want_title AND r.search_vector @@ q.ts
|
||||
AND coalesce(r.title, '') <> ''
|
||||
THEN ts_headline('simple', r.title, q.ts,
|
||||
'StartSel=,StopSel=,MaxFragments=0,MaxWords=25')
|
||||
WHEN :want_tag AND tm.recording_id IS NOT NULL THEN tm.snip
|
||||
WHEN :want_notes AND coalesce(r.notes, '') ILIKE '%' || :raw || '%'
|
||||
THEN left(r.notes, 200)
|
||||
WHEN :want_summary AND ls.search_vector @@ q.ts
|
||||
THEN ts_headline('simple',
|
||||
coalesce(ls.content->>'short', '') || ' ' || coalesce(ls.content->>'detailed', ''),
|
||||
q.ts, 'StartSel=,StopSel=,MaxFragments=1,MinWords=10,MaxWords=25')
|
||||
WHEN :want_transcript AND lt.search_vector @@ q.ts
|
||||
THEN ts_headline('simple', coalesce(lt.text, ''), q.ts,
|
||||
'StartSel=,StopSel=,MaxFragments=1,MinWords=10,MaxWords=25')
|
||||
END AS snippet
|
||||
FROM recordings r
|
||||
CROSS JOIN q
|
||||
LEFT JOIN lt ON lt.recording_id = r.id
|
||||
LEFT JOIN ls ON ls.recording_id = r.id
|
||||
LEFT JOIN tm ON tm.recording_id = r.id
|
||||
WHERE r.user_id = :uid AND r.deleted_at IS NULL
|
||||
)
|
||||
SELECT id, field, snippet FROM matched
|
||||
WHERE field IS NOT NULL
|
||||
ORDER BY CASE field WHEN 'title' THEN 0 WHEN 'tag' THEN 1 WHEN 'notes' THEN 2
|
||||
WHEN 'summary' THEN 3 ELSE 4 END,
|
||||
id
|
||||
LIMIT :limit OFFSET :offset
|
||||
"""
|
||||
|
||||
_PG_COUNT = """
|
||||
WITH q AS (SELECT plainto_tsquery('simple', :q) AS ts),
|
||||
lt AS (
|
||||
SELECT DISTINCT ON (recording_id) recording_id, search_vector
|
||||
FROM transcripts WHERE superseded_at IS NULL
|
||||
),
|
||||
ls AS (
|
||||
SELECT DISTINCT ON (recording_id) recording_id, search_vector
|
||||
FROM summaries WHERE superseded_at IS NULL
|
||||
),
|
||||
tm AS (
|
||||
SELECT DISTINCT rt.recording_id
|
||||
FROM recording_tags rt
|
||||
JOIN tags t ON t.id = rt.tag_id AND t.user_id = :uid, q
|
||||
WHERE t.search_vector @@ q.ts
|
||||
)
|
||||
SELECT count(*)
|
||||
FROM recordings r
|
||||
CROSS JOIN q
|
||||
LEFT JOIN lt ON lt.recording_id = r.id
|
||||
LEFT JOIN ls ON ls.recording_id = r.id
|
||||
LEFT JOIN tm ON tm.recording_id = r.id
|
||||
WHERE r.user_id = :uid AND r.deleted_at IS NULL
|
||||
AND (
|
||||
(:want_title AND r.search_vector @@ q.ts) OR
|
||||
(:want_transcript AND lt.search_vector @@ q.ts) OR
|
||||
(:want_summary AND ls.search_vector @@ q.ts) OR
|
||||
(:want_tag AND tm.recording_id IS NOT NULL) OR
|
||||
(:want_notes AND coalesce(r.notes, '') ILIKE '%' || :raw || '%')
|
||||
)
|
||||
"""
|
||||
|
||||
|
||||
async def _search_postgres(
|
||||
session: AsyncSession, user_id: uuid.UUID, q: str, scope: str,
|
||||
limit: int, offset: int,
|
||||
) -> tuple[list[SearchHit], int]:
|
||||
# NOTE: the title field matches anything the recordings vector hits
|
||||
# (title + notes); notes-only hits surface under 'title' headlines from
|
||||
# the title text. Acceptable precision tradeoff for a GIN-indexed path.
|
||||
params = {
|
||||
"uid": user_id,
|
||||
"q": q,
|
||||
"raw": q,
|
||||
"limit": limit,
|
||||
"offset": offset,
|
||||
"want_title": scope in ("all", "title"),
|
||||
"want_transcript": scope in ("all", "transcript"),
|
||||
"want_summary": scope in ("all", "summary"),
|
||||
"want_tag": scope in ("all", "tag"),
|
||||
"want_notes": scope in ("all", "notes"),
|
||||
}
|
||||
rows = (await session.execute(text(_PG_MATCHES), params)).all()
|
||||
total = await session.scalar(text(_PG_COUNT), params) or 0
|
||||
if not rows:
|
||||
return [], total
|
||||
ids = [r[0] for r in rows]
|
||||
by_id = {
|
||||
rec.id: rec
|
||||
for rec in (
|
||||
await session.scalars(select(Recording).where(Recording.id.in_(ids)))
|
||||
).all()
|
||||
}
|
||||
hits = [
|
||||
SearchHit(recording=by_id[r.id], field=r.field, snippet=r.snippet or "")
|
||||
for r in rows
|
||||
if r.id in by_id
|
||||
]
|
||||
return hits, total
|
||||
|
||||
|
||||
# --- public API ----------------------------------------------------------------
|
||||
|
||||
|
||||
async def search_recordings(
|
||||
session: AsyncSession,
|
||||
user_id: uuid.UUID,
|
||||
q: str,
|
||||
*,
|
||||
scope: str = "all",
|
||||
limit: int = 20,
|
||||
offset: int = 0,
|
||||
) -> tuple[list[SearchHit], int]:
|
||||
"""Returns (page of hits ordered by field rank, total). Empty q → no hits."""
|
||||
q = q.strip()
|
||||
if not q:
|
||||
return [], 0
|
||||
dialect = session.bind.dialect.name if session.bind else "sqlite"
|
||||
if dialect == "postgresql":
|
||||
return await _search_postgres(session, user_id, q, scope, limit, offset)
|
||||
return await _search_sqlite(session, user_id, q, scope, limit, offset)
|
||||
342
backend/shonar/services/uploads.py
Normal file
342
backend/shonar/services/uploads.py
Normal file
|
|
@ -0,0 +1,342 @@
|
|||
"""Chunked, resumable upload sessions.
|
||||
|
||||
Flow:
|
||||
1. POST /uploads -> session (uuid), chunk size, expiry
|
||||
2. PUT /uploads/{id}/chunks/{n} (idempotent per index; GET status lists
|
||||
received indexes so clients resume)
|
||||
3. POST /uploads/{id}/finalize -> validates size + magic bytes,
|
||||
assembles the object, creates the
|
||||
immutable original Asset and the
|
||||
Recording (or updates the existing
|
||||
recording for a retried
|
||||
client_recording_id).
|
||||
|
||||
Storage keys are server-generated UUID paths; clients never see them.
|
||||
Originals are immutable: finalize never overwrites an existing original.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import contextlib
|
||||
import hashlib
|
||||
import uuid
|
||||
from datetime import datetime, timedelta
|
||||
|
||||
from sqlalchemy import select
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from shonar.core.config import get_settings
|
||||
from shonar.db.models import (
|
||||
Asset,
|
||||
AssetKind,
|
||||
ProcessingStatus,
|
||||
Recording,
|
||||
UploadChunk,
|
||||
UploadSession,
|
||||
UploadSessionStatus,
|
||||
utcnow,
|
||||
)
|
||||
from shonar.services.media import is_compatible
|
||||
from shonar.storage import get_storage
|
||||
|
||||
SESSION_TTL = timedelta(hours=24)
|
||||
|
||||
|
||||
class UploadError(Exception):
|
||||
def __init__(self, message: str, status_code: int = 400):
|
||||
super().__init__(message)
|
||||
self.message = message
|
||||
self.status_code = status_code
|
||||
|
||||
|
||||
async def create_session(
|
||||
session: AsyncSession,
|
||||
user_id: uuid.UUID,
|
||||
declared_mime_type: str,
|
||||
declared_size_bytes: int,
|
||||
title: str | None,
|
||||
client_recording_id: str | None,
|
||||
transcription_model: str | None = None,
|
||||
) -> UploadSession:
|
||||
settings = get_settings()
|
||||
if declared_size_bytes <= 0 or declared_size_bytes > settings.max_upload_bytes:
|
||||
raise UploadError(
|
||||
f"Declared size must be between 1 and {settings.max_upload_bytes} bytes.", 413
|
||||
)
|
||||
allowed = [m.lower() for m in settings.allowed_audio_mime_types]
|
||||
if declared_mime_type.lower() not in allowed:
|
||||
raise UploadError(f"MIME type not allowed. Allowed: {', '.join(allowed)}", 415)
|
||||
|
||||
us = UploadSession(
|
||||
user_id=user_id,
|
||||
client_recording_id=client_recording_id,
|
||||
title=title,
|
||||
declared_mime_type=declared_mime_type.lower(),
|
||||
declared_size_bytes=declared_size_bytes,
|
||||
chunk_size_bytes=settings.max_chunk_bytes,
|
||||
expires_at=utcnow() + SESSION_TTL,
|
||||
transcription_model=(transcription_model.strip() if transcription_model else None),
|
||||
)
|
||||
session.add(us)
|
||||
await session.flush()
|
||||
return us
|
||||
|
||||
|
||||
async def get_owned_session(
|
||||
session: AsyncSession, user_id: uuid.UUID, session_id: uuid.UUID
|
||||
) -> UploadSession:
|
||||
us = await session.get(UploadSession, session_id)
|
||||
if us is None or us.user_id != user_id:
|
||||
raise UploadError("Upload session not found.", 404)
|
||||
if us.status == UploadSessionStatus.expired or (
|
||||
us.status != UploadSessionStatus.completed and us.expires_at < utcnow()
|
||||
):
|
||||
us.status = UploadSessionStatus.expired
|
||||
raise UploadError("Upload session expired. Start a new upload.", 410)
|
||||
return us
|
||||
|
||||
|
||||
async def put_chunk(
|
||||
session: AsyncSession,
|
||||
user_id: uuid.UUID,
|
||||
session_id: uuid.UUID,
|
||||
chunk_index: int,
|
||||
data: bytes,
|
||||
checksum_sha256: str | None,
|
||||
) -> UploadChunk:
|
||||
settings = get_settings()
|
||||
us = await get_owned_session(session, user_id, session_id)
|
||||
if us.status != UploadSessionStatus.open:
|
||||
raise UploadError(f"Session is {us.status.value}; cannot accept chunks.", 409)
|
||||
if chunk_index < 0:
|
||||
raise UploadError("Chunk index must be >= 0.", 400)
|
||||
if not data:
|
||||
raise UploadError("Empty chunk.", 400)
|
||||
if len(data) > settings.max_chunk_bytes:
|
||||
raise UploadError(f"Chunk exceeds max size {settings.max_chunk_bytes}.", 413)
|
||||
if checksum_sha256 and hashlib.sha256(data).hexdigest() != checksum_sha256.lower():
|
||||
raise UploadError("Chunk checksum mismatch.", 422)
|
||||
|
||||
existing = await session.scalar(
|
||||
select(UploadChunk).where(
|
||||
UploadChunk.session_id == us.id, UploadChunk.chunk_index == chunk_index
|
||||
)
|
||||
)
|
||||
if existing is not None:
|
||||
# Idempotent retry: same index re-sent replaces the stored bytes.
|
||||
if existing.size_bytes != len(data):
|
||||
await get_storage().delete(existing.storage_key)
|
||||
existing.size_bytes = len(data)
|
||||
existing.checksum_sha256 = hashlib.sha256(data).hexdigest()
|
||||
await get_storage().put(existing.storage_key, data)
|
||||
await session.flush()
|
||||
return existing
|
||||
|
||||
key = f"uploads/{us.id}/{chunk_index:06d}.part"
|
||||
await get_storage().put(key, data)
|
||||
chunk = UploadChunk(
|
||||
session_id=us.id,
|
||||
chunk_index=chunk_index,
|
||||
size_bytes=len(data),
|
||||
checksum_sha256=hashlib.sha256(data).hexdigest(),
|
||||
storage_key=key,
|
||||
)
|
||||
session.add(chunk)
|
||||
await session.flush()
|
||||
return chunk
|
||||
|
||||
|
||||
async def received_indexes(
|
||||
session: AsyncSession, user_id: uuid.UUID, session_id: uuid.UUID
|
||||
) -> list[int]:
|
||||
us = await get_owned_session(session, user_id, session_id)
|
||||
rows = await session.scalars(
|
||||
select(UploadChunk.chunk_index).where(UploadChunk.session_id == us.id)
|
||||
)
|
||||
return sorted(rows)
|
||||
|
||||
|
||||
async def finalize(
|
||||
session: AsyncSession,
|
||||
user_id: uuid.UUID,
|
||||
session_id: uuid.UUID,
|
||||
*,
|
||||
recorded_at: datetime | None,
|
||||
duration_seconds: float,
|
||||
latitude: float | None = None,
|
||||
longitude: float | None = None,
|
||||
location_accuracy_m: float | None = None,
|
||||
notes: str | None = None,
|
||||
transcription_model: str | None = None,
|
||||
) -> tuple[UploadSession, Recording, Asset]:
|
||||
"""Assemble chunks, validate, store the immutable original, and create or
|
||||
update the recording. Idempotent per client_recording_id."""
|
||||
from shonar.services.ai import ProviderConfigError
|
||||
from shonar.services.ai.model_registry import (
|
||||
effective_model,
|
||||
get_global_default_model,
|
||||
validate_model_name,
|
||||
)
|
||||
|
||||
us = await get_owned_session(session, user_id, session_id)
|
||||
# Resolve + validate the transcription model before touching audio.
|
||||
# Finalize body wins over the session's upload-screen choice; an
|
||||
# explicit override updates history, the default never rewrites it.
|
||||
override_raw = transcription_model or us.transcription_model
|
||||
override: str | None = None
|
||||
if override_raw:
|
||||
try:
|
||||
override = validate_model_name(override_raw)
|
||||
except ProviderConfigError as e:
|
||||
raise UploadError(str(e), 422) from None
|
||||
model = effective_model(override, await get_global_default_model(session))
|
||||
if us.status == UploadSessionStatus.completed and us.completed_asset_id:
|
||||
# Already finalized: return existing recording (retry-safe client).
|
||||
asset = await session.get(Asset, us.completed_asset_id)
|
||||
rec = await session.scalar(
|
||||
select(Recording).where(Recording.id == asset.recording_id)
|
||||
)
|
||||
if asset and rec:
|
||||
return us, rec, asset
|
||||
if us.status != UploadSessionStatus.open:
|
||||
raise UploadError(f"Session is {us.status.value}.", 409)
|
||||
|
||||
chunks = list(
|
||||
await session.scalars(
|
||||
select(UploadChunk).where(UploadChunk.session_id == us.id).order_by(
|
||||
UploadChunk.chunk_index
|
||||
)
|
||||
)
|
||||
)
|
||||
total = sum(c.size_bytes for c in chunks)
|
||||
if total != us.declared_size_bytes:
|
||||
raise UploadError(
|
||||
f"Size mismatch: received {total} of declared {us.declared_size_bytes} bytes. "
|
||||
"Upload missing chunks and retry.",
|
||||
422,
|
||||
)
|
||||
expected_indexes = list(range(len(chunks)))
|
||||
if [c.chunk_index for c in chunks] != expected_indexes:
|
||||
raise UploadError("Chunk sequence has gaps. Upload missing chunks and retry.", 422)
|
||||
|
||||
storage = get_storage()
|
||||
# Validate magic bytes from the first chunk.
|
||||
first = await storage.get(chunks[0].storage_key)
|
||||
if not is_compatible(us.declared_mime_type, first):
|
||||
raise UploadError(
|
||||
"File contents do not match the declared audio MIME type.", 415
|
||||
)
|
||||
|
||||
# Idempotency: same client_recording_id => update existing recording.
|
||||
recording: Recording | None = None
|
||||
if us.client_recording_id:
|
||||
recording = await session.scalar(
|
||||
select(Recording).where(
|
||||
Recording.user_id == user_id,
|
||||
Recording.client_recording_id == us.client_recording_id,
|
||||
)
|
||||
)
|
||||
|
||||
# NOTE(original-immutability): if the recording already has an original
|
||||
# asset we do NOT replace it; a re-upload with the same client id after
|
||||
# local edits updates metadata only, and the new audio is rejected as a
|
||||
# duplicate (the existing original is returned instead).
|
||||
if recording is not None:
|
||||
existing_original = await session.scalar(
|
||||
select(Asset).where(
|
||||
Asset.recording_id == recording.id, Asset.kind == AssetKind.original
|
||||
)
|
||||
)
|
||||
if existing_original is not None:
|
||||
us.status = UploadSessionStatus.completed
|
||||
us.completed_asset_id = existing_original.id
|
||||
recording.title = us.title or recording.title
|
||||
if override is not None:
|
||||
# Explicit re-choice replaces the saved model; the default
|
||||
# never rewrites history.
|
||||
recording.transcription_model = override
|
||||
await session.flush()
|
||||
from shonar.services import processing as _processing
|
||||
|
||||
# Same audio, maybe new metadata — and a failed pipeline deserves
|
||||
# another attempt. Idempotent: completed work is never redone.
|
||||
await _processing.enqueue_for_recording(session, recording)
|
||||
return us, recording, existing_original
|
||||
|
||||
# Assemble into the final object (streamed per chunk to bound memory).
|
||||
from shonar.services.media import extension_for
|
||||
|
||||
digest = hashlib.sha256()
|
||||
parts: list[bytes] = []
|
||||
for c in chunks:
|
||||
data = await storage.get(c.storage_key)
|
||||
digest.update(data)
|
||||
parts.append(data)
|
||||
checksum = digest.hexdigest()
|
||||
ext = extension_for(us.declared_mime_type, parts[0])
|
||||
final_key = f"recordings/{user_id}/{uuid.uuid4()}{ext}"
|
||||
blob = b"".join(parts)
|
||||
await storage.put(final_key, blob)
|
||||
|
||||
asset = Asset(
|
||||
recording_id=recording.id if recording else None,
|
||||
user_id=user_id,
|
||||
kind=AssetKind.original,
|
||||
storage_key=final_key,
|
||||
mime_type=us.declared_mime_type,
|
||||
size_bytes=total,
|
||||
checksum_sha256=checksum,
|
||||
)
|
||||
session.add(asset)
|
||||
|
||||
if recording is None:
|
||||
recording = Recording(
|
||||
user_id=user_id,
|
||||
client_recording_id=us.client_recording_id,
|
||||
title=us.title or "Untitled recording",
|
||||
recorded_at=recorded_at or utcnow(),
|
||||
duration_seconds=duration_seconds,
|
||||
notes=notes,
|
||||
latitude=latitude,
|
||||
longitude=longitude,
|
||||
location_accuracy_m=location_accuracy_m,
|
||||
processing_status=ProcessingStatus.uploaded,
|
||||
transcription_model=model,
|
||||
)
|
||||
session.add(recording)
|
||||
await session.flush()
|
||||
asset.recording_id = recording.id
|
||||
else:
|
||||
recording.title = us.title or recording.title
|
||||
recording.duration_seconds = duration_seconds or recording.duration_seconds
|
||||
recording.processing_status = ProcessingStatus.uploaded
|
||||
recording.processing_error = None
|
||||
|
||||
us.status = UploadSessionStatus.completed
|
||||
us.completed_asset_id = asset.id
|
||||
|
||||
# Clean up chunk parts (the assembled object is the source of truth).
|
||||
for c in chunks: # best-effort cleanup of chunk parts
|
||||
with contextlib.suppress(Exception):
|
||||
await storage.delete(c.storage_key)
|
||||
await session.flush()
|
||||
from shonar.services import processing as _processing
|
||||
|
||||
# New audio on disk: queue whatever AI stages apply (none configured =
|
||||
# ai_disabled, never an error).
|
||||
await _processing.enqueue_for_recording(session, recording)
|
||||
return us, recording, asset
|
||||
|
||||
|
||||
async def abort(session: AsyncSession, user_id: uuid.UUID, session_id: uuid.UUID) -> None:
|
||||
us = await get_owned_session(session, user_id, session_id)
|
||||
if us.status in (UploadSessionStatus.completed, UploadSessionStatus.aborted):
|
||||
return
|
||||
storage = get_storage()
|
||||
chunks = list(
|
||||
await session.scalars(select(UploadChunk).where(UploadChunk.session_id == us.id))
|
||||
)
|
||||
for c in chunks: # best-effort cleanup
|
||||
with contextlib.suppress(Exception):
|
||||
await storage.delete(c.storage_key)
|
||||
us.status = UploadSessionStatus.aborted
|
||||
Loading…
Add table
Add a link
Reference in a new issue