mritunjayk-ops's picture
Add LLM emotion classifier, result summary, and honest docs
f94ee13
Raw
History Blame Contribute Delete
6.77 kB
"""Label-isolated manifest parsing and per-file batch execution."""
from __future__ import annotations
import csv
from dataclasses import dataclass
from pathlib import Path
from typing import Callable, Generic, Literal, TypeVar
from .audio import AudioError, SUPPORTED_AUDIO_EXTENSIONS
DEFAULT_MAX_MANIFEST_BYTES = 1024 * 1024
T = TypeVar("T")
class ManifestError(ValueError):
"""A manifest-level validation failure safe to show to an authenticated user."""
def __init__(self, code: str, message: str) -> None:
super().__init__(message)
self.code = code
self.public_message = message
class BatchFileError(ValueError):
"""A path or file failure associated with one manifest row."""
def __init__(self, code: str, message: str) -> None:
super().__init__(message)
self.code = code
self.public_message = message
@dataclass(frozen=True)
class ManifestEntry:
"""The complete manifest contract visible to inference code."""
name: str
@dataclass(frozen=True)
class FileFailure:
code: str
message: str
@dataclass(frozen=True)
class BatchItemResult(Generic[T]):
name: str
status: Literal["ok", "error"]
value: T | None = None
failure: FileFailure | None = None
@dataclass(frozen=True)
class BatchReport(Generic[T]):
items: tuple[BatchItemResult[T], ...]
@property
def success_count(self) -> int:
return sum(item.status == "ok" for item in self.items)
@property
def failure_count(self) -> int:
return sum(item.status == "error" for item in self.items)
def _validate_manifest_name(name: str) -> str:
if not name or name != name.strip():
raise ManifestError("invalid_name", "Manifest filenames must be non-empty and trimmed.")
if "/" in name or "\\" in name or name in {".", ".."}:
raise ManifestError("unsafe_name", "Manifest filenames must be root-level basenames.")
if any(ord(character) < 32 for character in name):
raise ManifestError("invalid_name", "Manifest filenames contain invalid characters.")
if Path(name).suffix.lower() not in SUPPORTED_AUDIO_EXTENSIONS:
raise ManifestError("unsupported_extension", "Manifest contains an unsupported audio extension.")
return name
def parse_manifest_for_inference(
path: str | Path,
*,
max_bytes: int = DEFAULT_MAX_MANIFEST_BYTES,
) -> tuple[ManifestEntry, ...]:
"""Read only filenames; result_json is deliberately neither parsed nor retained."""
manifest = Path(path)
if manifest.is_symlink() or not manifest.is_file():
raise ManifestError("unsafe_manifest", "Manifest must be a regular file.")
if manifest.stat().st_size <= 0 or manifest.stat().st_size > max_bytes:
raise ManifestError("manifest_size", "Manifest is empty or exceeds the configured limit.")
try:
with manifest.open("r", encoding="utf-8-sig", newline="") as stream:
reader = csv.DictReader(stream)
if (
reader.fieldnames is None
or len(reader.fieldnames) != 2
or set(reader.fieldnames) != {"name", "result_json"}
):
raise ManifestError(
"invalid_columns",
"Manifest must contain exactly name and result_json columns.",
)
entries: list[ManifestEntry] = []
seen: set[str] = set()
for row in reader:
if None in row:
raise ManifestError("invalid_row", "Manifest contains a malformed CSV row.")
name = _validate_manifest_name(row.get("name", ""))
key = name.casefold()
if key in seen:
raise ManifestError("duplicate_name", "Manifest contains duplicate filenames.")
seen.add(key)
entries.append(ManifestEntry(name=name))
except UnicodeDecodeError as exc:
raise ManifestError("invalid_encoding", "Manifest must be UTF-8 encoded.") from exc
except csv.Error as exc:
raise ManifestError("invalid_csv", "Manifest is not valid CSV.") from exc
if not entries:
raise ManifestError("empty_manifest", "Manifest must contain at least one audio row.")
return tuple(entries)
def _resolve_batch_file(batch_root: Path, entry: ManifestEntry) -> Path:
candidate = batch_root / entry.name
if candidate.is_symlink():
raise BatchFileError("unsafe_file", "The listed audio file is unsafe.")
try:
resolved = candidate.resolve(strict=True)
except OSError as exc:
raise BatchFileError("missing_file", "The listed audio file is missing.") from exc
if resolved.parent != batch_root or not resolved.is_file():
raise BatchFileError("unsafe_file", "The listed audio file is unsafe.")
return resolved
def process_batch(
entries: tuple[ManifestEntry, ...],
batch_root: str | Path,
file_processor: Callable[[Path], T],
progress_callback: Callable[[int, int, str], None] | None = None,
) -> BatchReport[T]:
"""Process entries independently and never pass manifest labels to the processor."""
root = Path(batch_root)
if root.is_symlink() or not root.is_dir():
raise ManifestError("unsafe_batch_root", "Batch root must be a regular directory.")
root = root.resolve(strict=True)
results: list[BatchItemResult[T]] = []
total = len(entries)
for index, entry in enumerate(entries, start=1):
if progress_callback is not None:
progress_callback(index - 1, total, entry.name)
try:
path = _resolve_batch_file(root, entry)
value = file_processor(path)
results.append(BatchItemResult(name=entry.name, status="ok", value=value))
except AudioError as exc:
results.append(
BatchItemResult(
name=entry.name,
status="error",
failure=FileFailure(exc.code, exc.public_message),
)
)
except BatchFileError as exc:
results.append(
BatchItemResult(
name=entry.name,
status="error",
failure=FileFailure(exc.code, exc.public_message),
)
)
except Exception:
results.append(
BatchItemResult(
name=entry.name,
status="error",
failure=FileFailure("processing_failed", "File processing failed."),
)
)
if progress_callback is not None:
progress_callback(index, total, entry.name)
return BatchReport(items=tuple(results))