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
16
backend/README.md
Normal file
16
backend/README.md
Normal file
|
|
@ -0,0 +1,16 @@
|
|||
# S.H.O.N.A.R. backend
|
||||
|
||||
FastAPI + PostgreSQL backend for the S.H.O.N.A.R. Android app.
|
||||
|
||||
See the repository root [README](../README.md) and `docs/` for full documentation.
|
||||
|
||||
## Quick start (development)
|
||||
|
||||
```bash
|
||||
cd backend
|
||||
uv venv .venv && uv pip install -e ".[dev]"
|
||||
cp ../.env.example .env # then edit SHONAR_SECRET_KEY etc.
|
||||
uvicorn shonar.main:app --reload --port 8000
|
||||
```
|
||||
|
||||
Tests: `pytest` · Lint: `ruff check .` · Migrations: `alembic upgrade head`
|
||||
37
backend/alembic.ini
Normal file
37
backend/alembic.ini
Normal file
|
|
@ -0,0 +1,37 @@
|
|||
[alembic]
|
||||
script_location = migrations
|
||||
prepend_sys_path = .
|
||||
# URL is injected from settings in env.py; this is a placeholder.
|
||||
sqlalchemy.url =
|
||||
|
||||
[loggers]
|
||||
keys = root,sqlalchemy,alembic
|
||||
|
||||
[handlers]
|
||||
keys = console
|
||||
|
||||
[formatters]
|
||||
keys = generic
|
||||
|
||||
[logger_root]
|
||||
level = WARN
|
||||
handlers = console
|
||||
|
||||
[logger_sqlalchemy]
|
||||
level = WARN
|
||||
handlers =
|
||||
qualname = sqlalchemy.engine
|
||||
|
||||
[logger_alembic]
|
||||
level = INFO
|
||||
handlers =
|
||||
qualname = alembic
|
||||
|
||||
[handler_console]
|
||||
class = StreamHandler
|
||||
args = (sys.stderr,)
|
||||
level = NOTSET
|
||||
formatter = generic
|
||||
|
||||
[formatter_generic]
|
||||
format = %(levelname)-5.5s [%(name)s] %(message)s
|
||||
62
backend/migrations/env.py
Normal file
62
backend/migrations/env.py
Normal file
|
|
@ -0,0 +1,62 @@
|
|||
"""Alembic environment (async)."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
from logging.config import fileConfig
|
||||
|
||||
from alembic import context
|
||||
from sqlalchemy import pool
|
||||
from sqlalchemy.engine import Connection
|
||||
from sqlalchemy.ext.asyncio import async_engine_from_config
|
||||
|
||||
from shonar.core.config import get_settings
|
||||
from shonar.db import models # noqa: F401 (register tables)
|
||||
from shonar.db.base import Base
|
||||
|
||||
config = context.config
|
||||
if config.config_file_name is not None:
|
||||
fileConfig(config.config_file_name)
|
||||
|
||||
settings = get_settings()
|
||||
config.set_main_option("sqlalchemy.url", settings.database_url)
|
||||
|
||||
target_metadata = Base.metadata
|
||||
|
||||
|
||||
def run_migrations_offline() -> None:
|
||||
context.configure(
|
||||
url=settings.database_url,
|
||||
target_metadata=target_metadata,
|
||||
literal_binds=True,
|
||||
dialect_opts={"paramstyle": "named"},
|
||||
)
|
||||
with context.begin_transaction():
|
||||
context.run_migrations()
|
||||
|
||||
|
||||
def do_run_migrations(connection: Connection) -> None:
|
||||
context.configure(connection=connection, target_metadata=target_metadata)
|
||||
with context.begin_transaction():
|
||||
context.run_migrations()
|
||||
|
||||
|
||||
async def run_async_migrations() -> None:
|
||||
connectable = async_engine_from_config(
|
||||
config.get_section(config.config_ini_section, {}),
|
||||
prefix="sqlalchemy.",
|
||||
poolclass=pool.NullPool,
|
||||
)
|
||||
async with connectable.connect() as connection:
|
||||
await connection.run_sync(do_run_migrations)
|
||||
await connectable.dispose()
|
||||
|
||||
|
||||
def run_migrations_online() -> None:
|
||||
asyncio.run(run_async_migrations())
|
||||
|
||||
|
||||
if context.is_offline_mode():
|
||||
run_migrations_offline()
|
||||
else:
|
||||
run_migrations_online()
|
||||
26
backend/migrations/script.py.mako
Normal file
26
backend/migrations/script.py.mako
Normal file
|
|
@ -0,0 +1,26 @@
|
|||
"""${message}
|
||||
|
||||
Revision ID: ${up_revision}
|
||||
Revises: ${down_revision | comma,n}
|
||||
Create Date: ${create_date}
|
||||
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Sequence
|
||||
|
||||
import sqlalchemy as sa
|
||||
from alembic import op
|
||||
${imports if imports else ""}
|
||||
revision: str = ${repr(up_revision)}
|
||||
down_revision: str | None = ${repr(down_revision)}
|
||||
branch_labels: str | Sequence[str] | None = ${repr(branch_labels)}
|
||||
depends_on: str | Sequence[str] | None = ${repr(depends_on)}
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
${upgrades if upgrades else "pass"}
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
${downgrades if downgrades else "pass"}
|
||||
281
backend/migrations/versions/8d51af959ae4_initial_schema.py
Normal file
281
backend/migrations/versions/8d51af959ae4_initial_schema.py
Normal file
|
|
@ -0,0 +1,281 @@
|
|||
"""initial schema
|
||||
|
||||
Revision ID: 8d51af959ae4
|
||||
Revises:
|
||||
Create Date: 2026-09-08 13:07:10.455837
|
||||
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Sequence
|
||||
|
||||
import sqlalchemy as sa
|
||||
from alembic import op
|
||||
from sqlalchemy.dialects import postgresql
|
||||
|
||||
revision: str = '8d51af959ae4'
|
||||
down_revision: str | None = None
|
||||
branch_labels: str | Sequence[str] | None = None
|
||||
depends_on: str | Sequence[str] | None = None
|
||||
|
||||
|
||||
def _json_type() -> sa.types.TypeEngine:
|
||||
"""JSONB on PostgreSQL, plain JSON on SQLite (desktop bundled engine)."""
|
||||
if op.get_context().dialect.name == "postgresql":
|
||||
return postgresql.JSONB(astext_type=sa.Text())
|
||||
return sa.JSON()
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
# ### commands auto generated by Alembic - please adjust! ###
|
||||
op.create_table('users',
|
||||
sa.Column('email', sa.String(length=320), nullable=False),
|
||||
sa.Column('password_hash', sa.String(length=255), nullable=False),
|
||||
sa.Column('display_name', sa.String(length=120), nullable=True),
|
||||
sa.Column('is_active', sa.Boolean(), nullable=False),
|
||||
sa.Column('location_storage_enabled', sa.Boolean(), nullable=False),
|
||||
sa.Column('deleted_at', sa.DateTime(timezone=True), nullable=True),
|
||||
sa.Column('id', sa.Uuid(), nullable=False),
|
||||
sa.Column('created_at', sa.DateTime(timezone=True), nullable=False),
|
||||
sa.Column('updated_at', sa.DateTime(timezone=True), nullable=False),
|
||||
sa.PrimaryKeyConstraint('id')
|
||||
)
|
||||
op.create_index(op.f('ix_users_email'), 'users', ['email'], unique=True)
|
||||
op.create_table('devices',
|
||||
sa.Column('user_id', sa.Uuid(), nullable=False),
|
||||
sa.Column('name', sa.String(length=120), nullable=False),
|
||||
sa.Column('platform', sa.String(length=40), nullable=False),
|
||||
sa.Column('last_seen_at', sa.DateTime(timezone=True), nullable=False),
|
||||
sa.Column('revoked_at', sa.DateTime(timezone=True), nullable=True),
|
||||
sa.Column('id', sa.Uuid(), nullable=False),
|
||||
sa.Column('created_at', sa.DateTime(timezone=True), nullable=False),
|
||||
sa.Column('updated_at', sa.DateTime(timezone=True), nullable=False),
|
||||
sa.ForeignKeyConstraint(['user_id'], ['users.id'], ondelete='CASCADE'),
|
||||
sa.PrimaryKeyConstraint('id')
|
||||
)
|
||||
op.create_index(op.f('ix_devices_user_id'), 'devices', ['user_id'], unique=False)
|
||||
op.create_table('tags',
|
||||
sa.Column('user_id', sa.Uuid(), nullable=False),
|
||||
sa.Column('name', sa.String(length=80), nullable=False),
|
||||
sa.Column('id', sa.Uuid(), nullable=False),
|
||||
sa.Column('created_at', sa.DateTime(timezone=True), nullable=False),
|
||||
sa.Column('updated_at', sa.DateTime(timezone=True), nullable=False),
|
||||
sa.ForeignKeyConstraint(['user_id'], ['users.id'], ondelete='CASCADE'),
|
||||
sa.PrimaryKeyConstraint('id'),
|
||||
sa.UniqueConstraint('user_id', 'name')
|
||||
)
|
||||
op.create_index(op.f('ix_tags_user_id'), 'tags', ['user_id'], unique=False)
|
||||
op.create_table('recordings',
|
||||
sa.Column('user_id', sa.Uuid(), nullable=False),
|
||||
sa.Column('device_id', sa.Uuid(), nullable=True),
|
||||
sa.Column('client_recording_id', sa.String(length=64), nullable=True),
|
||||
sa.Column('title', sa.String(length=300), nullable=False),
|
||||
sa.Column('recorded_at', sa.DateTime(timezone=True), nullable=False),
|
||||
sa.Column('duration_seconds', sa.Float(), nullable=False),
|
||||
sa.Column('notes', sa.Text(), nullable=True),
|
||||
sa.Column('latitude', sa.Float(), nullable=True),
|
||||
sa.Column('longitude', sa.Float(), nullable=True),
|
||||
sa.Column('location_accuracy_m', sa.Float(), nullable=True),
|
||||
sa.Column('processing_status', sa.Enum('pending_upload', 'uploaded', 'queued', 'processing', 'completed', 'failed', 'ai_disabled', name='processing_status'), nullable=False),
|
||||
sa.Column('processing_error', sa.Text(), nullable=True),
|
||||
sa.Column('deleted_at', sa.DateTime(timezone=True), nullable=True),
|
||||
sa.Column('id', sa.Uuid(), nullable=False),
|
||||
sa.Column('created_at', sa.DateTime(timezone=True), nullable=False),
|
||||
sa.Column('updated_at', sa.DateTime(timezone=True), nullable=False),
|
||||
sa.ForeignKeyConstraint(['device_id'], ['devices.id'], ondelete='SET NULL'),
|
||||
sa.ForeignKeyConstraint(['user_id'], ['users.id'], ondelete='CASCADE'),
|
||||
sa.PrimaryKeyConstraint('id')
|
||||
)
|
||||
op.create_index(op.f('ix_recordings_client_recording_id'), 'recordings', ['client_recording_id'], unique=False)
|
||||
op.create_index(op.f('ix_recordings_processing_status'), 'recordings', ['processing_status'], unique=False)
|
||||
op.create_index('ix_recordings_user_client_id', 'recordings', ['user_id', 'client_recording_id'], unique=True, postgresql_where='client_recording_id IS NOT NULL', sqlite_where=sa.text('client_recording_id IS NOT NULL'))
|
||||
op.create_index(op.f('ix_recordings_user_id'), 'recordings', ['user_id'], unique=False)
|
||||
op.create_index('ix_recordings_user_recorded', 'recordings', ['user_id', 'recorded_at'], unique=False)
|
||||
op.create_table('refresh_tokens',
|
||||
sa.Column('user_id', sa.Uuid(), nullable=False),
|
||||
sa.Column('token_hash', sa.String(length=128), nullable=False),
|
||||
sa.Column('family', sa.Uuid(), nullable=False),
|
||||
sa.Column('device_id', sa.Uuid(), nullable=True),
|
||||
sa.Column('expires_at', sa.DateTime(timezone=True), nullable=False),
|
||||
sa.Column('revoked_at', sa.DateTime(timezone=True), nullable=True),
|
||||
sa.Column('replaced_by', sa.Uuid(), nullable=True),
|
||||
sa.Column('id', sa.Uuid(), nullable=False),
|
||||
sa.Column('created_at', sa.DateTime(timezone=True), nullable=False),
|
||||
sa.Column('updated_at', sa.DateTime(timezone=True), nullable=False),
|
||||
sa.ForeignKeyConstraint(['device_id'], ['devices.id'], ondelete='SET NULL'),
|
||||
sa.ForeignKeyConstraint(['replaced_by'], ['refresh_tokens.id'], ondelete='SET NULL'),
|
||||
sa.ForeignKeyConstraint(['user_id'], ['users.id'], ondelete='CASCADE'),
|
||||
sa.PrimaryKeyConstraint('id'),
|
||||
sa.UniqueConstraint('token_hash')
|
||||
)
|
||||
op.create_index(op.f('ix_refresh_tokens_family'), 'refresh_tokens', ['family'], unique=False)
|
||||
op.create_index(op.f('ix_refresh_tokens_user_id'), 'refresh_tokens', ['user_id'], unique=False)
|
||||
op.create_table('assets',
|
||||
sa.Column('recording_id', sa.Uuid(), nullable=True),
|
||||
sa.Column('user_id', sa.Uuid(), nullable=False),
|
||||
sa.Column('kind', sa.Enum('original', 'normalized', 'export', name='asset_kind'), nullable=False),
|
||||
sa.Column('storage_key', sa.String(length=500), nullable=False),
|
||||
sa.Column('mime_type', sa.String(length=100), nullable=False),
|
||||
sa.Column('size_bytes', sa.BigInteger(), nullable=False),
|
||||
sa.Column('checksum_sha256', sa.String(length=64), nullable=False),
|
||||
sa.Column('id', sa.Uuid(), nullable=False),
|
||||
sa.Column('created_at', sa.DateTime(timezone=True), nullable=False),
|
||||
sa.Column('updated_at', sa.DateTime(timezone=True), nullable=False),
|
||||
sa.ForeignKeyConstraint(['recording_id'], ['recordings.id'], ondelete='CASCADE'),
|
||||
sa.ForeignKeyConstraint(['user_id'], ['users.id'], ondelete='CASCADE'),
|
||||
sa.PrimaryKeyConstraint('id')
|
||||
)
|
||||
op.create_index(op.f('ix_assets_recording_id'), 'assets', ['recording_id'], unique=False)
|
||||
op.create_index(op.f('ix_assets_user_id'), 'assets', ['user_id'], unique=False)
|
||||
op.create_index('uq_assets_one_original', 'assets', ['recording_id'], unique=True, postgresql_where="kind = 'original'", sqlite_where=sa.text("kind = 'original'"))
|
||||
op.create_table('processing_jobs',
|
||||
sa.Column('recording_id', sa.Uuid(), nullable=False),
|
||||
sa.Column('job_type', sa.Enum('normalize_audio', 'transcribe', 'summarize', name='job_type'), nullable=False),
|
||||
sa.Column('status', sa.Enum('queued', 'running', 'succeeded', 'failed', 'skipped', name='job_status'), nullable=False),
|
||||
sa.Column('attempt', sa.Integer(), nullable=False),
|
||||
sa.Column('max_attempts', sa.Integer(), nullable=False),
|
||||
sa.Column('error', sa.Text(), nullable=True),
|
||||
sa.Column('started_at', sa.DateTime(timezone=True), nullable=True),
|
||||
sa.Column('finished_at', sa.DateTime(timezone=True), nullable=True),
|
||||
sa.Column('task_handle', sa.String(length=120), nullable=True),
|
||||
sa.Column('id', sa.Uuid(), nullable=False),
|
||||
sa.Column('created_at', sa.DateTime(timezone=True), nullable=False),
|
||||
sa.Column('updated_at', sa.DateTime(timezone=True), nullable=False),
|
||||
sa.ForeignKeyConstraint(['recording_id'], ['recordings.id'], ondelete='CASCADE'),
|
||||
sa.PrimaryKeyConstraint('id')
|
||||
)
|
||||
op.create_index(op.f('ix_processing_jobs_recording_id'), 'processing_jobs', ['recording_id'], unique=False)
|
||||
op.create_index('ix_processing_jobs_recording_type', 'processing_jobs', ['recording_id', 'job_type'], unique=False)
|
||||
op.create_index(op.f('ix_processing_jobs_status'), 'processing_jobs', ['status'], unique=False)
|
||||
op.create_table('recording_tags',
|
||||
sa.Column('recording_id', sa.Uuid(), nullable=False),
|
||||
sa.Column('tag_id', sa.Uuid(), nullable=False),
|
||||
sa.ForeignKeyConstraint(['recording_id'], ['recordings.id'], ondelete='CASCADE'),
|
||||
sa.ForeignKeyConstraint(['tag_id'], ['tags.id'], ondelete='CASCADE'),
|
||||
sa.PrimaryKeyConstraint('recording_id', 'tag_id')
|
||||
)
|
||||
op.create_table('summaries',
|
||||
sa.Column('recording_id', sa.Uuid(), nullable=False),
|
||||
sa.Column('version', sa.Integer(), nullable=False),
|
||||
sa.Column('superseded_at', sa.DateTime(timezone=True), nullable=True),
|
||||
sa.Column('provider', sa.String(length=80), nullable=False),
|
||||
sa.Column('model', sa.String(length=120), nullable=True),
|
||||
sa.Column('content', _json_type(), nullable=False),
|
||||
sa.Column('edited_by_user', sa.Boolean(), nullable=False),
|
||||
sa.Column('id', sa.Uuid(), nullable=False),
|
||||
sa.Column('created_at', sa.DateTime(timezone=True), nullable=False),
|
||||
sa.Column('updated_at', sa.DateTime(timezone=True), nullable=False),
|
||||
sa.ForeignKeyConstraint(['recording_id'], ['recordings.id'], ondelete='CASCADE'),
|
||||
sa.PrimaryKeyConstraint('id')
|
||||
)
|
||||
op.create_index(op.f('ix_summaries_recording_id'), 'summaries', ['recording_id'], unique=False)
|
||||
op.create_table('transcripts',
|
||||
sa.Column('recording_id', sa.Uuid(), nullable=False),
|
||||
sa.Column('version', sa.Integer(), nullable=False),
|
||||
sa.Column('superseded_at', sa.DateTime(timezone=True), nullable=True),
|
||||
sa.Column('language', sa.String(length=16), nullable=True),
|
||||
sa.Column('provider', sa.String(length=80), nullable=False),
|
||||
sa.Column('model', sa.String(length=120), nullable=True),
|
||||
sa.Column('text', sa.Text(), nullable=False),
|
||||
sa.Column('segments', _json_type(), nullable=True),
|
||||
sa.Column('edited_by_user', sa.Boolean(), nullable=False),
|
||||
sa.Column('id', sa.Uuid(), nullable=False),
|
||||
sa.Column('created_at', sa.DateTime(timezone=True), nullable=False),
|
||||
sa.Column('updated_at', sa.DateTime(timezone=True), nullable=False),
|
||||
sa.ForeignKeyConstraint(['recording_id'], ['recordings.id'], ondelete='CASCADE'),
|
||||
sa.PrimaryKeyConstraint('id')
|
||||
)
|
||||
op.create_index(op.f('ix_transcripts_recording_id'), 'transcripts', ['recording_id'], unique=False)
|
||||
op.create_index('ix_transcripts_recording_version', 'transcripts', ['recording_id', 'version'], unique=False)
|
||||
op.create_table('export_jobs',
|
||||
sa.Column('user_id', sa.Uuid(), nullable=False),
|
||||
sa.Column('recording_id', sa.Uuid(), nullable=False),
|
||||
sa.Column('export_type', sa.String(length=40), nullable=False),
|
||||
sa.Column('status', sa.Enum('queued', 'running', 'succeeded', 'failed', 'skipped', name='export_job_status'), nullable=False),
|
||||
sa.Column('asset_id', sa.Uuid(), nullable=True),
|
||||
sa.Column('error', sa.Text(), nullable=True),
|
||||
sa.Column('id', sa.Uuid(), nullable=False),
|
||||
sa.Column('created_at', sa.DateTime(timezone=True), nullable=False),
|
||||
sa.Column('updated_at', sa.DateTime(timezone=True), nullable=False),
|
||||
sa.ForeignKeyConstraint(['asset_id'], ['assets.id'], ondelete='SET NULL'),
|
||||
sa.ForeignKeyConstraint(['recording_id'], ['recordings.id'], ondelete='CASCADE'),
|
||||
sa.ForeignKeyConstraint(['user_id'], ['users.id'], ondelete='CASCADE'),
|
||||
sa.PrimaryKeyConstraint('id')
|
||||
)
|
||||
op.create_index(op.f('ix_export_jobs_recording_id'), 'export_jobs', ['recording_id'], unique=False)
|
||||
op.create_index(op.f('ix_export_jobs_user_id'), 'export_jobs', ['user_id'], unique=False)
|
||||
op.create_table('upload_sessions',
|
||||
sa.Column('user_id', sa.Uuid(), nullable=False),
|
||||
sa.Column('recording_id', sa.Uuid(), nullable=True),
|
||||
sa.Column('client_recording_id', sa.String(length=64), nullable=True),
|
||||
sa.Column('title', sa.String(length=300), nullable=True),
|
||||
sa.Column('declared_mime_type', sa.String(length=100), nullable=False),
|
||||
sa.Column('declared_size_bytes', sa.BigInteger(), nullable=False),
|
||||
sa.Column('chunk_size_bytes', sa.BigInteger(), nullable=False),
|
||||
sa.Column('status', sa.Enum('open', 'finalizing', 'completed', 'aborted', 'expired', name='upload_session_status'), nullable=False),
|
||||
sa.Column('expires_at', sa.DateTime(timezone=True), nullable=False),
|
||||
sa.Column('completed_asset_id', sa.Uuid(), nullable=True),
|
||||
sa.Column('id', sa.Uuid(), nullable=False),
|
||||
sa.Column('created_at', sa.DateTime(timezone=True), nullable=False),
|
||||
sa.Column('updated_at', sa.DateTime(timezone=True), nullable=False),
|
||||
sa.ForeignKeyConstraint(['completed_asset_id'], ['assets.id'], ondelete='SET NULL'),
|
||||
sa.ForeignKeyConstraint(['recording_id'], ['recordings.id'], ondelete='CASCADE'),
|
||||
sa.ForeignKeyConstraint(['user_id'], ['users.id'], ondelete='CASCADE'),
|
||||
sa.PrimaryKeyConstraint('id')
|
||||
)
|
||||
op.create_index(op.f('ix_upload_sessions_user_id'), 'upload_sessions', ['user_id'], unique=False)
|
||||
op.create_table('upload_chunks',
|
||||
sa.Column('id', sa.Integer(), nullable=False),
|
||||
sa.Column('session_id', sa.Uuid(), nullable=False),
|
||||
sa.Column('chunk_index', sa.Integer(), nullable=False),
|
||||
sa.Column('size_bytes', sa.BigInteger(), nullable=False),
|
||||
sa.Column('checksum_sha256', sa.String(length=64), nullable=False),
|
||||
sa.Column('storage_key', sa.String(length=500), nullable=False),
|
||||
sa.Column('created_at', sa.DateTime(timezone=True), nullable=False),
|
||||
sa.ForeignKeyConstraint(['session_id'], ['upload_sessions.id'], ondelete='CASCADE'),
|
||||
sa.PrimaryKeyConstraint('id'),
|
||||
sa.UniqueConstraint('session_id', 'chunk_index')
|
||||
)
|
||||
op.create_index(op.f('ix_upload_chunks_session_id'), 'upload_chunks', ['session_id'], unique=False)
|
||||
# ### end Alembic commands ###
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
# ### commands auto generated by Alembic - please adjust! ###
|
||||
op.drop_index(op.f('ix_upload_chunks_session_id'), table_name='upload_chunks')
|
||||
op.drop_table('upload_chunks')
|
||||
op.drop_index(op.f('ix_upload_sessions_user_id'), table_name='upload_sessions')
|
||||
op.drop_table('upload_sessions')
|
||||
op.drop_index(op.f('ix_export_jobs_user_id'), table_name='export_jobs')
|
||||
op.drop_index(op.f('ix_export_jobs_recording_id'), table_name='export_jobs')
|
||||
op.drop_table('export_jobs')
|
||||
op.drop_index('ix_transcripts_recording_version', table_name='transcripts')
|
||||
op.drop_index(op.f('ix_transcripts_recording_id'), table_name='transcripts')
|
||||
op.drop_table('transcripts')
|
||||
op.drop_index(op.f('ix_summaries_recording_id'), table_name='summaries')
|
||||
op.drop_table('summaries')
|
||||
op.drop_table('recording_tags')
|
||||
op.drop_index(op.f('ix_processing_jobs_status'), table_name='processing_jobs')
|
||||
op.drop_index('ix_processing_jobs_recording_type', table_name='processing_jobs')
|
||||
op.drop_index(op.f('ix_processing_jobs_recording_id'), table_name='processing_jobs')
|
||||
op.drop_table('processing_jobs')
|
||||
op.drop_index('uq_assets_one_original', table_name='assets', postgresql_where="kind = 'original'")
|
||||
op.drop_index(op.f('ix_assets_user_id'), table_name='assets')
|
||||
op.drop_index(op.f('ix_assets_recording_id'), table_name='assets')
|
||||
op.drop_table('assets')
|
||||
op.drop_index(op.f('ix_refresh_tokens_user_id'), table_name='refresh_tokens')
|
||||
op.drop_index(op.f('ix_refresh_tokens_family'), table_name='refresh_tokens')
|
||||
op.drop_table('refresh_tokens')
|
||||
op.drop_index('ix_recordings_user_recorded', table_name='recordings')
|
||||
op.drop_index(op.f('ix_recordings_user_id'), table_name='recordings')
|
||||
op.drop_index('ix_recordings_user_client_id', table_name='recordings', postgresql_where='client_recording_id IS NOT NULL')
|
||||
op.drop_index(op.f('ix_recordings_processing_status'), table_name='recordings')
|
||||
op.drop_index(op.f('ix_recordings_client_recording_id'), table_name='recordings')
|
||||
op.drop_table('recordings')
|
||||
op.drop_index(op.f('ix_tags_user_id'), table_name='tags')
|
||||
op.drop_table('tags')
|
||||
op.drop_index(op.f('ix_devices_user_id'), table_name='devices')
|
||||
op.drop_table('devices')
|
||||
op.drop_index(op.f('ix_users_email'), table_name='users')
|
||||
op.drop_table('users')
|
||||
# ### end Alembic commands ###
|
||||
86
backend/migrations/versions/fts0000000001_fts_columns.py
Normal file
86
backend/migrations/versions/fts0000000001_fts_columns.py
Normal file
|
|
@ -0,0 +1,86 @@
|
|||
"""full-text search columns (PostgreSQL tsvector)
|
||||
|
||||
Revision ID: fts0000000001
|
||||
Revises: 0c938663a363
|
||||
|
||||
Generated (STORED) tsvector columns + GIN indexes for search across title,
|
||||
transcript text, summary content, tags, and action items. The SearchBackend
|
||||
protocol keeps this swappable for Meilisearch/OpenSearch later.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Sequence
|
||||
|
||||
from alembic import op
|
||||
|
||||
revision: str = "fts0000000001"
|
||||
down_revision: str | None = "8d51af959ae4"
|
||||
branch_labels: str | Sequence[str] | None = None
|
||||
depends_on: str | Sequence[str] | None = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
# tsvector is PostgreSQL-only. On SQLite (desktop bundled engine) full-
|
||||
# text search is handled outside the DB (the app scans report files),
|
||||
# so this migration is a no-op.
|
||||
if op.get_context().dialect.name != "postgresql":
|
||||
return
|
||||
# recordings: title + notes + tag names are indexed; transcript/summary
|
||||
# contribute via their own tables (joined at query time).
|
||||
op.execute(
|
||||
"""
|
||||
ALTER TABLE recordings
|
||||
ADD COLUMN search_vector tsvector
|
||||
GENERATED ALWAYS AS (
|
||||
setweight(to_tsvector('simple', coalesce(title, '')), 'A') ||
|
||||
setweight(to_tsvector('simple', coalesce(notes, '')), 'B')
|
||||
) STORED;
|
||||
"""
|
||||
)
|
||||
op.execute("CREATE INDEX ix_recordings_search ON recordings USING GIN (search_vector);")
|
||||
|
||||
op.execute(
|
||||
"""
|
||||
ALTER TABLE transcripts
|
||||
ADD COLUMN search_vector tsvector
|
||||
GENERATED ALWAYS AS (
|
||||
setweight(to_tsvector('simple', coalesce(text, '')), 'B')
|
||||
) STORED;
|
||||
"""
|
||||
)
|
||||
op.execute("CREATE INDEX ix_transcripts_search ON transcripts USING GIN (search_vector);")
|
||||
|
||||
# summary content JSON -> extractive text for FTS. Generated columns
|
||||
# forbid subqueries/set-returning functions, so we index the JSON body
|
||||
# with punctuation stripped (covers key_points/decisions/action_items/
|
||||
# questions) plus weighted short/detailed fields.
|
||||
op.execute(
|
||||
"""
|
||||
ALTER TABLE summaries
|
||||
ADD COLUMN search_vector tsvector
|
||||
GENERATED ALWAYS AS (
|
||||
setweight(to_tsvector('simple', coalesce(content->>'short', '')), 'A') ||
|
||||
setweight(to_tsvector('simple', coalesce(content->>'detailed', '')), 'B') ||
|
||||
setweight(
|
||||
to_tsvector('simple',
|
||||
regexp_replace(coalesce(content::text, ''), '[\\[\\]{}"]', ' ', 'g')), 'C')
|
||||
) STORED;
|
||||
"""
|
||||
)
|
||||
op.execute("CREATE INDEX ix_summaries_search ON summaries USING GIN (search_vector);")
|
||||
|
||||
op.execute(
|
||||
"""
|
||||
ALTER TABLE tags ADD COLUMN IF NOT EXISTS search_vector tsvector
|
||||
GENERATED ALWAYS AS (to_tsvector('simple', coalesce(name, ''))) STORED;
|
||||
"""
|
||||
)
|
||||
op.execute("CREATE INDEX ix_tags_search ON tags USING GIN (search_vector);")
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
if op.get_context().dialect.name != "postgresql":
|
||||
return
|
||||
for table in ("tags", "summaries", "transcripts", "recordings"):
|
||||
op.execute(f"DROP INDEX IF EXISTS ix_{table}_search;")
|
||||
op.execute(f"ALTER TABLE {table} DROP COLUMN IF EXISTS search_vector;")
|
||||
69
backend/migrations/versions/m8models000001_model_support.py
Normal file
69
backend/migrations/versions/m8models000001_model_support.py
Normal file
|
|
@ -0,0 +1,69 @@
|
|||
"""per-recording transcription model + job progress + app settings
|
||||
|
||||
Revision ID: m8models000001
|
||||
Revises: fts0000000001
|
||||
|
||||
- recordings.transcription_model (nullable; NULL = server default)
|
||||
- processing_jobs.stage / processing_jobs.progress (display-only)
|
||||
- app_settings table (global default model, …)
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Sequence
|
||||
|
||||
import sqlalchemy as sa
|
||||
from alembic import op
|
||||
from sqlalchemy.dialects import postgresql
|
||||
|
||||
revision: str = 'm8models000001'
|
||||
down_revision: str | None = 'fts0000000001'
|
||||
branch_labels: str | Sequence[str] | None = None
|
||||
depends_on: str | Sequence[str] | None = None
|
||||
|
||||
|
||||
def _json_type() -> sa.types.TypeEngine:
|
||||
"""JSONB on PostgreSQL, plain JSON on SQLite (desktop bundled engine)."""
|
||||
if op.get_context().dialect.name == "postgresql":
|
||||
return postgresql.JSONB(astext_type=sa.Text())
|
||||
return sa.JSON()
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
op.add_column(
|
||||
"recordings",
|
||||
sa.Column("transcription_model", sa.String(length=32), nullable=True),
|
||||
)
|
||||
op.add_column(
|
||||
"upload_sessions",
|
||||
sa.Column("transcription_model", sa.String(length=32), nullable=True),
|
||||
)
|
||||
op.add_column(
|
||||
"processing_jobs",
|
||||
sa.Column("stage", sa.String(length=32), nullable=True),
|
||||
)
|
||||
op.add_column(
|
||||
"processing_jobs",
|
||||
sa.Column("progress", sa.Integer(), nullable=True),
|
||||
)
|
||||
op.create_table(
|
||||
"app_settings",
|
||||
sa.Column("key", sa.String(length=120), nullable=False),
|
||||
sa.Column(
|
||||
"value",
|
||||
_json_type(),
|
||||
nullable=False,
|
||||
),
|
||||
sa.Column("id", sa.Uuid(), nullable=False),
|
||||
sa.Column("created_at", sa.DateTime(timezone=True), nullable=False),
|
||||
sa.Column("updated_at", sa.DateTime(timezone=True), nullable=False),
|
||||
sa.PrimaryKeyConstraint("id"),
|
||||
sa.UniqueConstraint("key"),
|
||||
)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
op.drop_table("app_settings")
|
||||
op.drop_column("processing_jobs", "progress")
|
||||
op.drop_column("processing_jobs", "stage")
|
||||
op.drop_column("upload_sessions", "transcription_model")
|
||||
op.drop_column("recordings", "transcription_model")
|
||||
64
backend/pyproject.toml
Normal file
64
backend/pyproject.toml
Normal file
|
|
@ -0,0 +1,64 @@
|
|||
[project]
|
||||
name = "shonar-backend"
|
||||
version = "0.1.0"
|
||||
description = "S.H.O.N.A.R. — Self-hosted Oral Notes and Audio Recorder backend"
|
||||
readme = "README.md"
|
||||
requires-python = ">=3.11"
|
||||
license = { text = "Apache-2.0" }
|
||||
dependencies = [
|
||||
"fastapi>=0.115",
|
||||
"uvicorn[standard]>=0.30",
|
||||
"sqlalchemy[asyncio]>=2.0",
|
||||
"asyncpg>=0.29",
|
||||
"aiosqlite>=0.20", # tests only in practice; harmless dep
|
||||
"alembic>=1.13",
|
||||
"pydantic>=2.8",
|
||||
"pydantic-settings>=2.4",
|
||||
"email-validator>=2.0",
|
||||
"argon2-cffi>=23.1",
|
||||
"pyjwt>=2.9",
|
||||
"python-multipart>=0.0.9",
|
||||
"slowapi>=0.1.9",
|
||||
"redis>=5.0",
|
||||
"arq>=0.26",
|
||||
]
|
||||
|
||||
[project.optional-dependencies]
|
||||
s3 = ["boto3>=1.34"]
|
||||
faster-whisper = ["faster-whisper>=1.0"]
|
||||
dev = [
|
||||
"pytest>=8.0",
|
||||
"pytest-asyncio>=0.23",
|
||||
"httpx>=0.27",
|
||||
"ruff>=0.5",
|
||||
"mypy>=1.10",
|
||||
]
|
||||
|
||||
[build-system]
|
||||
requires = ["hatchling"]
|
||||
build-backend = "hatchling.build"
|
||||
|
||||
[tool.hatch.build.targets.wheel]
|
||||
packages = ["shonar"]
|
||||
|
||||
[tool.pytest.ini_options]
|
||||
asyncio_mode = "auto"
|
||||
asyncio_default_fixture_loop_scope = "session"
|
||||
asyncio_default_test_loop_scope = "session"
|
||||
testpaths = ["tests"]
|
||||
|
||||
[tool.ruff]
|
||||
line-length = 100
|
||||
target-version = "py311"
|
||||
exclude = ["migrations"]
|
||||
|
||||
[tool.ruff.lint]
|
||||
select = ["E", "F", "I", "UP", "B", "SIM"]
|
||||
|
||||
[tool.ruff.lint.per-file-ignores]
|
||||
# FastAPI idiom: Query(...)/Depends() as parameter defaults.
|
||||
"shonar/api/**" = ["B008"]
|
||||
|
||||
[tool.mypy]
|
||||
python_version = "3.11"
|
||||
ignore_missing_imports = true
|
||||
3
backend/shonar/__init__.py
Normal file
3
backend/shonar/__init__.py
Normal file
|
|
@ -0,0 +1,3 @@
|
|||
"""S.H.O.N.A.R. — Self-hosted Oral Notes and Audio Recorder."""
|
||||
|
||||
__version__ = "0.1.0"
|
||||
0
backend/shonar/api/__init__.py
Normal file
0
backend/shonar/api/__init__.py
Normal file
49
backend/shonar/api/deps.py
Normal file
49
backend/shonar/api/deps.py
Normal file
|
|
@ -0,0 +1,49 @@
|
|||
"""FastAPI dependencies: current user, session, rate limiting."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import uuid
|
||||
from typing import Annotated
|
||||
|
||||
from fastapi import Depends, HTTPException, Request, status
|
||||
from fastapi.security import HTTPAuthorizationCredentials, HTTPBearer
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from shonar.core.security import TokenError, decode_access_token
|
||||
from shonar.db.models import User
|
||||
from shonar.db.session import get_session
|
||||
|
||||
SessionDep = Annotated[AsyncSession, Depends(get_session)]
|
||||
|
||||
_bearer = HTTPBearer(auto_error=False)
|
||||
|
||||
|
||||
async def get_current_user(
|
||||
request: Request,
|
||||
session: SessionDep,
|
||||
creds: Annotated[HTTPAuthorizationCredentials | None, Depends(_bearer)] = None,
|
||||
) -> User:
|
||||
if creds is None or not creds.credentials:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_401_UNAUTHORIZED,
|
||||
detail="Not authenticated",
|
||||
headers={"WWW-Authenticate": "Bearer"},
|
||||
) from None
|
||||
try:
|
||||
payload = decode_access_token(creds.credentials)
|
||||
except TokenError:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_401_UNAUTHORIZED,
|
||||
detail="Invalid or expired token",
|
||||
headers={"WWW-Authenticate": "Bearer"},
|
||||
) from None
|
||||
user = await session.get(User, uuid.UUID(payload["sub"]))
|
||||
if user is None or not user.is_active or user.deleted_at is not None:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_401_UNAUTHORIZED, detail="Account unavailable"
|
||||
) from None
|
||||
request.state.user_id = user.id
|
||||
return user
|
||||
|
||||
|
||||
CurrentUser = Annotated[User, Depends(get_current_user)]
|
||||
69
backend/shonar/api/schemas_common.py
Normal file
69
backend/shonar/api/schemas_common.py
Normal file
|
|
@ -0,0 +1,69 @@
|
|||
"""Shared Pydantic schemas."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import uuid
|
||||
from datetime import datetime
|
||||
|
||||
from pydantic import BaseModel, ConfigDict, EmailStr, Field
|
||||
|
||||
|
||||
class ORMModel(BaseModel):
|
||||
model_config = ConfigDict(from_attributes=True)
|
||||
|
||||
|
||||
# --- auth ---
|
||||
|
||||
|
||||
class RegisterRequest(BaseModel):
|
||||
email: EmailStr
|
||||
password: str = Field(min_length=10, max_length=256)
|
||||
display_name: str | None = Field(default=None, max_length=120)
|
||||
|
||||
|
||||
class LoginRequest(BaseModel):
|
||||
email: EmailStr
|
||||
password: str
|
||||
device_name: str | None = Field(default=None, max_length=120)
|
||||
platform: str = Field(default="android", max_length=40)
|
||||
|
||||
|
||||
class TokenPair(BaseModel):
|
||||
access_token: str
|
||||
token_type: str = "bearer"
|
||||
expires_in: int
|
||||
refresh_token: str
|
||||
device_id: uuid.UUID | None = None
|
||||
|
||||
|
||||
class RefreshRequest(BaseModel):
|
||||
refresh_token: str
|
||||
|
||||
|
||||
class LogoutRequest(BaseModel):
|
||||
refresh_token: str
|
||||
|
||||
|
||||
class UserOut(ORMModel):
|
||||
id: uuid.UUID
|
||||
email: EmailStr
|
||||
display_name: str | None
|
||||
location_storage_enabled: bool
|
||||
created_at: datetime
|
||||
|
||||
|
||||
class UserUpdate(BaseModel):
|
||||
display_name: str | None = Field(default=None, max_length=120)
|
||||
location_storage_enabled: bool | None = None
|
||||
|
||||
|
||||
class DeviceOut(ORMModel):
|
||||
id: uuid.UUID
|
||||
name: str
|
||||
platform: str
|
||||
last_seen_at: datetime
|
||||
revoked_at: datetime | None
|
||||
|
||||
|
||||
class DeleteAccountRequest(BaseModel):
|
||||
password: str
|
||||
149
backend/shonar/api/schemas_recordings.py
Normal file
149
backend/shonar/api/schemas_recordings.py
Normal file
|
|
@ -0,0 +1,149 @@
|
|||
"""Recording / upload schemas."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import uuid
|
||||
from datetime import datetime
|
||||
|
||||
from pydantic import BaseModel, ConfigDict, Field
|
||||
|
||||
|
||||
class ORMModel(BaseModel):
|
||||
model_config = ConfigDict(from_attributes=True)
|
||||
|
||||
|
||||
class UploadSessionCreate(BaseModel):
|
||||
declared_mime_type: str = Field(max_length=100)
|
||||
declared_size_bytes: int = Field(gt=0)
|
||||
title: str | None = Field(default=None, max_length=300)
|
||||
client_recording_id: str | None = Field(default=None, max_length=64)
|
||||
# Optional per-recording transcription model override ("Use default"
|
||||
# when omitted). Validated against the model registry at finalize time.
|
||||
transcription_model: str | None = Field(default=None, max_length=32)
|
||||
|
||||
|
||||
class UploadSessionOut(ORMModel):
|
||||
id: uuid.UUID
|
||||
status: str
|
||||
chunk_size_bytes: int
|
||||
declared_mime_type: str
|
||||
declared_size_bytes: int
|
||||
expires_at: datetime
|
||||
|
||||
|
||||
class UploadStatusOut(BaseModel):
|
||||
id: uuid.UUID
|
||||
status: str
|
||||
received_chunk_indexes: list[int]
|
||||
|
||||
|
||||
class RecordingFinalize(BaseModel):
|
||||
recorded_at: datetime | None = None
|
||||
duration_seconds: float = Field(default=0.0, ge=0)
|
||||
# Location accepted ONLY if the user has location storage enabled
|
||||
# (enforced in the route; silently dropped otherwise).
|
||||
latitude: float | None = Field(default=None, ge=-90, le=90)
|
||||
longitude: float | None = Field(default=None, ge=-180, le=180)
|
||||
location_accuracy_m: float | None = Field(default=None, ge=0)
|
||||
notes: str | None = Field(default=None, max_length=20000)
|
||||
# Optional per-recording transcription model override. When both the
|
||||
# session and the finalize body specify one, the finalize body wins.
|
||||
transcription_model: str | None = Field(default=None, max_length=32)
|
||||
|
||||
|
||||
class RecordingOut(ORMModel):
|
||||
id: uuid.UUID
|
||||
title: str
|
||||
recorded_at: datetime
|
||||
duration_seconds: float
|
||||
notes: str | None
|
||||
latitude: float | None
|
||||
longitude: float | None
|
||||
processing_status: str
|
||||
processing_error: str | None
|
||||
# Effective transcription model saved at finalize time (override or the
|
||||
# global default then in force). None for pre-model rows.
|
||||
transcription_model: str | None
|
||||
has_audio: bool
|
||||
mime_type: str | None
|
||||
size_bytes: int | None
|
||||
tags: list[str]
|
||||
created_at: datetime
|
||||
updated_at: datetime
|
||||
|
||||
|
||||
class RecordingUpdate(BaseModel):
|
||||
title: str | None = Field(default=None, min_length=1, max_length=300)
|
||||
notes: str | None = Field(default=None, max_length=20000)
|
||||
tags: list[str] | None = Field(default=None, max_length=20)
|
||||
latitude: float | None = Field(default=None, ge=-90, le=90)
|
||||
longitude: float | None = Field(default=None, ge=-180, le=180)
|
||||
|
||||
|
||||
class RecordingListOut(BaseModel):
|
||||
items: list[RecordingOut]
|
||||
total: int
|
||||
limit: int
|
||||
offset: int
|
||||
|
||||
|
||||
class SegmentOut(BaseModel):
|
||||
start: float
|
||||
end: float
|
||||
text: str
|
||||
speaker: str | None = None
|
||||
|
||||
|
||||
class TranscriptOut(ORMModel):
|
||||
version: int
|
||||
language: str | None
|
||||
provider: str
|
||||
model: str | None
|
||||
text: str
|
||||
segments: list[SegmentOut]
|
||||
edited_by_user: bool
|
||||
created_at: datetime
|
||||
updated_at: datetime
|
||||
|
||||
|
||||
class SummaryOut(ORMModel):
|
||||
version: int
|
||||
provider: str
|
||||
model: str | None
|
||||
content: dict
|
||||
edited_by_user: bool
|
||||
created_at: datetime
|
||||
updated_at: datetime
|
||||
|
||||
|
||||
class ProcessingJobOut(ORMModel):
|
||||
job_type: str
|
||||
status: str
|
||||
attempt: int
|
||||
max_attempts: int
|
||||
error: str | None
|
||||
# Display-only phase within a running job ("loading-model",
|
||||
# "transcribing") plus 0-100 progress when known.
|
||||
stage: str | None
|
||||
progress: int | None
|
||||
started_at: datetime | None
|
||||
finished_at: datetime | None
|
||||
|
||||
|
||||
class TranscriptUpdate(BaseModel):
|
||||
"""User edit: stored as a new version with edited_by_user=True.
|
||||
|
||||
The AI pipeline never overwrites a newest user-edited row, so edits
|
||||
are verdicts. Segments are optional — when omitted the previous
|
||||
segments are dropped (plain-text correction).
|
||||
"""
|
||||
|
||||
text: str = Field(min_length=1, max_length=200000)
|
||||
segments: list[SegmentOut] | None = Field(default=None, max_length=2000)
|
||||
language: str | None = Field(default=None, max_length=16)
|
||||
|
||||
|
||||
class SummaryUpdate(BaseModel):
|
||||
"""User edit for the structured summary (same versioning rule)."""
|
||||
|
||||
content: dict = Field(max_length=100)
|
||||
23
backend/shonar/api/schemas_search.py
Normal file
23
backend/shonar/api/schemas_search.py
Normal file
|
|
@ -0,0 +1,23 @@
|
|||
"""Search + export schemas (M9)."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import uuid
|
||||
|
||||
from pydantic import BaseModel
|
||||
|
||||
|
||||
class SearchHitOut(BaseModel):
|
||||
id: uuid.UUID
|
||||
title: str
|
||||
field: str # title | tag | notes | summary | transcript
|
||||
snippet: str
|
||||
|
||||
|
||||
class SearchOut(BaseModel):
|
||||
query: str
|
||||
scope: str
|
||||
items: list[SearchHitOut]
|
||||
total: int
|
||||
limit: int
|
||||
offset: int
|
||||
25
backend/shonar/api/v1/__init__.py
Normal file
25
backend/shonar/api/v1/__init__.py
Normal file
|
|
@ -0,0 +1,25 @@
|
|||
"""API v1 router aggregation."""
|
||||
|
||||
from fastapi import APIRouter
|
||||
|
||||
from shonar.api.v1 import (
|
||||
auth,
|
||||
health,
|
||||
models,
|
||||
provider_info,
|
||||
recordings,
|
||||
search_exports,
|
||||
users,
|
||||
)
|
||||
|
||||
api_router = APIRouter(prefix="/api/v1")
|
||||
api_router.include_router(health.router)
|
||||
api_router.include_router(auth.router)
|
||||
api_router.include_router(users.router)
|
||||
api_router.include_router(recordings.router)
|
||||
api_router.include_router(models.router)
|
||||
api_router.include_router(provider_info.router)
|
||||
api_router.include_router(search_exports.router)
|
||||
|
||||
# Included as later milestones land:
|
||||
# - tags (M9 follow-on)
|
||||
91
backend/shonar/api/v1/auth.py
Normal file
91
backend/shonar/api/v1/auth.py
Normal file
|
|
@ -0,0 +1,91 @@
|
|||
"""Auth endpoints (rate-limited)."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from fastapi import APIRouter, HTTPException, Request
|
||||
|
||||
from shonar.api.deps import CurrentUser, SessionDep
|
||||
from shonar.api.schemas_common import (
|
||||
DeleteAccountRequest,
|
||||
LoginRequest,
|
||||
LogoutRequest,
|
||||
RefreshRequest,
|
||||
RegisterRequest,
|
||||
TokenPair,
|
||||
UserOut,
|
||||
)
|
||||
from shonar.core.config import get_settings
|
||||
from shonar.core.ratelimit import auth_limit
|
||||
from shonar.core.security import create_access_token
|
||||
from shonar.services import auth as auth_service
|
||||
from shonar.services.auth import AuthError
|
||||
|
||||
router = APIRouter(tags=["auth"])
|
||||
|
||||
|
||||
@router.post("/auth/register", response_model=TokenPair, status_code=201)
|
||||
@auth_limit
|
||||
async def register(body: RegisterRequest, request: Request, session: SessionDep):
|
||||
settings = get_settings()
|
||||
if not settings.allow_registration:
|
||||
raise HTTPException(403, "Registration is disabled on this server.") from None
|
||||
try:
|
||||
user = await auth_service.register_user(
|
||||
session, body.email, body.password, body.display_name
|
||||
)
|
||||
except AuthError as e:
|
||||
raise HTTPException(e.status_code, e.message) from None
|
||||
access, ttl = create_access_token(user.id)
|
||||
from shonar.services.auth import issue_refresh_token
|
||||
|
||||
refresh, _ = await issue_refresh_token(session, user.id, None, None)
|
||||
return TokenPair(access_token=access, expires_in=ttl, refresh_token=refresh)
|
||||
|
||||
|
||||
@router.post("/auth/login", response_model=TokenPair)
|
||||
@auth_limit
|
||||
async def login(body: LoginRequest, request: Request, session: SessionDep):
|
||||
try:
|
||||
user, refresh, device = await auth_service.login(
|
||||
session, body.email, body.password, body.device_name, body.platform
|
||||
)
|
||||
except AuthError as e:
|
||||
raise HTTPException(e.status_code, e.message) from None
|
||||
access, ttl = create_access_token(user.id, device.id)
|
||||
return TokenPair(
|
||||
access_token=access, expires_in=ttl, refresh_token=refresh, device_id=device.id
|
||||
)
|
||||
|
||||
|
||||
@router.post("/auth/refresh", response_model=TokenPair)
|
||||
@auth_limit
|
||||
async def refresh(body: RefreshRequest, request: Request, session: SessionDep):
|
||||
try:
|
||||
user, new_refresh, device_id = await auth_service.rotate_refresh_token(
|
||||
session, body.refresh_token
|
||||
)
|
||||
except AuthError as e:
|
||||
raise HTTPException(e.status_code, e.message) from None
|
||||
access, ttl = create_access_token(user.id, device_id)
|
||||
return TokenPair(
|
||||
access_token=access, expires_in=ttl, refresh_token=new_refresh, device_id=device_id
|
||||
)
|
||||
|
||||
|
||||
@router.post("/auth/logout", status_code=204)
|
||||
async def logout(body: LogoutRequest, session: SessionDep):
|
||||
await auth_service.logout(session, body.refresh_token)
|
||||
|
||||
|
||||
@router.get("/auth/me", response_model=UserOut)
|
||||
async def me(user: CurrentUser):
|
||||
return user
|
||||
|
||||
|
||||
@router.post("/auth/delete-account", status_code=202)
|
||||
async def delete_account(body: DeleteAccountRequest, user: CurrentUser, session: SessionDep):
|
||||
try:
|
||||
await auth_service.delete_account(session, user, body.password)
|
||||
except AuthError as e:
|
||||
raise HTTPException(e.status_code, e.message) from None
|
||||
return {"detail": "Account scheduled for deletion.", "grace_days": 30}
|
||||
65
backend/shonar/api/v1/health.py
Normal file
65
backend/shonar/api/v1/health.py
Normal file
|
|
@ -0,0 +1,65 @@
|
|||
"""Health and system endpoints."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import time
|
||||
|
||||
from fastapi import APIRouter
|
||||
from sqlalchemy import text
|
||||
|
||||
from shonar.api.deps import SessionDep
|
||||
from shonar.core.config import get_settings
|
||||
|
||||
router = APIRouter(tags=["health"])
|
||||
|
||||
_STARTED = time.monotonic()
|
||||
|
||||
|
||||
@router.get("/healthz")
|
||||
async def healthz() -> dict:
|
||||
"""Liveness: process up. No auth, no dependencies."""
|
||||
return {"status": "ok", "uptime_seconds": round(time.monotonic() - _STARTED, 1)}
|
||||
|
||||
|
||||
@router.get("/readyz")
|
||||
async def readyz(session: SessionDep) -> dict:
|
||||
"""Readiness: database reachable. Config warnings surfaced for admins
|
||||
via /api/v1/system/status instead of failing readiness."""
|
||||
try:
|
||||
await session.execute(text("SELECT 1"))
|
||||
db_ok = True
|
||||
except Exception:
|
||||
db_ok = False
|
||||
status_code_body = {"status": "ok" if db_ok else "degraded", "database": db_ok}
|
||||
return status_code_body
|
||||
|
||||
|
||||
@router.get("/system/status")
|
||||
async def system_status() -> dict:
|
||||
"""Public-ish status: which AI features are enabled (never any secrets).
|
||||
The app uses this to show honest AI-processing state to users."""
|
||||
settings = get_settings()
|
||||
return {
|
||||
"app": settings.app_name,
|
||||
"registration_enabled": settings.allow_registration,
|
||||
"ai": {
|
||||
"transcription_provider": settings.transcription_provider,
|
||||
"transcription_enabled": settings.transcription_provider != "none",
|
||||
"llm_provider": settings.llm_provider,
|
||||
"llm_enabled": settings.llm_provider != "none",
|
||||
# Explicit disclosure: are any external (non-local) calls made?
|
||||
"external_ai_in_use": (
|
||||
settings.transcription_provider == "whisper_http"
|
||||
and "localhost" not in settings.transcription_base_url
|
||||
and "127.0.0.1" not in settings.transcription_base_url
|
||||
)
|
||||
or (
|
||||
settings.llm_provider == "openai_compat"
|
||||
and "localhost" not in settings.llm_base_url
|
||||
and "127.0.0.1" not in settings.llm_base_url
|
||||
),
|
||||
},
|
||||
"storage_backend": settings.storage_backend,
|
||||
"audio_conversion_enabled": settings.audio_conversion_enabled,
|
||||
"config_warnings": settings.validate_production(),
|
||||
}
|
||||
132
backend/shonar/api/v1/models.py
Normal file
132
backend/shonar/api/v1/models.py
Normal file
|
|
@ -0,0 +1,132 @@
|
|||
"""Transcription model registry + global default (Stage 1).
|
||||
|
||||
- GET /models — supported models with display metadata, download and
|
||||
availability status, and which is the default.
|
||||
- GET /models/default — the current global default model name.
|
||||
- PUT /models/default — change the global default (affects future
|
||||
recordings only; saved per-recording models are never rewritten).
|
||||
- POST /models/{name}/download — fetch a model into the local cache
|
||||
(needs internet once; runs synchronously and may take minutes for
|
||||
large models).
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
|
||||
from fastapi import APIRouter, HTTPException
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
from shonar.api.deps import CurrentUser, SessionDep
|
||||
from shonar.services.ai import ProviderConfigError
|
||||
from shonar.services.ai.model_registry import (
|
||||
SUPPORTED_TRANSCRIPTION_MODELS,
|
||||
get_global_default_model,
|
||||
is_faster_whisper_installed,
|
||||
is_model_downloaded,
|
||||
set_global_default_model,
|
||||
validate_model_name,
|
||||
)
|
||||
|
||||
router = APIRouter(tags=["models"])
|
||||
|
||||
|
||||
class TranscriptionModelOut(BaseModel):
|
||||
name: str
|
||||
display_name: str
|
||||
description: str
|
||||
params: str
|
||||
approx_memory: str
|
||||
relative_speed: str
|
||||
is_default: bool
|
||||
downloaded: bool
|
||||
available: bool
|
||||
|
||||
|
||||
class ModelsOut(BaseModel):
|
||||
default_model: str
|
||||
faster_whisper_installed: bool
|
||||
models: list[TranscriptionModelOut]
|
||||
|
||||
|
||||
class DefaultModelUpdate(BaseModel):
|
||||
model: str = Field(min_length=1, max_length=32)
|
||||
|
||||
|
||||
class DefaultModelOut(BaseModel):
|
||||
default_model: str
|
||||
|
||||
|
||||
async def _models_out(session: SessionDep) -> ModelsOut:
|
||||
default = await get_global_default_model(session)
|
||||
installed = is_faster_whisper_installed()
|
||||
return ModelsOut(
|
||||
default_model=default,
|
||||
faster_whisper_installed=installed,
|
||||
models=[
|
||||
TranscriptionModelOut(
|
||||
name=info.name,
|
||||
display_name=info.display_name,
|
||||
description=info.description,
|
||||
params=info.params,
|
||||
approx_memory=info.approx_memory,
|
||||
relative_speed=info.relative_speed,
|
||||
is_default=info.name == default,
|
||||
downloaded=is_model_downloaded(info.name),
|
||||
available=installed and is_model_downloaded(info.name),
|
||||
)
|
||||
for info in SUPPORTED_TRANSCRIPTION_MODELS.values()
|
||||
],
|
||||
)
|
||||
|
||||
|
||||
@router.get("/models", response_model=ModelsOut)
|
||||
async def list_models(user: CurrentUser, session: SessionDep):
|
||||
return await _models_out(session)
|
||||
|
||||
|
||||
@router.get("/models/default", response_model=DefaultModelOut)
|
||||
async def get_default_model(user: CurrentUser, session: SessionDep):
|
||||
return DefaultModelOut(default_model=await get_global_default_model(session))
|
||||
|
||||
|
||||
@router.put("/models/default", response_model=DefaultModelOut)
|
||||
async def put_default_model(body: DefaultModelUpdate, user: CurrentUser, session: SessionDep):
|
||||
try:
|
||||
name = await set_global_default_model(session, body.model)
|
||||
except ProviderConfigError as e:
|
||||
raise HTTPException(422, str(e)) from None
|
||||
return DefaultModelOut(default_model=name)
|
||||
|
||||
|
||||
@router.post("/models/{name}/download", response_model=TranscriptionModelOut)
|
||||
async def download_model(name: str, user: CurrentUser, session: SessionDep):
|
||||
try:
|
||||
clean = validate_model_name(name)
|
||||
except ProviderConfigError as e:
|
||||
raise HTTPException(422, str(e)) from None
|
||||
if not is_faster_whisper_installed():
|
||||
raise HTTPException(
|
||||
501,
|
||||
"faster-whisper is not installed on this server "
|
||||
"(pip install shonar-backend[faster-whisper]).",
|
||||
)
|
||||
from shonar.services.ai.faster_whisper import _load_model
|
||||
|
||||
try:
|
||||
await asyncio.to_thread(_load_model, clean)
|
||||
except Exception as e: # noqa: BLE001 — surface download failures plainly
|
||||
raise HTTPException(500, f"Model download failed: {e}") from None
|
||||
default = await get_global_default_model(session)
|
||||
info = SUPPORTED_TRANSCRIPTION_MODELS[clean]
|
||||
return TranscriptionModelOut(
|
||||
name=info.name,
|
||||
display_name=info.display_name,
|
||||
description=info.description,
|
||||
params=info.params,
|
||||
approx_memory=info.approx_memory,
|
||||
relative_speed=info.relative_speed,
|
||||
is_default=info.name == default,
|
||||
downloaded=is_model_downloaded(clean),
|
||||
available=is_model_downloaded(clean),
|
||||
)
|
||||
38
backend/shonar/api/v1/provider_info.py
Normal file
38
backend/shonar/api/v1/provider_info.py
Normal file
|
|
@ -0,0 +1,38 @@
|
|||
"""Public provider handshake.
|
||||
|
||||
`GET /api/v1/provider-info` lets SHONAR clients (and platform probes on
|
||||
Start9/Umbrel) identify this server as SHONAR-compatible and learn its
|
||||
capabilities BEFORE any credentials are involved. Returns static deployment
|
||||
metadata only — never user data, never configuration secrets.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from fastapi import APIRouter
|
||||
|
||||
from shonar import __version__
|
||||
from shonar.core.config import get_settings
|
||||
|
||||
router = APIRouter(tags=["provider"])
|
||||
|
||||
|
||||
@router.get("/provider-info")
|
||||
async def provider_info() -> dict:
|
||||
"""Static, unauthenticated deployment identity for client probing."""
|
||||
settings = get_settings()
|
||||
# Capability flags reflect what this deployment can actually do right now.
|
||||
transcription = settings.transcription_provider != "none"
|
||||
llm = settings.llm_provider != "none"
|
||||
return {
|
||||
"kind": "shonar",
|
||||
"version": __version__,
|
||||
"api_version": "v1",
|
||||
"capabilities": {
|
||||
"chunked_upload": True,
|
||||
"server_transcription": transcription,
|
||||
"server_summary": llm,
|
||||
"account_deletion": True,
|
||||
},
|
||||
# Storage backend family (never a path): "local" | "s3"
|
||||
"storage_backend": settings.storage_backend,
|
||||
}
|
||||
528
backend/shonar/api/v1/recordings.py
Normal file
528
backend/shonar/api/v1/recordings.py
Normal file
|
|
@ -0,0 +1,528 @@
|
|||
"""Upload sessions + recordings CRUD.
|
||||
|
||||
Every endpoint enforces ownership server-side; missing/foreign resources
|
||||
return 404 (no existence leaks). Storage keys are never exposed.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import contextlib
|
||||
import uuid
|
||||
from datetime import datetime
|
||||
|
||||
from fastapi import APIRouter, Header, HTTPException, Query, Request, Response
|
||||
from sqlalchemy import func, select
|
||||
|
||||
from shonar.api.deps import CurrentUser, SessionDep
|
||||
from shonar.api.schemas_recordings import (
|
||||
ProcessingJobOut,
|
||||
RecordingFinalize,
|
||||
RecordingListOut,
|
||||
RecordingOut,
|
||||
RecordingUpdate,
|
||||
SummaryOut,
|
||||
SummaryUpdate,
|
||||
TranscriptOut,
|
||||
TranscriptUpdate,
|
||||
UploadSessionCreate,
|
||||
UploadSessionOut,
|
||||
UploadStatusOut,
|
||||
)
|
||||
from shonar.db.models import (
|
||||
Asset,
|
||||
AssetKind,
|
||||
ProcessingJob,
|
||||
Recording,
|
||||
RecordingTag,
|
||||
Summary,
|
||||
Tag,
|
||||
Transcript,
|
||||
utcnow,
|
||||
)
|
||||
from shonar.services import uploads as up
|
||||
|
||||
router = APIRouter(tags=["uploads", "recordings"])
|
||||
|
||||
|
||||
async def _recording_out(session, rec: Recording) -> RecordingOut: # noqa: ANN001
|
||||
original = await session.scalar(
|
||||
select(Asset).where(Asset.recording_id == rec.id, Asset.kind == AssetKind.original)
|
||||
)
|
||||
tags = await session.scalars(
|
||||
select(Tag.name)
|
||||
.join(RecordingTag, RecordingTag.tag_id == Tag.id)
|
||||
.where(RecordingTag.recording_id == rec.id)
|
||||
.order_by(Tag.name)
|
||||
)
|
||||
return RecordingOut(
|
||||
id=rec.id,
|
||||
title=rec.title,
|
||||
recorded_at=rec.recorded_at,
|
||||
duration_seconds=rec.duration_seconds,
|
||||
notes=rec.notes,
|
||||
latitude=rec.latitude,
|
||||
longitude=rec.longitude,
|
||||
processing_status=rec.processing_status.value,
|
||||
processing_error=rec.processing_error,
|
||||
transcription_model=rec.transcription_model,
|
||||
has_audio=original is not None,
|
||||
mime_type=original.mime_type if original else None,
|
||||
size_bytes=original.size_bytes if original else None,
|
||||
tags=list(tags),
|
||||
created_at=rec.created_at,
|
||||
updated_at=rec.updated_at,
|
||||
)
|
||||
|
||||
|
||||
# --- upload sessions ---------------------------------------------------------
|
||||
|
||||
|
||||
@router.post("/uploads", response_model=UploadSessionOut, status_code=201)
|
||||
async def create_upload(body: UploadSessionCreate, user: CurrentUser, session: SessionDep):
|
||||
try:
|
||||
us = await up.create_session(
|
||||
session, user.id, body.declared_mime_type, body.declared_size_bytes,
|
||||
body.title, body.client_recording_id, body.transcription_model,
|
||||
)
|
||||
except up.UploadError as e:
|
||||
raise HTTPException(e.status_code, e.message) from None
|
||||
return us
|
||||
|
||||
|
||||
@router.get("/uploads/{session_id}", response_model=UploadStatusOut)
|
||||
async def upload_status(session_id: uuid.UUID, user: CurrentUser, session: SessionDep):
|
||||
try:
|
||||
us = await up.get_owned_session(session, user.id, session_id)
|
||||
indexes = await up.received_indexes(session, user.id, session_id)
|
||||
except up.UploadError as e:
|
||||
raise HTTPException(e.status_code, e.message) from None
|
||||
return UploadStatusOut(id=us.id, status=us.status.value, received_chunk_indexes=indexes)
|
||||
|
||||
|
||||
@router.put("/uploads/{session_id}/chunks/{chunk_index}", status_code=201)
|
||||
async def put_chunk(
|
||||
session_id: uuid.UUID,
|
||||
chunk_index: int,
|
||||
request: Request,
|
||||
user: CurrentUser,
|
||||
session: SessionDep,
|
||||
x_chunk_sha256: str | None = Header(default=None),
|
||||
):
|
||||
data = await request.body()
|
||||
try:
|
||||
chunk = await up.put_chunk(session, user.id, session_id, chunk_index, data, x_chunk_sha256)
|
||||
except up.UploadError as e:
|
||||
raise HTTPException(e.status_code, e.message) from None
|
||||
return {"chunk_index": chunk.chunk_index, "size_bytes": chunk.size_bytes}
|
||||
|
||||
|
||||
@router.post("/uploads/{session_id}/finalize", response_model=RecordingOut, status_code=201)
|
||||
async def finalize_upload(
|
||||
session_id: uuid.UUID, body: RecordingFinalize, user: CurrentUser, session: SessionDep
|
||||
):
|
||||
# Location is stored ONLY with explicit per-user consent.
|
||||
lat = body.latitude if user.location_storage_enabled else None
|
||||
lon = body.longitude if user.location_storage_enabled else None
|
||||
acc = body.location_accuracy_m if user.location_storage_enabled else None
|
||||
try:
|
||||
_us, rec, _asset = await up.finalize(
|
||||
session, user.id, session_id,
|
||||
recorded_at=body.recorded_at, duration_seconds=body.duration_seconds,
|
||||
latitude=lat, longitude=lon, location_accuracy_m=acc, notes=body.notes,
|
||||
transcription_model=body.transcription_model,
|
||||
)
|
||||
except up.UploadError as e:
|
||||
raise HTTPException(e.status_code, e.message) from None
|
||||
return await _recording_out(session, rec)
|
||||
|
||||
|
||||
@router.delete("/uploads/{session_id}", status_code=204)
|
||||
async def abort_upload(session_id: uuid.UUID, user: CurrentUser, session: SessionDep):
|
||||
try:
|
||||
await up.abort(session, user.id, session_id)
|
||||
except up.UploadError as e:
|
||||
raise HTTPException(e.status_code, e.message) from None
|
||||
|
||||
|
||||
# --- recordings ----------------------------------------------------------------
|
||||
|
||||
|
||||
@router.get("/recordings", response_model=RecordingListOut)
|
||||
async def list_recordings(
|
||||
user: CurrentUser,
|
||||
session: SessionDep,
|
||||
limit: int = Query(default=50, ge=1, le=200),
|
||||
offset: int = Query(default=0, ge=0),
|
||||
sort: str = Query(
|
||||
default="recorded_at",
|
||||
pattern="^(recorded_at|created_at|duration_seconds|title)$",
|
||||
),
|
||||
order: str = Query(default="desc", pattern="^(asc|desc)$"),
|
||||
# M9 filters (AND-combined).
|
||||
tag: str | None = Query(default=None, max_length=80),
|
||||
status: str | None = Query(default=None, max_length=20),
|
||||
from_date: datetime | None = Query(default=None, description="recorded_at >= (UTC)"),
|
||||
to_date: datetime | None = Query(default=None, description="recorded_at <= (UTC)"),
|
||||
):
|
||||
where = [Recording.user_id == user.id, Recording.deleted_at.is_(None)]
|
||||
if tag:
|
||||
from shonar.db.models import RecordingTag as RT
|
||||
from shonar.db.models import Tag as T
|
||||
|
||||
where.append(
|
||||
Recording.id.in_(
|
||||
select(RT.recording_id)
|
||||
.join(T, T.id == RT.tag_id)
|
||||
.where(T.user_id == user.id, T.name == tag.strip().lower())
|
||||
)
|
||||
)
|
||||
if status:
|
||||
from shonar.db.models import ProcessingStatus
|
||||
|
||||
try:
|
||||
where.append(Recording.processing_status == ProcessingStatus(status))
|
||||
except ValueError:
|
||||
raise HTTPException(422, f"Unknown status: {status}") from None
|
||||
if from_date is not None:
|
||||
where.append(Recording.recorded_at >= from_date)
|
||||
if to_date is not None:
|
||||
where.append(Recording.recorded_at <= to_date)
|
||||
total = await session.scalar(select(func.count(Recording.id)).where(*where))
|
||||
col = getattr(Recording, sort)
|
||||
col = col.desc() if order == "desc" else col.asc()
|
||||
rows = await session.scalars(
|
||||
select(Recording).where(*where).order_by(col).limit(limit).offset(offset)
|
||||
)
|
||||
items = [await _recording_out(session, r) for r in rows]
|
||||
return RecordingListOut(items=items, total=total or 0, limit=limit, offset=offset)
|
||||
|
||||
|
||||
@router.get("/recordings/{recording_id}", response_model=RecordingOut)
|
||||
async def get_recording(recording_id: uuid.UUID, user: CurrentUser, session: SessionDep):
|
||||
rec = await session.get(Recording, recording_id)
|
||||
if rec is None or rec.user_id != user.id or rec.deleted_at is not None:
|
||||
raise HTTPException(404, "Recording not found")
|
||||
return await _recording_out(session, rec)
|
||||
|
||||
|
||||
@router.patch("/recordings/{recording_id}", response_model=RecordingOut)
|
||||
async def update_recording(
|
||||
recording_id: uuid.UUID, body: RecordingUpdate, user: CurrentUser, session: SessionDep
|
||||
):
|
||||
rec = await session.get(Recording, recording_id)
|
||||
if rec is None or rec.user_id != user.id or rec.deleted_at is not None:
|
||||
raise HTTPException(404, "Recording not found")
|
||||
if body.title is not None:
|
||||
rec.title = body.title
|
||||
if body.notes is not None:
|
||||
rec.notes = body.notes
|
||||
if body.latitude is not None and user.location_storage_enabled:
|
||||
rec.latitude = body.latitude
|
||||
if body.longitude is not None and user.location_storage_enabled:
|
||||
rec.longitude = body.longitude
|
||||
if body.tags is not None:
|
||||
# Replace tag set. Tags are per-user, created on demand.
|
||||
names = sorted(dict.fromkeys(t.strip().lower() for t in body.tags if t.strip()))[:20]
|
||||
existing = list(
|
||||
await session.scalars(select(Tag).where(Tag.user_id == user.id, Tag.name.in_(names)))
|
||||
)
|
||||
by_name = {t.name: t for t in existing}
|
||||
for name in names:
|
||||
if name not in by_name:
|
||||
t = Tag(user_id=user.id, name=name)
|
||||
session.add(t)
|
||||
await session.flush()
|
||||
by_name[name] = t
|
||||
# Deterministic replace: drop all links for this recording, re-add.
|
||||
from sqlalchemy import delete as sql_delete
|
||||
|
||||
await session.execute(
|
||||
sql_delete(RecordingTag).where(RecordingTag.recording_id == rec.id)
|
||||
)
|
||||
for name in names:
|
||||
session.add(RecordingTag(recording_id=rec.id, tag_id=by_name[name].id))
|
||||
await session.flush()
|
||||
return await _recording_out(session, rec)
|
||||
|
||||
|
||||
@router.delete("/recordings/{recording_id}", status_code=204)
|
||||
async def delete_recording(
|
||||
recording_id: uuid.UUID,
|
||||
user: CurrentUser,
|
||||
session: SessionDep,
|
||||
purge: bool = Query(default=False, description="true also deletes stored audio"),
|
||||
):
|
||||
"""Soft-delete by default; ?purge=true removes rows + stored files now."""
|
||||
rec = await session.get(Recording, recording_id)
|
||||
if rec is None or rec.user_id != user.id or rec.deleted_at is not None:
|
||||
raise HTTPException(404, "Recording not found")
|
||||
if not purge:
|
||||
rec.deleted_at = utcnow()
|
||||
await session.flush()
|
||||
return Response(status_code=204)
|
||||
|
||||
from shonar.storage import get_storage
|
||||
|
||||
storage = get_storage()
|
||||
assets = list(
|
||||
await session.scalars(select(Asset).where(Asset.recording_id == rec.id))
|
||||
)
|
||||
await session.delete(rec) # cascades to assets/transcripts/summaries/jobs
|
||||
await session.flush()
|
||||
for a in assets:
|
||||
with contextlib.suppress(Exception): # best effort
|
||||
await storage.delete(a.storage_key)
|
||||
return Response(status_code=204)
|
||||
|
||||
|
||||
@router.get("/recordings/{recording_id}/audio")
|
||||
async def download_audio(recording_id: uuid.UUID, user: CurrentUser, session: SessionDep):
|
||||
rec = await session.get(Recording, recording_id)
|
||||
if rec is None or rec.user_id != user.id or rec.deleted_at is not None:
|
||||
raise HTTPException(404, "Recording not found")
|
||||
original = await session.scalar(
|
||||
select(Asset).where(Asset.recording_id == rec.id, Asset.kind == AssetKind.original)
|
||||
)
|
||||
if original is None:
|
||||
raise HTTPException(404, "No audio stored for this recording")
|
||||
from fastapi.responses import Response as RawResponse
|
||||
|
||||
from shonar.storage import get_storage
|
||||
|
||||
data = await get_storage().get(original.storage_key)
|
||||
ext = original.storage_key[original.storage_key.rfind(".") :]
|
||||
filename = f"{rec.recorded_at:%Y%m%d-%H%M%S}{ext}"
|
||||
return RawResponse(
|
||||
content=data,
|
||||
media_type=original.mime_type,
|
||||
headers={
|
||||
"Content-Disposition": f'attachment; filename="{filename}"',
|
||||
"Cache-Control": "private, no-store",
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
# --- AI outputs (M7): latest transcript / summary / job states --------------
|
||||
|
||||
|
||||
async def _owned_recording(
|
||||
session: SessionDep, user: CurrentUser, recording_id: uuid.UUID
|
||||
) -> Recording:
|
||||
rec = await session.get(Recording, recording_id)
|
||||
if rec is None or rec.user_id != user.id or rec.deleted_at is not None:
|
||||
raise HTTPException(404, "Recording not found")
|
||||
return rec
|
||||
|
||||
|
||||
@router.get("/recordings/{recording_id}/transcript", response_model=TranscriptOut)
|
||||
async def get_transcript(
|
||||
recording_id: uuid.UUID, user: CurrentUser, session: SessionDep
|
||||
):
|
||||
rec = await _owned_recording(session, user, recording_id)
|
||||
row = await session.scalar(
|
||||
select(Transcript)
|
||||
.where(
|
||||
Transcript.recording_id == rec.id,
|
||||
Transcript.superseded_at.is_(None),
|
||||
)
|
||||
.order_by(Transcript.version.desc())
|
||||
)
|
||||
if row is None:
|
||||
raise HTTPException(404, "No transcript yet")
|
||||
return _transcript_out(row)
|
||||
|
||||
|
||||
@router.get("/recordings/{recording_id}/summary", response_model=SummaryOut)
|
||||
async def get_summary(recording_id: uuid.UUID, user: CurrentUser, session: SessionDep):
|
||||
rec = await _owned_recording(session, user, recording_id)
|
||||
row = await session.scalar(
|
||||
select(Summary)
|
||||
.where(
|
||||
Summary.recording_id == rec.id,
|
||||
Summary.superseded_at.is_(None),
|
||||
)
|
||||
.order_by(Summary.version.desc())
|
||||
)
|
||||
if row is None:
|
||||
raise HTTPException(404, "No summary yet")
|
||||
return row
|
||||
|
||||
|
||||
@router.get("/recordings/{recording_id}/jobs", response_model=list[ProcessingJobOut])
|
||||
async def list_jobs(recording_id: uuid.UUID, user: CurrentUser, session: SessionDep):
|
||||
rec = await _owned_recording(session, user, recording_id)
|
||||
rows = await session.scalars(
|
||||
select(ProcessingJob)
|
||||
.where(ProcessingJob.recording_id == rec.id)
|
||||
.order_by(ProcessingJob.id)
|
||||
)
|
||||
return list(rows)
|
||||
|
||||
|
||||
@router.post("/recordings/{recording_id}/reprocess", response_model=list[ProcessingJobOut])
|
||||
async def reprocess_recording(
|
||||
recording_id: uuid.UUID,
|
||||
user: CurrentUser,
|
||||
session: SessionDep,
|
||||
job: str = Query(default="summarize", pattern="^(transcribe|summarize)$"),
|
||||
model: str | None = Query(default=None, max_length=64),
|
||||
):
|
||||
"""Force one pipeline stage to run again (Summarize / Re-transcribe).
|
||||
|
||||
Unlike the enqueue-on-finalize path, this ignores prior success: a
|
||||
summary the user wants regenerated (better model, new prompt) is a
|
||||
deliberate request. Running jobs are left alone (409 instead of a
|
||||
duplicate).
|
||||
"""
|
||||
from shonar.db.models import JobStatus, ProcessingStatus
|
||||
from shonar.db.models import JobType as JT
|
||||
from shonar.services import processing as proc
|
||||
|
||||
rec = await _owned_recording(session, user, recording_id)
|
||||
job_type = JT(job)
|
||||
if job_type is JT.summarize and await proc.latest_transcript_text(
|
||||
session, rec.id
|
||||
) is None:
|
||||
raise HTTPException(409, "Transcribe first — there is nothing to summarize.")
|
||||
existing = await session.scalar(
|
||||
select(ProcessingJob)
|
||||
.where(
|
||||
ProcessingJob.recording_id == rec.id,
|
||||
ProcessingJob.job_type == job_type,
|
||||
)
|
||||
.order_by(ProcessingJob.id.desc())
|
||||
)
|
||||
if existing is not None and existing.status in (
|
||||
JobStatus.queued,
|
||||
JobStatus.running,
|
||||
):
|
||||
raise HTTPException(409, "That stage is already running.")
|
||||
if existing is None:
|
||||
session.add(ProcessingJob(recording_id=rec.id, job_type=job_type))
|
||||
else:
|
||||
existing.status = JobStatus.queued
|
||||
existing.attempt = 0
|
||||
existing.error = None
|
||||
existing.stage = None
|
||||
existing.progress = None
|
||||
existing.started_at = None
|
||||
existing.finished_at = None
|
||||
if job_type is JT.transcribe:
|
||||
if model is not None:
|
||||
# A re-transcribe may switch models; the saved per-recording
|
||||
# override is what the worker reads, so persist it here.
|
||||
from shonar.services.ai import ProviderConfigError
|
||||
from shonar.services.ai.model_registry import validate_model_name
|
||||
|
||||
try:
|
||||
rec.transcription_model = validate_model_name(model)
|
||||
except ProviderConfigError as e:
|
||||
raise HTTPException(422, str(e)) from e
|
||||
rec.processing_status = ProcessingStatus.processing
|
||||
rec.processing_error = None
|
||||
await session.flush()
|
||||
await proc.transport_enqueue(job_type, rec.id)
|
||||
rows = await session.scalars(
|
||||
select(ProcessingJob)
|
||||
.where(ProcessingJob.recording_id == rec.id)
|
||||
.order_by(ProcessingJob.id)
|
||||
)
|
||||
return list(rows)
|
||||
|
||||
|
||||
# --- M8: user edits (new version, edited_by_user=True; pipeline won't clobber)
|
||||
def _transcript_out(row: Transcript) -> TranscriptOut:
|
||||
return TranscriptOut(
|
||||
version=row.version,
|
||||
language=row.language,
|
||||
provider=row.provider,
|
||||
model=row.model,
|
||||
text=row.text,
|
||||
segments=[
|
||||
s for s in (row.segments or []) if isinstance(s, dict)
|
||||
],
|
||||
edited_by_user=row.edited_by_user,
|
||||
created_at=row.created_at,
|
||||
updated_at=row.updated_at,
|
||||
)
|
||||
|
||||
|
||||
@router.put("/recordings/{recording_id}/transcript", response_model=TranscriptOut)
|
||||
async def update_transcript(
|
||||
recording_id: uuid.UUID, body: TranscriptUpdate, user: CurrentUser, session: SessionDep
|
||||
):
|
||||
from sqlalchemy import func as sql_func
|
||||
|
||||
rec = await _owned_recording(session, user, recording_id)
|
||||
existing = list(
|
||||
await session.scalars(
|
||||
select(Transcript)
|
||||
.where(
|
||||
Transcript.recording_id == rec.id,
|
||||
Transcript.superseded_at.is_(None),
|
||||
)
|
||||
.order_by(Transcript.version.desc())
|
||||
)
|
||||
)
|
||||
now = utcnow()
|
||||
max_version = await session.scalar(
|
||||
select(sql_func.max(Transcript.version)).where(Transcript.recording_id == rec.id)
|
||||
)
|
||||
for row in existing:
|
||||
row.superseded_at = now
|
||||
row = Transcript(
|
||||
recording_id=rec.id,
|
||||
version=(max_version or 0) + 1,
|
||||
language=body.language,
|
||||
provider="user",
|
||||
model=None,
|
||||
text=body.text,
|
||||
segments=(
|
||||
[
|
||||
{"start": s.start, "end": s.end, "text": s.text, "speaker": s.speaker}
|
||||
for s in (body.segments or [])
|
||||
]
|
||||
if body.segments is not None
|
||||
else None
|
||||
),
|
||||
edited_by_user=True,
|
||||
)
|
||||
session.add(row)
|
||||
await session.flush()
|
||||
return _transcript_out(row)
|
||||
|
||||
|
||||
@router.put("/recordings/{recording_id}/summary", response_model=SummaryOut)
|
||||
async def update_summary(
|
||||
recording_id: uuid.UUID, body: SummaryUpdate, user: CurrentUser, session: SessionDep
|
||||
):
|
||||
from sqlalchemy import func as sql_func
|
||||
|
||||
rec = await _owned_recording(session, user, recording_id)
|
||||
existing = list(
|
||||
await session.scalars(
|
||||
select(Summary)
|
||||
.where(
|
||||
Summary.recording_id == rec.id,
|
||||
Summary.superseded_at.is_(None),
|
||||
)
|
||||
.order_by(Summary.version.desc())
|
||||
)
|
||||
)
|
||||
now = utcnow()
|
||||
max_version = await session.scalar(
|
||||
select(sql_func.max(Summary.version)).where(Summary.recording_id == rec.id)
|
||||
)
|
||||
for row in existing:
|
||||
row.superseded_at = now
|
||||
row = Summary(
|
||||
recording_id=rec.id,
|
||||
version=(max_version or 0) + 1,
|
||||
provider="user",
|
||||
model=None,
|
||||
content=body.content,
|
||||
edited_by_user=True,
|
||||
)
|
||||
session.add(row)
|
||||
await session.flush()
|
||||
return row
|
||||
81
backend/shonar/api/v1/search_exports.py
Normal file
81
backend/shonar/api/v1/search_exports.py
Normal file
|
|
@ -0,0 +1,81 @@
|
|||
"""Search + exports endpoints (M9).
|
||||
|
||||
Ownership is enforced everywhere: a search only ever sees the caller's own
|
||||
non-deleted recordings, and an export 404s on anything the caller doesn't
|
||||
own (no existence leak).
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import uuid
|
||||
|
||||
from fastapi import APIRouter, HTTPException, Query
|
||||
from fastapi.responses import Response
|
||||
|
||||
from shonar.api.deps import CurrentUser, SessionDep
|
||||
from shonar.api.schemas_search import SearchHitOut, SearchOut
|
||||
from shonar.db.models import Recording
|
||||
from shonar.services import exports as ex
|
||||
from shonar.services import search as sr
|
||||
|
||||
router = APIRouter(tags=["search", "exports"])
|
||||
|
||||
|
||||
# --- search -------------------------------------------------------------------
|
||||
|
||||
|
||||
@router.get("/search", response_model=SearchOut)
|
||||
async def search(
|
||||
user: CurrentUser,
|
||||
session: SessionDep,
|
||||
q: str = Query(..., min_length=1, max_length=200, description="Search text"),
|
||||
scope: str = Query(
|
||||
default="all",
|
||||
pattern="^(all|title|notes|transcript|summary|tag)$",
|
||||
description="Restrict the search to one field (default: all)",
|
||||
),
|
||||
limit: int = Query(default=20, ge=1, le=100),
|
||||
offset: int = Query(default=0, ge=0),
|
||||
):
|
||||
hits, total = await sr.search_recordings(
|
||||
session, user.id, q, scope=scope, limit=limit, offset=offset
|
||||
)
|
||||
items = [
|
||||
SearchHitOut(
|
||||
id=h.recording.id,
|
||||
title=h.recording.title,
|
||||
field=h.field,
|
||||
snippet=h.snippet,
|
||||
)
|
||||
for h in hits
|
||||
]
|
||||
return SearchOut(query=q, scope=scope, items=items, total=total, limit=limit, offset=offset)
|
||||
|
||||
|
||||
# --- exports ------------------------------------------------------------------
|
||||
|
||||
|
||||
@router.get("/recordings/{recording_id}/export")
|
||||
async def export_recording(
|
||||
recording_id: uuid.UUID,
|
||||
user: CurrentUser,
|
||||
session: SessionDep,
|
||||
fmt: str = Query(default="zip", pattern="^(audio|txt|md|zip)$"),
|
||||
):
|
||||
"""Return one export artifact inline (audio / txt / md / zip bundle)."""
|
||||
rec = await session.get(Recording, recording_id)
|
||||
if rec is None or rec.user_id != user.id or rec.deleted_at is not None:
|
||||
raise HTTPException(404, "Recording not found")
|
||||
try:
|
||||
result = await ex.build_export(session, rec, fmt)
|
||||
except ex.ExportError as e:
|
||||
raise HTTPException(e.status_code, e.message) from None
|
||||
await ex.record_export(session, user.id, rec, fmt, result)
|
||||
return Response(
|
||||
content=result.data,
|
||||
media_type=result.mime_type,
|
||||
headers={
|
||||
"Content-Disposition": f'attachment; filename="{result.filename}"',
|
||||
"Cache-Control": "private, no-store",
|
||||
},
|
||||
)
|
||||
54
backend/shonar/api/v1/users.py
Normal file
54
backend/shonar/api/v1/users.py
Normal file
|
|
@ -0,0 +1,54 @@
|
|||
"""Users and devices."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import uuid
|
||||
|
||||
from fastapi import APIRouter, HTTPException
|
||||
from sqlalchemy import select, update
|
||||
|
||||
from shonar.api.deps import CurrentUser, SessionDep
|
||||
from shonar.api.schemas_common import DeviceOut, UserOut, UserUpdate
|
||||
from shonar.db.models import Device, RefreshToken, utcnow
|
||||
|
||||
router = APIRouter(tags=["users", "devices"])
|
||||
|
||||
|
||||
@router.get("/users/me", response_model=UserOut)
|
||||
async def get_me(user: CurrentUser):
|
||||
return user
|
||||
|
||||
|
||||
@router.patch("/users/me", response_model=UserOut)
|
||||
async def update_me(body: UserUpdate, user: CurrentUser, session: SessionDep):
|
||||
if body.display_name is not None:
|
||||
user.display_name = body.display_name
|
||||
if body.location_storage_enabled is not None:
|
||||
user.location_storage_enabled = body.location_storage_enabled
|
||||
await session.flush()
|
||||
return user
|
||||
|
||||
|
||||
@router.get("/devices", response_model=list[DeviceOut])
|
||||
async def list_devices(user: CurrentUser, session: SessionDep):
|
||||
rows = await session.scalars(
|
||||
select(Device).where(Device.user_id == user.id).order_by(Device.last_seen_at.desc())
|
||||
)
|
||||
return list(rows)
|
||||
|
||||
|
||||
@router.delete("/devices/{device_id}", status_code=204)
|
||||
async def revoke_device(device_id: uuid.UUID, user: CurrentUser, session: SessionDep):
|
||||
device = await session.get(Device, device_id)
|
||||
# Ownership check — no cross-user access, and 404 (not 403) to avoid
|
||||
# leaking existence.
|
||||
if device is None or device.user_id != user.id:
|
||||
raise HTTPException(404, "Device not found")
|
||||
device.revoked_at = utcnow()
|
||||
# Revoke this device's live refresh tokens.
|
||||
await session.execute(
|
||||
update(RefreshToken)
|
||||
.where(RefreshToken.device_id == device.id, RefreshToken.revoked_at.is_(None))
|
||||
.values(revoked_at=utcnow())
|
||||
)
|
||||
return None
|
||||
0
backend/shonar/core/__init__.py
Normal file
0
backend/shonar/core/__init__.py
Normal file
138
backend/shonar/core/config.py
Normal file
138
backend/shonar/core/config.py
Normal file
|
|
@ -0,0 +1,138 @@
|
|||
"""Application settings.
|
||||
|
||||
All configuration comes from environment variables (and optionally an
|
||||
admin config file pointed at by SHONAR_CONFIG_FILE). API keys and other
|
||||
secrets must NEVER be hard-coded.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from functools import lru_cache
|
||||
|
||||
from pydantic import Field, field_validator
|
||||
from pydantic_settings import BaseSettings, SettingsConfigDict
|
||||
|
||||
|
||||
class Settings(BaseSettings):
|
||||
model_config = SettingsConfigDict(
|
||||
env_prefix="SHONAR_",
|
||||
env_file=".env",
|
||||
env_file_encoding="utf-8",
|
||||
extra="ignore",
|
||||
)
|
||||
|
||||
# --- Core -------------------------------------------------------------
|
||||
app_name: str = "S.H.O.N.A.R."
|
||||
debug: bool = False
|
||||
# Secret used to sign access tokens. MUST be set in production.
|
||||
secret_key: str = Field(default="", repr=False)
|
||||
access_token_ttl_minutes: int = 15
|
||||
refresh_token_ttl_days: int = 30
|
||||
# Comma-separated list of allowed registration modes: "open", "invite"
|
||||
allow_registration: bool = True
|
||||
|
||||
# --- Database -----------------------------------------------------------
|
||||
database_url: str = "postgresql+asyncpg://shonar:shonar@localhost:5432/shonar"
|
||||
db_pool_size: int = 5
|
||||
db_max_overflow: int = 10
|
||||
|
||||
# --- Storage ------------------------------------------------------------
|
||||
# "local" or "s3"
|
||||
storage_backend: str = "local"
|
||||
storage_path: str = "./data/storage"
|
||||
# S3 (used when storage_backend == "s3")
|
||||
s3_endpoint_url: str = ""
|
||||
s3_bucket: str = ""
|
||||
s3_region: str = "us-east-1"
|
||||
s3_access_key_id: str = Field(default="", repr=False)
|
||||
s3_secret_access_key: str = Field(default="", repr=False)
|
||||
|
||||
# --- Upload limits --------------------------------------------------------
|
||||
max_upload_bytes: int = 2 * 1024 * 1024 * 1024 # 2 GiB
|
||||
max_chunk_bytes: int = 16 * 1024 * 1024
|
||||
allowed_audio_mime_types: list[str] = [
|
||||
"audio/mp4",
|
||||
"audio/m4a",
|
||||
"audio/aac",
|
||||
"audio/wav",
|
||||
"audio/x-wav",
|
||||
"audio/ogg",
|
||||
"audio/opus",
|
||||
"audio/webm",
|
||||
"audio/mpeg",
|
||||
]
|
||||
|
||||
# --- Worker / queue -------------------------------------------------------
|
||||
redis_url: str = "redis://localhost:6379/0"
|
||||
# "arq" = Redis transport (server deployments). "inline" = in-process
|
||||
# asyncio runner (desktop bundled-lite engine: no Redis, one job at a
|
||||
# time). DB rows are the queue of record either way.
|
||||
queue_backend: str = "arq" # arq | inline
|
||||
# Apply 'alembic upgrade head' at startup. The desktop bundled engine
|
||||
# owns its SQLite file and must self-migrate; server deployments run
|
||||
# migrations in their deploy flow instead.
|
||||
auto_migrate: bool = False
|
||||
|
||||
# --- Audio processing -----------------------------------------------------
|
||||
# Optional server-side conversion via FFmpeg. Off by default.
|
||||
audio_conversion_enabled: bool = False
|
||||
ffmpeg_bin: str = "ffmpeg"
|
||||
|
||||
# --- AI providers -----------------------------------------------------
|
||||
# "none" disables all AI processing (recording/sync/playback still work).
|
||||
transcription_provider: str = "none" # none | whisper_http | faster_whisper
|
||||
transcription_model: str = "base"
|
||||
transcription_base_url: str = "" # whisper-compatible HTTP server
|
||||
transcription_api_key: str = Field(default="", repr=False)
|
||||
|
||||
llm_provider: str = "none" # none | openai_compat | ollama
|
||||
llm_model: str = ""
|
||||
llm_base_url: str = ""
|
||||
llm_api_key: str = Field(default="", repr=False)
|
||||
|
||||
# --- Search ---------------------------------------------------------
|
||||
search_backend: str = "postgres_fts" # postgres_fts (meilisearch: TODO)
|
||||
|
||||
# --- Retention --------------------------------------------------------
|
||||
# Grace window before the sweep hard-deletes soft-deleted recordings
|
||||
# and accounts (rows + stored files). Cancellations are valid until
|
||||
# the sweep fires.
|
||||
retention_grace_days: int = 30
|
||||
|
||||
# --- Misc -------------------------------------------------------------
|
||||
rate_limit_auth: str = "10/minute"
|
||||
rate_limit_default: str = "120/minute"
|
||||
|
||||
@field_validator("allowed_audio_mime_types", mode="before")
|
||||
@classmethod
|
||||
def _split_mime(cls, v): # noqa: ANN001, ANN206
|
||||
if isinstance(v, str):
|
||||
return [item.strip() for item in v.split(",") if item.strip()]
|
||||
return v
|
||||
|
||||
@property
|
||||
def ai_enabled(self) -> bool:
|
||||
return self.transcription_provider != "none" or self.llm_provider != "none"
|
||||
|
||||
def validate_production(self) -> list[str]:
|
||||
"""Return a list of configuration warnings (empty == OK)."""
|
||||
warnings: list[str] = []
|
||||
if not self.secret_key or len(self.secret_key) < 32:
|
||||
warnings.append(
|
||||
"SHONAR_SECRET_KEY is missing or shorter than 32 characters. "
|
||||
"Set a strong random secret in production."
|
||||
)
|
||||
if self.storage_backend == "s3" and not self.s3_bucket:
|
||||
warnings.append("storage_backend=s3 but SHONAR_S3_BUCKET is empty.")
|
||||
if self.transcription_provider == "whisper_http" and not self.transcription_base_url:
|
||||
warnings.append(
|
||||
"transcription_provider=whisper_http requires SHONAR_TRANSCRIPTION_BASE_URL."
|
||||
)
|
||||
if self.llm_provider == "openai_compat" and not self.llm_base_url:
|
||||
warnings.append("llm_provider=openai_compat requires SHONAR_LLM_BASE_URL.")
|
||||
return warnings
|
||||
|
||||
|
||||
@lru_cache
|
||||
def get_settings() -> Settings:
|
||||
return Settings()
|
||||
19
backend/shonar/core/ratelimit.py
Normal file
19
backend/shonar/core/ratelimit.py
Normal file
|
|
@ -0,0 +1,19 @@
|
|||
"""Rate limiting (slowapi). A single shared limiter instance so counters are
|
||||
consistent across endpoints; limit strings come from settings/env."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from slowapi import Limiter
|
||||
from slowapi.util import get_remote_address
|
||||
|
||||
from shonar.core.config import get_settings
|
||||
|
||||
_settings = get_settings()
|
||||
|
||||
limiter = Limiter(
|
||||
key_func=get_remote_address,
|
||||
default_limits=[],
|
||||
enabled=True,
|
||||
)
|
||||
|
||||
auth_limit = limiter.limit(_settings.rate_limit_auth)
|
||||
98
backend/shonar/core/security.py
Normal file
98
backend/shonar/core/security.py
Normal file
|
|
@ -0,0 +1,98 @@
|
|||
"""Security primitives: password hashing, access/refresh tokens, rate limit
|
||||
keying. Secrets come exclusively from settings/env."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import hashlib
|
||||
import hmac
|
||||
import secrets
|
||||
import uuid
|
||||
from datetime import UTC, datetime, timedelta
|
||||
from typing import Any
|
||||
|
||||
import jwt
|
||||
from argon2 import PasswordHasher
|
||||
from argon2.exceptions import InvalidHashError, VerificationError, VerifyMismatchError
|
||||
|
||||
from shonar.core.config import get_settings
|
||||
|
||||
_ph = PasswordHasher()
|
||||
|
||||
ACCESS_TOKEN_TYPE = "access"
|
||||
REFRESH_TOKEN_TYPE = "refresh"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Passwords
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def hash_password(password: str) -> str:
|
||||
return _ph.hash(password)
|
||||
|
||||
|
||||
def verify_password(password_hash: str, password: str) -> bool:
|
||||
try:
|
||||
return _ph.verify(password_hash, password)
|
||||
except (VerifyMismatchError, VerificationError, InvalidHashError):
|
||||
return False
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Access tokens (JWT)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def create_access_token(user_id: uuid.UUID, device_id: uuid.UUID | None = None) -> tuple[str, int]:
|
||||
"""Returns (token, ttl_seconds)."""
|
||||
settings = get_settings()
|
||||
ttl = settings.access_token_ttl_minutes * 60
|
||||
now = datetime.now(UTC)
|
||||
payload: dict[str, Any] = {
|
||||
"sub": str(user_id),
|
||||
"typ": ACCESS_TOKEN_TYPE,
|
||||
"iat": now,
|
||||
"exp": now + timedelta(seconds=ttl),
|
||||
"jti": uuid.uuid4().hex,
|
||||
}
|
||||
if device_id is not None:
|
||||
payload["dev"] = str(device_id)
|
||||
token = jwt.encode(payload, settings.secret_key, algorithm="HS256")
|
||||
return token, ttl
|
||||
|
||||
|
||||
class TokenError(Exception):
|
||||
"""Invalid or expired token."""
|
||||
|
||||
|
||||
def decode_access_token(token: str) -> dict[str, Any]:
|
||||
settings = get_settings()
|
||||
try:
|
||||
payload = jwt.decode(token, settings.secret_key, algorithms=["HS256"])
|
||||
except jwt.PyJWTError as exc:
|
||||
raise TokenError("invalid access token") from exc
|
||||
if payload.get("typ") != ACCESS_TOKEN_TYPE:
|
||||
raise TokenError("wrong token type")
|
||||
return payload
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Refresh tokens (opaque, rotated, hashed at rest)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def generate_refresh_token() -> str:
|
||||
"""Opaque high-entropy token; only its SHA-256 hash is ever stored."""
|
||||
return secrets.token_urlsafe(48)
|
||||
|
||||
|
||||
def hash_refresh_token(token: str) -> str:
|
||||
return hashlib.sha256(token.encode("utf-8")).hexdigest()
|
||||
|
||||
|
||||
def refresh_token_ttl() -> timedelta:
|
||||
return timedelta(days=get_settings().refresh_token_ttl_days)
|
||||
|
||||
|
||||
def constant_time_equals(a: str, b: str) -> bool:
|
||||
return hmac.compare_digest(a, b)
|
||||
3
backend/shonar/db/__init__.py
Normal file
3
backend/shonar/db/__init__.py
Normal file
|
|
@ -0,0 +1,3 @@
|
|||
"""DB package. Importing this registers every model on Base.metadata."""
|
||||
|
||||
from shonar.db import models # noqa: F401
|
||||
50
backend/shonar/db/base.py
Normal file
50
backend/shonar/db/base.py
Normal file
|
|
@ -0,0 +1,50 @@
|
|||
"""Public SQLAlchemy model base."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import uuid
|
||||
from datetime import UTC, datetime
|
||||
|
||||
from sqlalchemy import DateTime, TypeDecorator, Uuid
|
||||
from sqlalchemy.orm import DeclarativeBase, Mapped, mapped_column
|
||||
|
||||
|
||||
class UTCDT(TypeDecorator):
|
||||
"""Timezone-aware datetime that survives SQLite round-trips.
|
||||
|
||||
Postgres returns aware values (impl is a no-op there); SQLite has no
|
||||
tz info and hands back naive datetimes, which crash comparisons with
|
||||
``utcnow()``. Attach UTC on read whenever the driver lost it.
|
||||
"""
|
||||
|
||||
impl = DateTime(timezone=True)
|
||||
cache_ok = True
|
||||
|
||||
def process_result_value(self, value, dialect): # noqa: ANN001, ANN201
|
||||
if value is not None and value.tzinfo is None:
|
||||
return value.replace(tzinfo=UTC)
|
||||
return value
|
||||
|
||||
|
||||
def utcnow() -> datetime:
|
||||
return datetime.now(UTC)
|
||||
|
||||
|
||||
def new_uuid() -> uuid.UUID:
|
||||
return uuid.uuid4()
|
||||
|
||||
|
||||
class Base(DeclarativeBase):
|
||||
pass
|
||||
|
||||
|
||||
class PublicIdMixin:
|
||||
"""Gives every table a UUID public id; integer PKs stay internal."""
|
||||
|
||||
id: Mapped[uuid.UUID] = mapped_column(Uuid, primary_key=True, default=new_uuid)
|
||||
created_at: Mapped[datetime] = mapped_column(UTCDT(), default=utcnow, nullable=False)
|
||||
# default= as well as onupdate=: SQLite enforces NOT NULL on insert
|
||||
# (Postgres silently stored NULL for rows that never got an UPDATE).
|
||||
updated_at: Mapped[datetime] = mapped_column(
|
||||
UTCDT(), default=utcnow, onupdate=utcnow, nullable=False
|
||||
)
|
||||
33
backend/shonar/db/migrate.py
Normal file
33
backend/shonar/db/migrate.py
Normal file
|
|
@ -0,0 +1,33 @@
|
|||
"""Programmatic schema migrations for the bundled desktop engine.
|
||||
|
||||
Server deployments run ``alembic upgrade head`` in their deploy flow;
|
||||
the desktop app owns its SQLite file end-to-end, so the API applies
|
||||
migrations itself at startup (``SHONAR_AUTO_MIGRATE=1``).
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import logging
|
||||
from pathlib import Path
|
||||
|
||||
from alembic import command
|
||||
from alembic.config import Config
|
||||
|
||||
logger = logging.getLogger("shonar.migrate")
|
||||
|
||||
# backend/ dir: migrations/ + alembic.ini live here relative to this file.
|
||||
_BACKEND_DIR = Path(__file__).resolve().parents[2]
|
||||
|
||||
|
||||
async def upgrade_head(database_url: str) -> None:
|
||||
"""Run 'alembic upgrade head' for `database_url` without blocking the loop."""
|
||||
|
||||
def _run() -> None:
|
||||
cfg = Config(str(_BACKEND_DIR / "alembic.ini"))
|
||||
cfg.set_main_option("script_location", str(_BACKEND_DIR / "migrations"))
|
||||
cfg.set_main_option("sqlalchemy.url", database_url)
|
||||
command.upgrade(cfg, "head")
|
||||
|
||||
await asyncio.to_thread(_run)
|
||||
logger.info("database schema migrated to head")
|
||||
453
backend/shonar/db/models.py
Normal file
453
backend/shonar/db/models.py
Normal file
|
|
@ -0,0 +1,453 @@
|
|||
"""All persistent models.
|
||||
|
||||
Design rules:
|
||||
- Every client-visible identifier is a UUID (``PublicIdMixin.id``).
|
||||
- Filesystem/storage paths are NEVER exposed to clients.
|
||||
- Original uploads are immutable; processed audio lives in separate
|
||||
``Asset`` rows and never replaces an original.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import enum
|
||||
import uuid
|
||||
from datetime import datetime
|
||||
|
||||
# JSON that renders JSONB on PostgreSQL and plain JSON on SQLite (the
|
||||
# desktop bundled-lite engine runs on SQLite; JSONB does not compile there).
|
||||
from sqlalchemy import JSON as _JSON
|
||||
from sqlalchemy import (
|
||||
BigInteger,
|
||||
Boolean,
|
||||
Enum,
|
||||
Float,
|
||||
ForeignKey,
|
||||
Index,
|
||||
String,
|
||||
Text,
|
||||
UniqueConstraint,
|
||||
)
|
||||
from sqlalchemy.dialects.postgresql import JSONB
|
||||
from sqlalchemy.orm import Mapped, mapped_column, relationship
|
||||
|
||||
from shonar.db.base import UTCDT, Base, PublicIdMixin, utcnow
|
||||
|
||||
JSONType = _JSON().with_variant(JSONB(), "postgresql")
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Users & auth
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class User(Base, PublicIdMixin):
|
||||
__tablename__ = "users"
|
||||
|
||||
email: Mapped[str] = mapped_column(String(320), unique=True, index=True, nullable=False)
|
||||
password_hash: Mapped[str] = mapped_column(String(255), nullable=False)
|
||||
display_name: Mapped[str | None] = mapped_column(String(120))
|
||||
is_active: Mapped[bool] = mapped_column(Boolean, default=True, nullable=False)
|
||||
# Feature switches the user controls (mirror of app settings, server truth
|
||||
# for e.g. whether location metadata may be stored at all).
|
||||
location_storage_enabled: Mapped[bool] = mapped_column(Boolean, default=False, nullable=False)
|
||||
deleted_at: Mapped[datetime | None] = mapped_column(UTCDT())
|
||||
|
||||
devices: Mapped[list[Device]] = relationship(back_populates="user")
|
||||
recordings: Mapped[list[Recording]] = relationship(back_populates="user")
|
||||
|
||||
|
||||
class Device(Base, PublicIdMixin):
|
||||
__tablename__ = "devices"
|
||||
|
||||
user_id: Mapped[uuid.UUID] = mapped_column(
|
||||
ForeignKey("users.id", ondelete="CASCADE"), index=True, nullable=False
|
||||
)
|
||||
name: Mapped[str] = mapped_column(String(120), nullable=False)
|
||||
platform: Mapped[str] = mapped_column(String(40), default="android", nullable=False)
|
||||
last_seen_at: Mapped[datetime] = mapped_column(UTCDT(), default=utcnow)
|
||||
revoked_at: Mapped[datetime | None] = mapped_column(UTCDT())
|
||||
|
||||
user: Mapped[User] = relationship(back_populates="devices")
|
||||
|
||||
|
||||
class RefreshToken(Base, PublicIdMixin):
|
||||
"""Rotating refresh tokens, stored hashed, grouped into families for
|
||||
reuse detection. A refresh consumes one row and issues its replacement
|
||||
with the same ``family``."""
|
||||
|
||||
__tablename__ = "refresh_tokens"
|
||||
|
||||
user_id: Mapped[uuid.UUID] = mapped_column(
|
||||
ForeignKey("users.id", ondelete="CASCADE"), index=True, nullable=False
|
||||
)
|
||||
token_hash: Mapped[str] = mapped_column(String(128), unique=True, nullable=False)
|
||||
family: Mapped[uuid.UUID] = mapped_column(index=True, nullable=False)
|
||||
device_id: Mapped[uuid.UUID | None] = mapped_column(
|
||||
ForeignKey("devices.id", ondelete="SET NULL")
|
||||
)
|
||||
expires_at: Mapped[datetime] = mapped_column(UTCDT(), nullable=False)
|
||||
revoked_at: Mapped[datetime | None] = mapped_column(UTCDT())
|
||||
replaced_by: Mapped[uuid.UUID | None] = mapped_column(
|
||||
ForeignKey("refresh_tokens.id", ondelete="SET NULL")
|
||||
)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Recordings, assets, uploads
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class ProcessingStatus(enum.StrEnum):
|
||||
pending_upload = "pending_upload"
|
||||
uploaded = "uploaded"
|
||||
queued = "queued"
|
||||
processing = "processing"
|
||||
completed = "completed"
|
||||
failed = "failed"
|
||||
# No AI configured / requested — audio-only recording, fully usable.
|
||||
ai_disabled = "ai_disabled"
|
||||
|
||||
|
||||
class Recording(Base, PublicIdMixin):
|
||||
__tablename__ = "recordings"
|
||||
|
||||
user_id: Mapped[uuid.UUID] = mapped_column(
|
||||
ForeignKey("users.id", ondelete="CASCADE"), index=True, nullable=False
|
||||
)
|
||||
device_id: Mapped[uuid.UUID | None] = mapped_column(
|
||||
ForeignKey("devices.id", ondelete="SET NULL")
|
||||
)
|
||||
# Client-generated idempotency id so retried uploads update, not duplicate.
|
||||
client_recording_id: Mapped[str | None] = mapped_column(String(64), index=True)
|
||||
|
||||
title: Mapped[str] = mapped_column(String(300), nullable=False, default="Untitled recording")
|
||||
recorded_at: Mapped[datetime] = mapped_column(UTCDT(), nullable=False)
|
||||
duration_seconds: Mapped[float] = mapped_column(Float, default=0.0, nullable=False)
|
||||
notes: Mapped[str | None] = mapped_column(Text)
|
||||
# Location is stored ONLY when the user has explicitly enabled it.
|
||||
latitude: Mapped[float | None] = mapped_column(Float)
|
||||
longitude: Mapped[float | None] = mapped_column(Float)
|
||||
location_accuracy_m: Mapped[float | None] = mapped_column(Float)
|
||||
|
||||
processing_status: Mapped[ProcessingStatus] = mapped_column(
|
||||
Enum(
|
||||
ProcessingStatus,
|
||||
name="processing_status",
|
||||
values_callable=lambda e: [m.value for m in e],
|
||||
),
|
||||
default=ProcessingStatus.pending_upload,
|
||||
nullable=False,
|
||||
index=True,
|
||||
)
|
||||
processing_error: Mapped[str | None] = mapped_column(Text)
|
||||
# Effective faster-whisper model for this recording (e.g. "base").
|
||||
# Set at finalize time from the per-recording override or the global
|
||||
# default; NULL means "server default at processing time" (pre-model
|
||||
# rows). Never rewritten: changing the default affects future rows only.
|
||||
transcription_model: Mapped[str | None] = mapped_column(String(32))
|
||||
# The original audio is the Asset row with kind=original for this
|
||||
# recording (at most one, enforced by a partial unique index). Keeping
|
||||
# the pointer one-directional avoids a recordings<->assets FK cycle.
|
||||
deleted_at: Mapped[datetime | None] = mapped_column(UTCDT())
|
||||
|
||||
user: Mapped[User] = relationship(back_populates="recordings")
|
||||
assets: Mapped[list[Asset]] = relationship(
|
||||
back_populates="recording", cascade="all, delete-orphan"
|
||||
)
|
||||
transcripts: Mapped[list[Transcript]] = relationship(
|
||||
back_populates="recording", cascade="all, delete-orphan"
|
||||
)
|
||||
summaries: Mapped[list[Summary]] = relationship(
|
||||
back_populates="recording", cascade="all, delete-orphan"
|
||||
)
|
||||
tags: Mapped[list[Tag]] = relationship(secondary="recording_tags", back_populates="recordings")
|
||||
|
||||
__table_args__ = (
|
||||
Index("ix_recordings_user_recorded", "user_id", "recorded_at"),
|
||||
Index(
|
||||
"ix_recordings_user_client_id",
|
||||
"user_id",
|
||||
"client_recording_id",
|
||||
unique=True,
|
||||
postgresql_where="client_recording_id IS NOT NULL",
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
class AssetKind(enum.StrEnum):
|
||||
original = "original" # user's uploaded file — NEVER modified
|
||||
normalized = "normalized" # derivative for processing (ffmpeg)
|
||||
export = "export" # generated export bundle
|
||||
|
||||
|
||||
class Asset(Base, PublicIdMixin):
|
||||
"""A stored binary object. ``storage_key`` is server-internal only."""
|
||||
|
||||
__tablename__ = "assets"
|
||||
|
||||
recording_id: Mapped[uuid.UUID | None] = mapped_column(
|
||||
ForeignKey("recordings.id", ondelete="CASCADE"), index=True
|
||||
)
|
||||
user_id: Mapped[uuid.UUID] = mapped_column(
|
||||
ForeignKey("users.id", ondelete="CASCADE"), index=True, nullable=False
|
||||
)
|
||||
kind: Mapped[AssetKind] = mapped_column(
|
||||
Enum(AssetKind, name="asset_kind", values_callable=lambda e: [m.value for m in e]),
|
||||
nullable=False,
|
||||
)
|
||||
storage_key: Mapped[str] = mapped_column(String(500), nullable=False)
|
||||
mime_type: Mapped[str] = mapped_column(String(100), nullable=False)
|
||||
size_bytes: Mapped[int] = mapped_column(BigInteger, nullable=False, default=0)
|
||||
checksum_sha256: Mapped[str] = mapped_column(String(64), nullable=False)
|
||||
|
||||
recording: Mapped[Recording | None] = relationship(back_populates="assets")
|
||||
|
||||
__table_args__ = (
|
||||
# At most one immutable "original" per recording.
|
||||
Index(
|
||||
"uq_assets_one_original",
|
||||
"recording_id",
|
||||
unique=True,
|
||||
postgresql_where="kind = 'original'",
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
class UploadSessionStatus(enum.StrEnum):
|
||||
open = "open"
|
||||
finalizing = "finalizing"
|
||||
completed = "completed"
|
||||
aborted = "aborted"
|
||||
expired = "expired"
|
||||
|
||||
|
||||
class UploadSession(Base, PublicIdMixin):
|
||||
"""Chunked, resumable upload session."""
|
||||
|
||||
__tablename__ = "upload_sessions"
|
||||
|
||||
user_id: Mapped[uuid.UUID] = mapped_column(
|
||||
ForeignKey("users.id", ondelete="CASCADE"), index=True, nullable=False
|
||||
)
|
||||
recording_id: Mapped[uuid.UUID | None] = mapped_column(
|
||||
ForeignKey("recordings.id", ondelete="CASCADE")
|
||||
)
|
||||
client_recording_id: Mapped[str | None] = mapped_column(String(64))
|
||||
title: Mapped[str | None] = mapped_column(String(300))
|
||||
# Optional per-recording transcription model override chosen on the
|
||||
# upload screen. Validated at finalize time; the finalize body wins when
|
||||
# both specify one.
|
||||
transcription_model: Mapped[str | None] = mapped_column(String(32))
|
||||
declared_mime_type: Mapped[str] = mapped_column(String(100), nullable=False)
|
||||
declared_size_bytes: Mapped[int] = mapped_column(BigInteger, nullable=False)
|
||||
chunk_size_bytes: Mapped[int] = mapped_column(
|
||||
BigInteger, nullable=False, default=8 * 1024 * 1024
|
||||
)
|
||||
status: Mapped[UploadSessionStatus] = mapped_column(
|
||||
Enum(
|
||||
UploadSessionStatus,
|
||||
name="upload_session_status",
|
||||
values_callable=lambda e: [m.value for m in e],
|
||||
),
|
||||
default=UploadSessionStatus.open,
|
||||
nullable=False,
|
||||
)
|
||||
expires_at: Mapped[datetime] = mapped_column(UTCDT(), nullable=False)
|
||||
completed_asset_id: Mapped[uuid.UUID | None] = mapped_column(
|
||||
ForeignKey("assets.id", ondelete="SET NULL")
|
||||
)
|
||||
|
||||
chunks: Mapped[list[UploadChunk]] = relationship(
|
||||
back_populates="session", cascade="all, delete-orphan"
|
||||
)
|
||||
|
||||
|
||||
class UploadChunk(Base):
|
||||
__tablename__ = "upload_chunks"
|
||||
|
||||
id: Mapped[int] = mapped_column(primary_key=True)
|
||||
session_id: Mapped[uuid.UUID] = mapped_column(
|
||||
ForeignKey("upload_sessions.id", ondelete="CASCADE"), index=True, nullable=False
|
||||
)
|
||||
chunk_index: Mapped[int] = mapped_column(nullable=False)
|
||||
size_bytes: Mapped[int] = mapped_column(BigInteger, nullable=False)
|
||||
checksum_sha256: Mapped[str] = mapped_column(String(64), nullable=False)
|
||||
storage_key: Mapped[str] = mapped_column(String(500), nullable=False)
|
||||
created_at: Mapped[datetime] = mapped_column(UTCDT(), default=utcnow)
|
||||
|
||||
session: Mapped[UploadSession] = relationship(back_populates="chunks")
|
||||
|
||||
__table_args__ = (UniqueConstraint("session_id", "chunk_index"),)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Transcript / summary / tags
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class Transcript(Base, PublicIdMixin):
|
||||
__tablename__ = "transcripts"
|
||||
|
||||
recording_id: Mapped[uuid.UUID] = mapped_column(
|
||||
ForeignKey("recordings.id", ondelete="CASCADE"), index=True, nullable=False
|
||||
)
|
||||
# Versioned: regenerated transcripts supersede older rows; the newest
|
||||
# non-superseded row is authoritative. ``edited_by_user`` rows win.
|
||||
version: Mapped[int] = mapped_column(default=1, nullable=False)
|
||||
superseded_at: Mapped[datetime | None] = mapped_column(UTCDT())
|
||||
language: Mapped[str | None] = mapped_column(String(16))
|
||||
provider: Mapped[str] = mapped_column(String(80), nullable=False, default="manual")
|
||||
model: Mapped[str | None] = mapped_column(String(120))
|
||||
# Full text, plus segments: [{"start": 0.0, "end": 2.5, "text": "...",
|
||||
# "speaker": "S1"|null}, ...]
|
||||
text: Mapped[str] = mapped_column(Text, nullable=False, default="")
|
||||
segments: Mapped[list | None] = mapped_column(JSONType)
|
||||
edited_by_user: Mapped[bool] = mapped_column(Boolean, default=False, nullable=False)
|
||||
|
||||
recording: Mapped[Recording] = relationship(back_populates="transcripts")
|
||||
|
||||
__table_args__ = (Index("ix_transcripts_recording_version", "recording_id", "version"),)
|
||||
|
||||
|
||||
class Summary(Base, PublicIdMixin):
|
||||
"""Structured AI summary — user-editable.
|
||||
|
||||
JSON shape of ``content``:
|
||||
{
|
||||
"short": str,
|
||||
"detailed": str,
|
||||
"key_points": [str],
|
||||
"decisions": [str],
|
||||
"action_items": [str],
|
||||
"questions": [str]
|
||||
}
|
||||
"""
|
||||
|
||||
__tablename__ = "summaries"
|
||||
|
||||
recording_id: Mapped[uuid.UUID] = mapped_column(
|
||||
ForeignKey("recordings.id", ondelete="CASCADE"), index=True, nullable=False
|
||||
)
|
||||
version: Mapped[int] = mapped_column(default=1, nullable=False)
|
||||
superseded_at: Mapped[datetime | None] = mapped_column(UTCDT())
|
||||
provider: Mapped[str] = mapped_column(String(80), nullable=False, default="manual")
|
||||
model: Mapped[str | None] = mapped_column(String(120))
|
||||
content: Mapped[dict] = mapped_column(JSONType, nullable=False, default=dict)
|
||||
edited_by_user: Mapped[bool] = mapped_column(Boolean, default=False, nullable=False)
|
||||
|
||||
recording: Mapped[Recording] = relationship(back_populates="summaries")
|
||||
|
||||
|
||||
class Tag(Base, PublicIdMixin):
|
||||
__tablename__ = "tags"
|
||||
|
||||
user_id: Mapped[uuid.UUID] = mapped_column(
|
||||
ForeignKey("users.id", ondelete="CASCADE"), index=True, nullable=False
|
||||
)
|
||||
name: Mapped[str] = mapped_column(String(80), nullable=False)
|
||||
|
||||
recordings: Mapped[list[Recording]] = relationship(
|
||||
secondary="recording_tags", back_populates="tags"
|
||||
)
|
||||
|
||||
__table_args__ = (UniqueConstraint("user_id", "name"),)
|
||||
|
||||
|
||||
class RecordingTag(Base):
|
||||
__tablename__ = "recording_tags"
|
||||
|
||||
recording_id: Mapped[uuid.UUID] = mapped_column(
|
||||
ForeignKey("recordings.id", ondelete="CASCADE"), primary_key=True
|
||||
)
|
||||
tag_id: Mapped[uuid.UUID] = mapped_column(
|
||||
ForeignKey("tags.id", ondelete="CASCADE"), primary_key=True
|
||||
)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Processing jobs
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class JobStatus(enum.StrEnum):
|
||||
queued = "queued"
|
||||
running = "running"
|
||||
succeeded = "succeeded"
|
||||
failed = "failed"
|
||||
skipped = "skipped" # e.g. AI disabled, or user-edited output protected
|
||||
|
||||
|
||||
class JobType(enum.StrEnum):
|
||||
normalize_audio = "normalize_audio"
|
||||
transcribe = "transcribe"
|
||||
summarize = "summarize"
|
||||
|
||||
|
||||
class ProcessingJob(Base, PublicIdMixin):
|
||||
"""One pipeline step for one recording. Idempotent: reruns overwrite
|
||||
derived outputs (unless user-edited) and never touch originals."""
|
||||
|
||||
__tablename__ = "processing_jobs"
|
||||
|
||||
recording_id: Mapped[uuid.UUID] = mapped_column(
|
||||
ForeignKey("recordings.id", ondelete="CASCADE"), index=True, nullable=False
|
||||
)
|
||||
job_type: Mapped[JobType] = mapped_column(
|
||||
Enum(JobType, name="job_type", values_callable=lambda e: [m.value for m in e]),
|
||||
nullable=False,
|
||||
)
|
||||
status: Mapped[JobStatus] = mapped_column(
|
||||
Enum(JobStatus, name="job_status", values_callable=lambda e: [m.value for m in e]),
|
||||
default=JobStatus.queued,
|
||||
nullable=False,
|
||||
index=True,
|
||||
)
|
||||
attempt: Mapped[int] = mapped_column(default=0, nullable=False)
|
||||
max_attempts: Mapped[int] = mapped_column(default=3, nullable=False)
|
||||
error: Mapped[str | None] = mapped_column(Text)
|
||||
# Fine-grained phase for progress display (e.g. transcribe jobs report
|
||||
# "loading-model" then "transcribing"). Nullable: older rows predate it.
|
||||
# Never parsed by pipeline logic — display only.
|
||||
stage: Mapped[str | None] = mapped_column(String(32))
|
||||
# 0-100 work estimate within the current stage, when known.
|
||||
progress: Mapped[int | None] = mapped_column()
|
||||
started_at: Mapped[datetime | None] = mapped_column(UTCDT())
|
||||
finished_at: Mapped[datetime | None] = mapped_column(UTCDT())
|
||||
# Opaque arq task handle for observability.
|
||||
task_handle: Mapped[str | None] = mapped_column(String(120))
|
||||
|
||||
__table_args__ = (Index("ix_processing_jobs_recording_type", "recording_id", "job_type"),)
|
||||
|
||||
|
||||
class ExportJob(Base, PublicIdMixin):
|
||||
__tablename__ = "export_jobs"
|
||||
|
||||
user_id: Mapped[uuid.UUID] = mapped_column(
|
||||
ForeignKey("users.id", ondelete="CASCADE"), index=True, nullable=False
|
||||
)
|
||||
recording_id: Mapped[uuid.UUID] = mapped_column(
|
||||
ForeignKey("recordings.id", ondelete="CASCADE"), index=True, nullable=False
|
||||
)
|
||||
# "audio", "transcript_txt", "notes_md", "bundle_zip"
|
||||
export_type: Mapped[str] = mapped_column(String(40), nullable=False)
|
||||
status: Mapped[JobStatus] = mapped_column(
|
||||
Enum(JobStatus, name="export_job_status", values_callable=lambda e: [m.value for m in e]),
|
||||
default=JobStatus.queued,
|
||||
nullable=False,
|
||||
)
|
||||
asset_id: Mapped[uuid.UUID | None] = mapped_column(ForeignKey("assets.id", ondelete="SET NULL"))
|
||||
error: Mapped[str | None] = mapped_column(Text)
|
||||
|
||||
|
||||
# Ensure full-text search columns exist on Postgres (added via migration as
|
||||
# tsvector generated columns; see migrations/versions/*_fts.py).
|
||||
|
||||
|
||||
class AppSetting(Base, PublicIdMixin):
|
||||
"""Server-wide settings editable at runtime (global transcription model
|
||||
default, …). Single row per key; the desktop Settings page writes here."""
|
||||
|
||||
__tablename__ = "app_settings"
|
||||
|
||||
key: Mapped[str] = mapped_column(String(120), nullable=False, unique=True)
|
||||
value: Mapped[dict] = mapped_column(JSONType, nullable=False, default=dict)
|
||||
76
backend/shonar/db/session.py
Normal file
76
backend/shonar/db/session.py
Normal file
|
|
@ -0,0 +1,76 @@
|
|||
"""Async SQLAlchemy engine/session management."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import AsyncIterator
|
||||
|
||||
from sqlalchemy.ext.asyncio import AsyncSession, async_sessionmaker, create_async_engine
|
||||
|
||||
from shonar.core.config import get_settings
|
||||
|
||||
_engine = None
|
||||
_session_factory: async_sessionmaker[AsyncSession] | None = None
|
||||
|
||||
|
||||
def get_engine():
|
||||
global _engine, _session_factory
|
||||
if _engine is None:
|
||||
settings = get_settings()
|
||||
kwargs: dict = {"pool_pre_ping": True}
|
||||
url = settings.database_url
|
||||
# SQLite (tests, desktop bundled engine) does not support the pg
|
||||
# pool sizing kwargs.
|
||||
if url.startswith("sqlite"):
|
||||
kwargs = {}
|
||||
else:
|
||||
kwargs.update(pool_size=settings.db_pool_size, max_overflow=settings.db_max_overflow)
|
||||
_engine = create_async_engine(url, **kwargs)
|
||||
if url.startswith("sqlite"):
|
||||
_configure_sqlite(_engine)
|
||||
_session_factory = async_sessionmaker(_engine, expire_on_commit=False)
|
||||
return _engine
|
||||
|
||||
|
||||
def _configure_sqlite(engine) -> None: # noqa: ANN001
|
||||
"""Per-connection SQLite pragmas.
|
||||
|
||||
foreign_keys is OFF by default in SQLite; the schema leans on
|
||||
ON DELETE CASCADE, so every connection must enable it. WAL lets the
|
||||
inline-queue writer and API readers coexist without SQLITE_BUSY
|
||||
storms on the desktop box.
|
||||
"""
|
||||
from sqlalchemy import event
|
||||
|
||||
@event.listens_for(engine.sync_engine, "connect")
|
||||
def _pragmas(dbapi_conn, _record): # noqa: ANN001
|
||||
cur = dbapi_conn.cursor()
|
||||
cur.execute("PRAGMA foreign_keys=ON")
|
||||
cur.execute("PRAGMA journal_mode=WAL")
|
||||
cur.execute("PRAGMA busy_timeout=5000")
|
||||
cur.close()
|
||||
|
||||
|
||||
async def dispose_engine() -> None:
|
||||
global _engine, _session_factory
|
||||
if _engine is not None:
|
||||
await _engine.dispose()
|
||||
_engine = None
|
||||
_session_factory = None
|
||||
|
||||
|
||||
async def get_session() -> AsyncIterator[AsyncSession]:
|
||||
"""FastAPI dependency yielding a database session."""
|
||||
assert _session_factory is not None, "engine not initialised"
|
||||
async with _session_factory() as session:
|
||||
try:
|
||||
yield session
|
||||
await session.commit()
|
||||
except Exception:
|
||||
await session.rollback()
|
||||
raise
|
||||
|
||||
|
||||
def session_factory() -> async_sessionmaker[AsyncSession]:
|
||||
"""Shareable session factory for the worker (outside requests)."""
|
||||
assert _session_factory is not None, "engine not initialised"
|
||||
return _session_factory
|
||||
79
backend/shonar/main.py
Normal file
79
backend/shonar/main.py
Normal file
|
|
@ -0,0 +1,79 @@
|
|||
"""S.H.O.N.A.R. FastAPI application entrypoint.
|
||||
|
||||
Run (dev): uvicorn shonar.main:app --reload --port 8000
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
from contextlib import asynccontextmanager
|
||||
|
||||
from fastapi import FastAPI, Request
|
||||
from fastapi.responses import JSONResponse
|
||||
from slowapi import _rate_limit_exceeded_handler
|
||||
from slowapi.errors import RateLimitExceeded
|
||||
|
||||
from shonar import __version__
|
||||
from shonar.api.v1 import api_router
|
||||
from shonar.core.config import get_settings
|
||||
from shonar.core.ratelimit import limiter
|
||||
from shonar.db.session import dispose_engine, get_engine
|
||||
|
||||
logging.basicConfig(level=logging.INFO)
|
||||
logger = logging.getLogger("shonar")
|
||||
|
||||
|
||||
@asynccontextmanager
|
||||
async def lifespan(app: FastAPI):
|
||||
settings = get_settings()
|
||||
get_engine() # validate URL parses; connections are lazy
|
||||
for warning in settings.validate_production():
|
||||
logger.warning("CONFIG: %s", warning)
|
||||
if settings.auto_migrate:
|
||||
from shonar.db.migrate import upgrade_head
|
||||
|
||||
await upgrade_head(settings.database_url)
|
||||
inline = settings.queue_backend.strip().lower() == "inline"
|
||||
if inline:
|
||||
from shonar.services import inline_queue
|
||||
|
||||
await inline_queue.start()
|
||||
yield
|
||||
if inline:
|
||||
from shonar.services import inline_queue
|
||||
|
||||
await inline_queue.stop()
|
||||
await dispose_engine()
|
||||
|
||||
|
||||
app = FastAPI(
|
||||
title="SHONAR API",
|
||||
description=(
|
||||
"SHONAR — Self-hosted Oral Notes and Audio Recorder. REST API. "
|
||||
"All data stays on your server."
|
||||
),
|
||||
version=__version__,
|
||||
lifespan=lifespan,
|
||||
# Interactive docs at the conventional paths.
|
||||
docs_url="/docs",
|
||||
redoc_url="/redoc",
|
||||
openapi_url="/openapi.json",
|
||||
)
|
||||
|
||||
app.state.limiter = limiter
|
||||
app.add_exception_handler(RateLimitExceeded, _rate_limit_exceeded_handler)
|
||||
|
||||
|
||||
@app.exception_handler(Exception)
|
||||
async def unhandled_exception_handler(request: Request, exc: Exception):
|
||||
"""Safe error surface: log details server-side, never leak internals."""
|
||||
logger.exception("Unhandled error on %s %s", request.method, request.url.path)
|
||||
return JSONResponse(status_code=500, content={"detail": "Internal server error"})
|
||||
|
||||
|
||||
app.include_router(api_router)
|
||||
|
||||
|
||||
@app.get("/")
|
||||
async def root():
|
||||
return {"app": "SHONAR", "docs": "/docs", "health": "/api/v1/healthz"}
|
||||
0
backend/shonar/services/__init__.py
Normal file
0
backend/shonar/services/__init__.py
Normal file
160
backend/shonar/services/ai/__init__.py
Normal file
160
backend/shonar/services/ai/__init__.py
Normal file
|
|
@ -0,0 +1,160 @@
|
|||
"""AI provider interfaces (M7).
|
||||
|
||||
Two independent axes, both optional and both configured only through
|
||||
environment variables (never hard-coded keys):
|
||||
|
||||
- transcription: none | whisper_http | faster_whisper
|
||||
- LLM (summary/action items): none | openai_compat | ollama
|
||||
|
||||
"none" is a first-class choice: recording, sync, playback, and manual
|
||||
transcripts work with no AI configured at all. The pipeline treats a
|
||||
missing provider as "skip this stage", never as an error.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Protocol
|
||||
|
||||
from shonar.core.config import Settings
|
||||
|
||||
|
||||
class AIError(Exception):
|
||||
"""Base for AI failures. Messages must be user-safe: they surface in
|
||||
``processing_error`` and therefore on screens."""
|
||||
|
||||
|
||||
class ProviderConfigError(AIError):
|
||||
"""Persistent misconfiguration (bad credentials, unknown model, missing
|
||||
dependency). Fails the job immediately — retrying cannot help."""
|
||||
|
||||
|
||||
class ProviderTransientError(AIError):
|
||||
"""May succeed on retry (timeouts, 429/5xx). The worker requeues these
|
||||
up to the job's max attempts."""
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class Segment:
|
||||
start: float
|
||||
end: float
|
||||
text: str
|
||||
speaker: str | None = None
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class TranscriptResult:
|
||||
text: str
|
||||
language: str | None
|
||||
segments: list[Segment] = field(default_factory=list)
|
||||
model: str = ""
|
||||
|
||||
|
||||
class TranscriptionProvider(Protocol):
|
||||
name: str
|
||||
|
||||
async def transcribe(
|
||||
self,
|
||||
audio: bytes,
|
||||
mime: str,
|
||||
*,
|
||||
language_hint: str | None = None,
|
||||
on_progress=None, # optional Callable[[int], None], 0..99
|
||||
) -> TranscriptResult: ...
|
||||
|
||||
|
||||
SUMMARY_KEYS = ("short", "detailed", "key_points", "decisions", "action_items", "questions")
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class SummaryResult:
|
||||
short: str = ""
|
||||
detailed: str = ""
|
||||
key_points: tuple[str, ...] = ()
|
||||
decisions: tuple[str, ...] = ()
|
||||
action_items: tuple[str, ...] = ()
|
||||
questions: tuple[str, ...] = ()
|
||||
model: str = ""
|
||||
|
||||
def to_dict(self) -> dict:
|
||||
return {
|
||||
"short": self.short,
|
||||
"detailed": self.detailed,
|
||||
"key_points": list(self.key_points),
|
||||
"decisions": list(self.decisions),
|
||||
"action_items": list(self.action_items),
|
||||
"questions": list(self.questions),
|
||||
}
|
||||
|
||||
@classmethod
|
||||
def from_dict(cls, raw: dict, model: str = "") -> SummaryResult:
|
||||
def text(key: str) -> str:
|
||||
v = raw.get(key)
|
||||
return v if isinstance(v, str) else ""
|
||||
|
||||
def strs(key: str) -> tuple[str, ...]:
|
||||
v = raw.get(key)
|
||||
if not isinstance(v, list):
|
||||
return ()
|
||||
return tuple(s for s in v if isinstance(s, str) and s.strip())
|
||||
|
||||
return cls(
|
||||
short=text("short"),
|
||||
detailed=text("detailed"),
|
||||
key_points=strs("key_points"),
|
||||
decisions=strs("decisions"),
|
||||
action_items=strs("action_items"),
|
||||
questions=strs("questions"),
|
||||
model=model,
|
||||
)
|
||||
|
||||
|
||||
class LlmProvider(Protocol):
|
||||
name: str
|
||||
|
||||
async def summarize(self, transcript: str, *, title: str | None = None) -> SummaryResult: ...
|
||||
|
||||
|
||||
def get_transcription_provider(settings: Settings) -> TranscriptionProvider | None:
|
||||
"""None means "transcription stage skipped", never an error."""
|
||||
kind = settings.transcription_provider.strip().lower()
|
||||
if kind in ("", "none"):
|
||||
return None
|
||||
if kind == "whisper_http":
|
||||
from shonar.services.ai.whisper_http import WhisperHttpProvider
|
||||
|
||||
return WhisperHttpProvider(
|
||||
base_url=settings.transcription_base_url,
|
||||
model=settings.transcription_model,
|
||||
api_key=settings.transcription_api_key,
|
||||
)
|
||||
if kind == "faster_whisper":
|
||||
from shonar.services.ai.faster_whisper import FasterWhisperProvider
|
||||
|
||||
return FasterWhisperProvider(model=settings.transcription_model)
|
||||
raise ProviderConfigError(
|
||||
f"Unknown transcription provider: {settings.transcription_provider!r}"
|
||||
)
|
||||
|
||||
|
||||
def get_llm_provider(settings: Settings) -> LlmProvider | None:
|
||||
"""None means "summary stage skipped", never an error."""
|
||||
kind = settings.llm_provider.strip().lower()
|
||||
if kind in ("", "none"):
|
||||
return None
|
||||
if kind == "openai_compat":
|
||||
from shonar.services.ai.openai_compat import OpenAICompatProvider
|
||||
|
||||
return OpenAICompatProvider(
|
||||
base_url=settings.llm_base_url,
|
||||
model=settings.llm_model,
|
||||
api_key=settings.llm_api_key,
|
||||
)
|
||||
if kind == "ollama":
|
||||
from shonar.services.ai.ollama import OllamaProvider
|
||||
|
||||
return OllamaProvider(
|
||||
base_url=settings.llm_base_url,
|
||||
model=settings.llm_model,
|
||||
)
|
||||
raise ProviderConfigError(f"Unknown LLM provider: {settings.llm_provider!r}")
|
||||
35
backend/shonar/services/ai/_llm.py
Normal file
35
backend/shonar/services/ai/_llm.py
Normal file
|
|
@ -0,0 +1,35 @@
|
|||
"""Shared summary contract: one system prompt, one JSON shape.
|
||||
|
||||
The model replies with JSON only; partial replies are accepted and missing
|
||||
keys default to empty (a terse-but-valid summary beats a failed job).
|
||||
Transcript input is truncated to bound context — a way in, not a report.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from shonar.services.ai import SummaryResult
|
||||
|
||||
SYSTEM_PROMPT = (
|
||||
"You summarize voice recordings for the speaker's own later reference. "
|
||||
"Reply with JSON only, exactly these keys: "
|
||||
'{"short": "1-2 sentences", "detailed": "a faithful paragraph", '
|
||||
'"key_points": [], "decisions": [], "action_items": [], "questions": []}. '
|
||||
"Empty arrays when absent. Never invent names, dates, or commitments "
|
||||
"not stated in the transcript."
|
||||
)
|
||||
|
||||
MAX_TRANSCRIPT_CHARS = 12_000
|
||||
|
||||
|
||||
def build_user_message(transcript: str, title: str | None) -> str:
|
||||
text = transcript[:MAX_TRANSCRIPT_CHARS]
|
||||
if len(transcript) > MAX_TRANSCRIPT_CHARS:
|
||||
text += f"\n\n[truncated from {len(transcript)} chars]"
|
||||
head = f'Title: "{title}"\n\n' if title else ""
|
||||
return head + "Transcript:\n" + text
|
||||
|
||||
|
||||
def parse_summary(data: object, model: str) -> SummaryResult:
|
||||
if not isinstance(data, dict):
|
||||
return SummaryResult(model=model)
|
||||
return SummaryResult.from_dict(data, model=model)
|
||||
108
backend/shonar/services/ai/faster_whisper.py
Normal file
108
backend/shonar/services/ai/faster_whisper.py
Normal file
|
|
@ -0,0 +1,108 @@
|
|||
"""Local transcription via faster-whisper (optional dependency).
|
||||
|
||||
Runs fully on this machine: audio never leaves the server for this stage.
|
||||
The import is lazy so the base install (and every test run) works without
|
||||
the heavyweight dependency.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import tempfile
|
||||
import threading
|
||||
from pathlib import Path
|
||||
|
||||
from shonar.services.ai import ProviderConfigError, Segment, TranscriptResult
|
||||
from shonar.services.ai.model_registry import (
|
||||
download_instructions,
|
||||
is_model_downloaded,
|
||||
validate_model_name,
|
||||
)
|
||||
|
||||
# One loaded model per name, shared across jobs in this worker process.
|
||||
# ctranslate2 inference is thread-safe; a lock serializes first-load only.
|
||||
_MODEL_CACHE: dict[str, object] = {}
|
||||
_MODEL_CACHE_LOCK = threading.Lock()
|
||||
|
||||
|
||||
def _load_model(name: str) -> object:
|
||||
from faster_whisper import WhisperModel
|
||||
|
||||
with _MODEL_CACHE_LOCK:
|
||||
model = _MODEL_CACHE.get(name)
|
||||
if model is None:
|
||||
model = WhisperModel(name, device="auto")
|
||||
_MODEL_CACHE[name] = model
|
||||
return model
|
||||
|
||||
|
||||
def ensure_model_available(name: str) -> None:
|
||||
"""Fail fast with download instructions instead of triggering a surprise
|
||||
multi-GB download inside a transcription job."""
|
||||
if not is_model_downloaded(name):
|
||||
raise ProviderConfigError(download_instructions(name))
|
||||
|
||||
|
||||
class FasterWhisperProvider:
|
||||
name = "faster_whisper"
|
||||
|
||||
def __init__(self, model: str = "base") -> None:
|
||||
try:
|
||||
import faster_whisper # noqa: F401
|
||||
except ImportError as e:
|
||||
raise ProviderConfigError(
|
||||
"faster-whisper is not installed (pip install shonar-backend[faster-whisper])."
|
||||
) from e
|
||||
self.model = validate_model_name(model or "base")
|
||||
|
||||
async def transcribe(
|
||||
self,
|
||||
audio: bytes,
|
||||
mime: str,
|
||||
*,
|
||||
language_hint: str | None = None,
|
||||
on_progress=None, # Callable[[int], None] | None — 0..99 percent
|
||||
) -> TranscriptResult:
|
||||
# faster-whisper is blocking CPU work: keep it off the event loop.
|
||||
return await asyncio.to_thread(self._run, audio, language_hint, on_progress)
|
||||
|
||||
def _run(
|
||||
self,
|
||||
audio: bytes,
|
||||
language_hint: str | None,
|
||||
on_progress=None,
|
||||
) -> TranscriptResult:
|
||||
path: Path | None = None
|
||||
try:
|
||||
with tempfile.NamedTemporaryFile(suffix=".m4a", delete=False) as f:
|
||||
f.write(audio)
|
||||
path = Path(f.name)
|
||||
ensure_model_available(self.model)
|
||||
model = _load_model(self.model)
|
||||
segments_iter, info = model.transcribe( # type: ignore[union-attr]
|
||||
str(path),
|
||||
beam_size=5,
|
||||
language=language_hint,
|
||||
)
|
||||
duration = float(getattr(info, "duration", 0.0) or 0.0)
|
||||
segments = []
|
||||
last_pct = -1
|
||||
for s in segments_iter:
|
||||
segments.append(Segment(start=s.start, end=s.end, text=s.text.strip()))
|
||||
if on_progress is not None and duration > 0:
|
||||
# Throttle: report only on whole-percent gains. Capped
|
||||
# at 99 — the caller commits 100 when the row finishes.
|
||||
pct = min(99, int(s.end / duration * 100))
|
||||
if pct > last_pct:
|
||||
last_pct = pct
|
||||
on_progress(pct)
|
||||
text = " ".join(s.text for s in segments).strip()
|
||||
return TranscriptResult(
|
||||
text=text,
|
||||
language=getattr(info, "language", None),
|
||||
segments=segments,
|
||||
model=self.model,
|
||||
)
|
||||
finally:
|
||||
if path is not None:
|
||||
path.unlink(missing_ok=True)
|
||||
165
backend/shonar/services/ai/model_registry.py
Normal file
165
backend/shonar/services/ai/model_registry.py
Normal file
|
|
@ -0,0 +1,165 @@
|
|||
"""Transcription model registry (Stage 1).
|
||||
|
||||
The supported faster-whisper sizes, their display metadata, validation, and
|
||||
local availability checks. Nothing here downloads anything: faster-whisper
|
||||
fetches from HuggingFace on first use, so "downloaded" is answered by
|
||||
inspecting the HF hub cache, and "available" additionally requires the
|
||||
faster-whisper package itself.
|
||||
|
||||
No silent substitution anywhere: unknown names are rejected, missing
|
||||
downloads fail fast with instructions.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
from dataclasses import dataclass
|
||||
from pathlib import Path
|
||||
|
||||
from sqlalchemy import select
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from shonar.db.models import AppSetting
|
||||
from shonar.services.ai import ProviderConfigError
|
||||
|
||||
DEFAULT_MODEL = "base"
|
||||
DEFAULT_MODEL_KEY = "transcription.default_model"
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class TranscriptionModelInfo:
|
||||
name: str
|
||||
display_name: str
|
||||
description: str
|
||||
params: str
|
||||
approx_memory: str
|
||||
relative_speed: str
|
||||
|
||||
|
||||
SUPPORTED_TRANSCRIPTION_MODELS: dict[str, TranscriptionModelInfo] = {
|
||||
"tiny": TranscriptionModelInfo(
|
||||
name="tiny",
|
||||
display_name="Tiny",
|
||||
description="Fastest and lightest. Good for quick drafts and slow machines.",
|
||||
params="~39M",
|
||||
approx_memory="~1 GB RAM",
|
||||
relative_speed="~10x real-time (CPU)",
|
||||
),
|
||||
"base": TranscriptionModelInfo(
|
||||
name="base",
|
||||
display_name="Base (default)",
|
||||
description="Balanced default. Works reasonably well on ordinary computers.",
|
||||
params="~74M",
|
||||
approx_memory="~1 GB RAM",
|
||||
relative_speed="~7x real-time (CPU)",
|
||||
),
|
||||
"small": TranscriptionModelInfo(
|
||||
name="small",
|
||||
display_name="Small",
|
||||
description="Better accuracy with higher resource usage.",
|
||||
params="~244M",
|
||||
approx_memory="~2 GB RAM",
|
||||
relative_speed="~4x real-time (CPU)",
|
||||
),
|
||||
"medium": TranscriptionModelInfo(
|
||||
name="medium",
|
||||
display_name="Medium",
|
||||
description="Higher accuracy and slower performance.",
|
||||
params="~769M",
|
||||
approx_memory="~5 GB RAM",
|
||||
relative_speed="~2x real-time (CPU)",
|
||||
),
|
||||
"large-v3": TranscriptionModelInfo(
|
||||
name="large-v3",
|
||||
display_name="Large v3",
|
||||
description="Highest accuracy and greatest resource requirements.",
|
||||
params="~1.5B",
|
||||
approx_memory="~10 GB RAM",
|
||||
relative_speed="~1x real-time (CPU)",
|
||||
),
|
||||
}
|
||||
|
||||
|
||||
def normalize_model_name(raw: str | None) -> str:
|
||||
"""Case/whitespace-tolerant normalization. Never maps one model to another."""
|
||||
return (raw or "").strip().lower()
|
||||
|
||||
|
||||
def validate_model_name(raw: str | None) -> str:
|
||||
"""Return the normalized name, or raise with a helpful message."""
|
||||
name = normalize_model_name(raw)
|
||||
if name in SUPPORTED_TRANSCRIPTION_MODELS:
|
||||
return name
|
||||
supported = ", ".join(sorted(SUPPORTED_TRANSCRIPTION_MODELS))
|
||||
raise ProviderConfigError(
|
||||
f"Unsupported transcription model {raw!r}. Supported models: {supported}. "
|
||||
"Check the spelling — a different model is never substituted silently."
|
||||
)
|
||||
|
||||
|
||||
def _hf_hub_cache() -> Path:
|
||||
try:
|
||||
from huggingface_hub.constants import HF_HUB_CACHE
|
||||
|
||||
return Path(HF_HUB_CACHE)
|
||||
except ImportError:
|
||||
return Path(
|
||||
os.environ.get("HF_HUB_CACHE", str(Path.home() / ".cache" / "huggingface" / "hub"))
|
||||
)
|
||||
|
||||
|
||||
def is_model_downloaded(name: str) -> bool:
|
||||
"""True when a non-empty faster-whisper snapshot for `name` sits in the
|
||||
HuggingFace hub cache (repo Systran/faster-whisper-<name>)."""
|
||||
repo_dir = _hf_hub_cache() / f"models--Systran--faster-whisper-{name}"
|
||||
snapshots = repo_dir / "snapshots"
|
||||
if not snapshots.is_dir():
|
||||
return False
|
||||
return any(s.is_dir() and any(s.iterdir()) for s in snapshots.iterdir())
|
||||
|
||||
|
||||
def is_faster_whisper_installed() -> bool:
|
||||
try:
|
||||
import faster_whisper # noqa: F401
|
||||
|
||||
return True
|
||||
except ImportError:
|
||||
return False
|
||||
|
||||
|
||||
def download_instructions(name: str) -> str:
|
||||
return (
|
||||
f'Model "{name}" is not downloaded. Download it with: '
|
||||
f"POST /api/v1/models/{name}/download "
|
||||
"(needs internet once), or run any transcription with that model selected — "
|
||||
"faster-whisper fetches it from HuggingFace automatically."
|
||||
)
|
||||
|
||||
|
||||
async def get_global_default_model(session: AsyncSession) -> str:
|
||||
row = await session.scalar(select(AppSetting).where(AppSetting.key == DEFAULT_MODEL_KEY))
|
||||
if row is None:
|
||||
return DEFAULT_MODEL
|
||||
model = row.value.get("model") if isinstance(row.value, dict) else None
|
||||
return model if model in SUPPORTED_TRANSCRIPTION_MODELS else DEFAULT_MODEL
|
||||
|
||||
|
||||
async def set_global_default_model(session: AsyncSession, raw: str) -> str:
|
||||
"""Validate + persist the global default. Affects future recordings only —
|
||||
existing rows keep their saved model."""
|
||||
name = validate_model_name(raw)
|
||||
row = await session.scalar(select(AppSetting).where(AppSetting.key == DEFAULT_MODEL_KEY))
|
||||
if row is None:
|
||||
row = AppSetting(key=DEFAULT_MODEL_KEY, value={"model": name})
|
||||
session.add(row)
|
||||
else:
|
||||
row.value = {"model": name}
|
||||
await session.flush()
|
||||
return name
|
||||
|
||||
|
||||
def effective_model(recording_model: str | None, global_default: str) -> str:
|
||||
"""Per-recording override wins; otherwise the global default."""
|
||||
if recording_model and recording_model in SUPPORTED_TRANSCRIPTION_MODELS:
|
||||
return recording_model
|
||||
return global_default
|
||||
78
backend/shonar/services/ai/ollama.py
Normal file
78
backend/shonar/services/ai/ollama.py
Normal file
|
|
@ -0,0 +1,78 @@
|
|||
"""Summaries via a local Ollama server (`/api/chat`, JSON mode)."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import AsyncIterator
|
||||
from contextlib import asynccontextmanager
|
||||
|
||||
import httpx
|
||||
|
||||
from shonar.services.ai import ProviderConfigError, ProviderTransientError, SummaryResult
|
||||
from shonar.services.ai._llm import SYSTEM_PROMPT, build_user_message, parse_summary
|
||||
|
||||
|
||||
class OllamaProvider:
|
||||
name = "ollama"
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
base_url: str,
|
||||
model: str = "",
|
||||
timeout_s: float = 900.0,
|
||||
http_client: httpx.AsyncClient | None = None,
|
||||
) -> None:
|
||||
if not base_url.strip():
|
||||
raise ProviderConfigError("ollama needs SHONAR_LLM_BASE_URL.")
|
||||
if not model.strip():
|
||||
raise ProviderConfigError("ollama needs SHONAR_LLM_MODEL.")
|
||||
self.base_url = base_url.rstrip("/")
|
||||
self.model = model
|
||||
self.timeout_s = timeout_s
|
||||
self.http_client = http_client
|
||||
|
||||
@asynccontextmanager
|
||||
async def _client(self) -> AsyncIterator[httpx.AsyncClient]:
|
||||
if self.http_client is not None:
|
||||
yield self.http_client
|
||||
else:
|
||||
async with httpx.AsyncClient(timeout=self.timeout_s) as client:
|
||||
yield client
|
||||
|
||||
async def summarize(self, transcript: str, *, title: str | None = None) -> SummaryResult:
|
||||
payload = {
|
||||
"model": self.model,
|
||||
"stream": False,
|
||||
"format": "json",
|
||||
# qwen3-family models "think" by default: a long reasoning
|
||||
# chain before the JSON answer, brutally slow on CPU and it
|
||||
# does not improve the summary. Ask for the answer directly
|
||||
# (ignored by non-thinking models).
|
||||
"think": False,
|
||||
"options": {"num_ctx": 8192},
|
||||
"messages": [
|
||||
{"role": "system", "content": SYSTEM_PROMPT},
|
||||
{"role": "user", "content": build_user_message(transcript, title)},
|
||||
],
|
||||
}
|
||||
try:
|
||||
async with self._client() as client:
|
||||
resp = await client.post(f"{self.base_url}/api/chat", json=payload)
|
||||
except (httpx.TimeoutException, httpx.TransportError) as e:
|
||||
raise ProviderTransientError(f"Ollama unreachable: {type(e).__name__}") from e
|
||||
if resp.status_code == 404:
|
||||
# Missing model and missing route both 404 here; both are
|
||||
# configuration, not weather.
|
||||
raise ProviderConfigError("Ollama has no such model or route (HTTP 404).")
|
||||
if resp.status_code != 200:
|
||||
raise ProviderTransientError(f"Summarization failed (HTTP {resp.status_code}).")
|
||||
try:
|
||||
content = resp.json()["message"]["content"]
|
||||
except (ValueError, KeyError, TypeError) as e:
|
||||
raise ProviderTransientError("Ollama sent an unreadable reply.") from e
|
||||
import json as _json
|
||||
|
||||
try:
|
||||
data = _json.loads(content)
|
||||
except ValueError as e:
|
||||
raise ProviderTransientError("Ollama reply was not JSON.") from e
|
||||
return parse_summary(data, self.model)
|
||||
82
backend/shonar/services/ai/openai_compat.py
Normal file
82
backend/shonar/services/ai/openai_compat.py
Normal file
|
|
@ -0,0 +1,82 @@
|
|||
"""Summaries via any OpenAI-compatible chat endpoint (self-hosted
|
||||
vLLM/llama.cpp server, commercial API, …) with JSON mode.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import AsyncIterator
|
||||
from contextlib import asynccontextmanager
|
||||
|
||||
import httpx
|
||||
|
||||
from shonar.services.ai import ProviderConfigError, ProviderTransientError, SummaryResult
|
||||
from shonar.services.ai._llm import SYSTEM_PROMPT, build_user_message, parse_summary
|
||||
|
||||
|
||||
class OpenAICompatProvider:
|
||||
name = "openai_compat"
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
base_url: str,
|
||||
model: str = "",
|
||||
api_key: str = "",
|
||||
timeout_s: float = 180.0,
|
||||
http_client: httpx.AsyncClient | None = None,
|
||||
) -> None:
|
||||
if not base_url.strip():
|
||||
raise ProviderConfigError("openai_compat needs SHONAR_LLM_BASE_URL.")
|
||||
if not model.strip():
|
||||
raise ProviderConfigError("openai_compat needs SHONAR_LLM_MODEL.")
|
||||
self.base_url = base_url.rstrip("/")
|
||||
self.model = model
|
||||
self.api_key = api_key
|
||||
self.timeout_s = timeout_s
|
||||
self.http_client = http_client
|
||||
|
||||
@asynccontextmanager
|
||||
async def _client(self) -> AsyncIterator[httpx.AsyncClient]:
|
||||
if self.http_client is not None:
|
||||
yield self.http_client
|
||||
else:
|
||||
async with httpx.AsyncClient(timeout=self.timeout_s) as client:
|
||||
yield client
|
||||
|
||||
async def summarize(self, transcript: str, *, title: str | None = None) -> SummaryResult:
|
||||
headers = (
|
||||
{"Authorization": f"Bearer {self.api_key}"} if self.api_key else {}
|
||||
)
|
||||
payload = {
|
||||
"model": self.model,
|
||||
"temperature": 0.2,
|
||||
"response_format": {"type": "json_object"},
|
||||
"messages": [
|
||||
{"role": "system", "content": SYSTEM_PROMPT},
|
||||
{"role": "user", "content": build_user_message(transcript, title)},
|
||||
],
|
||||
}
|
||||
try:
|
||||
async with self._client() as client:
|
||||
resp = await client.post(
|
||||
f"{self.base_url}/v1/chat/completions", headers=headers, json=payload
|
||||
)
|
||||
except (httpx.TimeoutException, httpx.TransportError) as e:
|
||||
raise ProviderTransientError(f"LLM unreachable: {type(e).__name__}") from e
|
||||
if resp.status_code in (401, 403, 404):
|
||||
raise ProviderConfigError(f"LLM refused the request (HTTP {resp.status_code}).")
|
||||
if resp.status_code == 429 or resp.status_code >= 500:
|
||||
raise ProviderTransientError(f"LLM busy (HTTP {resp.status_code}).")
|
||||
if resp.status_code != 200:
|
||||
raise ProviderTransientError(f"Summarization failed (HTTP {resp.status_code}).")
|
||||
try:
|
||||
body = resp.json()
|
||||
content = body["choices"][0]["message"]["content"]
|
||||
except (ValueError, KeyError, IndexError, TypeError) as e:
|
||||
raise ProviderTransientError("LLM sent an unreadable reply.") from e
|
||||
import json as _json
|
||||
|
||||
try:
|
||||
data = _json.loads(content)
|
||||
except ValueError as e:
|
||||
raise ProviderTransientError("LLM reply was not JSON.") from e
|
||||
return parse_summary(data, self.model)
|
||||
128
backend/shonar/services/ai/whisper_http.py
Normal file
128
backend/shonar/services/ai/whisper_http.py
Normal file
|
|
@ -0,0 +1,128 @@
|
|||
"""Transcription via any OpenAI-compatible `/v1/audio/transcriptions`
|
||||
endpoint (self-hosted whisper.cpp server, commercial Whisper API, …).
|
||||
Sends `verbose_json` so segment timings come back with the text.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import AsyncIterator
|
||||
from contextlib import asynccontextmanager
|
||||
|
||||
import httpx
|
||||
|
||||
from shonar.services.ai import (
|
||||
ProviderConfigError,
|
||||
ProviderTransientError,
|
||||
Segment,
|
||||
TranscriptResult,
|
||||
)
|
||||
|
||||
|
||||
class WhisperHttpProvider:
|
||||
name = "whisper_http"
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
base_url: str,
|
||||
model: str = "base",
|
||||
api_key: str = "",
|
||||
timeout_s: float = 300.0,
|
||||
http_client: httpx.AsyncClient | None = None,
|
||||
) -> None:
|
||||
if not base_url.strip():
|
||||
raise ProviderConfigError(
|
||||
"whisper_http needs SHONAR_TRANSCRIPTION_BASE_URL."
|
||||
)
|
||||
self.base_url = base_url.rstrip("/")
|
||||
self.model = model
|
||||
self.api_key = api_key
|
||||
self.timeout_s = timeout_s
|
||||
self.http_client = http_client
|
||||
|
||||
@asynccontextmanager
|
||||
async def _client(self) -> AsyncIterator[httpx.AsyncClient]:
|
||||
if self.http_client is not None:
|
||||
yield self.http_client
|
||||
else:
|
||||
async with httpx.AsyncClient(timeout=self.timeout_s) as client:
|
||||
yield client
|
||||
|
||||
async def transcribe(
|
||||
self,
|
||||
audio: bytes,
|
||||
mime: str,
|
||||
*,
|
||||
language_hint: str | None = None,
|
||||
on_progress=None, # accepted for protocol parity; not reported
|
||||
) -> TranscriptResult:
|
||||
headers = (
|
||||
{"Authorization": f"Bearer {self.api_key}"} if self.api_key else {}
|
||||
)
|
||||
data: dict[str, str] = {"model": self.model, "response_format": "verbose_json"}
|
||||
if language_hint:
|
||||
data["language"] = language_hint
|
||||
files = {"file": (f"audio.{_ext(mime)}", audio, mime or "application/octet-stream")}
|
||||
try:
|
||||
async with self._client() as client:
|
||||
resp = await client.post(
|
||||
f"{self.base_url}/v1/audio/transcriptions",
|
||||
headers=headers,
|
||||
data=data,
|
||||
files=files,
|
||||
)
|
||||
except (httpx.TimeoutException, httpx.TransportError) as e:
|
||||
raise ProviderTransientError(
|
||||
f"Transcription service unreachable: {type(e).__name__}"
|
||||
) from e
|
||||
if resp.status_code in (401, 403, 404):
|
||||
raise ProviderConfigError(
|
||||
f"Transcription service refused the request (HTTP {resp.status_code})."
|
||||
)
|
||||
if resp.status_code == 429 or resp.status_code >= 500:
|
||||
raise ProviderTransientError(
|
||||
f"Transcription service busy (HTTP {resp.status_code})."
|
||||
)
|
||||
if resp.status_code != 200:
|
||||
raise ProviderTransientError(
|
||||
f"Transcription failed (HTTP {resp.status_code})."
|
||||
)
|
||||
try:
|
||||
body = resp.json()
|
||||
except ValueError as e:
|
||||
raise ProviderTransientError("Transcription service sent no JSON.") from e
|
||||
segments = []
|
||||
raw_segs = body.get("segments")
|
||||
if isinstance(raw_segs, list):
|
||||
for s in raw_segs:
|
||||
if not isinstance(s, dict):
|
||||
continue
|
||||
try:
|
||||
segments.append(
|
||||
Segment(
|
||||
start=float(s.get("start", 0.0)),
|
||||
end=float(s.get("end", 0.0)),
|
||||
text=str(s.get("text", "")),
|
||||
)
|
||||
)
|
||||
except (TypeError, ValueError):
|
||||
continue
|
||||
text = body.get("text")
|
||||
return TranscriptResult(
|
||||
text=text if isinstance(text, str) else "",
|
||||
language=body.get("language") if isinstance(body.get("language"), str) else None,
|
||||
segments=segments,
|
||||
model=self.model,
|
||||
)
|
||||
|
||||
|
||||
def _ext(mime: str) -> str:
|
||||
return {
|
||||
"audio/mp4": "m4a",
|
||||
"audio/m4a": "m4a",
|
||||
"audio/wav": "wav",
|
||||
"audio/x-wav": "wav",
|
||||
"audio/ogg": "ogg",
|
||||
"audio/opus": "ogg",
|
||||
"audio/webm": "webm",
|
||||
"audio/mpeg": "mp3",
|
||||
}.get(mime.lower().split(";")[0].strip(), "bin")
|
||||
167
backend/shonar/services/auth.py
Normal file
167
backend/shonar/services/auth.py
Normal file
|
|
@ -0,0 +1,167 @@
|
|||
"""Authentication service: registration, login, rotating refresh tokens with
|
||||
reuse detection, logout, account deletion."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import uuid
|
||||
from datetime import UTC, datetime
|
||||
|
||||
from sqlalchemy import select, update
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from shonar.core.security import (
|
||||
generate_refresh_token,
|
||||
hash_password,
|
||||
hash_refresh_token,
|
||||
refresh_token_ttl,
|
||||
verify_password,
|
||||
)
|
||||
from shonar.db.models import Device, RefreshToken, User, utcnow
|
||||
|
||||
|
||||
class AuthError(Exception):
|
||||
"""Safe, client-displayable auth failure (never leaks which half failed
|
||||
beyond what the flow requires)."""
|
||||
|
||||
def __init__(self, message: str, status_code: int = 401):
|
||||
super().__init__(message)
|
||||
self.message = message
|
||||
self.status_code = status_code
|
||||
|
||||
|
||||
async def register_user(
|
||||
session: AsyncSession, email: str, password: str, display_name: str | None
|
||||
) -> User:
|
||||
email = email.strip().lower()
|
||||
existing = await session.scalar(select(User).where(User.email == email))
|
||||
if existing is not None:
|
||||
# Use a generic message; do not reveal whether the account exists in
|
||||
# flows where that matters. For self-hosted registration the UX cost
|
||||
# of "email already registered" is acceptable and helpful.
|
||||
raise AuthError("An account with this email already exists.", 409)
|
||||
user = User(
|
||||
email=email,
|
||||
password_hash=hash_password(password),
|
||||
display_name=display_name,
|
||||
)
|
||||
session.add(user)
|
||||
await session.flush()
|
||||
return user
|
||||
|
||||
|
||||
async def issue_refresh_token(
|
||||
session: AsyncSession,
|
||||
user_id: uuid.UUID,
|
||||
family: uuid.UUID | None,
|
||||
device_id: uuid.UUID | None,
|
||||
) -> tuple[str, RefreshToken]:
|
||||
token = generate_refresh_token()
|
||||
rt = RefreshToken(
|
||||
user_id=user_id,
|
||||
token_hash=hash_refresh_token(token),
|
||||
family=family or uuid.uuid4(),
|
||||
device_id=device_id,
|
||||
expires_at=datetime.now(UTC) + refresh_token_ttl(),
|
||||
)
|
||||
session.add(rt)
|
||||
await session.flush()
|
||||
return token, rt
|
||||
|
||||
|
||||
async def login(
|
||||
session: AsyncSession,
|
||||
email: str,
|
||||
password: str,
|
||||
device_name: str | None,
|
||||
platform: str,
|
||||
) -> tuple[User, str, Device]:
|
||||
"""Returns (user, refresh_token, device). Raises AuthError safely."""
|
||||
email = email.strip().lower()
|
||||
user = await session.scalar(select(User).where(User.email == email))
|
||||
if user is None or user.deleted_at is not None or not user.is_active:
|
||||
raise AuthError("Invalid email or password.")
|
||||
if not verify_password(user.password_hash, password):
|
||||
raise AuthError("Invalid email or password.")
|
||||
|
||||
device = Device(user_id=user.id, name=device_name or "Android device", platform=platform)
|
||||
session.add(device)
|
||||
await session.flush()
|
||||
|
||||
refresh_token, _ = await issue_refresh_token(session, user.id, None, device.id)
|
||||
return user, refresh_token, device
|
||||
|
||||
|
||||
async def rotate_refresh_token(
|
||||
session: AsyncSession, presented_token: str
|
||||
) -> tuple[User, str, uuid.UUID | None]:
|
||||
"""Consume a refresh token and issue a replacement in the same family.
|
||||
|
||||
Reuse detection: presenting an already-consumed/revoked token revokes the
|
||||
entire family (an attacker's stolen token dies along with the real one).
|
||||
"""
|
||||
token_hash = hash_refresh_token(presented_token)
|
||||
rt = await session.scalar(select(RefreshToken).where(RefreshToken.token_hash == token_hash))
|
||||
now = utcnow()
|
||||
|
||||
if rt is None:
|
||||
raise AuthError("Invalid refresh token.")
|
||||
|
||||
if rt.revoked_at is not None or rt.replaced_by is not None:
|
||||
# REUSE DETECTED — revoke the whole family. Commit BEFORE raising:
|
||||
# the request's transaction would otherwise roll back on the 401 and
|
||||
# silently undo the security-revocation.
|
||||
await session.execute(
|
||||
update(RefreshToken)
|
||||
.where(RefreshToken.family == rt.family, RefreshToken.revoked_at.is_(None))
|
||||
.values(revoked_at=now)
|
||||
)
|
||||
await session.commit()
|
||||
raise AuthError("Refresh token reuse detected. Please log in again.", 401)
|
||||
|
||||
if rt.expires_at < now:
|
||||
raise AuthError("Refresh token expired.", 401)
|
||||
|
||||
user = await session.get(User, rt.user_id)
|
||||
if user is None or user.deleted_at is not None or not user.is_active:
|
||||
raise AuthError("Account unavailable.", 401)
|
||||
|
||||
new_token, new_rt = await issue_refresh_token(session, user.id, rt.family, rt.device_id)
|
||||
rt.revoked_at = now
|
||||
rt.replaced_by = new_rt.id
|
||||
|
||||
if rt.device_id is not None:
|
||||
device = await session.get(Device, rt.device_id)
|
||||
if device is not None:
|
||||
device.last_seen_at = now
|
||||
await session.flush()
|
||||
return user, new_token, rt.device_id
|
||||
|
||||
|
||||
async def logout(session: AsyncSession, presented_token: str) -> None:
|
||||
"""Revoke the presented token's whole family (logs the device out)."""
|
||||
token_hash = hash_refresh_token(presented_token)
|
||||
rt = await session.scalar(select(RefreshToken).where(RefreshToken.token_hash == token_hash))
|
||||
if rt is None:
|
||||
return # idempotent
|
||||
await session.execute(
|
||||
update(RefreshToken)
|
||||
.where(RefreshToken.family == rt.family, RefreshToken.revoked_at.is_(None))
|
||||
.values(revoked_at=utcnow())
|
||||
)
|
||||
|
||||
|
||||
async def delete_account(session: AsyncSession, user: User, password: str) -> None:
|
||||
if not verify_password(user.password_hash, password):
|
||||
raise AuthError("Invalid password.", 403)
|
||||
now = utcnow()
|
||||
user.deleted_at = now
|
||||
user.is_active = False
|
||||
# Revoke every refresh token for the user.
|
||||
await session.execute(
|
||||
update(RefreshToken)
|
||||
.where(RefreshToken.user_id == user.id, RefreshToken.revoked_at.is_(None))
|
||||
.values(revoked_at=now)
|
||||
)
|
||||
# NOTE: hard deletion of rows/files is performed by the retention sweep
|
||||
# (services/retention.py) so an accidental deletion can be cancelled
|
||||
# within the grace window (see docs/security.md).
|
||||
251
backend/shonar/services/exports.py
Normal file
251
backend/shonar/services/exports.py
Normal file
|
|
@ -0,0 +1,251 @@
|
|||
"""Exports (M9): audio, transcript txt, notes markdown, bundle zip.
|
||||
|
||||
Synchronous generation — every artifact is small enough (text, or one
|
||||
audio file) that a background job adds failure modes, not speed. Each
|
||||
successful export records an ExportJob row and stores the produced bytes
|
||||
as an Asset(kind=export) so the audit trail exists; the response is the
|
||||
file itself (no separate download-asset round trip).
|
||||
|
||||
Formats:
|
||||
audio — the original upload, byte-identical, original mime/extension
|
||||
txt — current transcript text
|
||||
md — notes.md: title, metadata, notes, summary sections, transcript
|
||||
zip — bundle: original audio + transcript.txt + notes.md
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import io
|
||||
import uuid
|
||||
import zipfile
|
||||
from dataclasses import dataclass
|
||||
|
||||
from sqlalchemy import select
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from shonar.db.models import (
|
||||
Asset,
|
||||
AssetKind,
|
||||
ExportJob,
|
||||
JobStatus,
|
||||
Recording,
|
||||
Summary,
|
||||
Transcript,
|
||||
)
|
||||
|
||||
|
||||
class ExportError(Exception):
|
||||
def __init__(self, status_code: int, message: str):
|
||||
super().__init__(message)
|
||||
self.status_code = status_code
|
||||
self.message = message
|
||||
|
||||
|
||||
EXPORT_FORMATS = ("audio", "txt", "md", "zip")
|
||||
|
||||
|
||||
@dataclass
|
||||
class ExportResult:
|
||||
filename: str
|
||||
mime_type: str
|
||||
data: bytes
|
||||
|
||||
|
||||
def _safe_filename(rec: Recording) -> str:
|
||||
"""Slug the title; fall back to the recorded_at stamp. Never leaks ids."""
|
||||
base = "".join(
|
||||
c if (c.isalnum() or c in "-_ ") else " " for c in (rec.title or "")
|
||||
).strip()
|
||||
if not base:
|
||||
base = f"recording-{rec.recorded_at:%Y%m%d-%H%M%S}"
|
||||
return base[:120]
|
||||
|
||||
|
||||
async def _current_transcript(session: AsyncSession, rec_id: uuid.UUID) -> Transcript | None:
|
||||
return await session.scalar(
|
||||
select(Transcript)
|
||||
.where(Transcript.recording_id == rec_id, Transcript.superseded_at.is_(None))
|
||||
.order_by(Transcript.version.desc())
|
||||
)
|
||||
|
||||
|
||||
async def _current_summary(session: AsyncSession, rec_id: uuid.UUID) -> Summary | None:
|
||||
return await session.scalar(
|
||||
select(Summary)
|
||||
.where(Summary.recording_id == rec_id, Summary.superseded_at.is_(None))
|
||||
.order_by(Summary.version.desc())
|
||||
)
|
||||
|
||||
|
||||
async def _original_asset(session: AsyncSession, rec_id: uuid.UUID) -> Asset | None:
|
||||
return await session.scalar(
|
||||
select(Asset).where(Asset.recording_id == rec_id, Asset.kind == AssetKind.original)
|
||||
)
|
||||
|
||||
|
||||
def _summary_md(summary: Summary | None) -> str:
|
||||
if summary is None:
|
||||
return ""
|
||||
c = summary.content or {}
|
||||
out = ["## Summary\n"]
|
||||
if c.get("short"):
|
||||
out.append(f"{c['short']}\n")
|
||||
if c.get("detailed"):
|
||||
out.append(f"### Detailed\n\n{c['detailed']}\n")
|
||||
for key, header in (
|
||||
("key_points", "Key Points"),
|
||||
("decisions", "Decisions"),
|
||||
("action_items", "Action Items"),
|
||||
("questions", "Questions"),
|
||||
):
|
||||
items = c.get(key)
|
||||
if isinstance(items, list) and items:
|
||||
out.append(f"### {header}\n")
|
||||
out.extend(f"- {x}" for x in items)
|
||||
out.append("")
|
||||
return "\n".join(out)
|
||||
|
||||
|
||||
def _transcript_md(t: Transcript | None) -> str:
|
||||
if t is None:
|
||||
return ""
|
||||
lines = ["## Transcript\n"]
|
||||
segs = [s for s in (t.segments or []) if isinstance(s, dict)]
|
||||
if segs:
|
||||
for s in segs:
|
||||
start = float(s.get("start", 0.0))
|
||||
stamp = f"{int(start // 60):02d}:{start % 60:04.1f}"
|
||||
speaker = f"**{s['speaker']}**: " if s.get("speaker") else ""
|
||||
lines.append(f"- `[{stamp}]` {speaker}{s.get('text', '').strip()}")
|
||||
else:
|
||||
lines.append(t.text or "")
|
||||
return "\n".join(lines) + "\n"
|
||||
|
||||
|
||||
def _notes_md(
|
||||
rec: Recording, t: Transcript | None, summary: Summary | None,
|
||||
tag_names: list[str] | None = None,
|
||||
) -> str:
|
||||
parts = [
|
||||
f"# {rec.title}\n",
|
||||
f"- Recorded: {rec.recorded_at:%Y-%m-%d %H:%M} UTC",
|
||||
f"- Duration: {rec.duration_seconds:.1f}s",
|
||||
]
|
||||
if tag_names:
|
||||
parts.append("- Tags: " + ", ".join(f"`{x}`" for x in tag_names))
|
||||
parts.append("")
|
||||
if rec.notes:
|
||||
parts.append(f"## Notes\n\n{rec.notes}\n")
|
||||
s = _summary_md(summary)
|
||||
if s:
|
||||
parts.append(s)
|
||||
tr = _transcript_md(t)
|
||||
if tr:
|
||||
parts.append(tr)
|
||||
return "\n".join(parts)
|
||||
|
||||
|
||||
def _zip(files: list[tuple[str, bytes]]) -> bytes:
|
||||
buf = io.BytesIO()
|
||||
with zipfile.ZipFile(buf, "w", zipfile.ZIP_DEFLATED) as z:
|
||||
for name, data in files:
|
||||
z.writestr(name, data)
|
||||
return buf.getvalue()
|
||||
|
||||
|
||||
async def build_export(
|
||||
session: AsyncSession, rec: Recording, export_format: str
|
||||
) -> ExportResult:
|
||||
"""Build one export artifact for an owned, non-deleted recording."""
|
||||
from shonar.storage import get_storage
|
||||
|
||||
if export_format not in EXPORT_FORMATS:
|
||||
raise ExportError(422, f"Unknown export format. Use one of: {', '.join(EXPORT_FORMATS)}")
|
||||
|
||||
original = await _original_asset(session, rec.id)
|
||||
t = await _current_transcript(session, rec.id)
|
||||
summary = await _current_summary(session, rec.id)
|
||||
from shonar.services.search import tag_names
|
||||
|
||||
tags = await tag_names(session, rec.id)
|
||||
base = _safe_filename(rec)
|
||||
|
||||
if export_format == "audio":
|
||||
if original is None:
|
||||
raise ExportError(404, "No audio stored for this recording")
|
||||
data = await get_storage().get(original.storage_key)
|
||||
ext = original.storage_key[original.storage_key.rfind(".") :]
|
||||
return ExportResult(filename=f"{base}{ext}", mime_type=original.mime_type, data=data)
|
||||
|
||||
if export_format == "txt":
|
||||
if t is None:
|
||||
raise ExportError(404, "No transcript yet — transcribe first")
|
||||
return ExportResult(
|
||||
filename=f"{base}.txt", mime_type="text/plain; charset=utf-8",
|
||||
data=(t.text or "").encode("utf-8"),
|
||||
)
|
||||
|
||||
if export_format == "md":
|
||||
if t is None and summary is None and not rec.notes:
|
||||
raise ExportError(
|
||||
404, "Nothing to export — this recording has no notes, transcript, or summary"
|
||||
)
|
||||
return ExportResult(
|
||||
filename=f"{base}.md", mime_type="text/markdown; charset=utf-8",
|
||||
data=_notes_md(rec, t, summary, tags).encode("utf-8"),
|
||||
)
|
||||
|
||||
# zip bundle: whatever exists, always at least the audio when present.
|
||||
if original is None and t is None and summary is None and not rec.notes:
|
||||
raise ExportError(404, "Nothing to export for this recording")
|
||||
files: list[tuple[str, bytes]] = []
|
||||
if original is not None:
|
||||
audio = await get_storage().get(original.storage_key)
|
||||
ext = original.storage_key[original.storage_key.rfind(".") :]
|
||||
files.append((f"{base}{ext}", audio))
|
||||
if t is not None:
|
||||
files.append(("transcript.txt", (t.text or "").encode("utf-8")))
|
||||
files.append(("notes.md", _notes_md(rec, t, summary, tags).encode("utf-8")))
|
||||
return ExportResult(
|
||||
filename=f"{base}.zip", mime_type="application/zip", data=_zip(files)
|
||||
)
|
||||
|
||||
|
||||
async def record_export(
|
||||
session: AsyncSession, user_id: uuid.UUID, rec: Recording,
|
||||
export_format: str, result: ExportResult,
|
||||
) -> None:
|
||||
"""Persist the audit trail: ExportJob(succeeded) + Asset(kind=export).
|
||||
|
||||
Best-effort storage of the artifact bytes; a storage failure never
|
||||
fails the download the user already received.
|
||||
"""
|
||||
from shonar.storage import get_storage
|
||||
|
||||
job = ExportJob(
|
||||
user_id=user_id,
|
||||
recording_id=rec.id,
|
||||
export_type=export_format,
|
||||
status=JobStatus.succeeded,
|
||||
)
|
||||
session.add(job)
|
||||
try:
|
||||
key = f"exports/{user_id}/{rec.id}/{export_format}-{uuid.uuid4().hex}"
|
||||
await get_storage().put(key, result.data)
|
||||
import hashlib
|
||||
|
||||
asset = Asset(
|
||||
recording_id=rec.id,
|
||||
user_id=user_id,
|
||||
kind=AssetKind.export,
|
||||
storage_key=key,
|
||||
mime_type=result.mime_type,
|
||||
size_bytes=len(result.data),
|
||||
checksum_sha256=hashlib.sha256(result.data).hexdigest(),
|
||||
)
|
||||
session.add(asset)
|
||||
await session.flush()
|
||||
job.asset_id = asset.id
|
||||
except Exception: # noqa: BLE001 — audit copy is best-effort
|
||||
pass
|
||||
await session.flush()
|
||||
141
backend/shonar/services/inline_queue.py
Normal file
141
backend/shonar/services/inline_queue.py
Normal file
|
|
@ -0,0 +1,141 @@
|
|||
"""In-process job runner (desktop bundled-lite engine).
|
||||
|
||||
``queue_backend=inline`` replaces the arq/Redis transport with a single
|
||||
asyncio consumer inside the uvicorn process: DB rows stay the source of
|
||||
truth (ProcessingJob), this just runs the work. One job at a time —
|
||||
local faster-whisper is multi-GB per pass, mirroring the worker's
|
||||
``max_jobs=1`` rule.
|
||||
|
||||
Transient failures retry in-process with backoff up to MAX_TRIES (the
|
||||
same budget arq gives via ``max_tries``); the final failure is recorded
|
||||
by the task body itself (see processing._fail).
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import contextlib
|
||||
import logging
|
||||
|
||||
from shonar.db.models import JobType
|
||||
from shonar.services import processing
|
||||
|
||||
logger = logging.getLogger("shonar.inline_queue")
|
||||
|
||||
RETRY_DELAY_SECONDS = 5.0
|
||||
# Desktop engine: hard-delete expired soft-deletes once the app has been up
|
||||
# for a day, then daily. Delayed first run keeps startup snappy.
|
||||
RETENTION_INTERVAL_SECONDS = 24 * 3600.0
|
||||
|
||||
_queue: asyncio.Queue[tuple[str, str, int]] | None = None
|
||||
_consumer: asyncio.Task | None = None
|
||||
_retention: asyncio.Task | None = None
|
||||
|
||||
|
||||
def _get_queue() -> asyncio.Queue[tuple[str, str, int]]:
|
||||
global _queue
|
||||
if _queue is None:
|
||||
_queue = asyncio.Queue()
|
||||
return _queue
|
||||
|
||||
|
||||
async def start() -> None:
|
||||
"""Start the consumer and re-run anything the DB says is pending."""
|
||||
global _consumer, _retention
|
||||
_get_queue()
|
||||
# Inline engine = single process: any row still marked `running` at
|
||||
# startup is a corpse from the previous process (the worker died with
|
||||
# it). sweep_stale's 2h live-worker grace — correct for multi-worker
|
||||
# arq deployments — would starve these jobs, so reclaim them first.
|
||||
from sqlalchemy import update
|
||||
|
||||
from shonar.db.models import JobStatus, ProcessingJob
|
||||
from shonar.db.session import session_factory
|
||||
|
||||
async with session_factory()() as s:
|
||||
await s.execute(
|
||||
update(ProcessingJob)
|
||||
.where(ProcessingJob.status == JobStatus.running)
|
||||
.values(status=JobStatus.queued, started_at=None)
|
||||
)
|
||||
await s.commit()
|
||||
if _consumer is None or _consumer.done():
|
||||
_consumer = asyncio.create_task(_consume(), name="shonar-inline-queue")
|
||||
if _retention is None or _retention.done():
|
||||
_retention = asyncio.create_task(_retention_loop(), name="shonar-retention")
|
||||
# Crash recovery: queued rows (and orphaned running rows requeued by
|
||||
# the sweep) go back on the in-process queue.
|
||||
count = await processing.sweep_stale()
|
||||
if count:
|
||||
logger.info("inline queue startup sweep requeued %d jobs", count)
|
||||
|
||||
|
||||
async def _retention_loop() -> None:
|
||||
"""Daily hard-delete sweep for the desktop engine (no arq cron here).
|
||||
|
||||
First pass a few minutes after start (the app may only run for hours
|
||||
at a time, so a full-day initial sleep could starve the sweep), then
|
||||
once a day while running.
|
||||
"""
|
||||
from shonar.services import retention
|
||||
|
||||
await asyncio.sleep(120.0)
|
||||
while True:
|
||||
try:
|
||||
purged = await retention.sweep_deleted()
|
||||
if purged["recordings"] or purged["users"]:
|
||||
logger.info("inline retention sweep: %s", purged)
|
||||
except asyncio.CancelledError:
|
||||
raise
|
||||
except Exception: # noqa: BLE001 — the loop must survive any failure
|
||||
logger.exception("inline retention sweep failed")
|
||||
await asyncio.sleep(RETENTION_INTERVAL_SECONDS)
|
||||
|
||||
|
||||
async def stop() -> None:
|
||||
global _consumer, _retention
|
||||
if _consumer is not None:
|
||||
_consumer.cancel()
|
||||
with contextlib.suppress(BaseException): # noqa: BLE001 — shutdown is best-effort
|
||||
await _consumer
|
||||
_consumer = None
|
||||
if _retention is not None:
|
||||
_retention.cancel()
|
||||
with contextlib.suppress(BaseException): # noqa: BLE001
|
||||
await _retention
|
||||
_retention = None
|
||||
|
||||
|
||||
async def enqueue(job_type: JobType, recording_id: str) -> None:
|
||||
await _get_queue().put((job_type.value, str(recording_id), 1))
|
||||
|
||||
|
||||
async def _consume() -> None:
|
||||
q = _get_queue()
|
||||
while True:
|
||||
job_value, recording_id, attempt = await q.get()
|
||||
ctx = {"job_try": attempt}
|
||||
try:
|
||||
if job_value == JobType.transcribe.value:
|
||||
await processing.run_transcribe(ctx, recording_id)
|
||||
else:
|
||||
await processing.run_summarize(ctx, recording_id)
|
||||
except processing.ProviderTransientError as e:
|
||||
if attempt < processing.MAX_TRIES:
|
||||
logger.warning(
|
||||
"inline job %s:%s transient failure (%s); retry %d/%d",
|
||||
job_value, recording_id, e, attempt + 1, processing.MAX_TRIES,
|
||||
)
|
||||
await asyncio.sleep(RETRY_DELAY_SECONDS)
|
||||
await q.put((job_value, recording_id, attempt + 1))
|
||||
else:
|
||||
logger.error(
|
||||
"inline job %s:%s failed after %d tries: %s",
|
||||
job_value, recording_id, attempt, e,
|
||||
)
|
||||
except asyncio.CancelledError:
|
||||
raise
|
||||
except Exception: # noqa: BLE001 — consumer must survive any task crash
|
||||
logger.exception("inline job %s:%s crashed", job_value, recording_id)
|
||||
finally:
|
||||
q.task_done()
|
||||
74
backend/shonar/services/media.py
Normal file
74
backend/shonar/services/media.py
Normal file
|
|
@ -0,0 +1,74 @@
|
|||
"""Audio format validation: declared MIME type vs magic bytes.
|
||||
|
||||
Families decouple declared MIME aliases (audio/mp4 vs audio/m4a) from what
|
||||
the bytes actually are. The original file is stored exactly as uploaded —
|
||||
validation never rewrites it.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
FAMILY_BY_MIME = {
|
||||
"audio/mp4": "mp4",
|
||||
"audio/m4a": "mp4",
|
||||
"audio/aac": "aac",
|
||||
"audio/wav": "wav",
|
||||
"audio/x-wav": "wav",
|
||||
"audio/ogg": "ogg",
|
||||
"audio/opus": "ogg",
|
||||
"audio/webm": "webm",
|
||||
"audio/mpeg": "mpeg",
|
||||
}
|
||||
|
||||
EXT_BY_FAMILY = {
|
||||
"mp4": ".m4a",
|
||||
"aac": ".aac",
|
||||
"wav": ".wav",
|
||||
"ogg": ".ogg",
|
||||
"webm": ".webm",
|
||||
"mpeg": ".mp3",
|
||||
}
|
||||
|
||||
|
||||
def sniff_audio_family(data: bytes) -> str | None:
|
||||
"""Return the audio family from magic bytes, or None if unrecognized."""
|
||||
if len(data) >= 12:
|
||||
if data[4:8] == b"ftyp":
|
||||
return "mp4"
|
||||
if data[:4] == b"RIFF" and data[8:12] == b"WAVE":
|
||||
return "wav"
|
||||
if data[:4] == b"OggS":
|
||||
return "ogg"
|
||||
if data[:4] == b"\x1a\x45\xdf\xa3":
|
||||
return "webm"
|
||||
if data[:3] == b"ID3":
|
||||
return "mpeg"
|
||||
if len(data) >= 2 and data[0] == 0xFF and (data[1] & 0xF6) == 0xF0:
|
||||
# ADTS frame sync: AAC (also accepted as mpeg-family audio)
|
||||
return "aac"
|
||||
return None
|
||||
|
||||
|
||||
def declared_family(mime_type: str) -> str | None:
|
||||
return FAMILY_BY_MIME.get(mime_type.lower().split(";")[0].strip())
|
||||
|
||||
|
||||
def is_compatible(mime_type: str, data: bytes) -> bool:
|
||||
"""True when the declared MIME matches the sniffed magic bytes.
|
||||
|
||||
mpeg and aac are treated as one family: Android records AAC in ADTS or
|
||||
in MP4 containers and MIME reporting around these is inconsistent.
|
||||
"""
|
||||
declared = declared_family(mime_type)
|
||||
sniffed = sniff_audio_family(data)
|
||||
if declared is None or sniffed is None:
|
||||
return False
|
||||
if {declared, sniffed} == {"mpeg", "aac"}:
|
||||
return True
|
||||
return declared == sniffed
|
||||
|
||||
|
||||
def extension_for(mime_type: str, data: bytes) -> str:
|
||||
sniffed = sniff_audio_family(data)
|
||||
if sniffed is not None:
|
||||
return EXT_BY_FAMILY[sniffed]
|
||||
return EXT_BY_FAMILY.get(declared_family(mime_type) or "", ".bin")
|
||||
591
backend/shonar/services/processing.py
Normal file
591
backend/shonar/services/processing.py
Normal file
|
|
@ -0,0 +1,591 @@
|
|||
"""AI processing pipeline (M7): uploaded -> transcribed -> summarized.
|
||||
|
||||
State lives in the database (ProcessingJob rows); arq/redis is transport
|
||||
only. That ordering is deliberate: if redis is down, uploads still succeed
|
||||
(the jobs sit queued) and the worker sweep picks them up. The only race is
|
||||
a task running before the API transaction commits (recording invisible) —
|
||||
missing rows are transient failures, so arq retry absorbs it.
|
||||
|
||||
Entry points:
|
||||
- ``enqueue_for_recording``: called from upload finalize AND re-callable
|
||||
later (manual transcripts in M8 re-enter here). Idempotent.
|
||||
- ``run_transcribe`` / ``run_summarize``: arq task bodies. ``ctx`` is an
|
||||
arq context in production and a plain dict in tests; only
|
||||
``ctx.get("job_try", 1)`` is read.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import logging
|
||||
import uuid
|
||||
from datetime import timedelta
|
||||
|
||||
from sqlalchemy import func, select
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from shonar.core.config import get_settings
|
||||
from shonar.db.models import (
|
||||
Asset,
|
||||
AssetKind,
|
||||
JobStatus,
|
||||
JobType,
|
||||
ProcessingJob,
|
||||
ProcessingStatus,
|
||||
Recording,
|
||||
Summary,
|
||||
Transcript,
|
||||
utcnow,
|
||||
)
|
||||
from shonar.services.ai import (
|
||||
AIError,
|
||||
ProviderTransientError,
|
||||
get_llm_provider,
|
||||
get_transcription_provider,
|
||||
)
|
||||
from shonar.storage import get_storage
|
||||
|
||||
logger = logging.getLogger("shonar.processing")
|
||||
|
||||
|
||||
def _accepts_on_progress(provider) -> bool:
|
||||
"""True when the provider's transcribe() accepts on_progress=."""
|
||||
import inspect
|
||||
|
||||
try:
|
||||
return "on_progress" in inspect.signature(provider.transcribe).parameters
|
||||
except (TypeError, ValueError): # builtins / exotic callables
|
||||
return False
|
||||
|
||||
|
||||
def _thread_progress_reporter(job_id):
|
||||
"""Callback safe to invoke from a worker thread (faster-whisper runs
|
||||
via asyncio.to_thread): schedules a tiny session update on the loop.
|
||||
|
||||
Progress writes are best-effort display state — failures are logged
|
||||
and swallowed, never allowed to disturb the transcription itself."""
|
||||
loop = asyncio.get_running_loop()
|
||||
|
||||
def report(pct: int) -> None:
|
||||
async def _write() -> None:
|
||||
from sqlalchemy import update
|
||||
|
||||
from shonar.db.session import session_factory
|
||||
|
||||
try:
|
||||
async with session_factory()() as s:
|
||||
await s.execute(
|
||||
update(ProcessingJob)
|
||||
.where(ProcessingJob.id == job_id)
|
||||
.values(progress=int(pct))
|
||||
)
|
||||
await s.commit()
|
||||
except Exception: # pragma: no cover - display state only
|
||||
logger.debug("progress write failed for job %s", job_id, exc_info=True)
|
||||
|
||||
try:
|
||||
asyncio.run_coroutine_threadsafe(_write(), loop)
|
||||
except RuntimeError: # loop already gone (shutdown race)
|
||||
logger.debug("progress dropped for job %s (loop gone)", job_id)
|
||||
|
||||
return report
|
||||
|
||||
MAX_TRIES = 3
|
||||
|
||||
# A `running` job younger than this is treated as live work, not a crash
|
||||
# relic — long transcriptions legitimately outrun the 5-minute sweep, and
|
||||
# re-enqueuing them mid-flight caused a copy storm (one whisper per copy).
|
||||
STALE_RUNNING_AFTER = timedelta(hours=2)
|
||||
|
||||
|
||||
# --- entry -----------------------------------------------------------------
|
||||
|
||||
|
||||
async def enqueue_for_recording(session: AsyncSession, rec: Recording) -> list[JobType]:
|
||||
"""Queue whatever AI stages apply. Safe to call repeatedly: completed
|
||||
work is never redone, failed work is reset for another attempt."""
|
||||
if rec.deleted_at is not None:
|
||||
return []
|
||||
settings = get_settings()
|
||||
tprov = get_transcription_provider(settings)
|
||||
lprov = get_llm_provider(settings)
|
||||
queued: list[JobType] = []
|
||||
if tprov is not None and await reset_or_create(session, rec, JobType.transcribe):
|
||||
queued.append(JobType.transcribe)
|
||||
if (
|
||||
lprov is not None
|
||||
and await latest_transcript_text(session, rec.id) is not None
|
||||
and await reset_or_create(session, rec, JobType.summarize)
|
||||
):
|
||||
queued.append(JobType.summarize)
|
||||
if tprov is None and lprov is None:
|
||||
if rec.processing_status != ProcessingStatus.ai_disabled:
|
||||
rec.processing_status = ProcessingStatus.ai_disabled
|
||||
rec.processing_error = None
|
||||
elif queued and rec.processing_status not in (
|
||||
ProcessingStatus.processing,
|
||||
ProcessingStatus.completed,
|
||||
):
|
||||
rec.processing_status = ProcessingStatus.queued
|
||||
rec.processing_error = None
|
||||
await session.flush()
|
||||
for jt in queued:
|
||||
await transport_enqueue(jt, rec.id)
|
||||
return queued
|
||||
|
||||
|
||||
async def reset_or_create(
|
||||
session: AsyncSession, rec: Recording, job_type: JobType
|
||||
) -> bool:
|
||||
"""Ensure a queued job row. Returns True when (re)queued now: new rows,
|
||||
plus failed/skipped rows (a re-upload or re-entry deserves another
|
||||
attempt). Queued/running/succeeded rows are left alone."""
|
||||
existing = await session.scalar(
|
||||
select(ProcessingJob)
|
||||
.where(
|
||||
ProcessingJob.recording_id == rec.id,
|
||||
ProcessingJob.job_type == job_type,
|
||||
)
|
||||
.order_by(ProcessingJob.id.desc())
|
||||
)
|
||||
if existing is None:
|
||||
session.add(ProcessingJob(recording_id=rec.id, job_type=job_type))
|
||||
return True
|
||||
if existing.status in (JobStatus.queued, JobStatus.running, JobStatus.succeeded):
|
||||
return False
|
||||
existing.status = JobStatus.queued
|
||||
existing.attempt = 0
|
||||
existing.error = None
|
||||
existing.stage = None
|
||||
existing.progress = None
|
||||
existing.started_at = None
|
||||
existing.finished_at = None
|
||||
return True
|
||||
|
||||
|
||||
async def latest_transcript_text(
|
||||
session: AsyncSession, recording_id: uuid.UUID
|
||||
) -> str | None:
|
||||
"""Newest non-superseded transcript; user-edited rows win over newer
|
||||
machine rows (an edit is a verdict, not a draft)."""
|
||||
rows = (
|
||||
await session.scalars(
|
||||
select(Transcript)
|
||||
.where(
|
||||
Transcript.recording_id == recording_id,
|
||||
Transcript.superseded_at.is_(None),
|
||||
)
|
||||
.order_by(Transcript.version.desc())
|
||||
)
|
||||
).all()
|
||||
if not rows:
|
||||
return None
|
||||
for r in rows:
|
||||
if r.edited_by_user and r.text.strip():
|
||||
return r.text
|
||||
text = rows[0].text
|
||||
return text if text.strip() else None
|
||||
|
||||
|
||||
async def transport_enqueue(job_type: JobType, recording_id: uuid.UUID) -> None:
|
||||
"""Best-effort trigger for the configured backend ('arq' or 'inline').
|
||||
Failure only logs: the DB rows are the real queue and the sweep picks
|
||||
up anything the transport missed."""
|
||||
if get_settings().queue_backend.strip().lower() == "inline":
|
||||
from shonar.services import inline_queue
|
||||
|
||||
await inline_queue.enqueue(job_type, str(recording_id))
|
||||
return
|
||||
from arq import create_pool
|
||||
from arq.connections import RedisSettings
|
||||
from arq.constants import result_key_prefix
|
||||
|
||||
job_id = f"{job_type.value}:{recording_id}"
|
||||
try:
|
||||
pool = await create_pool(RedisSettings.from_dsn(get_settings().redis_url))
|
||||
try:
|
||||
# Deterministic _job_id: arq refuses a duplicate while a copy of
|
||||
# (job_type, recording) is queued or running, so sweep pokes and
|
||||
# retried enqueues can never pile up concurrent copies of the
|
||||
# same work.
|
||||
#
|
||||
# arq dedupes on job key OR result key, and failed runs leave a
|
||||
# result behind for keep_result days — that would silently swallow
|
||||
# deliberate re-runs of previously-finished work. Drop only the
|
||||
# stale RESULT key first (never the job key: that is exactly what
|
||||
# shields a live copy from duplicate enqueue).
|
||||
await pool.delete(result_key_prefix + job_id)
|
||||
await pool.enqueue_job(
|
||||
"run_" + job_type.value,
|
||||
str(recording_id),
|
||||
_job_id=job_id,
|
||||
)
|
||||
finally:
|
||||
await pool.aclose()
|
||||
except Exception as e: # noqa: BLE001 — transport must never break uploads
|
||||
logger.warning("arq enqueue failed (%s); worker sweep will pick it up", e)
|
||||
|
||||
|
||||
# --- tasks -----------------------------------------------------------------
|
||||
|
||||
|
||||
def _job_try(ctx: dict) -> int:
|
||||
try:
|
||||
return int(ctx.get("job_try", 1))
|
||||
except (TypeError, ValueError):
|
||||
return 1
|
||||
|
||||
|
||||
async def _load(session: AsyncSession, recording_id: str) -> Recording | None:
|
||||
try:
|
||||
rid = uuid.UUID(recording_id)
|
||||
except ValueError:
|
||||
return None
|
||||
return await session.get(Recording, rid)
|
||||
|
||||
|
||||
async def _job(
|
||||
session: AsyncSession, rec: Recording, job_type: JobType
|
||||
) -> ProcessingJob:
|
||||
job = await session.scalar(
|
||||
select(ProcessingJob)
|
||||
.where(
|
||||
ProcessingJob.recording_id == rec.id,
|
||||
ProcessingJob.job_type == job_type,
|
||||
)
|
||||
.order_by(ProcessingJob.id.desc())
|
||||
)
|
||||
if job is None:
|
||||
job = ProcessingJob(recording_id=rec.id, job_type=job_type)
|
||||
session.add(job)
|
||||
await session.flush()
|
||||
return job
|
||||
|
||||
|
||||
async def _fail(
|
||||
session: AsyncSession,
|
||||
rec: Recording,
|
||||
job: ProcessingJob,
|
||||
message: str,
|
||||
ctx: dict,
|
||||
exc: AIError | None = None,
|
||||
) -> None:
|
||||
"""Config errors fail now; transient errors fail only on the last try
|
||||
(returning normally), otherwise they raise for arq retry."""
|
||||
transient = exc is None or isinstance(exc, ProviderTransientError)
|
||||
if transient and _job_try(ctx) < MAX_TRIES:
|
||||
# Running state was committed before the long phase; put the row
|
||||
# back to queued for the retry and persist that (a raise no longer
|
||||
# rolls the pre-phase commit back).
|
||||
job.status = JobStatus.queued
|
||||
job.attempt = _job_try(ctx)
|
||||
job.stage = None
|
||||
await session.commit()
|
||||
raise ProviderTransientError(message)
|
||||
job.status = JobStatus.failed
|
||||
job.error = message
|
||||
job.stage = None
|
||||
job.finished_at = utcnow()
|
||||
if job.job_type == JobType.summarize and await latest_transcript_text(
|
||||
session, rec.id
|
||||
):
|
||||
# The transcript is usable; a summary timeout must not mark the
|
||||
# whole recording failed (the summary can be re-run separately).
|
||||
rec.processing_status = ProcessingStatus.completed
|
||||
rec.processing_error = f"Summary failed: {message}"
|
||||
else:
|
||||
rec.processing_status = ProcessingStatus.failed
|
||||
rec.processing_error = message
|
||||
await session.flush()
|
||||
|
||||
|
||||
async def run_transcribe(ctx: dict, recording_id: str) -> None:
|
||||
"""Transcribe the original audio; chain into summarization when an LLM
|
||||
is configured."""
|
||||
from shonar.db.session import session_factory
|
||||
|
||||
settings = get_settings()
|
||||
async with session_factory()() as session:
|
||||
rec = await _load(session, recording_id)
|
||||
if rec is None or rec.deleted_at is not None:
|
||||
# Finalize race: the API transaction may not have committed yet.
|
||||
raise ProviderTransientError("Recording not ready; retrying.")
|
||||
job = await _job(session, rec, JobType.transcribe)
|
||||
provider = get_transcription_provider(settings)
|
||||
if provider is None:
|
||||
job.status = JobStatus.skipped
|
||||
await session.flush()
|
||||
await session.commit()
|
||||
await _maybe_chain_summarize(session, rec)
|
||||
return
|
||||
# The recording's saved model is authoritative: a per-recording
|
||||
# override wins, otherwise the global default in force at finalize
|
||||
# time (changing the default never rewrites history).
|
||||
from shonar.services.ai.model_registry import (
|
||||
effective_model,
|
||||
get_global_default_model,
|
||||
)
|
||||
|
||||
model = effective_model(rec.transcription_model, await get_global_default_model(session))
|
||||
if provider.name == "faster_whisper":
|
||||
from shonar.services.ai.faster_whisper import FasterWhisperProvider
|
||||
|
||||
try:
|
||||
provider = FasterWhisperProvider(model=model)
|
||||
except AIError as e:
|
||||
await _fail(session, rec, job, str(e), ctx, e)
|
||||
await session.commit()
|
||||
return
|
||||
original = await session.scalar(
|
||||
select(Asset).where(
|
||||
Asset.recording_id == rec.id, Asset.kind == AssetKind.original
|
||||
)
|
||||
)
|
||||
if original is None:
|
||||
raise ProviderTransientError("Audio not ready; retrying.")
|
||||
job.status = JobStatus.running
|
||||
job.attempt = _job_try(ctx)
|
||||
job.started_at = utcnow()
|
||||
job.stage = "loading-model"
|
||||
job.progress = None
|
||||
rec.processing_status = ProcessingStatus.processing
|
||||
rec.processing_error = None
|
||||
# Commit before the long CPU phase: an open transaction is invisible
|
||||
# to other readers (and on SQLite it locks out the progress writer).
|
||||
await session.commit()
|
||||
try:
|
||||
audio = await get_storage().get(original.storage_key)
|
||||
job.stage = "transcribing"
|
||||
await session.commit()
|
||||
kwargs = {}
|
||||
if _accepts_on_progress(provider):
|
||||
kwargs["on_progress"] = _thread_progress_reporter(job.id)
|
||||
result = await provider.transcribe(audio, original.mime_type, **kwargs)
|
||||
except AIError as e:
|
||||
await _fail(session, rec, job, str(e), ctx, e)
|
||||
await session.commit()
|
||||
return
|
||||
await store_transcript(
|
||||
session, rec, result.text, result.segments, result.language,
|
||||
provider.name, getattr(result, "model", ""),
|
||||
)
|
||||
job.status = JobStatus.succeeded
|
||||
job.stage = None
|
||||
job.progress = 100
|
||||
job.finished_at = utcnow()
|
||||
await session.flush()
|
||||
await _maybe_chain_summarize(session, rec)
|
||||
await session.commit()
|
||||
|
||||
|
||||
async def _maybe_chain_summarize(session: AsyncSession, rec: Recording) -> None:
|
||||
"""After transcription (or a skip): summarize when possible, else finish."""
|
||||
settings = get_settings()
|
||||
if get_llm_provider(settings) is None:
|
||||
if rec.processing_status != ProcessingStatus.completed:
|
||||
rec.processing_status = ProcessingStatus.completed
|
||||
rec.processing_error = None
|
||||
await session.flush()
|
||||
return
|
||||
if await latest_transcript_text(session, rec.id) is None:
|
||||
# LLM configured but nothing to summarize (e.g. empty transcript).
|
||||
existing = await session.scalar(
|
||||
select(ProcessingJob).where(
|
||||
ProcessingJob.recording_id == rec.id,
|
||||
ProcessingJob.job_type == JobType.summarize,
|
||||
)
|
||||
)
|
||||
if existing is not None and existing.status == JobStatus.queued:
|
||||
existing.status = JobStatus.skipped
|
||||
if rec.processing_status != ProcessingStatus.completed:
|
||||
rec.processing_status = ProcessingStatus.completed
|
||||
await session.flush()
|
||||
return
|
||||
if await reset_or_create(session, rec, JobType.summarize):
|
||||
await transport_enqueue(JobType.summarize, rec.id)
|
||||
# Status stays `processing` until the summarize task lands.
|
||||
|
||||
|
||||
async def run_summarize(ctx: dict, recording_id: str) -> None:
|
||||
"""Summarize the latest transcript into the structured summary shape."""
|
||||
from shonar.db.session import session_factory
|
||||
|
||||
settings = get_settings()
|
||||
async with session_factory()() as session:
|
||||
rec = await _load(session, recording_id)
|
||||
if rec is None or rec.deleted_at is not None:
|
||||
raise ProviderTransientError("Recording not ready; retrying.")
|
||||
job = await _job(session, rec, JobType.summarize)
|
||||
provider = get_llm_provider(settings)
|
||||
if provider is None:
|
||||
job.status = JobStatus.skipped
|
||||
await session.flush()
|
||||
await session.commit()
|
||||
return
|
||||
text = await latest_transcript_text(session, rec.id)
|
||||
if text is None:
|
||||
job.status = JobStatus.skipped
|
||||
await session.flush()
|
||||
if rec.processing_status != ProcessingStatus.completed:
|
||||
rec.processing_status = ProcessingStatus.completed
|
||||
await session.commit()
|
||||
return
|
||||
job.status = JobStatus.running
|
||||
job.attempt = _job_try(ctx)
|
||||
job.started_at = utcnow()
|
||||
job.stage = "summarizing"
|
||||
rec.processing_status = ProcessingStatus.processing
|
||||
await session.commit() # visible before the long LLM call
|
||||
try:
|
||||
result = await provider.summarize(text, title=rec.title)
|
||||
except AIError as e:
|
||||
await _fail(session, rec, job, str(e), ctx, e)
|
||||
await session.commit()
|
||||
return
|
||||
await store_summary(session, rec, result.to_dict(), provider.name, result.model)
|
||||
job.status = JobStatus.succeeded
|
||||
job.stage = None
|
||||
job.progress = 100
|
||||
job.finished_at = utcnow()
|
||||
rec.processing_status = ProcessingStatus.completed
|
||||
rec.processing_error = None
|
||||
await session.flush()
|
||||
await session.commit()
|
||||
|
||||
|
||||
async def store_transcript(
|
||||
session: AsyncSession,
|
||||
rec: Recording,
|
||||
text: str,
|
||||
segments: list,
|
||||
language: str | None,
|
||||
provider_name: str,
|
||||
model: str,
|
||||
) -> None:
|
||||
"""Insert a new auto version; a newest user-edited row wins instead and
|
||||
nothing is inserted (edits are verdicts)."""
|
||||
existing = (
|
||||
await session.scalars(
|
||||
select(Transcript)
|
||||
.where(
|
||||
Transcript.recording_id == rec.id,
|
||||
Transcript.superseded_at.is_(None),
|
||||
)
|
||||
.order_by(Transcript.version.desc())
|
||||
)
|
||||
).all()
|
||||
if existing and existing[0].edited_by_user:
|
||||
return
|
||||
now = utcnow()
|
||||
max_version = await session.scalar(
|
||||
select(func.max(Transcript.version)).where(Transcript.recording_id == rec.id)
|
||||
)
|
||||
for row in existing:
|
||||
row.superseded_at = now
|
||||
session.add(
|
||||
Transcript(
|
||||
recording_id=rec.id,
|
||||
version=(max_version or 0) + 1,
|
||||
language=language,
|
||||
provider=provider_name,
|
||||
model=model,
|
||||
text=text,
|
||||
segments=[
|
||||
{"start": s.start, "end": s.end, "text": s.text, "speaker": s.speaker}
|
||||
for s in segments
|
||||
],
|
||||
edited_by_user=False,
|
||||
)
|
||||
)
|
||||
await session.flush()
|
||||
|
||||
|
||||
async def store_summary(
|
||||
session: AsyncSession,
|
||||
rec: Recording,
|
||||
content: dict,
|
||||
provider_name: str,
|
||||
model: str,
|
||||
) -> None:
|
||||
existing = (
|
||||
await session.scalars(
|
||||
select(Summary)
|
||||
.where(
|
||||
Summary.recording_id == rec.id,
|
||||
Summary.superseded_at.is_(None),
|
||||
)
|
||||
.order_by(Summary.version.desc())
|
||||
)
|
||||
).all()
|
||||
if existing and existing[0].edited_by_user:
|
||||
return
|
||||
now = utcnow()
|
||||
max_version = await session.scalar(
|
||||
select(func.max(Summary.version)).where(Summary.recording_id == rec.id)
|
||||
)
|
||||
for row in existing:
|
||||
row.superseded_at = now
|
||||
session.add(
|
||||
Summary(
|
||||
recording_id=rec.id,
|
||||
version=(max_version or 0) + 1,
|
||||
provider=provider_name,
|
||||
model=model,
|
||||
content=content,
|
||||
edited_by_user=False,
|
||||
)
|
||||
)
|
||||
await session.flush()
|
||||
|
||||
|
||||
async def sweep_stale(limit: int = 100) -> int:
|
||||
"""Crash recovery + transport-loss backstop: requeue jobs stuck running
|
||||
or sitting queued, newest catastrophe first. Returns jobs re-enqueued."""
|
||||
from shonar.db.session import session_factory
|
||||
|
||||
count = 0
|
||||
async with session_factory()() as session:
|
||||
rows = (
|
||||
await session.scalars(
|
||||
select(ProcessingJob)
|
||||
.where(ProcessingJob.status.in_([JobStatus.queued, JobStatus.running]))
|
||||
.order_by(ProcessingJob.id.desc())
|
||||
.limit(limit)
|
||||
)
|
||||
).all()
|
||||
for job in rows:
|
||||
rec = await session.get(Recording, job.recording_id)
|
||||
if rec is None or rec.deleted_at is not None:
|
||||
job.status = JobStatus.skipped
|
||||
continue
|
||||
if job.status == JobStatus.running:
|
||||
if job.attempt >= MAX_TRIES:
|
||||
job.status = JobStatus.failed
|
||||
job.error = "Worker died too many times."
|
||||
job.finished_at = utcnow()
|
||||
continue
|
||||
# Live work, not a crash relic: a running job that started
|
||||
# recently belongs to a worker still chewing it (long audio
|
||||
# outruns the sweep cadence). Only orphaned runs — no
|
||||
# started_at, or older than the grace window — get requeued.
|
||||
if job.started_at is not None and (
|
||||
utcnow() - job.started_at
|
||||
) < STALE_RUNNING_AFTER:
|
||||
continue
|
||||
job.status = JobStatus.queued
|
||||
job.error = None
|
||||
count += 1
|
||||
await session.commit()
|
||||
# Transport outside the transaction: rows are the queue, this just pokes.
|
||||
async with session_factory()() as session:
|
||||
rows = (
|
||||
await session.scalars(
|
||||
select(ProcessingJob)
|
||||
.where(ProcessingJob.status == JobStatus.queued)
|
||||
.order_by(ProcessingJob.id.desc())
|
||||
.limit(limit)
|
||||
)
|
||||
).all()
|
||||
for job in rows:
|
||||
await transport_enqueue(job.job_type, job.recording_id)
|
||||
return count
|
||||
101
backend/shonar/services/retention.py
Normal file
101
backend/shonar/services/retention.py
Normal file
|
|
@ -0,0 +1,101 @@
|
|||
"""Retention sweep (M9): hard-delete what passed its grace window.
|
||||
|
||||
Two policies, both driven by ``deleted_at``:
|
||||
|
||||
* **Recordings** soft-deleted more than ``retention_grace_days`` ago have
|
||||
their rows removed (cascade cleans transcripts/summaries/jobs/tags) and
|
||||
every stored asset file deleted best-effort.
|
||||
* **Accounts** deleted more than ``retention_grace_days`` ago are hard
|
||||
deleted (user cascade takes their recordings/assets/devices/tokens);
|
||||
their storage files are collected the same way.
|
||||
|
||||
Running inside the grace window is a no-op, so an accidental delete stays
|
||||
cancellable until the sweep actually fires. The sweep is idempotent and
|
||||
safe to run on any cadence (worker cron + inline-queue timer).
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import contextlib
|
||||
import logging
|
||||
from datetime import timedelta
|
||||
|
||||
from sqlalchemy import delete as sql_delete
|
||||
from sqlalchemy import select
|
||||
|
||||
from shonar.core.config import get_settings
|
||||
from shonar.db.models import Asset, Recording, User, utcnow
|
||||
from shonar.db.session import session_factory
|
||||
from shonar.storage import get_storage
|
||||
|
||||
logger = logging.getLogger("shonar.retention")
|
||||
|
||||
|
||||
async def sweep_deleted(limit: int = 200) -> dict[str, int]:
|
||||
"""Hard-purge expired recordings and accounts. Returns counts."""
|
||||
grace = timedelta(days=get_settings().retention_grace_days)
|
||||
cutoff = utcnow() - grace
|
||||
purged = {"recordings": 0, "users": 0, "files": 0}
|
||||
storage = get_storage()
|
||||
|
||||
async with session_factory()() as session:
|
||||
# --- recordings (skip rows whose account is also expiring: the
|
||||
# user cascade below collects their files in one pass) ---
|
||||
expiring_users = select(User.id).where(
|
||||
User.deleted_at.is_not(None), User.deleted_at < cutoff
|
||||
)
|
||||
recs = list(
|
||||
await session.scalars(
|
||||
select(Recording)
|
||||
.where(
|
||||
Recording.deleted_at.is_not(None),
|
||||
Recording.deleted_at < cutoff,
|
||||
Recording.user_id.notin_(expiring_users),
|
||||
)
|
||||
.limit(limit)
|
||||
)
|
||||
)
|
||||
for rec in recs:
|
||||
assets = list(
|
||||
await session.scalars(select(Asset).where(Asset.recording_id == rec.id))
|
||||
)
|
||||
await session.delete(rec)
|
||||
await session.flush()
|
||||
for a in assets:
|
||||
with contextlib.suppress(Exception): # best effort; DB row is gone
|
||||
await storage.delete(a.storage_key)
|
||||
purged["files"] += 1
|
||||
purged["recordings"] += 1
|
||||
|
||||
# --- accounts ---
|
||||
users = list(
|
||||
await session.scalars(
|
||||
select(User)
|
||||
.where(User.deleted_at.is_not(None), User.deleted_at < cutoff)
|
||||
.limit(limit)
|
||||
)
|
||||
)
|
||||
for user in users:
|
||||
assets = list(
|
||||
await session.scalars(select(Asset).where(Asset.user_id == user.id))
|
||||
)
|
||||
keys = [a.storage_key for a in assets]
|
||||
# DB-level delete: the ORM would null the NOT NULL FKs of the
|
||||
# user's recordings before the ON DELETE CASCADE could fire.
|
||||
await session.execute(sql_delete(User).where(User.id == user.id))
|
||||
await session.flush()
|
||||
for key in keys:
|
||||
with contextlib.suppress(Exception):
|
||||
await storage.delete(key)
|
||||
purged["files"] += 1
|
||||
purged["users"] += 1
|
||||
|
||||
await session.commit()
|
||||
|
||||
if purged["recordings"] or purged["users"]:
|
||||
logger.info(
|
||||
"retention sweep purged %d recordings, %d accounts, %d files (grace %dd)",
|
||||
purged["recordings"], purged["users"], purged["files"],
|
||||
get_settings().retention_grace_days,
|
||||
)
|
||||
return purged
|
||||
292
backend/shonar/services/search.py
Normal file
292
backend/shonar/services/search.py
Normal file
|
|
@ -0,0 +1,292 @@
|
|||
"""Full-text search across recordings (M9).
|
||||
|
||||
Postgres uses the tsvector columns from migration ``fts0000000001``
|
||||
(title/notes, transcript text, summary content, tag names). SQLite (the
|
||||
desktop bundled-lite engine) falls back to a substring scan; libraries
|
||||
there are single-user and small, and the endpoint contract is identical.
|
||||
|
||||
This module is the whole SearchBackend seam — a Meilisearch/OpenSearch
|
||||
implementation would replace it, not the callers.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import uuid
|
||||
from dataclasses import dataclass
|
||||
|
||||
from sqlalchemy import select, text
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from shonar.db.models import Recording, RecordingTag, Summary, Tag, Transcript
|
||||
|
||||
# ``scope`` values accepted by the endpoint.
|
||||
SCOPES = ("all", "title", "notes", "transcript", "summary", "tag")
|
||||
|
||||
# Hit ordering: a title hit outranks a transcript hit.
|
||||
_FIELD_RANK = {"title": 0, "tag": 1, "notes": 2, "summary": 3, "transcript": 4}
|
||||
|
||||
|
||||
@dataclass
|
||||
class SearchHit:
|
||||
recording: Recording
|
||||
field: str # where the best match landed
|
||||
snippet: str
|
||||
|
||||
|
||||
def _snippet_around(value: str, at: int, width: int = 160) -> str:
|
||||
"""~``width`` chars centred on ``at``, word-bounded, with ellipses."""
|
||||
half = width // 2
|
||||
start = max(0, at - half)
|
||||
end = min(len(value), at + half)
|
||||
if start > 0:
|
||||
start = value.rfind(" ", 0, start) + 1 or start
|
||||
if end < len(value):
|
||||
nxt = value.find(" ", end)
|
||||
end = nxt if nxt != -1 else end
|
||||
prefix = "…" if start > 0 else ""
|
||||
suffix = "…" if end < len(value) else ""
|
||||
return f"{prefix}{value[start:end].strip()}{suffix}"
|
||||
|
||||
|
||||
async def tag_names(session: AsyncSession, recording_id: uuid.UUID) -> list[str]:
|
||||
return list(
|
||||
await session.scalars(
|
||||
select(Tag.name)
|
||||
.join(RecordingTag, RecordingTag.tag_id == Tag.id)
|
||||
.where(RecordingTag.recording_id == recording_id)
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
async def _current_transcript_text(session: AsyncSession, recording_id: uuid.UUID) -> str:
|
||||
row = await session.scalar(
|
||||
select(Transcript.text)
|
||||
.where(Transcript.recording_id == recording_id, Transcript.superseded_at.is_(None))
|
||||
.order_by(Transcript.version.desc())
|
||||
)
|
||||
return row or ""
|
||||
|
||||
|
||||
def _summary_flat(content: dict | None) -> str:
|
||||
if not content:
|
||||
return ""
|
||||
parts = [str(content.get("short", "")), str(content.get("detailed", ""))]
|
||||
for key in ("key_points", "decisions", "action_items", "questions"):
|
||||
v = content.get(key)
|
||||
if isinstance(v, list):
|
||||
parts.extend(str(x) for x in v)
|
||||
return " ".join(p for p in parts if p)
|
||||
|
||||
|
||||
# --- SQLite fallback ----------------------------------------------------------
|
||||
|
||||
|
||||
async def _search_sqlite(
|
||||
session: AsyncSession, user_id: uuid.UUID, q: str, scope: str,
|
||||
limit: int, offset: int,
|
||||
) -> tuple[list[SearchHit], int]:
|
||||
needles = [t for t in q.lower().split() if t]
|
||||
if not needles:
|
||||
return [], 0
|
||||
recs = list(
|
||||
await session.scalars(
|
||||
select(Recording)
|
||||
.where(Recording.user_id == user_id, Recording.deleted_at.is_(None))
|
||||
.order_by(Recording.recorded_at.desc())
|
||||
)
|
||||
)
|
||||
hits: list[SearchHit] = []
|
||||
for rec in recs:
|
||||
fields: list[tuple[str, str]] = []
|
||||
if scope in ("all", "title"):
|
||||
fields.append(("title", rec.title or ""))
|
||||
if scope in ("all", "notes"):
|
||||
fields.append(("notes", rec.notes or ""))
|
||||
if scope in ("all", "tag"):
|
||||
fields.append(("tag", " ".join(await tag_names(session, rec.id))))
|
||||
if scope in ("all", "transcript"):
|
||||
fields.append(("transcript", await _current_transcript_text(session, rec.id)))
|
||||
if scope in ("all", "summary"):
|
||||
content = await session.scalar(
|
||||
select(Summary.content)
|
||||
.where(Summary.recording_id == rec.id, Summary.superseded_at.is_(None))
|
||||
.order_by(Summary.version.desc())
|
||||
)
|
||||
fields.append(("summary", _summary_flat(content)))
|
||||
# A field matches when EVERY needle appears in it (AND semantics,
|
||||
# matching plainto_tsquery on the Postgres side).
|
||||
best: SearchHit | None = None
|
||||
for field, value in fields:
|
||||
low = value.lower()
|
||||
if not all(n in low for n in needles):
|
||||
continue
|
||||
at = low.find(needles[0])
|
||||
cand = SearchHit(
|
||||
recording=rec, field=field, snippet=_snippet_around(value, max(at, 0))
|
||||
)
|
||||
if best is None or _FIELD_RANK[field] < _FIELD_RANK[best.field]:
|
||||
best = cand
|
||||
if best is not None:
|
||||
hits.append(best)
|
||||
hits.sort(key=lambda h: (_FIELD_RANK[h.field], h.recording.recorded_at), reverse=False)
|
||||
return hits[offset : offset + limit], len(hits)
|
||||
|
||||
|
||||
# --- PostgreSQL tsvector path -------------------------------------------------
|
||||
|
||||
# CTE ``q`` carries the parsed tsquery so it is computed once. Notes are not
|
||||
# in the recordings search_vector weights the way we want headlines, so notes
|
||||
# match by substring like the SQLite path (title/notes share the vector; the
|
||||
# field classifier prefers 'title' when the vector hits).
|
||||
_PG_MATCHES = """
|
||||
WITH q AS (SELECT plainto_tsquery('simple', :q) AS ts),
|
||||
lt AS (
|
||||
SELECT DISTINCT ON (recording_id) recording_id, search_vector, text
|
||||
FROM transcripts WHERE superseded_at IS NULL
|
||||
ORDER BY recording_id, version DESC
|
||||
),
|
||||
ls AS (
|
||||
SELECT DISTINCT ON (recording_id) recording_id, content, search_vector
|
||||
FROM summaries WHERE superseded_at IS NULL
|
||||
ORDER BY recording_id, version DESC
|
||||
),
|
||||
tm AS (
|
||||
SELECT DISTINCT rt.recording_id,
|
||||
ts_headline('simple', t.name, q.ts,
|
||||
'StartSel=,StopSel=,MaxFragments=0') AS snip
|
||||
FROM recording_tags rt
|
||||
JOIN tags t ON t.id = rt.tag_id AND t.user_id = :uid, q
|
||||
),
|
||||
matched AS (
|
||||
SELECT r.id AS id,
|
||||
CASE
|
||||
WHEN :want_title AND r.search_vector @@ q.ts
|
||||
AND coalesce(r.title, '') <> ''
|
||||
THEN 'title'
|
||||
WHEN :want_tag AND tm.recording_id IS NOT NULL THEN 'tag'
|
||||
WHEN :want_notes AND coalesce(r.notes, '') ILIKE '%' || :raw || '%' THEN 'notes'
|
||||
WHEN :want_summary AND ls.search_vector @@ q.ts THEN 'summary'
|
||||
WHEN :want_transcript AND lt.search_vector @@ q.ts THEN 'transcript'
|
||||
END AS field,
|
||||
CASE
|
||||
WHEN :want_title AND r.search_vector @@ q.ts
|
||||
AND coalesce(r.title, '') <> ''
|
||||
THEN ts_headline('simple', r.title, q.ts,
|
||||
'StartSel=,StopSel=,MaxFragments=0,MaxWords=25')
|
||||
WHEN :want_tag AND tm.recording_id IS NOT NULL THEN tm.snip
|
||||
WHEN :want_notes AND coalesce(r.notes, '') ILIKE '%' || :raw || '%'
|
||||
THEN left(r.notes, 200)
|
||||
WHEN :want_summary AND ls.search_vector @@ q.ts
|
||||
THEN ts_headline('simple',
|
||||
coalesce(ls.content->>'short', '') || ' ' || coalesce(ls.content->>'detailed', ''),
|
||||
q.ts, 'StartSel=,StopSel=,MaxFragments=1,MinWords=10,MaxWords=25')
|
||||
WHEN :want_transcript AND lt.search_vector @@ q.ts
|
||||
THEN ts_headline('simple', coalesce(lt.text, ''), q.ts,
|
||||
'StartSel=,StopSel=,MaxFragments=1,MinWords=10,MaxWords=25')
|
||||
END AS snippet
|
||||
FROM recordings r
|
||||
CROSS JOIN q
|
||||
LEFT JOIN lt ON lt.recording_id = r.id
|
||||
LEFT JOIN ls ON ls.recording_id = r.id
|
||||
LEFT JOIN tm ON tm.recording_id = r.id
|
||||
WHERE r.user_id = :uid AND r.deleted_at IS NULL
|
||||
)
|
||||
SELECT id, field, snippet FROM matched
|
||||
WHERE field IS NOT NULL
|
||||
ORDER BY CASE field WHEN 'title' THEN 0 WHEN 'tag' THEN 1 WHEN 'notes' THEN 2
|
||||
WHEN 'summary' THEN 3 ELSE 4 END,
|
||||
id
|
||||
LIMIT :limit OFFSET :offset
|
||||
"""
|
||||
|
||||
_PG_COUNT = """
|
||||
WITH q AS (SELECT plainto_tsquery('simple', :q) AS ts),
|
||||
lt AS (
|
||||
SELECT DISTINCT ON (recording_id) recording_id, search_vector
|
||||
FROM transcripts WHERE superseded_at IS NULL
|
||||
),
|
||||
ls AS (
|
||||
SELECT DISTINCT ON (recording_id) recording_id, search_vector
|
||||
FROM summaries WHERE superseded_at IS NULL
|
||||
),
|
||||
tm AS (
|
||||
SELECT DISTINCT rt.recording_id
|
||||
FROM recording_tags rt
|
||||
JOIN tags t ON t.id = rt.tag_id AND t.user_id = :uid, q
|
||||
WHERE t.search_vector @@ q.ts
|
||||
)
|
||||
SELECT count(*)
|
||||
FROM recordings r
|
||||
CROSS JOIN q
|
||||
LEFT JOIN lt ON lt.recording_id = r.id
|
||||
LEFT JOIN ls ON ls.recording_id = r.id
|
||||
LEFT JOIN tm ON tm.recording_id = r.id
|
||||
WHERE r.user_id = :uid AND r.deleted_at IS NULL
|
||||
AND (
|
||||
(:want_title AND r.search_vector @@ q.ts) OR
|
||||
(:want_transcript AND lt.search_vector @@ q.ts) OR
|
||||
(:want_summary AND ls.search_vector @@ q.ts) OR
|
||||
(:want_tag AND tm.recording_id IS NOT NULL) OR
|
||||
(:want_notes AND coalesce(r.notes, '') ILIKE '%' || :raw || '%')
|
||||
)
|
||||
"""
|
||||
|
||||
|
||||
async def _search_postgres(
|
||||
session: AsyncSession, user_id: uuid.UUID, q: str, scope: str,
|
||||
limit: int, offset: int,
|
||||
) -> tuple[list[SearchHit], int]:
|
||||
# NOTE: the title field matches anything the recordings vector hits
|
||||
# (title + notes); notes-only hits surface under 'title' headlines from
|
||||
# the title text. Acceptable precision tradeoff for a GIN-indexed path.
|
||||
params = {
|
||||
"uid": user_id,
|
||||
"q": q,
|
||||
"raw": q,
|
||||
"limit": limit,
|
||||
"offset": offset,
|
||||
"want_title": scope in ("all", "title"),
|
||||
"want_transcript": scope in ("all", "transcript"),
|
||||
"want_summary": scope in ("all", "summary"),
|
||||
"want_tag": scope in ("all", "tag"),
|
||||
"want_notes": scope in ("all", "notes"),
|
||||
}
|
||||
rows = (await session.execute(text(_PG_MATCHES), params)).all()
|
||||
total = await session.scalar(text(_PG_COUNT), params) or 0
|
||||
if not rows:
|
||||
return [], total
|
||||
ids = [r[0] for r in rows]
|
||||
by_id = {
|
||||
rec.id: rec
|
||||
for rec in (
|
||||
await session.scalars(select(Recording).where(Recording.id.in_(ids)))
|
||||
).all()
|
||||
}
|
||||
hits = [
|
||||
SearchHit(recording=by_id[r.id], field=r.field, snippet=r.snippet or "")
|
||||
for r in rows
|
||||
if r.id in by_id
|
||||
]
|
||||
return hits, total
|
||||
|
||||
|
||||
# --- public API ----------------------------------------------------------------
|
||||
|
||||
|
||||
async def search_recordings(
|
||||
session: AsyncSession,
|
||||
user_id: uuid.UUID,
|
||||
q: str,
|
||||
*,
|
||||
scope: str = "all",
|
||||
limit: int = 20,
|
||||
offset: int = 0,
|
||||
) -> tuple[list[SearchHit], int]:
|
||||
"""Returns (page of hits ordered by field rank, total). Empty q → no hits."""
|
||||
q = q.strip()
|
||||
if not q:
|
||||
return [], 0
|
||||
dialect = session.bind.dialect.name if session.bind else "sqlite"
|
||||
if dialect == "postgresql":
|
||||
return await _search_postgres(session, user_id, q, scope, limit, offset)
|
||||
return await _search_sqlite(session, user_id, q, scope, limit, offset)
|
||||
342
backend/shonar/services/uploads.py
Normal file
342
backend/shonar/services/uploads.py
Normal file
|
|
@ -0,0 +1,342 @@
|
|||
"""Chunked, resumable upload sessions.
|
||||
|
||||
Flow:
|
||||
1. POST /uploads -> session (uuid), chunk size, expiry
|
||||
2. PUT /uploads/{id}/chunks/{n} (idempotent per index; GET status lists
|
||||
received indexes so clients resume)
|
||||
3. POST /uploads/{id}/finalize -> validates size + magic bytes,
|
||||
assembles the object, creates the
|
||||
immutable original Asset and the
|
||||
Recording (or updates the existing
|
||||
recording for a retried
|
||||
client_recording_id).
|
||||
|
||||
Storage keys are server-generated UUID paths; clients never see them.
|
||||
Originals are immutable: finalize never overwrites an existing original.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import contextlib
|
||||
import hashlib
|
||||
import uuid
|
||||
from datetime import datetime, timedelta
|
||||
|
||||
from sqlalchemy import select
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from shonar.core.config import get_settings
|
||||
from shonar.db.models import (
|
||||
Asset,
|
||||
AssetKind,
|
||||
ProcessingStatus,
|
||||
Recording,
|
||||
UploadChunk,
|
||||
UploadSession,
|
||||
UploadSessionStatus,
|
||||
utcnow,
|
||||
)
|
||||
from shonar.services.media import is_compatible
|
||||
from shonar.storage import get_storage
|
||||
|
||||
SESSION_TTL = timedelta(hours=24)
|
||||
|
||||
|
||||
class UploadError(Exception):
|
||||
def __init__(self, message: str, status_code: int = 400):
|
||||
super().__init__(message)
|
||||
self.message = message
|
||||
self.status_code = status_code
|
||||
|
||||
|
||||
async def create_session(
|
||||
session: AsyncSession,
|
||||
user_id: uuid.UUID,
|
||||
declared_mime_type: str,
|
||||
declared_size_bytes: int,
|
||||
title: str | None,
|
||||
client_recording_id: str | None,
|
||||
transcription_model: str | None = None,
|
||||
) -> UploadSession:
|
||||
settings = get_settings()
|
||||
if declared_size_bytes <= 0 or declared_size_bytes > settings.max_upload_bytes:
|
||||
raise UploadError(
|
||||
f"Declared size must be between 1 and {settings.max_upload_bytes} bytes.", 413
|
||||
)
|
||||
allowed = [m.lower() for m in settings.allowed_audio_mime_types]
|
||||
if declared_mime_type.lower() not in allowed:
|
||||
raise UploadError(f"MIME type not allowed. Allowed: {', '.join(allowed)}", 415)
|
||||
|
||||
us = UploadSession(
|
||||
user_id=user_id,
|
||||
client_recording_id=client_recording_id,
|
||||
title=title,
|
||||
declared_mime_type=declared_mime_type.lower(),
|
||||
declared_size_bytes=declared_size_bytes,
|
||||
chunk_size_bytes=settings.max_chunk_bytes,
|
||||
expires_at=utcnow() + SESSION_TTL,
|
||||
transcription_model=(transcription_model.strip() if transcription_model else None),
|
||||
)
|
||||
session.add(us)
|
||||
await session.flush()
|
||||
return us
|
||||
|
||||
|
||||
async def get_owned_session(
|
||||
session: AsyncSession, user_id: uuid.UUID, session_id: uuid.UUID
|
||||
) -> UploadSession:
|
||||
us = await session.get(UploadSession, session_id)
|
||||
if us is None or us.user_id != user_id:
|
||||
raise UploadError("Upload session not found.", 404)
|
||||
if us.status == UploadSessionStatus.expired or (
|
||||
us.status != UploadSessionStatus.completed and us.expires_at < utcnow()
|
||||
):
|
||||
us.status = UploadSessionStatus.expired
|
||||
raise UploadError("Upload session expired. Start a new upload.", 410)
|
||||
return us
|
||||
|
||||
|
||||
async def put_chunk(
|
||||
session: AsyncSession,
|
||||
user_id: uuid.UUID,
|
||||
session_id: uuid.UUID,
|
||||
chunk_index: int,
|
||||
data: bytes,
|
||||
checksum_sha256: str | None,
|
||||
) -> UploadChunk:
|
||||
settings = get_settings()
|
||||
us = await get_owned_session(session, user_id, session_id)
|
||||
if us.status != UploadSessionStatus.open:
|
||||
raise UploadError(f"Session is {us.status.value}; cannot accept chunks.", 409)
|
||||
if chunk_index < 0:
|
||||
raise UploadError("Chunk index must be >= 0.", 400)
|
||||
if not data:
|
||||
raise UploadError("Empty chunk.", 400)
|
||||
if len(data) > settings.max_chunk_bytes:
|
||||
raise UploadError(f"Chunk exceeds max size {settings.max_chunk_bytes}.", 413)
|
||||
if checksum_sha256 and hashlib.sha256(data).hexdigest() != checksum_sha256.lower():
|
||||
raise UploadError("Chunk checksum mismatch.", 422)
|
||||
|
||||
existing = await session.scalar(
|
||||
select(UploadChunk).where(
|
||||
UploadChunk.session_id == us.id, UploadChunk.chunk_index == chunk_index
|
||||
)
|
||||
)
|
||||
if existing is not None:
|
||||
# Idempotent retry: same index re-sent replaces the stored bytes.
|
||||
if existing.size_bytes != len(data):
|
||||
await get_storage().delete(existing.storage_key)
|
||||
existing.size_bytes = len(data)
|
||||
existing.checksum_sha256 = hashlib.sha256(data).hexdigest()
|
||||
await get_storage().put(existing.storage_key, data)
|
||||
await session.flush()
|
||||
return existing
|
||||
|
||||
key = f"uploads/{us.id}/{chunk_index:06d}.part"
|
||||
await get_storage().put(key, data)
|
||||
chunk = UploadChunk(
|
||||
session_id=us.id,
|
||||
chunk_index=chunk_index,
|
||||
size_bytes=len(data),
|
||||
checksum_sha256=hashlib.sha256(data).hexdigest(),
|
||||
storage_key=key,
|
||||
)
|
||||
session.add(chunk)
|
||||
await session.flush()
|
||||
return chunk
|
||||
|
||||
|
||||
async def received_indexes(
|
||||
session: AsyncSession, user_id: uuid.UUID, session_id: uuid.UUID
|
||||
) -> list[int]:
|
||||
us = await get_owned_session(session, user_id, session_id)
|
||||
rows = await session.scalars(
|
||||
select(UploadChunk.chunk_index).where(UploadChunk.session_id == us.id)
|
||||
)
|
||||
return sorted(rows)
|
||||
|
||||
|
||||
async def finalize(
|
||||
session: AsyncSession,
|
||||
user_id: uuid.UUID,
|
||||
session_id: uuid.UUID,
|
||||
*,
|
||||
recorded_at: datetime | None,
|
||||
duration_seconds: float,
|
||||
latitude: float | None = None,
|
||||
longitude: float | None = None,
|
||||
location_accuracy_m: float | None = None,
|
||||
notes: str | None = None,
|
||||
transcription_model: str | None = None,
|
||||
) -> tuple[UploadSession, Recording, Asset]:
|
||||
"""Assemble chunks, validate, store the immutable original, and create or
|
||||
update the recording. Idempotent per client_recording_id."""
|
||||
from shonar.services.ai import ProviderConfigError
|
||||
from shonar.services.ai.model_registry import (
|
||||
effective_model,
|
||||
get_global_default_model,
|
||||
validate_model_name,
|
||||
)
|
||||
|
||||
us = await get_owned_session(session, user_id, session_id)
|
||||
# Resolve + validate the transcription model before touching audio.
|
||||
# Finalize body wins over the session's upload-screen choice; an
|
||||
# explicit override updates history, the default never rewrites it.
|
||||
override_raw = transcription_model or us.transcription_model
|
||||
override: str | None = None
|
||||
if override_raw:
|
||||
try:
|
||||
override = validate_model_name(override_raw)
|
||||
except ProviderConfigError as e:
|
||||
raise UploadError(str(e), 422) from None
|
||||
model = effective_model(override, await get_global_default_model(session))
|
||||
if us.status == UploadSessionStatus.completed and us.completed_asset_id:
|
||||
# Already finalized: return existing recording (retry-safe client).
|
||||
asset = await session.get(Asset, us.completed_asset_id)
|
||||
rec = await session.scalar(
|
||||
select(Recording).where(Recording.id == asset.recording_id)
|
||||
)
|
||||
if asset and rec:
|
||||
return us, rec, asset
|
||||
if us.status != UploadSessionStatus.open:
|
||||
raise UploadError(f"Session is {us.status.value}.", 409)
|
||||
|
||||
chunks = list(
|
||||
await session.scalars(
|
||||
select(UploadChunk).where(UploadChunk.session_id == us.id).order_by(
|
||||
UploadChunk.chunk_index
|
||||
)
|
||||
)
|
||||
)
|
||||
total = sum(c.size_bytes for c in chunks)
|
||||
if total != us.declared_size_bytes:
|
||||
raise UploadError(
|
||||
f"Size mismatch: received {total} of declared {us.declared_size_bytes} bytes. "
|
||||
"Upload missing chunks and retry.",
|
||||
422,
|
||||
)
|
||||
expected_indexes = list(range(len(chunks)))
|
||||
if [c.chunk_index for c in chunks] != expected_indexes:
|
||||
raise UploadError("Chunk sequence has gaps. Upload missing chunks and retry.", 422)
|
||||
|
||||
storage = get_storage()
|
||||
# Validate magic bytes from the first chunk.
|
||||
first = await storage.get(chunks[0].storage_key)
|
||||
if not is_compatible(us.declared_mime_type, first):
|
||||
raise UploadError(
|
||||
"File contents do not match the declared audio MIME type.", 415
|
||||
)
|
||||
|
||||
# Idempotency: same client_recording_id => update existing recording.
|
||||
recording: Recording | None = None
|
||||
if us.client_recording_id:
|
||||
recording = await session.scalar(
|
||||
select(Recording).where(
|
||||
Recording.user_id == user_id,
|
||||
Recording.client_recording_id == us.client_recording_id,
|
||||
)
|
||||
)
|
||||
|
||||
# NOTE(original-immutability): if the recording already has an original
|
||||
# asset we do NOT replace it; a re-upload with the same client id after
|
||||
# local edits updates metadata only, and the new audio is rejected as a
|
||||
# duplicate (the existing original is returned instead).
|
||||
if recording is not None:
|
||||
existing_original = await session.scalar(
|
||||
select(Asset).where(
|
||||
Asset.recording_id == recording.id, Asset.kind == AssetKind.original
|
||||
)
|
||||
)
|
||||
if existing_original is not None:
|
||||
us.status = UploadSessionStatus.completed
|
||||
us.completed_asset_id = existing_original.id
|
||||
recording.title = us.title or recording.title
|
||||
if override is not None:
|
||||
# Explicit re-choice replaces the saved model; the default
|
||||
# never rewrites history.
|
||||
recording.transcription_model = override
|
||||
await session.flush()
|
||||
from shonar.services import processing as _processing
|
||||
|
||||
# Same audio, maybe new metadata — and a failed pipeline deserves
|
||||
# another attempt. Idempotent: completed work is never redone.
|
||||
await _processing.enqueue_for_recording(session, recording)
|
||||
return us, recording, existing_original
|
||||
|
||||
# Assemble into the final object (streamed per chunk to bound memory).
|
||||
from shonar.services.media import extension_for
|
||||
|
||||
digest = hashlib.sha256()
|
||||
parts: list[bytes] = []
|
||||
for c in chunks:
|
||||
data = await storage.get(c.storage_key)
|
||||
digest.update(data)
|
||||
parts.append(data)
|
||||
checksum = digest.hexdigest()
|
||||
ext = extension_for(us.declared_mime_type, parts[0])
|
||||
final_key = f"recordings/{user_id}/{uuid.uuid4()}{ext}"
|
||||
blob = b"".join(parts)
|
||||
await storage.put(final_key, blob)
|
||||
|
||||
asset = Asset(
|
||||
recording_id=recording.id if recording else None,
|
||||
user_id=user_id,
|
||||
kind=AssetKind.original,
|
||||
storage_key=final_key,
|
||||
mime_type=us.declared_mime_type,
|
||||
size_bytes=total,
|
||||
checksum_sha256=checksum,
|
||||
)
|
||||
session.add(asset)
|
||||
|
||||
if recording is None:
|
||||
recording = Recording(
|
||||
user_id=user_id,
|
||||
client_recording_id=us.client_recording_id,
|
||||
title=us.title or "Untitled recording",
|
||||
recorded_at=recorded_at or utcnow(),
|
||||
duration_seconds=duration_seconds,
|
||||
notes=notes,
|
||||
latitude=latitude,
|
||||
longitude=longitude,
|
||||
location_accuracy_m=location_accuracy_m,
|
||||
processing_status=ProcessingStatus.uploaded,
|
||||
transcription_model=model,
|
||||
)
|
||||
session.add(recording)
|
||||
await session.flush()
|
||||
asset.recording_id = recording.id
|
||||
else:
|
||||
recording.title = us.title or recording.title
|
||||
recording.duration_seconds = duration_seconds or recording.duration_seconds
|
||||
recording.processing_status = ProcessingStatus.uploaded
|
||||
recording.processing_error = None
|
||||
|
||||
us.status = UploadSessionStatus.completed
|
||||
us.completed_asset_id = asset.id
|
||||
|
||||
# Clean up chunk parts (the assembled object is the source of truth).
|
||||
for c in chunks: # best-effort cleanup of chunk parts
|
||||
with contextlib.suppress(Exception):
|
||||
await storage.delete(c.storage_key)
|
||||
await session.flush()
|
||||
from shonar.services import processing as _processing
|
||||
|
||||
# New audio on disk: queue whatever AI stages apply (none configured =
|
||||
# ai_disabled, never an error).
|
||||
await _processing.enqueue_for_recording(session, recording)
|
||||
return us, recording, asset
|
||||
|
||||
|
||||
async def abort(session: AsyncSession, user_id: uuid.UUID, session_id: uuid.UUID) -> None:
|
||||
us = await get_owned_session(session, user_id, session_id)
|
||||
if us.status in (UploadSessionStatus.completed, UploadSessionStatus.aborted):
|
||||
return
|
||||
storage = get_storage()
|
||||
chunks = list(
|
||||
await session.scalars(select(UploadChunk).where(UploadChunk.session_id == us.id))
|
||||
)
|
||||
for c in chunks: # best-effort cleanup
|
||||
with contextlib.suppress(Exception):
|
||||
await storage.delete(c.storage_key)
|
||||
us.status = UploadSessionStatus.aborted
|
||||
179
backend/shonar/storage/__init__.py
Normal file
179
backend/shonar/storage/__init__.py
Normal file
|
|
@ -0,0 +1,179 @@
|
|||
"""Storage abstraction.
|
||||
|
||||
Backends:
|
||||
- ``LocalStorage``: filesystem under ``SHONAR_STORAGE_PATH`` (dev + default).
|
||||
- ``S3Storage``: any S3-compatible object store (extra: ``pip install
|
||||
shonar-backend[s3]``).
|
||||
|
||||
Storage keys are server-internal and validated against path traversal:
|
||||
they are UUID-based by construction and never derived from user input.
|
||||
At-rest encryption is NOT implemented; see docs/security.md for the
|
||||
documented optional design.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import shutil
|
||||
from pathlib import Path, PurePosixPath
|
||||
from typing import Protocol
|
||||
|
||||
from shonar.core.config import get_settings
|
||||
|
||||
_KEY_CHARSET = set("abcdefghijklmnopqrstuvwxyz0123456789-./")
|
||||
|
||||
|
||||
class StorageError(Exception):
|
||||
pass
|
||||
|
||||
|
||||
def validate_storage_key(key: str) -> str:
|
||||
"""Reject anything that could escape the storage root."""
|
||||
if not key or len(key) > 500:
|
||||
raise StorageError("invalid storage key")
|
||||
pure = PurePosixPath(key)
|
||||
if pure.is_absolute() or ".." in pure.parts:
|
||||
raise StorageError("invalid storage key")
|
||||
if not set(key) <= _KEY_CHARSET:
|
||||
raise StorageError("invalid storage key")
|
||||
return key
|
||||
|
||||
|
||||
class StorageBackend(Protocol):
|
||||
async def put(self, key: str, data: bytes) -> int: ...
|
||||
async def put_file(self, key: str, src_path: Path) -> int: ...
|
||||
async def get(self, key: str) -> bytes: ...
|
||||
async def open_path(self, key: str) -> Path | None:
|
||||
"""Local file path if the backend can provide one, else None."""
|
||||
...
|
||||
|
||||
async def delete(self, key: str) -> None: ...
|
||||
async def exists(self, key: str) -> bool: ...
|
||||
|
||||
|
||||
class LocalStorage:
|
||||
def __init__(self, root: str | Path):
|
||||
self.root = Path(root).resolve()
|
||||
self.root.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
def _path(self, key: str) -> Path:
|
||||
validate_storage_key(key)
|
||||
path = (self.root / key).resolve()
|
||||
# Defense in depth: resolved path must stay under root.
|
||||
if not path.is_relative_to(self.root):
|
||||
raise StorageError("storage key escapes storage root")
|
||||
return path
|
||||
|
||||
async def put(self, key: str, data: bytes) -> int:
|
||||
path = self._path(key)
|
||||
path.parent.mkdir(parents=True, exist_ok=True)
|
||||
await asyncio.to_thread(path.write_bytes, data)
|
||||
return len(data)
|
||||
|
||||
async def put_file(self, key: str, src_path: Path) -> int:
|
||||
path = self._path(key)
|
||||
path.parent.mkdir(parents=True, exist_ok=True)
|
||||
await asyncio.to_thread(shutil.copyfile, src_path, path)
|
||||
return path.stat().st_size
|
||||
|
||||
async def get(self, key: str) -> bytes:
|
||||
path = self._path(key)
|
||||
if not path.exists():
|
||||
raise StorageError("object not found")
|
||||
return await asyncio.to_thread(path.read_bytes)
|
||||
|
||||
async def open_path(self, key: str) -> Path | None:
|
||||
path = self._path(key)
|
||||
return path if path.exists() else None
|
||||
|
||||
async def delete(self, key: str) -> None:
|
||||
path = self._path(key)
|
||||
if path.exists():
|
||||
await asyncio.to_thread(path.unlink)
|
||||
|
||||
async def exists(self, key: str) -> bool:
|
||||
return self._path(key).exists()
|
||||
|
||||
|
||||
class S3Storage: # pragma: no cover - requires boto3 + a real/mini endpoint
|
||||
def __init__(
|
||||
self, endpoint_url: str, bucket: str, region: str, access_key: str, secret_key: str
|
||||
):
|
||||
import boto3 # optional extra
|
||||
|
||||
self.bucket = bucket
|
||||
self.s3 = boto3.client(
|
||||
"s3",
|
||||
endpoint_url=endpoint_url or None,
|
||||
region_name=region,
|
||||
aws_access_key_id=access_key or None,
|
||||
aws_secret_access_key=secret_key or None,
|
||||
)
|
||||
|
||||
async def put(self, key: str, data: bytes) -> int:
|
||||
validate_storage_key(key)
|
||||
await asyncio.to_thread(self.s3.put_object, Bucket=self.bucket, Key=key, Body=data)
|
||||
return len(data)
|
||||
|
||||
async def put_file(self, key: str, src_path: Path) -> int:
|
||||
validate_storage_key(key)
|
||||
await asyncio.to_thread(self.s3.upload_file, str(src_path), self.bucket, key)
|
||||
return src_path.stat().st_size
|
||||
|
||||
async def get(self, key: str) -> bytes:
|
||||
validate_storage_key(key)
|
||||
|
||||
def _get() -> bytes:
|
||||
obj = self.s3.get_object(Bucket=self.bucket, Key=key)
|
||||
return obj["Body"].read()
|
||||
|
||||
return await asyncio.to_thread(_get)
|
||||
|
||||
async def open_path(self, key: str) -> Path | None:
|
||||
return None # callers must use get()/streaming
|
||||
|
||||
async def delete(self, key: str) -> None:
|
||||
validate_storage_key(key)
|
||||
await asyncio.to_thread(self.s3.delete_object, Bucket=self.bucket, Key=key)
|
||||
|
||||
async def exists(self, key: str) -> bool:
|
||||
validate_storage_key(key)
|
||||
|
||||
def _head() -> bool:
|
||||
from botocore.exceptions import ClientError
|
||||
|
||||
try:
|
||||
self.s3.head_object(Bucket=self.bucket, Key=key)
|
||||
return True
|
||||
except ClientError:
|
||||
return False
|
||||
|
||||
return await asyncio.to_thread(_head)
|
||||
|
||||
|
||||
_backend: StorageBackend | None = None
|
||||
|
||||
|
||||
def get_storage() -> StorageBackend:
|
||||
global _backend
|
||||
if _backend is None:
|
||||
settings = get_settings()
|
||||
if settings.storage_backend == "local":
|
||||
_backend = LocalStorage(settings.storage_path)
|
||||
elif settings.storage_backend == "s3":
|
||||
_backend = S3Storage(
|
||||
settings.s3_endpoint_url,
|
||||
settings.s3_bucket,
|
||||
settings.s3_region,
|
||||
settings.s3_access_key_id,
|
||||
settings.s3_secret_access_key,
|
||||
)
|
||||
else:
|
||||
raise StorageError(f"unknown storage backend: {settings.storage_backend}")
|
||||
return _backend
|
||||
|
||||
|
||||
def set_storage(backend: StorageBackend | None) -> None:
|
||||
"""Test seam."""
|
||||
global _backend
|
||||
_backend = backend
|
||||
79
backend/shonar/worker.py
Normal file
79
backend/shonar/worker.py
Normal file
|
|
@ -0,0 +1,79 @@
|
|||
"""arq worker entrypoint (M7): runs the AI pipeline tasks.
|
||||
|
||||
Run: arq shonar.worker.WorkerSettings
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
|
||||
from arq import cron
|
||||
from arq.connections import RedisSettings
|
||||
|
||||
from shonar.core.config import get_settings
|
||||
from shonar.db.session import dispose_engine, get_engine
|
||||
from shonar.services import processing
|
||||
|
||||
logger = logging.getLogger("shonar.worker")
|
||||
|
||||
|
||||
async def run_transcribe(ctx: dict, recording_id: str) -> None:
|
||||
await processing.run_transcribe(ctx, recording_id)
|
||||
|
||||
|
||||
async def run_summarize(ctx: dict, recording_id: str) -> None:
|
||||
await processing.run_summarize(ctx, recording_id)
|
||||
|
||||
|
||||
async def sweep(ctx: dict) -> None: # noqa: ARG001 — arq cron signature
|
||||
count = await processing.sweep_stale()
|
||||
if count:
|
||||
logger.info("sweep re-enqueued %d stale jobs", count)
|
||||
|
||||
|
||||
async def retention_sweep(ctx: dict) -> None: # noqa: ARG001 — arq cron signature
|
||||
from shonar.services import retention
|
||||
|
||||
purged = await retention.sweep_deleted()
|
||||
if purged["recordings"] or purged["users"]:
|
||||
logger.info("retention sweep: %s", purged)
|
||||
|
||||
|
||||
async def startup(ctx: dict) -> None:
|
||||
get_engine()
|
||||
# Crash recovery before accepting new work: jobs stuck `running` and
|
||||
# queued rows the transport missed go back through arq.
|
||||
count = await processing.sweep_stale()
|
||||
if count:
|
||||
logger.info("startup sweep re-enqueued %d stale jobs", count)
|
||||
|
||||
|
||||
async def shutdown(ctx: dict) -> None: # noqa: ARG001
|
||||
await dispose_engine()
|
||||
|
||||
|
||||
def _redis() -> RedisSettings:
|
||||
return RedisSettings.from_dsn(get_settings().redis_url)
|
||||
|
||||
|
||||
class WorkerSettings:
|
||||
functions = [run_transcribe, run_summarize, sweep, retention_sweep]
|
||||
# Transport-loss backstop beyond the startup sweep: anything still
|
||||
# queued (missed enqueue, dead worker between runs) goes back through
|
||||
# arq every 5 minutes. Rows are the queue; this just pokes.
|
||||
cron_jobs = [
|
||||
cron(sweep, minute={0, 5, 10, 15, 20, 25, 30, 35, 40, 45, 50, 55}),
|
||||
# Hard-delete expired soft-deletes once a day (3:17 local, off-peak).
|
||||
cron(retention_sweep, hour=3, minute=17),
|
||||
]
|
||||
on_startup = startup
|
||||
on_shutdown = shutdown
|
||||
redis_settings = _redis()
|
||||
# Retry budget for transient provider failures; the tasks themselves
|
||||
# mark jobs failed on the last try (see processing.MAX_TRIES).
|
||||
max_tries = 3
|
||||
# One job per worker process at a time. Local faster-whisper is
|
||||
# multi-GB per concurrent pass; default max_jobs=10 let one worker run
|
||||
# ~8 transcriptions at once and OOM-swap the box. Scale by running more
|
||||
# worker processes, never by raising this.
|
||||
max_jobs = 1
|
||||
0
backend/tests/__init__.py
Normal file
0
backend/tests/__init__.py
Normal file
113
backend/tests/conftest.py
Normal file
113
backend/tests/conftest.py
Normal file
|
|
@ -0,0 +1,113 @@
|
|||
"""Pytest fixtures.
|
||||
|
||||
Tests run against a real PostgreSQL (deploy/docker-compose.dev.yml) using a
|
||||
dedicated ``shonar_test`` database, plus a temporary local storage root.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
import tempfile
|
||||
from collections.abc import AsyncIterator
|
||||
from pathlib import Path
|
||||
|
||||
import pytest
|
||||
import pytest_asyncio
|
||||
|
||||
# Configure env BEFORE importing the app so Settings picks it up.
|
||||
TEST_DB = os.environ.get(
|
||||
"SHONAR_TEST_DATABASE_URL",
|
||||
"postgresql+asyncpg://shonar:shonar@localhost:5432/shonar_test",
|
||||
)
|
||||
os.environ["SHONAR_DATABASE_URL"] = TEST_DB
|
||||
os.environ["SHONAR_SECRET_KEY"] = "test-secret-key-0123456789abcdef0123456789abcdef"
|
||||
os.environ["SHONAR_STORAGE_BACKEND"] = "local"
|
||||
# Effectively disable the auth rate limit under test (dedicated tests cover
|
||||
# the limiter behaviour itself).
|
||||
os.environ["SHONAR_RATE_LIMIT_AUTH"] = "10000/minute"
|
||||
|
||||
_tmp_storage = tempfile.mkdtemp(prefix="shonar-test-storage-")
|
||||
os.environ["SHONAR_STORAGE_PATH"] = _tmp_storage
|
||||
|
||||
|
||||
@pytest_asyncio.fixture(scope="session", loop_scope="session")
|
||||
async def _setup_db() -> AsyncIterator[None]:
|
||||
from shonar.db import models # noqa: F401
|
||||
from shonar.db.base import Base
|
||||
from shonar.db.session import dispose_engine, get_engine
|
||||
|
||||
engine = get_engine()
|
||||
async with engine.begin() as conn:
|
||||
await conn.run_sync(Base.metadata.drop_all)
|
||||
await conn.run_sync(Base.metadata.create_all)
|
||||
# The tsvector generated columns live only in migration
|
||||
# fts0000000001 (not in the ORM models), so create_all misses them.
|
||||
# Apply the same DDL the migration applies (M9 search needs them).
|
||||
if engine.dialect.name == "postgresql":
|
||||
# Reuse the real migration's DDL (not a copy) via a sync
|
||||
# MigrationContext — op.execute() is synchronous there.
|
||||
import importlib.util
|
||||
from pathlib import Path
|
||||
|
||||
def _apply(sync_conn):
|
||||
from alembic.migration import MigrationContext
|
||||
from alembic.operations import Operations
|
||||
|
||||
spec = importlib.util.spec_from_file_location(
|
||||
"fts_migration",
|
||||
Path(__file__).resolve().parents[1]
|
||||
/ "migrations/versions/fts0000000001_fts_columns.py",
|
||||
)
|
||||
assert spec is not None and spec.loader is not None
|
||||
mod = importlib.util.module_from_spec(spec)
|
||||
spec.loader.exec_module(mod)
|
||||
ctx = MigrationContext.configure(sync_conn)
|
||||
with Operations.context(ctx):
|
||||
mod.upgrade()
|
||||
|
||||
await conn.run_sync(_apply)
|
||||
yield
|
||||
await dispose_engine()
|
||||
|
||||
|
||||
@pytest_asyncio.fixture(loop_scope="session", autouse=True)
|
||||
async def clean_db(_setup_db: None) -> AsyncIterator[None]:
|
||||
"""Truncate between tests for isolation."""
|
||||
yield
|
||||
from sqlalchemy import delete, text
|
||||
|
||||
from shonar.db.base import Base
|
||||
from shonar.db.session import _session_factory, get_engine # type: ignore[attr-defined]
|
||||
|
||||
assert _session_factory is not None
|
||||
async with _session_factory() as s:
|
||||
if get_engine().dialect.name == "postgresql":
|
||||
await s.execute(
|
||||
text(
|
||||
"TRUNCATE users, devices, refresh_tokens, recordings, assets, "
|
||||
"upload_sessions, upload_chunks, transcripts, summaries, tags, "
|
||||
"recording_tags, processing_jobs, export_jobs, app_settings "
|
||||
"RESTART IDENTITY CASCADE"
|
||||
)
|
||||
)
|
||||
else:
|
||||
# SQLite: delete child-first (reverse dependency order).
|
||||
for table in reversed(Base.metadata.sorted_tables):
|
||||
await s.execute(delete(table))
|
||||
await s.commit()
|
||||
|
||||
|
||||
@pytest_asyncio.fixture(loop_scope="session")
|
||||
async def client(_setup_db) -> AsyncIterator:
|
||||
from httpx import ASGITransport, AsyncClient
|
||||
|
||||
from shonar.main import app
|
||||
|
||||
transport = ASGITransport(app=app)
|
||||
async with AsyncClient(transport=transport, base_url="http://test") as c:
|
||||
yield c
|
||||
|
||||
|
||||
@pytest.fixture()
|
||||
def storage_root() -> Path:
|
||||
return Path(_tmp_storage)
|
||||
184
backend/tests/test_ai_adapters.py
Normal file
184
backend/tests/test_ai_adapters.py
Normal file
|
|
@ -0,0 +1,184 @@
|
|||
"""Adapter wire-protocol tests (M7): JSON shapes in/out, error mapping.
|
||||
|
||||
HTTP adapters take an injectable httpx client; faster-whisper is an
|
||||
optional dependency and is only exercised when installed.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import importlib.util
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
|
||||
from shonar.core.config import Settings
|
||||
from shonar.services import ai
|
||||
from shonar.services.ai import ProviderConfigError, SummaryResult
|
||||
from shonar.services.ai.ollama import OllamaProvider
|
||||
from shonar.services.ai.openai_compat import OpenAICompatProvider
|
||||
from shonar.services.ai.whisper_http import WhisperHttpProvider
|
||||
|
||||
|
||||
def mock_client(handler) -> httpx.AsyncClient:
|
||||
return httpx.AsyncClient(transport=httpx.MockTransport(handler))
|
||||
|
||||
|
||||
# --- factories -----------------------------------------------------------------
|
||||
|
||||
|
||||
def test_factories_none_means_skip():
|
||||
s = Settings(transcription_provider="none", llm_provider="none")
|
||||
assert ai.get_transcription_provider(s) is None
|
||||
assert ai.get_llm_provider(s) is None
|
||||
|
||||
|
||||
def test_factories_unknown_names_raise_config_error():
|
||||
s = Settings(transcription_provider="whisper-9k")
|
||||
with pytest.raises(ProviderConfigError):
|
||||
ai.get_transcription_provider(s)
|
||||
s = Settings(llm_provider="clippy")
|
||||
with pytest.raises(ProviderConfigError):
|
||||
ai.get_llm_provider(s)
|
||||
|
||||
|
||||
def test_factories_missing_urls_raise_config_error():
|
||||
s = Settings(transcription_provider="whisper_http", transcription_base_url="")
|
||||
with pytest.raises(ProviderConfigError):
|
||||
ai.get_transcription_provider(s)
|
||||
s = Settings(llm_provider="openai_compat", llm_base_url="http://x", llm_model="")
|
||||
with pytest.raises(ProviderConfigError):
|
||||
ai.get_llm_provider(s)
|
||||
|
||||
|
||||
def test_faster_whisper_missing_dep_is_config_error():
|
||||
if importlib.util.find_spec("faster_whisper") is not None:
|
||||
pytest.skip("faster-whisper installed; missing-dep path N/A")
|
||||
# Constructor validates eagerly so misconfiguration fails at startup,
|
||||
# not on the first recording.
|
||||
with pytest.raises(ProviderConfigError):
|
||||
ai.get_transcription_provider(Settings(transcription_provider="faster_whisper"))
|
||||
|
||||
|
||||
# --- whisper_http -----------------------------------------------------------------
|
||||
|
||||
|
||||
def whisper_ok(request: httpx.Request) -> httpx.Response:
|
||||
assert request.url.path == "/v1/audio/transcriptions"
|
||||
assert request.method == "POST"
|
||||
return httpx.Response(200, json={
|
||||
"text": "hello world",
|
||||
"language": "en",
|
||||
"segments": [{"start": 0.0, "end": 1.2, "text": "hello world"}],
|
||||
})
|
||||
|
||||
|
||||
async def test_whisper_http_happy_path():
|
||||
p = WhisperHttpProvider("http://stt:8000", model="small",
|
||||
http_client=mock_client(whisper_ok))
|
||||
res = await p.transcribe(b"\x00" * 16, "audio/wav")
|
||||
assert res.text == "hello world"
|
||||
assert res.language == "en"
|
||||
assert [(s.start, s.end, s.text) for s in res.segments] == [(0.0, 1.2, "hello world")]
|
||||
assert res.model == "small"
|
||||
|
||||
|
||||
async def test_whisper_http_401_is_config_error():
|
||||
async def denied(request: httpx.Request) -> httpx.Response:
|
||||
return httpx.Response(401, json={"detail": "nope"})
|
||||
p = WhisperHttpProvider("http://stt:8000", http_client=mock_client(denied))
|
||||
with pytest.raises(ProviderConfigError):
|
||||
await p.transcribe(b"\x00" * 16, "audio/wav")
|
||||
|
||||
|
||||
async def test_whisper_http_503_is_transient():
|
||||
async def busy(request: httpx.Request) -> httpx.Response:
|
||||
return httpx.Response(503, text="overloaded")
|
||||
from shonar.services.ai import ProviderTransientError
|
||||
|
||||
p = WhisperHttpProvider("http://stt:8000", http_client=mock_client(busy))
|
||||
with pytest.raises(ProviderTransientError):
|
||||
await p.transcribe(b"\x00" * 16, "audio/wav")
|
||||
|
||||
|
||||
# --- openai_compat ------------------------------------------------------------------
|
||||
|
||||
|
||||
def chat_ok(request: httpx.Request) -> httpx.Response:
|
||||
assert request.url.path == "/v1/chat/completions"
|
||||
body = {
|
||||
"short": "Standup.",
|
||||
"detailed": "The team met.",
|
||||
"key_points": ["a", "b"],
|
||||
"decisions": ["ship"],
|
||||
"action_items": [{"not": "a string"}, "call ana"],
|
||||
"questions": [],
|
||||
"extra_key": "ignored",
|
||||
}
|
||||
import json as _json
|
||||
|
||||
return httpx.Response(200, json={"choices": [{"message": {"content": _json.dumps(body)}}]})
|
||||
|
||||
|
||||
async def test_openai_compat_parses_and_sanitizes():
|
||||
p = OpenAICompatProvider("http://llm:8000", model="qwen",
|
||||
http_client=mock_client(chat_ok))
|
||||
res = await p.summarize("a very long meeting transcript", title="Standup")
|
||||
assert isinstance(res, SummaryResult)
|
||||
assert res.short == "Standup."
|
||||
assert res.action_items == ("call ana",) # non-strings dropped
|
||||
assert res.model == "qwen"
|
||||
|
||||
|
||||
async def test_openai_compat_partial_json_gets_defaults():
|
||||
import json as _json
|
||||
|
||||
async def partial(request: httpx.Request) -> httpx.Response:
|
||||
return httpx.Response(200, json={
|
||||
"choices": [{"message": {"content": _json.dumps({"short": "Hi."})}}]
|
||||
})
|
||||
p = OpenAICompatProvider("http://llm:8000", model="qwen",
|
||||
http_client=mock_client(partial))
|
||||
res = await p.summarize("hi")
|
||||
assert res.short == "Hi."
|
||||
assert res.detailed == "" and res.key_points == ()
|
||||
|
||||
|
||||
async def test_openai_compat_non_json_is_transient():
|
||||
async def garbage(request: httpx.Request) -> httpx.Response:
|
||||
return httpx.Response(200, json={
|
||||
"choices": [{"message": {"content": "Sure! Here it is..."}}]
|
||||
})
|
||||
from shonar.services.ai import ProviderTransientError
|
||||
|
||||
p = OpenAICompatProvider("http://llm:8000", model="qwen",
|
||||
http_client=mock_client(garbage))
|
||||
with pytest.raises(ProviderTransientError):
|
||||
await p.summarize("hi")
|
||||
|
||||
|
||||
# --- ollama ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
async def test_ollama_happy_path():
|
||||
import json as _json
|
||||
|
||||
async def ok(request: httpx.Request) -> httpx.Response:
|
||||
assert request.url.path == "/api/chat"
|
||||
payload = _json.loads(request.content.decode())
|
||||
assert payload["format"] == "json" and payload["stream"] is False
|
||||
return httpx.Response(200, json={
|
||||
"message": {"content": _json.dumps({"short": "S.", "detailed": "D."})}
|
||||
})
|
||||
p = OllamaProvider("http://ollama:11434", model="llama3",
|
||||
http_client=mock_client(ok))
|
||||
res = await p.summarize("meeting notes")
|
||||
assert res.short == "S." and res.model == "llama3"
|
||||
|
||||
|
||||
async def test_ollama_404_is_config_error():
|
||||
async def missing(request: httpx.Request) -> httpx.Response:
|
||||
return httpx.Response(404, text="model not found")
|
||||
p = OllamaProvider("http://ollama:11434", model="nope",
|
||||
http_client=mock_client(missing))
|
||||
with pytest.raises(ProviderConfigError):
|
||||
await p.summarize("hi")
|
||||
410
backend/tests/test_ai_pipeline.py
Normal file
410
backend/tests/test_ai_pipeline.py
Normal file
|
|
@ -0,0 +1,410 @@
|
|||
"""AI pipeline tests (M7): status flow, versioning, failure modes, endpoints.
|
||||
|
||||
Provider fakes stand in for real STT/LLM services (no network, no model
|
||||
downloads); adapter wire-protocol tests live in test_ai_adapters.py.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import struct
|
||||
import uuid
|
||||
|
||||
import pytest
|
||||
from sqlalchemy import func, select
|
||||
|
||||
from shonar.db import session as db_session
|
||||
from shonar.db.models import (
|
||||
JobStatus,
|
||||
JobType,
|
||||
ProcessingJob,
|
||||
Transcript,
|
||||
)
|
||||
from shonar.services import processing
|
||||
from shonar.services.ai import (
|
||||
ProviderConfigError,
|
||||
ProviderTransientError,
|
||||
Segment,
|
||||
SummaryResult,
|
||||
TranscriptResult,
|
||||
)
|
||||
|
||||
|
||||
def wav_bytes(payload_len: int = 64) -> bytes:
|
||||
data = bytes(range(payload_len % 256)) * (payload_len // 256 + 1)
|
||||
data = data[:payload_len]
|
||||
header = (
|
||||
b"RIFF" + struct.pack("<I", 36 + len(data)) + b"WAVE"
|
||||
+ b"fmt " + struct.pack("<IHHIIHH", 16, 1, 1, 8000, 8000, 1, 8)
|
||||
+ b"data" + struct.pack("<I", len(data))
|
||||
)
|
||||
return header + data
|
||||
|
||||
|
||||
AUTH = {"email": "m7@example.com", "password": "m7-test-passw0rd-123"}
|
||||
|
||||
|
||||
async def user_tokens(client, email=AUTH["email"], password=AUTH["password"]):
|
||||
r = await client.post("/api/v1/auth/register", json={"email": email, "password": password})
|
||||
assert r.status_code == 201, r.text
|
||||
return r.json()["access_token"]
|
||||
|
||||
|
||||
async def upload_recording(client, token, client_id=None, title="M7 standup"):
|
||||
h = {"Authorization": f"Bearer {token}"}
|
||||
data = wav_bytes()
|
||||
r = await client.post(
|
||||
"/api/v1/uploads",
|
||||
json={"declared_mime_type": "audio/wav", "declared_size_bytes": len(data),
|
||||
"client_recording_id": client_id, "title": title},
|
||||
headers=h,
|
||||
)
|
||||
assert r.status_code == 201, r.text
|
||||
sid = r.json()["id"]
|
||||
r = await client.put(f"/api/v1/uploads/{sid}/chunks/0", content=data,
|
||||
headers={**h, "content-type": "application/octet-stream"})
|
||||
assert r.status_code == 201, r.text
|
||||
r = await client.post(f"/api/v1/uploads/{sid}/finalize",
|
||||
json={"duration_seconds": 5.0}, headers=h)
|
||||
assert r.status_code == 201, r.text
|
||||
return r.json()
|
||||
|
||||
|
||||
class FakeTranscriber:
|
||||
name = "fake-stt"
|
||||
|
||||
def __init__(self, text="hello world from the meeting", fail=None):
|
||||
self.text = text
|
||||
self.fail = fail
|
||||
self.calls = 0
|
||||
|
||||
async def transcribe(self, audio, mime, *, language_hint=None):
|
||||
self.calls += 1
|
||||
assert len(audio) > 0 and mime == "audio/wav"
|
||||
if self.fail is not None:
|
||||
raise self.fail
|
||||
return TranscriptResult(
|
||||
text=self.text, language="en",
|
||||
segments=[Segment(0.0, 1.0, self.text)], model="fake-stt-1",
|
||||
)
|
||||
|
||||
|
||||
class FakeLlm:
|
||||
name = "fake-llm"
|
||||
|
||||
def __init__(self, fail=None):
|
||||
self.fail = fail
|
||||
self.seen = []
|
||||
|
||||
async def summarize(self, transcript, *, title=None):
|
||||
self.seen.append(transcript)
|
||||
if self.fail is not None:
|
||||
raise self.fail
|
||||
return SummaryResult(
|
||||
short="Standup happened.", detailed="The team met and spoke.",
|
||||
key_points=("a",), decisions=(), action_items=("ship it",),
|
||||
questions=(), model="fake-llm-1",
|
||||
)
|
||||
|
||||
|
||||
_UNSET = object()
|
||||
|
||||
|
||||
def use_fakes(monkeypatch, stt=_UNSET, llm=_UNSET):
|
||||
tprov = FakeTranscriber() if stt is _UNSET else stt
|
||||
lprov = FakeLlm() if llm is _UNSET else llm
|
||||
monkeypatch.setattr(processing, "get_transcription_provider", lambda settings: tprov)
|
||||
monkeypatch.setattr(processing, "get_llm_provider", lambda settings: lprov)
|
||||
|
||||
|
||||
async def jobs_for(recording_id):
|
||||
async with db_session._session_factory() as s:
|
||||
rows = (await s.scalars(
|
||||
select(ProcessingJob).where(ProcessingJob.recording_id == uuid.UUID(recording_id))
|
||||
)).all()
|
||||
return {j.job_type: j for j in rows}
|
||||
|
||||
|
||||
# --- no AI configured --------------------------------------------------------
|
||||
|
||||
|
||||
async def test_none_configured_marks_ai_disabled(client):
|
||||
token = await user_tokens(client)
|
||||
rec = await upload_recording(client, token)
|
||||
assert rec["processing_status"] == "ai_disabled"
|
||||
assert await jobs_for(rec["id"]) == {}
|
||||
# Status endpoints: nothing there yet, but the recording exists.
|
||||
h = {"Authorization": f"Bearer {token}"}
|
||||
r = await client.get(f"/api/v1/recordings/{rec['id']}/transcript", headers=h)
|
||||
assert r.status_code == 404
|
||||
r = await client.get(f"/api/v1/recordings/{rec['id']}/summary", headers=h)
|
||||
assert r.status_code == 404
|
||||
r = await client.get(f"/api/v1/recordings/{rec['id']}/jobs", headers=h)
|
||||
assert r.status_code == 200 and r.json() == []
|
||||
|
||||
|
||||
# --- happy path ----------------------------------------------------------------
|
||||
|
||||
|
||||
async def test_full_pipeline_transcribe_then_summarize(client, monkeypatch):
|
||||
stt, llm = FakeTranscriber(), FakeLlm()
|
||||
use_fakes(monkeypatch, stt, llm)
|
||||
token = await user_tokens(client)
|
||||
rec = await upload_recording(client, token, client_id="m7-full-1")
|
||||
assert rec["processing_status"] == "queued"
|
||||
jobs = await jobs_for(rec["id"])
|
||||
assert set(jobs) == {JobType.transcribe}
|
||||
assert jobs[JobType.transcribe].status == JobStatus.queued
|
||||
|
||||
await processing.run_transcribe({}, rec["id"])
|
||||
jobs = await jobs_for(rec["id"])
|
||||
assert jobs[JobType.transcribe].status == JobStatus.succeeded
|
||||
assert set(jobs) == {JobType.transcribe, JobType.summarize}
|
||||
h = {"Authorization": f"Bearer {token}"}
|
||||
r = await client.get(f"/api/v1/recordings/{rec['id']}", headers=h)
|
||||
assert r.json()["processing_status"] == "processing"
|
||||
|
||||
t = await client.get(f"/api/v1/recordings/{rec['id']}/transcript", headers=h)
|
||||
assert t.status_code == 200, t.text
|
||||
body = t.json()
|
||||
assert body["text"] == "hello world from the meeting"
|
||||
assert body["version"] == 1 and body["provider"] == "fake-stt"
|
||||
assert body["segments"][0]["text"].startswith("hello")
|
||||
assert llm.seen == []
|
||||
|
||||
await processing.run_summarize({}, rec["id"])
|
||||
jobs = await jobs_for(rec["id"])
|
||||
assert jobs[JobType.summarize].status == JobStatus.succeeded
|
||||
s = await client.get(f"/api/v1/recordings/{rec['id']}/summary", headers=h)
|
||||
assert s.status_code == 200, s.text
|
||||
content = s.json()["content"]
|
||||
assert content["short"] == "Standup happened."
|
||||
assert content["action_items"] == ["ship it"]
|
||||
assert set(content) == {"short", "detailed", "key_points", "decisions",
|
||||
"action_items", "questions"}
|
||||
r = await client.get(f"/api/v1/recordings/{rec['id']}", headers=h)
|
||||
assert r.json()["processing_status"] == "completed"
|
||||
assert llm.seen == ["hello world from the meeting"]
|
||||
|
||||
jobs_rows = await client.get(f"/api/v1/recordings/{rec['id']}/jobs", headers=h)
|
||||
assert jobs_rows.status_code == 200
|
||||
assert {j["job_type"] for j in jobs_rows.json()} == {"transcribe", "summarize"}
|
||||
|
||||
|
||||
async def test_enqueue_is_idempotent(client, monkeypatch):
|
||||
use_fakes(monkeypatch, FakeTranscriber(), FakeLlm())
|
||||
token = await user_tokens(client)
|
||||
rec = await upload_recording(client, token, client_id="m7-idem-1")
|
||||
await processing.run_transcribe({}, rec["id"])
|
||||
await processing.run_summarize({}, rec["id"])
|
||||
async with db_session._session_factory() as s:
|
||||
from shonar.db.models import Recording
|
||||
|
||||
row = await s.get(Recording, uuid.UUID(rec["id"]))
|
||||
queued = await processing.enqueue_for_recording(s, row)
|
||||
assert queued == []
|
||||
n = await s.scalar(
|
||||
select(func.count(ProcessingJob.id)).where(
|
||||
ProcessingJob.recording_id == uuid.UUID(rec["id"]))
|
||||
)
|
||||
assert n == 2
|
||||
await s.commit()
|
||||
|
||||
|
||||
# --- failure modes -------------------------------------------------------------
|
||||
|
||||
|
||||
async def test_transient_failure_retries_then_fails(client, monkeypatch):
|
||||
stt = FakeTranscriber(fail=ProviderTransientError("stt down"))
|
||||
use_fakes(monkeypatch, stt, FakeLlm())
|
||||
token = await user_tokens(client)
|
||||
rec = await upload_recording(client, token, client_id="m7-fail-1")
|
||||
# First tries raise for arq retry; the attempt rolls back with the
|
||||
# transaction, so the row still reads queued (arq tracks the tries).
|
||||
with pytest.raises(ProviderTransientError):
|
||||
await processing.run_transcribe({"job_try": 1}, rec["id"])
|
||||
jobs = await jobs_for(rec["id"])
|
||||
assert jobs[JobType.transcribe].status == JobStatus.queued
|
||||
# Last try marks the job (and recording) failed with a safe message.
|
||||
await processing.run_transcribe({"job_try": 3}, rec["id"])
|
||||
jobs = await jobs_for(rec["id"])
|
||||
assert jobs[JobType.transcribe].status == JobStatus.failed
|
||||
assert jobs[JobType.transcribe].error == "stt down"
|
||||
h = {"Authorization": f"Bearer {token}"}
|
||||
r = await client.get(f"/api/v1/recordings/{rec['id']}", headers=h)
|
||||
assert r.json()["processing_status"] == "failed"
|
||||
assert r.json()["processing_error"] == "stt down"
|
||||
|
||||
|
||||
async def test_config_error_fails_fast(client, monkeypatch):
|
||||
stt = FakeTranscriber(fail=ProviderConfigError("bad credentials"))
|
||||
use_fakes(monkeypatch, stt, FakeLlm())
|
||||
token = await user_tokens(client)
|
||||
rec = await upload_recording(client, token, client_id="m7-cfg-1")
|
||||
await processing.run_transcribe({"job_try": 1}, rec["id"]) # no raise
|
||||
jobs = await jobs_for(rec["id"])
|
||||
assert jobs[JobType.transcribe].status == JobStatus.failed
|
||||
assert stt.calls == 1
|
||||
|
||||
|
||||
# --- versioning -----------------------------------------------------------------
|
||||
|
||||
|
||||
async def test_rerun_supersedes_auto_but_not_user_edits(client, monkeypatch):
|
||||
use_fakes(monkeypatch, FakeTranscriber("v1 text"), FakeLlm())
|
||||
token = await user_tokens(client)
|
||||
rec = await upload_recording(client, token, client_id="m7-ver-1")
|
||||
await processing.run_transcribe({}, rec["id"])
|
||||
rid = uuid.UUID(rec["id"])
|
||||
|
||||
async def versions():
|
||||
async with db_session._session_factory() as s:
|
||||
rows = (await s.scalars(
|
||||
select(Transcript).where(Transcript.recording_id == rid)
|
||||
.order_by(Transcript.version))).all()
|
||||
return [(r.version, r.text, r.superseded_at is not None, r.edited_by_user)
|
||||
for r in rows]
|
||||
|
||||
assert await versions() == [(1, "v1 text", False, False)]
|
||||
|
||||
# Queue another transcription run manually (re-entry path).
|
||||
async with db_session._session_factory() as s:
|
||||
from shonar.db.models import Recording
|
||||
|
||||
row = await s.get(Recording, rid)
|
||||
await processing.enqueue_for_recording(s, row)
|
||||
await s.commit()
|
||||
use_fakes(monkeypatch, FakeTranscriber("v2 text"), FakeLlm())
|
||||
await processing.run_transcribe({}, rec["id"])
|
||||
assert await versions() == [(1, "v1 text", True, False), (2, "v2 text", False, False)]
|
||||
|
||||
# A user edit wins: the next auto run inserts nothing.
|
||||
async with db_session._session_factory() as s:
|
||||
v2 = await s.scalar(
|
||||
select(Transcript).where(Transcript.recording_id == rid,
|
||||
Transcript.version == 2))
|
||||
v2.edited_by_user = True
|
||||
v2.text = "user corrected text"
|
||||
from shonar.db.models import Recording
|
||||
|
||||
row = await s.get(Recording, rid)
|
||||
await processing.enqueue_for_recording(s, row)
|
||||
await s.commit()
|
||||
use_fakes(monkeypatch, FakeTranscriber("v3 text"), FakeLlm())
|
||||
await processing.run_transcribe({}, rec["id"])
|
||||
got = await versions()
|
||||
assert len(got) == 2 and got[1][1] == "user corrected text"
|
||||
|
||||
|
||||
# --- llm-only and skips ----------------------------------------------------------
|
||||
|
||||
|
||||
async def test_llm_only_without_transcript_queues_nothing(client, monkeypatch):
|
||||
use_fakes(monkeypatch, None, FakeLlm())
|
||||
token = await user_tokens(client)
|
||||
rec = await upload_recording(client, token, client_id="m7-llm-1")
|
||||
assert rec["processing_status"] == "uploaded"
|
||||
assert await jobs_for(rec["id"]) == {}
|
||||
|
||||
|
||||
async def test_llm_only_with_manual_transcript_summarizes(client, monkeypatch):
|
||||
use_fakes(monkeypatch, None, FakeLlm())
|
||||
token = await user_tokens(client)
|
||||
rec = await upload_recording(client, token, client_id="m7-llm-2")
|
||||
rid = uuid.UUID(rec["id"])
|
||||
async with db_session._session_factory() as s:
|
||||
from shonar.db.models import Recording
|
||||
|
||||
row = await s.get(Recording, rid)
|
||||
s.add(Transcript(recording_id=rid, version=1, provider="manual",
|
||||
text="handwritten notes", edited_by_user=True))
|
||||
await s.flush()
|
||||
queued = await processing.enqueue_for_recording(s, row)
|
||||
assert queued == [JobType.summarize]
|
||||
await s.commit()
|
||||
await processing.run_summarize({}, rec["id"])
|
||||
h = {"Authorization": f"Bearer {token}"}
|
||||
s = await client.get(f"/api/v1/recordings/{rec['id']}/summary", headers=h)
|
||||
assert s.status_code == 200
|
||||
r = await client.get(f"/api/v1/recordings/{rec['id']}", headers=h)
|
||||
assert r.json()["processing_status"] == "completed"
|
||||
|
||||
|
||||
# --- sweep ------------------------------------------------------------------------
|
||||
|
||||
|
||||
async def test_sweep_requeues_stale_jobs(client, monkeypatch):
|
||||
use_fakes(monkeypatch, FakeTranscriber(), FakeLlm())
|
||||
token = await user_tokens(client)
|
||||
rec = await upload_recording(client, token, client_id="m7-sweep-1")
|
||||
rid = uuid.UUID(rec["id"])
|
||||
async with db_session._session_factory() as s:
|
||||
from shonar.db.models import Recording
|
||||
|
||||
row = await s.get(Recording, rid)
|
||||
# Simulate a crashed worker + a lost transport, respectively.
|
||||
await processing.reset_or_create(s, row, JobType.transcribe)
|
||||
jobs = await jobs_for(rec["id"])
|
||||
jobs[JobType.transcribe].status = JobStatus.running
|
||||
await s.commit()
|
||||
# reset_or_create leaves running rows alone, so force the second shape:
|
||||
async with db_session._session_factory() as s:
|
||||
extra = ProcessingJob(recording_id=rid, job_type=JobType.summarize,
|
||||
status=JobStatus.queued)
|
||||
s.add(extra)
|
||||
await s.commit()
|
||||
count = await processing.sweep_stale()
|
||||
assert count == 2
|
||||
jobs = await jobs_for(rec["id"])
|
||||
assert jobs[JobType.transcribe].status == JobStatus.queued
|
||||
assert jobs[JobType.summarize].status == JobStatus.queued
|
||||
|
||||
|
||||
async def test_sweep_leaves_live_running_jobs_alone(client, monkeypatch):
|
||||
"""Regression: sweep re-enqueued every in-flight transcription every
|
||||
5 minutes; with local whisper each copy ran concurrently and one worker
|
||||
OOM-swap-swelled the box. A running job started recently is live work."""
|
||||
from datetime import UTC, datetime, timedelta
|
||||
|
||||
use_fakes(monkeypatch, FakeTranscriber(), FakeLlm())
|
||||
token = await user_tokens(client)
|
||||
rec = await upload_recording(client, token, client_id="m7-sweep-live")
|
||||
rid = uuid.UUID(rec["id"])
|
||||
async with db_session._session_factory() as s:
|
||||
from shonar.db.models import Recording
|
||||
|
||||
row = await s.get(Recording, rid)
|
||||
await processing.reset_or_create(s, row, JobType.transcribe)
|
||||
live = (await s.scalars(
|
||||
select(ProcessingJob).where(
|
||||
ProcessingJob.recording_id == rid,
|
||||
ProcessingJob.job_type == JobType.transcribe)
|
||||
)).one()
|
||||
live.status = JobStatus.running
|
||||
live.started_at = datetime.now(UTC) # worker is chewing it right now
|
||||
stale = ProcessingJob(recording_id=rid, job_type=JobType.summarize,
|
||||
status=JobStatus.running, attempt=1)
|
||||
stale.started_at = datetime.now(UTC) - timedelta(
|
||||
hours=processing.STALE_RUNNING_AFTER.total_seconds() / 3600 + 1)
|
||||
s.add(stale)
|
||||
await s.commit()
|
||||
count = await processing.sweep_stale()
|
||||
assert count == 1 # only the orphaned old run
|
||||
jobs = await jobs_for(rec["id"])
|
||||
assert jobs[JobType.transcribe].status == JobStatus.running
|
||||
assert jobs[JobType.summarize].status == JobStatus.queued
|
||||
|
||||
|
||||
# --- ownership ---------------------------------------------------------------------
|
||||
|
||||
|
||||
async def test_ai_endpoints_enforce_ownership(client, monkeypatch):
|
||||
use_fakes(monkeypatch, FakeTranscriber(), FakeLlm())
|
||||
token = await user_tokens(client)
|
||||
rec = await upload_recording(client, token, client_id="m7-own-1")
|
||||
await processing.run_transcribe({}, rec["id"])
|
||||
other = await user_tokens(client, email="m7-other@example.com",
|
||||
password="m7-test-passw0rd-456")
|
||||
h = {"Authorization": f"Bearer {other}"}
|
||||
for path in ("transcript", "summary", "jobs"):
|
||||
r = await client.get(f"/api/v1/recordings/{rec['id']}/{path}", headers=h)
|
||||
assert r.status_code == 404, path
|
||||
110
backend/tests/test_auth.py
Normal file
110
backend/tests/test_auth.py
Normal file
|
|
@ -0,0 +1,110 @@
|
|||
"""Auth flow tests: register, login, refresh rotation + reuse detection,
|
||||
logout, protected access, account guards."""
|
||||
|
||||
AUTH = {"email": "test@example.com", "password": "correct-horse-battery"}
|
||||
|
||||
|
||||
async def register(client, email=AUTH["email"], password=AUTH["password"]):
|
||||
return await client.post(
|
||||
"/api/v1/auth/register",
|
||||
json={"email": email, "password": password, "display_name": "Tester"},
|
||||
)
|
||||
|
||||
|
||||
async def test_register_returns_token_pair(client):
|
||||
r = await register(client)
|
||||
assert r.status_code == 201, r.text
|
||||
body = r.json()
|
||||
assert body["token_type"] == "bearer"
|
||||
assert body["expires_in"] == 15 * 60
|
||||
assert body["access_token"] and body["refresh_token"]
|
||||
|
||||
|
||||
async def test_register_rejects_duplicate_email(client):
|
||||
assert (await register(client)).status_code == 201
|
||||
r = await register(client)
|
||||
assert r.status_code == 409
|
||||
|
||||
|
||||
async def test_register_rejects_weak_password(client):
|
||||
r = await client.post(
|
||||
"/api/v1/auth/register", json={"email": "x@example.com", "password": "short"}
|
||||
)
|
||||
assert r.status_code == 422
|
||||
|
||||
|
||||
async def test_login_success_and_failure(client):
|
||||
await register(client)
|
||||
r = await client.post(
|
||||
"/api/v1/auth/login", json={"email": AUTH["email"], "password": AUTH["password"]}
|
||||
)
|
||||
assert r.status_code == 200
|
||||
assert r.json()["device_id"]
|
||||
|
||||
bad = await client.post("/api/v1/auth/login", json={"email": AUTH["email"], "password": "***"})
|
||||
assert bad.status_code == 401
|
||||
# Same generic message either way (no user enumeration via password check)
|
||||
bad2 = await client.post(
|
||||
"/api/v1/auth/login", json={"email": "nobody@example.com", "password": "***"}
|
||||
)
|
||||
assert bad2.status_code == 401
|
||||
assert bad.json()["detail"] == bad2.json()["detail"]
|
||||
|
||||
|
||||
async def test_me_requires_valid_token(client):
|
||||
r = await client.get("/api/v1/auth/me")
|
||||
assert r.status_code == 401
|
||||
tok = (await register(client)).json()["access_token"]
|
||||
r = await client.get("/api/v1/auth/me", headers={"Authorization": f"Bearer {tok}"})
|
||||
assert r.status_code == 200
|
||||
assert r.json()["email"] == AUTH["email"]
|
||||
|
||||
|
||||
async def test_refresh_rotates_and_detects_reuse(client):
|
||||
tok = (await register(client)).json()
|
||||
old_refresh = tok["refresh_token"]
|
||||
|
||||
r = await client.post("/api/v1/auth/refresh", json={"refresh_token": old_refresh})
|
||||
assert r.status_code == 200
|
||||
new_refresh = r.json()["refresh_token"]
|
||||
assert new_refresh != old_refresh
|
||||
|
||||
# Old token is dead; reuse kills the whole family.
|
||||
r2 = await client.post("/api/v1/auth/refresh", json={"refresh_token": old_refresh})
|
||||
assert r2.status_code == 401
|
||||
|
||||
# The replacement is also revoked (family revocation).
|
||||
r3 = await client.post("/api/v1/auth/refresh", json={"refresh_token": new_refresh})
|
||||
assert r3.status_code == 401
|
||||
|
||||
|
||||
async def test_logout_revokes_refresh(client):
|
||||
tok = (await register(client)).json()
|
||||
r = await client.post("/api/v1/auth/logout", json={"refresh_token": tok["refresh_token"]})
|
||||
assert r.status_code == 204
|
||||
r2 = await client.post("/api/v1/auth/refresh", json={"refresh_token": tok["refresh_token"]})
|
||||
assert r2.status_code == 401
|
||||
|
||||
|
||||
async def test_access_token_with_refresh_type_rejected(client):
|
||||
tok = (await register(client)).json()
|
||||
r = await client.get(
|
||||
"/api/v1/auth/me", headers={"Authorization": f"Bearer {tok['refresh_token']}"}
|
||||
)
|
||||
assert r.status_code == 401
|
||||
|
||||
|
||||
async def test_delete_account_requires_password(client):
|
||||
tok = (await register(client)).json()
|
||||
h = {"Authorization": f"Bearer {tok['access_token']}"}
|
||||
r = await client.post("/api/v1/auth/delete-account", json={"password": "***"}, headers=h)
|
||||
assert r.status_code == 403
|
||||
r = await client.post(
|
||||
"/api/v1/auth/delete-account", json={"password": AUTH["password"]}, headers=h
|
||||
)
|
||||
assert r.status_code == 202
|
||||
# Deleted account can no longer log in.
|
||||
r = await client.post(
|
||||
"/api/v1/auth/login", json={"email": AUTH["email"], "password": AUTH["password"]}
|
||||
)
|
||||
assert r.status_code == 401
|
||||
29
backend/tests/test_health.py
Normal file
29
backend/tests/test_health.py
Normal file
|
|
@ -0,0 +1,29 @@
|
|||
"""Health-check and system-status tests."""
|
||||
|
||||
|
||||
async def test_healthz(client):
|
||||
r = await client.get("/api/v1/healthz")
|
||||
assert r.status_code == 200
|
||||
body = r.json()
|
||||
assert body["status"] == "ok"
|
||||
assert body["uptime_seconds"] >= 0
|
||||
|
||||
|
||||
async def test_readyz_with_db(client):
|
||||
r = await client.get("/api/v1/readyz")
|
||||
assert r.status_code == 200
|
||||
assert r.json() == {"status": "ok", "database": True}
|
||||
|
||||
|
||||
async def test_system_status_reports_ai_state_honestly(client):
|
||||
r = await client.get("/api/v1/system/status")
|
||||
assert r.status_code == 200
|
||||
body = r.json()
|
||||
# Default test config: AI disabled, no external calls.
|
||||
assert body["ai"]["transcription_enabled"] is False
|
||||
assert body["ai"]["llm_enabled"] is False
|
||||
assert body["ai"]["external_ai_in_use"] is False
|
||||
# No secrets ever present in the public status payload.
|
||||
body_text = r.text.lower()
|
||||
for leak in ("api_key", "password", "secret_key"):
|
||||
assert leak not in body_text
|
||||
153
backend/tests/test_inline_queue.py
Normal file
153
backend/tests/test_inline_queue.py
Normal file
|
|
@ -0,0 +1,153 @@
|
|||
"""Inline queue tests (desktop bundled-lite engine): with
|
||||
``queue_backend=inline`` an upload runs transcribe→summarize automatically
|
||||
in-process — no arq, no Redis. Providers are fakes (see test_ai_pipeline).
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import uuid
|
||||
|
||||
from sqlalchemy import select
|
||||
|
||||
from shonar.core.config import get_settings
|
||||
from shonar.db import session as db_session
|
||||
from shonar.db.models import JobStatus, JobType, ProcessingJob
|
||||
from shonar.services import inline_queue, processing
|
||||
|
||||
from .test_ai_pipeline import FakeLlm, FakeTranscriber, upload_recording, use_fakes, user_tokens
|
||||
|
||||
|
||||
def use_inline(monkeypatch):
|
||||
"""Force queue_backend=inline for everything resolved via processing."""
|
||||
inline_settings = get_settings().model_copy(update={"queue_backend": "inline"})
|
||||
monkeypatch.setattr(processing, "get_settings", lambda: inline_settings)
|
||||
|
||||
|
||||
async def _wait_terminal(recording_id: str, want: int = 2, timeout: float = 20.0):
|
||||
"""Poll the DB until `want` jobs reach a terminal state."""
|
||||
deadline = asyncio.get_running_loop().time() + timeout
|
||||
rows: list[ProcessingJob] = []
|
||||
while asyncio.get_running_loop().time() < deadline:
|
||||
async with db_session._session_factory() as s:
|
||||
rows = list(
|
||||
(
|
||||
await s.scalars(
|
||||
select(ProcessingJob).where(
|
||||
ProcessingJob.recording_id == uuid.UUID(recording_id)
|
||||
)
|
||||
)
|
||||
).all()
|
||||
)
|
||||
if sum(
|
||||
1
|
||||
for j in rows
|
||||
if j.status in (JobStatus.succeeded, JobStatus.failed, JobStatus.skipped)
|
||||
) >= want:
|
||||
return {j.job_type: j.status for j in rows}
|
||||
await asyncio.sleep(0.1)
|
||||
raise AssertionError(
|
||||
f"jobs did not reach terminal state in {timeout}s: "
|
||||
f"{[(j.job_type, j.status) for j in rows]}"
|
||||
)
|
||||
|
||||
|
||||
async def test_inline_backend_processes_upload_without_arq(client, monkeypatch):
|
||||
use_fakes(monkeypatch, FakeTranscriber(), FakeLlm())
|
||||
use_inline(monkeypatch)
|
||||
await inline_queue.start()
|
||||
try:
|
||||
token = await user_tokens(client, email="inline@shonar.dev")
|
||||
rec = await upload_recording(client, token, client_id="inline-1")
|
||||
|
||||
statuses = await _wait_terminal(rec["id"], want=2)
|
||||
assert statuses[JobType.transcribe] == JobStatus.succeeded
|
||||
assert statuses[JobType.summarize] == JobStatus.succeeded
|
||||
|
||||
h = {"Authorization": f"Bearer {token}"}
|
||||
r = await client.get(f"/api/v1/recordings/{rec['id']}", headers=h)
|
||||
assert r.json()["processing_status"] == "completed"
|
||||
t = await client.get(f"/api/v1/recordings/{rec['id']}/transcript", headers=h)
|
||||
assert t.json()["text"] == "hello world from the meeting"
|
||||
s = await client.get(f"/api/v1/recordings/{rec['id']}/summary", headers=h)
|
||||
assert s.json()["content"]["short"] == "Standup happened."
|
||||
finally:
|
||||
await inline_queue.stop()
|
||||
|
||||
|
||||
async def test_inline_retries_transient_failure(client, monkeypatch):
|
||||
stt = FakeTranscriber()
|
||||
state = {"tries": 0}
|
||||
original = stt.transcribe
|
||||
|
||||
async def flaky(audio, mime, *, language_hint=None):
|
||||
state["tries"] += 1
|
||||
if state["tries"] == 1:
|
||||
raise processing.ProviderTransientError("temporary hiccup")
|
||||
return await original(audio, mime, language_hint=language_hint)
|
||||
|
||||
stt.transcribe = flaky
|
||||
use_fakes(monkeypatch, stt, FakeLlm())
|
||||
use_inline(monkeypatch)
|
||||
monkeypatch.setattr(inline_queue, "RETRY_DELAY_SECONDS", 0.05)
|
||||
await inline_queue.start()
|
||||
try:
|
||||
token = await user_tokens(client, email="inline2@shonar.dev")
|
||||
rec = await upload_recording(client, token, client_id="inline-2")
|
||||
statuses = await _wait_terminal(rec["id"], want=2)
|
||||
assert statuses[JobType.transcribe] == JobStatus.succeeded
|
||||
assert state["tries"] >= 2 # the retry actually happened
|
||||
finally:
|
||||
await inline_queue.stop()
|
||||
|
||||
|
||||
async def test_reprocess_reruns_in_place_with_model(client, monkeypatch):
|
||||
"""POST /reprocess?job=transcribe&model= re-runs the pipeline on the
|
||||
SAME recording (no re-upload) and persists the model override."""
|
||||
stt = FakeTranscriber()
|
||||
use_fakes(monkeypatch, stt, FakeLlm())
|
||||
use_inline(monkeypatch)
|
||||
await inline_queue.start()
|
||||
try:
|
||||
token = await user_tokens(client, email="reproc@shonar.dev")
|
||||
h = {"Authorization": f"Bearer {token}"}
|
||||
rec = await upload_recording(client, token, client_id="reproc-1")
|
||||
await _wait_terminal(rec["id"], want=2)
|
||||
calls_before = stt.calls
|
||||
|
||||
r = await client.post(
|
||||
f"/api/v1/recordings/{rec['id']}/reprocess?job=transcribe&model=small",
|
||||
headers=h,
|
||||
)
|
||||
assert r.status_code == 200, r.text
|
||||
statuses = await _wait_terminal(rec["id"], want=2)
|
||||
assert statuses[JobType.transcribe] == JobStatus.succeeded
|
||||
assert stt.calls == calls_before + 1 # re-ran, no new recording
|
||||
|
||||
# Same recording id, model override persisted on the row.
|
||||
r = await client.get(f"/api/v1/recordings/{rec['id']}", headers=h)
|
||||
assert r.json()["id"] == rec["id"]
|
||||
assert r.json()["transcription_model"] == "small"
|
||||
finally:
|
||||
await inline_queue.stop()
|
||||
|
||||
|
||||
async def test_reprocess_validation(client, monkeypatch):
|
||||
token = await user_tokens(client, email="reproc2@shonar.dev")
|
||||
h = {"Authorization": f"Bearer {token}"}
|
||||
rec = await upload_recording(client, token, client_id="reproc-2")
|
||||
# Nothing transcribed yet -> summarize refuses.
|
||||
r = await client.post(
|
||||
f"/api/v1/recordings/{rec['id']}/reprocess?job=summarize", headers=h)
|
||||
assert r.status_code == 409
|
||||
# Bad model name -> 422 (never silently substituted).
|
||||
r = await client.post(
|
||||
f"/api/v1/recordings/{rec['id']}/reprocess?job=transcribe&model=not-a-model",
|
||||
headers=h)
|
||||
assert r.status_code == 422
|
||||
# Not the owner -> 404.
|
||||
other = await user_tokens(client, email="reproc3@shonar.dev")
|
||||
r = await client.post(
|
||||
f"/api/v1/recordings/{rec['id']}/reprocess?job=transcribe",
|
||||
headers={"Authorization": f"Bearer {other}"})
|
||||
assert r.status_code == 404
|
||||
304
backend/tests/test_m9.py
Normal file
304
backend/tests/test_m9.py
Normal file
|
|
@ -0,0 +1,304 @@
|
|||
"""M9 tests: search, exports, retention sweep.
|
||||
|
||||
Search runs against the real backend dialect (Postgres in CI/dev, SQLite
|
||||
via SHONAR_TEST_DATABASE_URL) — both paths share the endpoint contract.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import io
|
||||
import zipfile
|
||||
|
||||
from tests.test_recordings import auth, user_tokens, wav_bytes
|
||||
|
||||
|
||||
async def make_recording(client, token, title, notes=None, tags=None, recorded_at=None):
|
||||
h = await auth(token)
|
||||
audio = wav_bytes()
|
||||
body = {"declared_mime_type": "audio/wav", "declared_size_bytes": len(audio), "title": title}
|
||||
r = await client.post("/api/v1/uploads", json=body, headers=h)
|
||||
sid = r.json()["id"]
|
||||
await client.put(
|
||||
f"/api/v1/uploads/{sid}/chunks/0", content=audio,
|
||||
headers={**h, "content-type": "application/octet-stream"},
|
||||
)
|
||||
fin = {"duration_seconds": 5.0}
|
||||
if recorded_at:
|
||||
fin["recorded_at"] = recorded_at
|
||||
if notes:
|
||||
fin["notes"] = notes
|
||||
r = await client.post(f"/api/v1/uploads/{sid}/finalize", json=fin, headers=h)
|
||||
assert r.status_code == 201, r.text
|
||||
rec = r.json()
|
||||
if tags:
|
||||
r = await client.patch(
|
||||
f"/api/v1/recordings/{rec['id']}", json={"tags": tags}, headers=h
|
||||
)
|
||||
assert r.status_code == 200
|
||||
return rec["id"]
|
||||
|
||||
|
||||
async def add_transcript(rec_id: str, text: str, segments=None):
|
||||
"""Insert a transcript row directly (no AI provider under test)."""
|
||||
from shonar.db.models import Transcript
|
||||
from shonar.db.session import session_factory
|
||||
|
||||
async with session_factory()() as s:
|
||||
s.add(Transcript(recording_id=rec_id, text=text, segments=segments, provider="test"))
|
||||
await s.commit()
|
||||
|
||||
|
||||
async def add_summary(rec_id: str, content: dict):
|
||||
from shonar.db.models import Summary
|
||||
from shonar.db.session import session_factory
|
||||
|
||||
async with session_factory()() as s:
|
||||
s.add(Summary(recording_id=rec_id, content=content, provider="test"))
|
||||
await s.commit()
|
||||
|
||||
|
||||
# --- search -------------------------------------------------------------------
|
||||
|
||||
|
||||
async def test_search_title_and_transcript(client):
|
||||
token = await user_tokens(client, email="m9s1@example.com")
|
||||
rid_t = await make_recording(client, token, "Quarterly budget review")
|
||||
rid_x = await make_recording(client, token, "Grocery list")
|
||||
await add_transcript(rid_x, "remember to buy kale chips and quinoa tonight")
|
||||
|
||||
r = await client.get("/api/v1/search", params={"q": "budget"}, headers=await auth(token))
|
||||
assert r.status_code == 200, r.text
|
||||
body = r.json()
|
||||
assert body["total"] == 1
|
||||
assert body["items"][0]["id"] == rid_t
|
||||
assert body["items"][0]["field"] == "title"
|
||||
|
||||
r = await client.get("/api/v1/search", params={"q": "quinoa"}, headers=await auth(token))
|
||||
body = r.json()
|
||||
assert body["total"] == 1
|
||||
assert body["items"][0]["id"] == rid_x
|
||||
assert body["items"][0]["field"] == "transcript"
|
||||
assert "quinoa" in body["items"][0]["snippet"].lower()
|
||||
|
||||
|
||||
async def test_search_scope_and_tag(client):
|
||||
token = await user_tokens(client, email="m9s2@example.com")
|
||||
rid = await make_recording(client, token, "Standup", tags=["daily"])
|
||||
await add_transcript(rid, "we discussed the daily standup format")
|
||||
|
||||
# tag scope finds by tag name
|
||||
r = await client.get(
|
||||
"/api/v1/search", params={"q": "daily", "scope": "tag"}, headers=await auth(token)
|
||||
)
|
||||
assert r.json()["total"] == 1
|
||||
assert r.json()["items"][0]["field"] == "tag"
|
||||
|
||||
# a transcript-only word does NOT match under scope=title
|
||||
r = await client.get(
|
||||
"/api/v1/search", params={"q": "discussed", "scope": "title"}, headers=await auth(token)
|
||||
)
|
||||
assert r.json()["total"] == 0
|
||||
r = await client.get(
|
||||
"/api/v1/search", params={"q": "discussed", "scope": "transcript"},
|
||||
headers=await auth(token),
|
||||
)
|
||||
assert r.json()["total"] == 1
|
||||
|
||||
|
||||
async def test_search_isolation_and_deleted(client):
|
||||
token_a = await user_tokens(client, email="m9s3a@example.com")
|
||||
token_b = await user_tokens(client, email="m9s3b@example.com")
|
||||
rid = await make_recording(client, token_a, "secret sauce recipe")
|
||||
|
||||
h_b = await auth(token_b)
|
||||
r = await client.get("/api/v1/search", params={"q": "secret"}, headers=h_b)
|
||||
assert r.json()["total"] == 0 # other users' data invisible
|
||||
|
||||
# soft-deleted rows drop out of search
|
||||
await client.delete(f"/api/v1/recordings/{rid}", headers=await auth(token_a))
|
||||
r = await client.get(
|
||||
"/api/v1/search", params={"q": "secret"}, headers=await auth(token_a)
|
||||
)
|
||||
assert r.json()["total"] == 0
|
||||
|
||||
|
||||
async def test_search_summary_and_notes(client):
|
||||
token = await user_tokens(client, email="m9s4@example.com")
|
||||
rid = await make_recording(client, token, "Meeting", notes="bring the projector cable")
|
||||
await add_summary(rid, {"short": "sprint retro", "action_items": ["fix flaky test"]})
|
||||
|
||||
r = await client.get("/api/v1/search", params={"q": "projector"}, headers=await auth(token))
|
||||
assert r.json()["total"] == 1
|
||||
r = await client.get("/api/v1/search", params={"q": "flaky"}, headers=await auth(token))
|
||||
assert r.json()["total"] == 1
|
||||
|
||||
|
||||
# --- list filters ---------------------------------------------------------------
|
||||
|
||||
|
||||
async def test_list_filters_tag_status_date(client):
|
||||
token = await user_tokens(client, email="m9f1@example.com")
|
||||
await make_recording(client, token, "Old one", tags=["keep"],
|
||||
recorded_at="2020-01-01T10:00:00Z")
|
||||
rid_new = await make_recording(client, token, "New one", tags=["keep"],
|
||||
recorded_at="2026-01-01T10:00:00Z")
|
||||
h = await auth(token)
|
||||
|
||||
r = await client.get("/api/v1/recordings", params={"tag": "keep"}, headers=h)
|
||||
assert r.json()["total"] == 2
|
||||
r = await client.get(
|
||||
"/api/v1/recordings", params={"tag": "keep", "from_date": "2025-06-01T00:00:00Z"},
|
||||
headers=h,
|
||||
)
|
||||
body = r.json()
|
||||
assert body["total"] == 1 and body["items"][0]["id"] == rid_new
|
||||
|
||||
r = await client.get("/api/v1/recordings", params={"status": "ai_disabled"}, headers=h)
|
||||
assert r.json()["total"] == 2
|
||||
r = await client.get("/api/v1/recordings", params={"status": "bogus"}, headers=h)
|
||||
assert r.status_code == 422
|
||||
r = await client.get("/api/v1/recordings", params={"tag": "nope"}, headers=h)
|
||||
assert r.json()["total"] == 0
|
||||
|
||||
|
||||
# --- exports --------------------------------------------------------------------
|
||||
|
||||
|
||||
async def test_export_formats(client):
|
||||
token = await user_tokens(client, email="m9e1@example.com")
|
||||
rid = await make_recording(client, token, "Retro & Planning", notes="retro notes here",
|
||||
tags=["team"])
|
||||
await add_transcript(
|
||||
rid, "first segment second segment",
|
||||
segments=[{"start": 0.0, "end": 2.5, "text": "first segment", "speaker": None},
|
||||
{"start": 2.5, "end": 5.0, "text": "second segment", "speaker": "S1"}],
|
||||
)
|
||||
await add_summary(rid, {"short": "one line", "action_items": ["do the thing"]})
|
||||
h = await auth(token)
|
||||
|
||||
r = await client.get(f"/api/v1/recordings/{rid}/export", params={"fmt": "txt"}, headers=h)
|
||||
assert r.status_code == 200
|
||||
assert r.text == "first segment second segment"
|
||||
assert "attachment" in r.headers["content-disposition"]
|
||||
assert ".txt" in r.headers["content-disposition"]
|
||||
|
||||
r = await client.get(f"/api/v1/recordings/{rid}/export", params={"fmt": "md"}, headers=h)
|
||||
assert r.status_code == 200
|
||||
assert "# Retro & Planning" in r.text
|
||||
assert "do the thing" in r.text
|
||||
assert "`[00:02.5]` **S1**: second segment" in r.text
|
||||
assert "`team`" in r.text
|
||||
|
||||
r = await client.get(f"/api/v1/recordings/{rid}/export", params={"fmt": "zip"}, headers=h)
|
||||
assert r.status_code == 200
|
||||
assert r.headers["content-type"] == "application/zip"
|
||||
with zipfile.ZipFile(io.BytesIO(r.content)) as z:
|
||||
names = z.namelist()
|
||||
assert "transcript.txt" in names and "notes.md" in names
|
||||
assert any(n.endswith(".wav") for n in names)
|
||||
|
||||
r = await client.get(f"/api/v1/recordings/{rid}/export", params={"fmt": "audio"}, headers=h)
|
||||
assert r.status_code == 200
|
||||
assert r.content.startswith(b"RIFF")
|
||||
|
||||
|
||||
async def test_export_missing_and_foreign(client):
|
||||
token = await user_tokens(client, email="m9e2@example.com")
|
||||
rid = await make_recording(client, token, "Bare") # no transcript/summary/notes
|
||||
token2 = await user_tokens(client, email="m9e2b@example.com")
|
||||
|
||||
r = await client.get(f"/api/v1/recordings/{rid}/export", params={"fmt": "txt"},
|
||||
headers=await auth(token))
|
||||
assert r.status_code == 404 # no transcript yet
|
||||
r = await client.get(f"/api/v1/recordings/{rid}/export", params={"fmt": "md"},
|
||||
headers=await auth(token))
|
||||
assert r.status_code == 404
|
||||
# zip still works with just the audio
|
||||
r = await client.get(f"/api/v1/recordings/{rid}/export", params={"fmt": "zip"},
|
||||
headers=await auth(token))
|
||||
assert r.status_code == 200
|
||||
|
||||
r = await client.get(f"/api/v1/recordings/{rid}/export", params={"fmt": "audio"},
|
||||
headers=await auth(token2))
|
||||
assert r.status_code == 404 # not yours
|
||||
|
||||
|
||||
# --- retention sweep --------------------------------------------------------------
|
||||
|
||||
|
||||
async def test_retention_sweep_purges_expired(client):
|
||||
from datetime import timedelta
|
||||
|
||||
from shonar.db.models import Asset, Recording, utcnow
|
||||
from shonar.db.session import session_factory
|
||||
from shonar.services import retention
|
||||
|
||||
token = await user_tokens(client, email="m9r1@example.com")
|
||||
rid = await make_recording(client, token, "Doomed")
|
||||
h = await auth(token)
|
||||
|
||||
# Soft-delete, then push deleted_at past the grace window directly.
|
||||
r = await client.delete(f"/api/v1/recordings/{rid}", headers=h)
|
||||
assert r.status_code == 204
|
||||
async with session_factory()() as s:
|
||||
rec = await s.get(Recording, rid)
|
||||
rec.deleted_at = utcnow() - timedelta(days=31)
|
||||
# remember the storage key before the row vanishes
|
||||
from shonar.db.models import Asset
|
||||
|
||||
asset = (
|
||||
await s.execute(
|
||||
Asset.__table__.select().where(Asset.recording_id == rec.id) # noqa: SLF001
|
||||
)
|
||||
).first()
|
||||
storage_key = asset.storage_key if asset else None
|
||||
await s.commit()
|
||||
assert storage_key is not None
|
||||
|
||||
purged = await retention.sweep_deleted()
|
||||
assert purged["recordings"] == 1
|
||||
assert purged["files"] >= 1
|
||||
|
||||
async with session_factory()() as s:
|
||||
assert await s.get(Recording, rid) is None
|
||||
from shonar.storage import get_storage
|
||||
|
||||
assert not await get_storage().exists(storage_key) # file gone too
|
||||
|
||||
# Within the window, nothing is purged (cancellable).
|
||||
rid2 = await make_recording(client, token, "Fresh delete")
|
||||
await client.delete(f"/api/v1/recordings/{rid2}", headers=h)
|
||||
purged = await retention.sweep_deleted()
|
||||
assert purged["recordings"] == 0
|
||||
async with session_factory()() as s:
|
||||
assert await s.get(Recording, rid2) is not None
|
||||
|
||||
|
||||
async def test_retention_sweep_account(client):
|
||||
from datetime import timedelta
|
||||
|
||||
from shonar.db.models import Recording, utcnow
|
||||
from shonar.db.session import session_factory
|
||||
from shonar.services import retention
|
||||
|
||||
token = await user_tokens(client, email="m9r2@example.com")
|
||||
rid = await make_recording(client, token, "Gone with user")
|
||||
|
||||
async with session_factory()() as s:
|
||||
user = await s.scalar(select_user("m9r2@example.com"))
|
||||
user.deleted_at = utcnow() - timedelta(days=31)
|
||||
await s.commit()
|
||||
|
||||
purged = await retention.sweep_deleted()
|
||||
assert purged["users"] == 1
|
||||
async with session_factory()() as s:
|
||||
assert await s.get(Recording, rid) is None # cascade
|
||||
assert await s.scalar(select_user("m9r2@example.com")) is None
|
||||
|
||||
|
||||
def select_user(email: str):
|
||||
from sqlalchemy import select
|
||||
|
||||
from shonar.db.models import User
|
||||
|
||||
return select(User).where(User.email == email)
|
||||
245
backend/tests/test_models.py
Normal file
245
backend/tests/test_models.py
Normal file
|
|
@ -0,0 +1,245 @@
|
|||
"""Stage 1: per-recording transcription models + models API.
|
||||
|
||||
Global default (base) with optional per-recording overrides; the worker
|
||||
uses the exact saved model; history is never rewritten; unknown models
|
||||
are rejected; missing downloads fail fast with instructions.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import struct
|
||||
|
||||
import pytest
|
||||
from sqlalchemy import select
|
||||
|
||||
from shonar.db import session as db_session
|
||||
from shonar.db.models import Recording
|
||||
from shonar.services import processing
|
||||
|
||||
|
||||
def wav_bytes(payload_len: int = 64) -> bytes:
|
||||
data = bytes(range(payload_len % 256)) * (payload_len // 256 + 1)
|
||||
data = data[:payload_len]
|
||||
return (
|
||||
b"RIFF" + struct.pack("<I", 36 + len(data)) + b"WAVE"
|
||||
+ b"fmt " + struct.pack("<IHHIIHH", 16, 1, 1, 8000, 8000, 1, 8)
|
||||
+ b"data" + struct.pack("<I", len(data))
|
||||
) + data
|
||||
|
||||
|
||||
AUTH = {"email": "models@example.com", "password": "models-test-passw0rd-123"}
|
||||
|
||||
|
||||
async def user_tokens(client):
|
||||
r = await client.post(
|
||||
"/api/v1/auth/register",
|
||||
json={"email": AUTH["email"], "password": AUTH["password"]},
|
||||
)
|
||||
assert r.status_code == 201, r.text
|
||||
return r.json()["access_token"]
|
||||
|
||||
|
||||
async def upload_recording(client, token, client_id, session_model=None, finalize_model=None):
|
||||
h = {"Authorization": f"Bearer {token}"}
|
||||
data = wav_bytes()
|
||||
body = {"declared_mime_type": "audio/wav", "declared_size_bytes": len(data),
|
||||
"client_recording_id": client_id}
|
||||
if session_model is not None:
|
||||
body["transcription_model"] = session_model
|
||||
r = await client.post("/api/v1/uploads", json=body, headers=h)
|
||||
assert r.status_code == 201, r.text
|
||||
sid = r.json()["id"]
|
||||
r = await client.put(f"/api/v1/uploads/{sid}/chunks/0", content=data,
|
||||
headers={**h, "content-type": "application/octet-stream"})
|
||||
assert r.status_code == 201, r.text
|
||||
fbody = {"duration_seconds": 5.0}
|
||||
if finalize_model is not None:
|
||||
fbody["transcription_model"] = finalize_model
|
||||
r = await client.post(f"/api/v1/uploads/{sid}/finalize", json=fbody, headers=h)
|
||||
return r
|
||||
|
||||
|
||||
async def test_models_endpoint_lists_five(client):
|
||||
token = await user_tokens(client)
|
||||
r = await client.get("/api/v1/models", headers={"Authorization": f"Bearer {token}"})
|
||||
assert r.status_code == 200, r.text
|
||||
body = r.json()
|
||||
assert body["default_model"] == "base"
|
||||
assert [m["name"] for m in body["models"]] == ["tiny", "base", "small", "medium", "large-v3"]
|
||||
base = next(m for m in body["models"] if m["name"] == "base")
|
||||
assert base["is_default"] is True
|
||||
assert "ordinary computers" in base["description"]
|
||||
for m in body["models"]:
|
||||
assert isinstance(m["downloaded"], bool)
|
||||
assert isinstance(m["available"], bool)
|
||||
|
||||
|
||||
async def test_default_get_put_validation(client):
|
||||
token = await user_tokens(client)
|
||||
h = {"Authorization": f"Bearer {token}"}
|
||||
r = await client.get("/api/v1/models/default", headers=h)
|
||||
assert r.json() == {"default_model": "base"}
|
||||
r = await client.put("/api/v1/models/default", json={"model": "small"}, headers=h)
|
||||
assert r.status_code == 200, r.text
|
||||
assert r.json() == {"default_model": "small"}
|
||||
r = await client.put("/api/v1/models/default", json={"model": "xxl-turbo"}, headers=h)
|
||||
assert r.status_code == 422
|
||||
assert "Supported models" in r.text
|
||||
# rejected change did not stick
|
||||
r = await client.get("/api/v1/models/default", headers=h)
|
||||
assert r.json() == {"default_model": "small"}
|
||||
|
||||
|
||||
async def test_finalize_override_and_default_history(client):
|
||||
token = await user_tokens(client)
|
||||
h = {"Authorization": f"Bearer {token}"}
|
||||
r = await upload_recording(client, token, "m-override-1", finalize_model="small")
|
||||
assert r.status_code == 201, r.text
|
||||
assert r.json()["transcription_model"] == "small"
|
||||
|
||||
r = await upload_recording(client, token, "m-default-1")
|
||||
assert r.json()["transcription_model"] == "base"
|
||||
|
||||
# Changing the default affects future rows only.
|
||||
r = await client.put("/api/v1/models/default", json={"model": "tiny"}, headers=h)
|
||||
assert r.status_code == 200
|
||||
r = await upload_recording(client, token, "m-default-2")
|
||||
assert r.json()["transcription_model"] == "tiny"
|
||||
|
||||
async with db_session._session_factory() as s:
|
||||
got = dict((await s.execute(
|
||||
select(Recording.client_recording_id, Recording.transcription_model)
|
||||
.where(Recording.client_recording_id.in_(
|
||||
["m-override-1", "m-default-1", "m-default-2"])))
|
||||
).all())
|
||||
assert got == {"m-override-1": "small", "m-default-1": "base", "m-default-2": "tiny"}
|
||||
|
||||
|
||||
async def test_session_override_and_finalize_wins(client):
|
||||
token = await user_tokens(client)
|
||||
r = await upload_recording(client, token, "m-sess-1", session_model="small")
|
||||
assert r.status_code == 201, r.text
|
||||
assert r.json()["transcription_model"] == "small"
|
||||
|
||||
r = await upload_recording(client, token, "m-sess-2", session_model="small",
|
||||
finalize_model="tiny")
|
||||
assert r.status_code == 201, r.text
|
||||
assert r.json()["transcription_model"] == "tiny"
|
||||
|
||||
|
||||
async def test_finalize_invalid_override_rejected(client):
|
||||
token = await user_tokens(client)
|
||||
r = await upload_recording(client, token, "m-bad-1", finalize_model="xxl-turbo")
|
||||
assert r.status_code == 422
|
||||
assert "Supported models" in r.text
|
||||
|
||||
|
||||
class SpyTranscriber:
|
||||
"""Stands in for FasterWhisperProvider; records the model it was built with."""
|
||||
|
||||
name = "faster_whisper"
|
||||
seen_models: list = []
|
||||
|
||||
def __init__(self, model: str = "base"):
|
||||
type(self).seen_models.append(model)
|
||||
self.model = model
|
||||
|
||||
async def transcribe(self, audio, mime, *, language_hint=None):
|
||||
from shonar.services.ai import Segment, TranscriptResult
|
||||
|
||||
return TranscriptResult(
|
||||
text="spy text", language="en",
|
||||
segments=[Segment(0.0, 1.0, "spy text")], model=self.model,
|
||||
)
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def _clear_spy():
|
||||
SpyTranscriber.seen_models = []
|
||||
yield
|
||||
SpyTranscriber.seen_models = []
|
||||
|
||||
|
||||
async def _run_with_spy(monkeypatch, recording_id):
|
||||
import shonar.services.ai.faster_whisper as fw
|
||||
|
||||
# The worker rebuilds a faster_whisper provider from the saved model;
|
||||
# the late import inside run_transcribe picks up this spy.
|
||||
monkeypatch.setattr(fw, "FasterWhisperProvider", SpyTranscriber)
|
||||
monkeypatch.setattr(
|
||||
processing, "get_transcription_provider",
|
||||
lambda settings: SpyTranscriber(model="ignored"),
|
||||
)
|
||||
await processing.run_transcribe({}, recording_id)
|
||||
|
||||
|
||||
async def test_worker_uses_saved_model(client, monkeypatch):
|
||||
token = await user_tokens(client)
|
||||
rec = (await upload_recording(client, token, "m-work-1", finalize_model="small")).json()
|
||||
await _run_with_spy(monkeypatch, rec["id"])
|
||||
# Last construction is the worker rebuild from the saved model.
|
||||
assert SpyTranscriber.seen_models[-1] == "small"
|
||||
|
||||
h = {"Authorization": f"Bearer {token}"}
|
||||
r = await client.get(f"/api/v1/recordings/{rec['id']}/transcript", headers=h)
|
||||
assert r.status_code == 200
|
||||
assert r.json()["model"] == "small"
|
||||
|
||||
r = await client.get(f"/api/v1/recordings/{rec['id']}/jobs", headers=h)
|
||||
job = next(j for j in r.json() if j["job_type"] == "transcribe")
|
||||
assert job["status"] == "succeeded"
|
||||
assert job["stage"] is None
|
||||
assert job["progress"] == 100
|
||||
|
||||
|
||||
async def test_worker_keeps_history_after_default_change(client, monkeypatch):
|
||||
token = await user_tokens(client)
|
||||
h = {"Authorization": f"Bearer {token}"}
|
||||
rec = (await upload_recording(client, token, "m-work-2")).json()
|
||||
assert rec["transcription_model"] == "base"
|
||||
r = await client.put("/api/v1/models/default", json={"model": "tiny"}, headers=h)
|
||||
assert r.status_code == 200
|
||||
await _run_with_spy(monkeypatch, rec["id"])
|
||||
assert SpyTranscriber.seen_models[-1] == "base"
|
||||
|
||||
|
||||
async def test_unavailable_model_fails_with_instructions(client, monkeypatch):
|
||||
import shonar.services.ai.faster_whisper as fw
|
||||
from shonar.services.ai.faster_whisper import FasterWhisperProvider
|
||||
|
||||
# Real provider class, but the model is (simulated) not downloaded.
|
||||
monkeypatch.setattr(fw, "is_model_downloaded", lambda name: False)
|
||||
monkeypatch.setattr(
|
||||
processing, "get_transcription_provider",
|
||||
lambda settings: FasterWhisperProvider(model="base"),
|
||||
)
|
||||
token = await user_tokens(client)
|
||||
h = {"Authorization": f"Bearer {token}"}
|
||||
rec = (await upload_recording(client, token, "m-work-3")).json()
|
||||
await processing.run_transcribe({}, rec["id"])
|
||||
|
||||
r = await client.get(f"/api/v1/recordings/{rec['id']}", headers=h)
|
||||
assert r.json()["processing_status"] == "failed"
|
||||
assert "not downloaded" in (r.json()["processing_error"] or "")
|
||||
r = await client.get(f"/api/v1/recordings/{rec['id']}/jobs", headers=h)
|
||||
job = next(j for j in r.json() if j["job_type"] == "transcribe")
|
||||
assert job["status"] == "failed"
|
||||
assert "download" in (job["error"] or "").lower()
|
||||
|
||||
|
||||
async def test_reupload_explicit_override_updates_model(client):
|
||||
token = await user_tokens(client)
|
||||
r = await upload_recording(client, token, "m-reup-1", finalize_model="small")
|
||||
assert r.json()["transcription_model"] == "small"
|
||||
# Same client id, new explicit choice: model history moves with it.
|
||||
r = await upload_recording(client, token, "m-reup-1", finalize_model="tiny")
|
||||
assert r.status_code == 201, r.text
|
||||
assert r.json()["transcription_model"] == "tiny"
|
||||
# Same client id, no choice: history untouched.
|
||||
r = await upload_recording(client, token, "m-reup-1")
|
||||
assert r.json()["transcription_model"] == "tiny"
|
||||
|
||||
|
||||
async def test_models_need_auth(client):
|
||||
r = await client.get("/api/v1/models")
|
||||
assert r.status_code in (401, 403)
|
||||
35
backend/tests/test_provider_info.py
Normal file
35
backend/tests/test_provider_info.py
Normal file
|
|
@ -0,0 +1,35 @@
|
|||
"""provider-info handshake tests (P2/P3 client probing depends on this)."""
|
||||
|
||||
|
||||
async def test_provider_info_identifies_shonar(client):
|
||||
r = await client.get("/api/v1/provider-info")
|
||||
assert r.status_code == 200
|
||||
body = r.json()
|
||||
assert body["kind"] == "shonar"
|
||||
assert body["api_version"] == "v1"
|
||||
caps = body["capabilities"]
|
||||
assert caps["chunked_upload"] is True
|
||||
assert caps["account_deletion"] is True
|
||||
# test env configures no AI providers -> flags must be false
|
||||
assert caps["server_transcription"] is False
|
||||
assert caps["server_summary"] is False
|
||||
assert body["storage_backend"] in ("local", "s3")
|
||||
|
||||
|
||||
async def test_provider_info_leaks_no_paths_or_secrets(client):
|
||||
r = await client.get("/api/v1/provider-info")
|
||||
body = r.json()
|
||||
# every string value must be a bare identifier, never a filesystem path
|
||||
def walk(v):
|
||||
if isinstance(v, str):
|
||||
assert not v.startswith("/"), f"path-like value in handshake: {v!r}"
|
||||
low = v.lower()
|
||||
for banned in ("secret", "token", "password", "key"):
|
||||
assert banned not in low
|
||||
elif isinstance(v, dict):
|
||||
for x in v.values():
|
||||
walk(x)
|
||||
elif isinstance(v, list):
|
||||
for x in v:
|
||||
walk(x)
|
||||
walk(body)
|
||||
264
backend/tests/test_recordings.py
Normal file
264
backend/tests/test_recordings.py
Normal file
|
|
@ -0,0 +1,264 @@
|
|||
"""Upload sessions + recordings CRUD tests (M2).
|
||||
|
||||
Uses real WAV magic bytes; validation is byte-level, so fakes would be
|
||||
testing the wrong thing.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import struct
|
||||
import uuid
|
||||
|
||||
|
||||
def wav_bytes(payload_len: int = 64) -> bytes:
|
||||
data = bytes(range(payload_len % 256)) * (payload_len // 256 + 1)
|
||||
data = data[:payload_len]
|
||||
header = (
|
||||
b"RIFF" + struct.pack("<I", 36 + len(data)) + b"WAVE"
|
||||
+ b"fmt " + struct.pack("<IHHIIHH", 16, 1, 1, 8000, 8000, 1, 8)
|
||||
+ b"data" + struct.pack("<I", len(data))
|
||||
)
|
||||
return header + data
|
||||
|
||||
|
||||
def mp4_bytes() -> bytes:
|
||||
return b"\x00\x00\x00 ftypM4A " + b"\x00" * 64
|
||||
|
||||
|
||||
AUTH = {"email": "m2@example.com",
|
||||
"password": "m2-test-" + "passw0rd-123"}
|
||||
|
||||
|
||||
async def user_tokens(client, email=AUTH["email"], password=AUTH["password"]):
|
||||
r = await client.post("/api/v1/auth/register", json={"email": email, "password": password})
|
||||
assert r.status_code == 201, r.text
|
||||
return r.json()["access_token"]
|
||||
|
||||
|
||||
async def auth(token: str) -> dict:
|
||||
return {"Authorization": f"Bearer {token}"}
|
||||
|
||||
|
||||
async def upload_full(client, token: str, data: bytes, mime="audio/wav",
|
||||
client_id=None, title=None):
|
||||
h = await auth(token)
|
||||
r = await client.post(
|
||||
"/api/v1/uploads",
|
||||
json={"declared_mime_type": mime, "declared_size_bytes": len(data),
|
||||
"client_recording_id": client_id, "title": title},
|
||||
headers=h,
|
||||
)
|
||||
assert r.status_code == 201, r.text
|
||||
sid = r.json()["id"]
|
||||
r = await client.put(f"/api/v1/uploads/{sid}/chunks/0", content=data,
|
||||
headers={**h, "content-type": "application/octet-stream"})
|
||||
assert r.status_code == 201, r.text
|
||||
r = await client.post(
|
||||
f"/api/v1/uploads/{sid}/finalize",
|
||||
json={"duration_seconds": 12.5},
|
||||
headers=h,
|
||||
)
|
||||
return sid, r
|
||||
|
||||
|
||||
# --- upload session ----------------------------------------------------------
|
||||
|
||||
|
||||
async def test_upload_happy_path_creates_recording(client):
|
||||
token = await user_tokens(client)
|
||||
data = wav_bytes()
|
||||
sid, r = await upload_full(client, token, data, title="Standup")
|
||||
assert r.status_code == 201, r.text
|
||||
rec = r.json()
|
||||
assert rec["title"] == "Standup"
|
||||
assert rec["has_audio"] is True
|
||||
# M7: with no AI configured the pipeline marks audio-only explicitly.
|
||||
assert rec["processing_status"] == "ai_disabled"
|
||||
assert rec["duration_seconds"] == 12.5
|
||||
# No storage keys or internals leak.
|
||||
assert "storage" not in r.text and "key" not in r.text.lower().replace("chunk", "")
|
||||
|
||||
|
||||
async def test_upload_rejects_bad_mime_declared(client):
|
||||
token = await user_tokens(client)
|
||||
h = await auth(token)
|
||||
r = await client.post("/api/v1/uploads",
|
||||
json={"declared_mime_type": "application/x-msdownload",
|
||||
"declared_size_bytes": 100}, headers=h)
|
||||
assert r.status_code == 415
|
||||
|
||||
|
||||
async def test_upload_rejects_oversize(client):
|
||||
token = await user_tokens(client)
|
||||
h = await auth(token)
|
||||
r = await client.post("/api/v1/uploads",
|
||||
json={"declared_mime_type": "audio/wav",
|
||||
"declared_size_bytes": 5 * 1024**3}, headers=h)
|
||||
assert r.status_code == 413
|
||||
|
||||
|
||||
async def test_finalize_rejects_bytes_not_matching_mime(client):
|
||||
token = await user_tokens(client)
|
||||
data = mp4_bytes()
|
||||
_sid, r = await upload_full(client, token, data, mime="audio/wav")
|
||||
assert r.status_code == 415
|
||||
|
||||
|
||||
async def test_finalize_rejects_size_mismatch(client):
|
||||
token = await user_tokens(client)
|
||||
h = await auth(token)
|
||||
data = wav_bytes()
|
||||
r = await client.post("/api/v1/uploads",
|
||||
json={"declared_mime_type": "audio/wav",
|
||||
"declared_size_bytes": len(data) + 10}, headers=h)
|
||||
sid = r.json()["id"]
|
||||
await client.put(f"/api/v1/uploads/{sid}/chunks/0", content=data,
|
||||
headers={**h, "content-type": "application/octet-stream"})
|
||||
r = await client.post(f"/api/v1/uploads/{sid}/finalize", json={}, headers=h)
|
||||
assert r.status_code == 422
|
||||
detail = r.json()["detail"].lower()
|
||||
assert "missing chunks" in detail or "size mismatch" in detail
|
||||
|
||||
|
||||
async def test_chunk_resume_status_and_idempotency(client):
|
||||
token = await user_tokens(client)
|
||||
h = await auth(token)
|
||||
data = wav_bytes()
|
||||
r = await client.post("/api/v1/uploads",
|
||||
json={"declared_mime_type": "audio/wav",
|
||||
"declared_size_bytes": len(data)}, headers=h)
|
||||
sid = r.json()["id"]
|
||||
r = await client.get(f"/api/v1/uploads/{sid}", headers=h)
|
||||
assert r.status_code == 200
|
||||
assert r.json()["received_chunk_indexes"] == []
|
||||
hdr = {**h, "content-type": "application/octet-stream"}
|
||||
await client.put(f"/api/v1/uploads/{sid}/chunks/0", content=data, headers=hdr)
|
||||
# Duplicate PUT of chunk 0 (retry) must not duplicate or corrupt.
|
||||
await client.put(f"/api/v1/uploads/{sid}/chunks/0", content=data, headers=hdr)
|
||||
r = await client.get(f"/api/v1/uploads/{sid}", headers=h)
|
||||
assert r.json()["received_chunk_indexes"] == [0]
|
||||
r = await client.post(f"/api/v1/uploads/{sid}/finalize", json={}, headers=h)
|
||||
assert r.status_code == 201
|
||||
|
||||
|
||||
async def test_chunk_checksum_enforced(client):
|
||||
token = await user_tokens(client)
|
||||
h = await auth(token)
|
||||
data = wav_bytes()
|
||||
r = await client.post("/api/v1/uploads",
|
||||
json={"declared_mime_type": "audio/wav",
|
||||
"declared_size_bytes": len(data)}, headers=h)
|
||||
sid = r.json()["id"]
|
||||
r = await client.put(f"/api/v1/uploads/{sid}/chunks/0", content=data,
|
||||
headers={**h, "content-type": "application/octet-stream",
|
||||
"x-chunk-sha256": "0" * 64})
|
||||
assert r.status_code == 422
|
||||
|
||||
|
||||
async def test_finalize_idempotent_per_client_recording_id(client):
|
||||
token = await user_tokens(client)
|
||||
cid = str(uuid.uuid4())
|
||||
_sid1, r1 = await upload_full(client, token, wav_bytes(64), client_id=cid)
|
||||
_sid2, r2 = await upload_full(client, token, wav_bytes(64), client_id=cid, title="Renamed")
|
||||
assert r1.status_code == 201 and r2.status_code == 201
|
||||
# Same recording id, original preserved, metadata updated.
|
||||
assert r1.json()["id"] == r2.json()["id"]
|
||||
assert r2.json()["title"] == "Renamed"
|
||||
|
||||
|
||||
# --- ownership ---------------------------------------------------------------
|
||||
|
||||
|
||||
async def test_cross_user_isolation(client):
|
||||
ta = await user_tokens(client, "a@example.com")
|
||||
tb = await user_tokens(client, "b@example.com")
|
||||
_sid, r = await upload_full(client, ta, wav_bytes())
|
||||
rec_id = r.json()["id"]
|
||||
r = await client.get(f"/api/v1/recordings/{rec_id}", headers=await auth(tb))
|
||||
assert r.status_code == 404
|
||||
r = await client.get(f"/api/v1/recordings/{rec_id}/audio", headers=await auth(tb))
|
||||
assert r.status_code == 404
|
||||
r = await client.get("/api/v1/recordings", headers=await auth(tb))
|
||||
assert r.json()["total"] == 0
|
||||
|
||||
|
||||
async def test_uploads_require_auth(client):
|
||||
r = await client.get("/api/v1/recordings")
|
||||
assert r.status_code == 401
|
||||
|
||||
|
||||
# --- recordings CRUD -----------------------------------------------------------
|
||||
|
||||
|
||||
async def test_update_metadata_and_tags(client):
|
||||
token = await user_tokens(client)
|
||||
h = await auth(token)
|
||||
_sid, r = await upload_full(client, token, wav_bytes())
|
||||
rec_id = r.json()["id"]
|
||||
r = await client.patch(f"/api/v1/recordings/{rec_id}",
|
||||
json={"title": "Sync meeting", "notes": "n1",
|
||||
"tags": ["Work", " meeting ", "work"]}, headers=h)
|
||||
assert r.status_code == 200
|
||||
body = r.json()
|
||||
assert body["title"] == "Sync meeting"
|
||||
assert body["tags"] == ["meeting", "work"] # normalized, deduped, sorted
|
||||
# listing shows same
|
||||
r = await client.get("/api/v1/recordings", headers=h)
|
||||
assert r.json()["total"] == 1
|
||||
assert r.json()["items"][0]["tags"] == ["meeting", "work"]
|
||||
|
||||
|
||||
async def test_location_dropped_without_consent(client):
|
||||
token = await user_tokens(client)
|
||||
h = await auth(token)
|
||||
_sid, r = await upload_full(client, token, wav_bytes())
|
||||
rec_id = r.json()["id"]
|
||||
r = await client.patch(f"/api/v1/recordings/{rec_id}",
|
||||
json={"latitude": 41.8, "longitude": -87.6}, headers=h)
|
||||
assert r.json()["latitude"] is None
|
||||
# enable consent
|
||||
await client.patch("/api/v1/users/me", json={"location_storage_enabled": True}, headers=h)
|
||||
_sid2, r2 = await upload_full(client, token, wav_bytes(128), client_id=str(uuid.uuid4()))
|
||||
rid2 = r2.json()["id"]
|
||||
r = await client.patch(f"/api/v1/recordings/{rid2}",
|
||||
json={"latitude": 41.8, "longitude": -87.6}, headers=h)
|
||||
assert r.json()["latitude"] == 41.8
|
||||
|
||||
|
||||
async def test_soft_then_purge_delete(client, storage_root):
|
||||
token = await user_tokens(client)
|
||||
h = await auth(token)
|
||||
_sid, r = await upload_full(client, token, wav_bytes())
|
||||
rec_id = r.json()["id"]
|
||||
|
||||
files_before = list(storage_root.rglob("*"))
|
||||
assert any(p.is_file() for p in files_before)
|
||||
|
||||
r = await client.delete(f"/api/v1/recordings/{rec_id}", headers=h)
|
||||
assert r.status_code == 204
|
||||
r = await client.get(f"/api/v1/recordings/{rec_id}", headers=h)
|
||||
assert r.status_code == 404
|
||||
|
||||
# Purge deletes rows AND stored files.
|
||||
token2 = await user_tokens(client, "p2@example.com")
|
||||
h2 = await auth(token2)
|
||||
_sid, r = await upload_full(client, token2, wav_bytes(96))
|
||||
rec2 = r.json()["id"]
|
||||
r = await client.delete(f"/api/v1/recordings/{rec2}?purge=true", headers=h2)
|
||||
assert r.status_code == 204
|
||||
remaining = [p for p in storage_root.rglob("*") if p.is_file() and f"{rec2}" in str(p)]
|
||||
assert remaining == []
|
||||
|
||||
|
||||
async def test_download_audio_roundtrip(client):
|
||||
token = await user_tokens(client)
|
||||
h = await auth(token)
|
||||
data = wav_bytes(128)
|
||||
_sid, r = await upload_full(client, token, data)
|
||||
rec_id = r.json()["id"]
|
||||
r = await client.get(f"/api/v1/recordings/{rec_id}/audio", headers=h)
|
||||
assert r.status_code == 200
|
||||
assert r.content == data
|
||||
assert r.headers["content-type"] == "audio/wav"
|
||||
assert "attachment" in r.headers["content-disposition"]
|
||||
assert "no-store" in r.headers["cache-control"]
|
||||
189
backend/tests/test_transcript_edits.py
Normal file
189
backend/tests/test_transcript_edits.py
Normal file
|
|
@ -0,0 +1,189 @@
|
|||
"""M8: user edits to transcript/summary via PUT endpoints.
|
||||
|
||||
Edits create a new version with edited_by_user=True; the AI pipeline
|
||||
must not overwrite them afterwards.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import struct
|
||||
import uuid
|
||||
|
||||
from sqlalchemy import select
|
||||
|
||||
from shonar.db import session as db_session
|
||||
from shonar.db.models import Transcript
|
||||
from shonar.services import processing
|
||||
|
||||
AUTH = {"email": "m8@example.com", "password": "m8-test-passw0rd-123"}
|
||||
|
||||
|
||||
async def user_tokens(client):
|
||||
r = await client.post(
|
||||
"/api/v1/auth/register",
|
||||
json={"email": AUTH["email"], "password": AUTH["password"]},
|
||||
)
|
||||
assert r.status_code == 201, r.text
|
||||
return r.json()["access_token"]
|
||||
|
||||
|
||||
def wav_bytes(payload_len: int = 64) -> bytes:
|
||||
data = bytes(range(payload_len % 256)) * (payload_len // 256 + 1)
|
||||
data = data[:payload_len]
|
||||
header = (
|
||||
b"RIFF" + struct.pack("<I", 36 + len(data)) + b"WAVE"
|
||||
+ b"fmt " + struct.pack("<IHHIIHH", 16, 1, 1, 8000, 8000, 1, 8)
|
||||
+ b"data" + struct.pack("<I", len(data))
|
||||
)
|
||||
return header + data
|
||||
|
||||
|
||||
async def upload_recording(client, token, client_id=None):
|
||||
h = {"Authorization": f"Bearer {token}"}
|
||||
data = wav_bytes()
|
||||
r = await client.post(
|
||||
"/api/v1/uploads",
|
||||
json={"declared_mime_type": "audio/wav", "declared_size_bytes": len(data),
|
||||
"client_recording_id": client_id},
|
||||
headers=h,
|
||||
)
|
||||
assert r.status_code == 201, r.text
|
||||
sid = r.json()["id"]
|
||||
r = await client.put(f"/api/v1/uploads/{sid}/chunks/0", content=data,
|
||||
headers={**h, "content-type": "application/octet-stream"})
|
||||
assert r.status_code == 201, r.text
|
||||
r = await client.post(f"/api/v1/uploads/{sid}/finalize",
|
||||
json={"duration_seconds": 5.0}, headers=h)
|
||||
assert r.status_code == 201, r.text
|
||||
return r.json()
|
||||
|
||||
|
||||
async def test_put_transcript_creates_user_version(client):
|
||||
token = await user_tokens(client)
|
||||
rec = await upload_recording(client, token, client_id="m8-t-1")
|
||||
h = {"Authorization": f"Bearer {token}"}
|
||||
rid = rec["id"]
|
||||
|
||||
r = await client.put(f"/api/v1/recordings/{rid}/transcript",
|
||||
json={"text": "hello corrected"}, headers=h)
|
||||
assert r.status_code == 200, r.text
|
||||
body = r.json()
|
||||
assert body["version"] == 1
|
||||
assert body["text"] == "hello corrected"
|
||||
assert body["edited_by_user"] is True
|
||||
assert body["provider"] == "user"
|
||||
|
||||
r = await client.get(f"/api/v1/recordings/{rid}/transcript", headers=h)
|
||||
assert r.status_code == 200
|
||||
assert r.json()["text"] == "hello corrected"
|
||||
|
||||
# Second edit supersedes the first.
|
||||
r = await client.put(f"/api/v1/recordings/{rid}/transcript",
|
||||
json={"text": "second pass",
|
||||
"segments": [{"start": 0.0, "end": 1.0, "text": "second pass"}]},
|
||||
headers=h)
|
||||
assert r.status_code == 200, r.text
|
||||
assert r.json()["version"] == 2
|
||||
assert r.json()["segments"][0]["text"] == "second pass"
|
||||
|
||||
async with db_session._session_factory() as s:
|
||||
rows = (await s.scalars(
|
||||
select(Transcript).where(Transcript.recording_id == uuid.UUID(rid))
|
||||
.order_by(Transcript.version))).all()
|
||||
assert [(r.version, r.superseded_at is not None, r.edited_by_user) for r in rows] == [
|
||||
(1, True, True), (2, False, True)]
|
||||
|
||||
|
||||
async def test_put_summary_creates_user_version(client):
|
||||
token = await user_tokens(client)
|
||||
rec = await upload_recording(client, token, client_id="m8-s-1")
|
||||
h = {"Authorization": f"Bearer {token}"}
|
||||
rid = rec["id"]
|
||||
|
||||
content = {"short": "s", "action_items": ["ship it"]}
|
||||
r = await client.put(f"/api/v1/recordings/{rid}/summary",
|
||||
json={"content": content}, headers=h)
|
||||
assert r.status_code == 200, r.text
|
||||
body = r.json()
|
||||
assert body["version"] == 1
|
||||
assert body["content"] == content
|
||||
assert body["edited_by_user"] is True
|
||||
|
||||
r = await client.get(f"/api/v1/recordings/{rid}/summary", headers=h)
|
||||
assert r.status_code == 200
|
||||
assert r.json()["content"]["action_items"] == ["ship it"]
|
||||
|
||||
|
||||
async def test_user_edit_survives_auto_pipeline(client, monkeypatch):
|
||||
"""An auto transcribe run after a user edit inserts nothing."""
|
||||
from shonar.services.ai import Segment as AiSegment
|
||||
from shonar.services.ai import SummaryResult, TranscriptResult
|
||||
|
||||
class FakeTranscriber:
|
||||
name = "fake-stt"
|
||||
|
||||
async def transcribe(self, audio, mime, *, language_hint=None):
|
||||
return TranscriptResult(
|
||||
text="auto text", language="en",
|
||||
segments=[AiSegment(0.0, 1.0, "auto text")], model="fake-stt-1",
|
||||
)
|
||||
|
||||
class FakeLlm:
|
||||
name = "fake-llm"
|
||||
|
||||
async def summarize(self, transcript, *, title=None):
|
||||
return SummaryResult(short="s", detailed="d", key_points=(),
|
||||
decisions=(), action_items=(), questions=(),
|
||||
model="fake-llm-1")
|
||||
|
||||
monkeypatch.setattr(
|
||||
processing, "get_transcription_provider", lambda settings: FakeTranscriber())
|
||||
monkeypatch.setattr(processing, "get_llm_provider", lambda settings: FakeLlm())
|
||||
|
||||
token = await user_tokens(client)
|
||||
rec = await upload_recording(client, token, client_id="m8-t-2")
|
||||
h = {"Authorization": f"Bearer {token}"}
|
||||
rid = rec["id"]
|
||||
|
||||
r = await client.put(f"/api/v1/recordings/{rid}/transcript",
|
||||
json={"text": "user verdict"}, headers=h)
|
||||
assert r.status_code == 200, r.text
|
||||
|
||||
monkeypatch.setattr(
|
||||
processing, "get_transcription_provider", lambda settings: FakeTranscriber())
|
||||
monkeypatch.setattr(processing, "get_llm_provider", lambda settings: FakeLlm())
|
||||
await processing.run_transcribe({}, rid)
|
||||
|
||||
r = await client.get(f"/api/v1/recordings/{rid}/transcript", headers=h)
|
||||
assert r.status_code == 200
|
||||
assert r.json()["text"] == "user verdict"
|
||||
|
||||
|
||||
async def test_edit_endpoints_enforce_ownership(client):
|
||||
token = await user_tokens(client)
|
||||
rec = await upload_recording(client, token, client_id="m8-t-3")
|
||||
rid = rec["id"]
|
||||
|
||||
r = await client.post("/api/v1/auth/register",
|
||||
json={"email": "m8-other@example.com", "password": "m8-other-passw0rd-1"})
|
||||
assert r.status_code == 201, r.text
|
||||
other = r.json()["access_token"]
|
||||
h2 = {"Authorization": f"Bearer {other}"}
|
||||
|
||||
r = await client.put(f"/api/v1/recordings/{rid}/transcript",
|
||||
json={"text": "hijack"}, headers=h2)
|
||||
assert r.status_code == 404
|
||||
r = await client.put(f"/api/v1/recordings/{rid}/summary",
|
||||
json={"content": {}}, headers=h2)
|
||||
assert r.status_code == 404
|
||||
|
||||
|
||||
async def test_edit_validation(client):
|
||||
token = await user_tokens(client)
|
||||
rec = await upload_recording(client, token, client_id="m8-t-4")
|
||||
h = {"Authorization": f"Bearer {token}"}
|
||||
rid = rec["id"]
|
||||
|
||||
r = await client.put(f"/api/v1/recordings/{rid}/transcript",
|
||||
json={"text": ""}, headers=h)
|
||||
assert r.status_code == 422
|
||||
Loading…
Add table
Add a link
Reference in a new issue