mirror of
https://github.com/jamiepine/voicebox.git
synced 2026-10-04 01:25:18 -07:00
delete_generations_by_profile (and the version sweep under it) committed on their own, so a failure between the generation sweep and the profile row left the profile listed with its history already gone. Both helpers take a commit flag now, and delete_profile passes commit=False so the single db.commit() at its end publishes the entire cascade. Other callers keep the default and are unchanged.
423 lines
13 KiB
Python
423 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
|
|
_delete_generation_children(generation_id, db)
|
|
|
|
# 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():
|
|
audio_path.unlink()
|
|
|
|
# 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.
|
|
_delete_generation_children(generation.id, db)
|
|
|
|
# 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
|
|
},
|
|
}
|