mirror of
https://github.com/jamiepine/voicebox.git
synced 2026-10-04 01:25:18 -07:00
Review follow-up: a rollback after a failed main-audio unlink restored version rows whose files were already removed. Treat a locked main file like the failed-generation sweep does: log and continue, leaving at worst an orphaned audio file rather than rows that point at nothing.
430 lines
13 KiB
Python
430 lines
13 KiB
Python
"""
|
|
Generation history management module.
|
|
"""
|
|
|
|
from typing import List, Optional, Tuple
|
|
from datetime import UTC, datetime
|
|
import logging
|
|
import uuid
|
|
import shutil
|
|
from pathlib import Path
|
|
from sqlalchemy.orm import Session
|
|
from sqlalchemy import or_
|
|
|
|
from ..models import GenerationRequest, GenerationResponse, HistoryQuery, HistoryResponse, HistoryListResponse, GenerationVersionResponse, EffectConfig
|
|
from ..database import Generation as DBGeneration, GenerationVersion as DBGenerationVersion, StoryItem as DBStoryItem, VoiceProfile as DBVoiceProfile
|
|
from .. import config
|
|
|
|
logger = logging.getLogger(__name__)
|
|
|
|
|
|
def _delete_generation_children(generation_id: str, db: Session, commit: bool = True) -> None:
|
|
"""Remove the rows that reference a generation, plus any version audio files.
|
|
|
|
Story items and versions both point at the generation by a non-null FK.
|
|
The story detail query inner-joins generations, so a leftover story item
|
|
vanishes from the timeline while staying in the table forever.
|
|
"""
|
|
from . import versions as versions_mod
|
|
|
|
db.query(DBStoryItem).filter_by(generation_id=generation_id).delete()
|
|
versions_mod.delete_versions_for_generation(generation_id, db, commit=commit)
|
|
|
|
|
|
def _get_versions_for_generations(generation_ids: list[str], db: Session) -> dict:
|
|
"""Fetch versions for many generations in a single query.
|
|
|
|
Returns a mapping of ``generation_id -> (versions, active_version_id)``
|
|
using the same shape as ``_get_versions_for_generation()``, so callers
|
|
can batch a whole page of generations without an N+1 query.
|
|
"""
|
|
import json
|
|
|
|
ids = list(dict.fromkeys(generation_ids))
|
|
if not ids:
|
|
return {}
|
|
|
|
versions_rows = (
|
|
db.query(DBGenerationVersion)
|
|
.filter(DBGenerationVersion.generation_id.in_(ids))
|
|
.order_by(DBGenerationVersion.created_at)
|
|
.all()
|
|
)
|
|
|
|
versions_by_generation: dict[str, list] = {}
|
|
active_by_generation: dict[str, Optional[str]] = {}
|
|
for v in versions_rows:
|
|
versions = versions_by_generation.setdefault(v.generation_id, [])
|
|
effects_chain = None
|
|
if v.effects_chain:
|
|
try:
|
|
raw = json.loads(v.effects_chain)
|
|
effects_chain = [EffectConfig(**e) for e in raw]
|
|
except Exception:
|
|
pass
|
|
versions.append(GenerationVersionResponse(
|
|
id=v.id,
|
|
generation_id=v.generation_id,
|
|
label=v.label,
|
|
audio_path=v.audio_path,
|
|
effects_chain=effects_chain,
|
|
is_default=v.is_default,
|
|
created_at=v.created_at,
|
|
))
|
|
if v.is_default:
|
|
active_by_generation[v.generation_id] = v.id
|
|
|
|
return {
|
|
generation_id: (
|
|
versions_by_generation.get(generation_id),
|
|
active_by_generation.get(generation_id),
|
|
)
|
|
for generation_id in ids
|
|
}
|
|
|
|
|
|
def _get_versions_for_generation(generation_id: str, db: Session) -> tuple:
|
|
"""Get versions list and active version ID for a single generation."""
|
|
return _get_versions_for_generations([generation_id], db)[generation_id]
|
|
|
|
|
|
async def create_generation(
|
|
profile_id: str,
|
|
text: str,
|
|
language: str,
|
|
audio_path: str,
|
|
duration: float,
|
|
seed: Optional[int],
|
|
db: Session,
|
|
instruct: Optional[str] = None,
|
|
generation_id: Optional[str] = None,
|
|
status: str = "completed",
|
|
engine: Optional[str] = "qwen",
|
|
model_size: Optional[str] = None,
|
|
source: str = "manual",
|
|
) -> GenerationResponse:
|
|
"""
|
|
Create a new generation history entry.
|
|
|
|
Args:
|
|
profile_id: Profile ID used for generation
|
|
text: Generated text
|
|
language: Language code
|
|
audio_path: Path where audio was saved
|
|
duration: Audio duration in seconds
|
|
seed: Random seed used (if any)
|
|
db: Database session
|
|
instruct: Natural language instruction used (if any)
|
|
generation_id: Pre-assigned ID (for async generation flow)
|
|
status: Generation status (generating, completed, failed)
|
|
engine: TTS engine used (qwen, luxtts, chatterbox, chatterbox_turbo)
|
|
model_size: Model size variant (1.7B, 0.6B) — only relevant for qwen
|
|
source: Origin marker stored on the row. ``"manual"`` for regular
|
|
/generate calls; ``"personality_speak"`` for rows created
|
|
by the /profiles/{id}/speak endpoint. Enables filtering the
|
|
history view for personality-driven output.
|
|
|
|
Returns:
|
|
Created generation entry
|
|
"""
|
|
db_generation = DBGeneration(
|
|
id=generation_id or str(uuid.uuid4()),
|
|
profile_id=profile_id,
|
|
text=text,
|
|
language=language,
|
|
audio_path=audio_path,
|
|
duration=duration,
|
|
seed=seed,
|
|
instruct=instruct,
|
|
engine=engine,
|
|
model_size=model_size,
|
|
status=status,
|
|
source=source,
|
|
created_at=datetime.now(UTC),
|
|
)
|
|
|
|
db.add(db_generation)
|
|
db.commit()
|
|
db.refresh(db_generation)
|
|
|
|
return GenerationResponse.model_validate(db_generation)
|
|
|
|
|
|
async def update_generation_status(
|
|
generation_id: str,
|
|
status: str,
|
|
db: Session,
|
|
audio_path: Optional[str] = None,
|
|
duration: Optional[float] = None,
|
|
error: Optional[str] = None,
|
|
) -> Optional[GenerationResponse]:
|
|
"""Update the status of a generation (used by async generation flow)."""
|
|
generation = db.query(DBGeneration).filter_by(id=generation_id).first()
|
|
if not generation:
|
|
return None
|
|
|
|
generation.status = status
|
|
if audio_path is not None:
|
|
generation.audio_path = audio_path
|
|
if duration is not None:
|
|
generation.duration = duration
|
|
if error is not None:
|
|
generation.error = error
|
|
|
|
db.commit()
|
|
db.refresh(generation)
|
|
return GenerationResponse.model_validate(generation)
|
|
|
|
|
|
async def get_generation(
|
|
generation_id: str,
|
|
db: Session,
|
|
) -> Optional[GenerationResponse]:
|
|
"""
|
|
Get a generation by ID.
|
|
|
|
Args:
|
|
generation_id: Generation ID
|
|
db: Database session
|
|
|
|
Returns:
|
|
Generation or None if not found
|
|
"""
|
|
generation = db.query(DBGeneration).filter_by(id=generation_id).first()
|
|
if not generation:
|
|
return None
|
|
|
|
return GenerationResponse.model_validate(generation)
|
|
|
|
|
|
async def list_generations(
|
|
query: HistoryQuery,
|
|
db: Session,
|
|
) -> HistoryListResponse:
|
|
"""
|
|
List generations with optional filters.
|
|
|
|
Args:
|
|
query: Query parameters (filters, pagination)
|
|
db: Database session
|
|
|
|
Returns:
|
|
HistoryListResponse with items and total count
|
|
"""
|
|
# Build base query with join to get profile name
|
|
q = db.query(
|
|
DBGeneration,
|
|
DBVoiceProfile.name.label('profile_name')
|
|
).join(
|
|
DBVoiceProfile,
|
|
DBGeneration.profile_id == DBVoiceProfile.id
|
|
)
|
|
|
|
# Apply profile filter
|
|
if query.profile_id:
|
|
q = q.filter(DBGeneration.profile_id == query.profile_id)
|
|
|
|
# Apply search filter (searches in text content)
|
|
if query.search:
|
|
search_pattern = f"%{query.search}%"
|
|
q = q.filter(DBGeneration.text.like(search_pattern))
|
|
|
|
# Get total count before pagination
|
|
total_count = q.count()
|
|
|
|
# Apply ordering (newest first)
|
|
q = q.order_by(DBGeneration.created_at.desc())
|
|
|
|
# Apply pagination
|
|
q = q.offset(query.offset).limit(query.limit)
|
|
|
|
# Execute query
|
|
results = q.all()
|
|
|
|
# Fetch versions for every generation on this page with a single
|
|
# query instead of one SELECT per generation (N+1).
|
|
versions_by_generation = _get_versions_for_generations(
|
|
[generation.id for generation, _ in results], db
|
|
)
|
|
|
|
# Convert to HistoryResponse with profile_name
|
|
items = []
|
|
for generation, profile_name in results:
|
|
versions, active_version_id = versions_by_generation[generation.id]
|
|
items.append(HistoryResponse(
|
|
id=generation.id,
|
|
profile_id=generation.profile_id,
|
|
profile_name=profile_name,
|
|
text=generation.text,
|
|
language=generation.language,
|
|
audio_path=generation.audio_path,
|
|
duration=generation.duration,
|
|
seed=generation.seed,
|
|
instruct=generation.instruct,
|
|
engine=generation.engine or "qwen",
|
|
model_size=generation.model_size,
|
|
status=generation.status or "completed",
|
|
error=generation.error,
|
|
is_favorited=bool(generation.is_favorited),
|
|
created_at=generation.created_at,
|
|
versions=versions,
|
|
active_version_id=active_version_id,
|
|
))
|
|
|
|
return HistoryListResponse(
|
|
items=items,
|
|
total=total_count,
|
|
)
|
|
|
|
|
|
async def delete_generation(
|
|
generation_id: str,
|
|
db: Session,
|
|
) -> bool:
|
|
"""
|
|
Delete a generation.
|
|
|
|
Args:
|
|
generation_id: Generation ID
|
|
db: Database session
|
|
|
|
Returns:
|
|
True if deleted, False if not found
|
|
"""
|
|
generation = db.query(DBGeneration).filter_by(id=generation_id).first()
|
|
if not generation:
|
|
return False
|
|
|
|
# Delete all version files and records; commit together with the
|
|
# generation row below so an unlink failure rolls everything back.
|
|
_delete_generation_children(generation_id, db, commit=False)
|
|
|
|
# Delete main audio file (if not already removed by version cleanup)
|
|
if generation.audio_path:
|
|
audio_path = config.resolve_storage_path(generation.audio_path)
|
|
if audio_path is not None and audio_path.exists():
|
|
try:
|
|
audio_path.unlink()
|
|
except OSError:
|
|
# Version files are already gone by now, so rolling back would
|
|
# leave rows pointing at missing audio. Mirror the sweep below:
|
|
# keep going and leave the locked file as an orphan instead.
|
|
logger.warning("Could not delete generation audio %s", audio_path)
|
|
|
|
# Delete from database
|
|
db.delete(generation)
|
|
db.commit()
|
|
|
|
return True
|
|
|
|
|
|
async def delete_failed_generations(db: Session) -> int:
|
|
"""
|
|
Delete every generation whose status is 'failed'.
|
|
|
|
Used by the "Clear failed" action in the UI so users can tidy up
|
|
history after the model wasn't loaded, the app was closed mid-run,
|
|
or a generation otherwise errored out (see issue #410).
|
|
|
|
Returns:
|
|
Number of generations deleted.
|
|
"""
|
|
failed = db.query(DBGeneration).filter(DBGeneration.status == "failed").all()
|
|
count = 0
|
|
for generation in failed:
|
|
# Clean up version files/rows first; one commit at the end.
|
|
_delete_generation_children(generation.id, db, commit=False)
|
|
|
|
# Remove the main audio file if it somehow made it to disk.
|
|
if generation.audio_path:
|
|
audio_path = config.resolve_storage_path(generation.audio_path)
|
|
if audio_path is not None and audio_path.exists():
|
|
try:
|
|
audio_path.unlink()
|
|
except OSError:
|
|
# Best-effort cleanup — don't abort the whole sweep
|
|
# if a single file can't be removed.
|
|
logger.warning("Could not delete generation audio %s", audio_path)
|
|
|
|
db.delete(generation)
|
|
count += 1
|
|
|
|
db.commit()
|
|
return count
|
|
|
|
|
|
async def delete_generations_by_profile(
|
|
profile_id: str,
|
|
db: Session,
|
|
commit: bool = True,
|
|
) -> int:
|
|
"""
|
|
Delete all generations for a profile.
|
|
|
|
Args:
|
|
profile_id: Profile ID
|
|
db: Database session
|
|
commit: Commit at the end. Pass False when the caller owns the
|
|
transaction (e.g. deleting the profile itself) so the whole
|
|
cascade lands in one commit.
|
|
|
|
Returns:
|
|
Number of generations deleted
|
|
"""
|
|
generations = db.query(DBGeneration).filter_by(profile_id=profile_id).all()
|
|
|
|
count = 0
|
|
for generation in generations:
|
|
# Delete associated version files and rows first
|
|
_delete_generation_children(generation.id, db, commit=commit)
|
|
|
|
# Delete audio file
|
|
audio_path = config.resolve_storage_path(generation.audio_path)
|
|
if audio_path is not None and audio_path.exists():
|
|
try:
|
|
audio_path.unlink()
|
|
except OSError:
|
|
# A file locked by playback shouldn't abort the whole sweep
|
|
# and leave the profile half-deleted.
|
|
logger.warning("Could not delete generation audio %s", audio_path)
|
|
|
|
# Delete from database
|
|
db.delete(generation)
|
|
count += 1
|
|
|
|
if commit:
|
|
db.commit()
|
|
|
|
return count
|
|
|
|
|
|
async def get_generation_stats(db: Session) -> dict:
|
|
"""
|
|
Get generation statistics.
|
|
|
|
Args:
|
|
db: Database session
|
|
|
|
Returns:
|
|
Statistics dictionary
|
|
"""
|
|
from sqlalchemy import func
|
|
|
|
total = db.query(func.count(DBGeneration.id)).scalar()
|
|
|
|
total_duration = db.query(func.sum(DBGeneration.duration)).scalar() or 0
|
|
|
|
# Get generations by profile
|
|
by_profile = db.query(
|
|
DBGeneration.profile_id,
|
|
func.count(DBGeneration.id).label('count')
|
|
).group_by(DBGeneration.profile_id).all()
|
|
|
|
return {
|
|
"total_generations": total,
|
|
"total_duration_seconds": total_duration,
|
|
"generations_by_profile": {
|
|
profile_id: count for profile_id, count in by_profile
|
|
},
|
|
}
|