254 lines
8.4 KiB
Python
254 lines
8.4 KiB
Python
"""Batch lifecycle: one trimmed range of one video, turned into frames.
|
||
|
||
A batch is the unit of work everything downstream hangs off — auto-annotation,
|
||
review, and the merge into the master dataset all address a batch. The same
|
||
video can produce many batches with different ranges (REQ-023).
|
||
|
||
extracting -> extracted -> labeling -> reviewing -> approved -> merged
|
||
\\-> failed
|
||
"""
|
||
|
||
import os
|
||
import time
|
||
from typing import List, Optional
|
||
|
||
from PIL import Image
|
||
|
||
from backend import config, db, jobs, library, projects, video
|
||
|
||
|
||
class BatchError(Exception):
|
||
pass
|
||
|
||
|
||
def batch_dir(project_slug: str, batch_id: int) -> str:
|
||
return os.path.join(config.project_dir(project_slug), "batches", str(batch_id))
|
||
|
||
|
||
def frames_dir(project_slug: str, batch_id: int) -> str:
|
||
return os.path.join(batch_dir(project_slug, batch_id), "frames")
|
||
|
||
|
||
def create(project_id: int, rel: str, start_sec: float, end_sec: float, fps: float) -> dict:
|
||
"""Register a batch and queue its extraction job (REQ-020…022)."""
|
||
project = projects.get(project_id)
|
||
if project is None:
|
||
raise BatchError("No such project")
|
||
|
||
try:
|
||
video_path = library.resolve(project["video_root"], rel)
|
||
except library.LibraryError as exc:
|
||
raise BatchError(str(exc))
|
||
|
||
try:
|
||
info = video.probe(video_path)
|
||
except video.VideoError as exc:
|
||
raise BatchError(str(exc))
|
||
end_sec = min(end_sec, info["duration"]) if info["duration"] else end_sec
|
||
if end_sec <= start_sec:
|
||
raise BatchError("The end of the range must be after its start")
|
||
if fps <= 0:
|
||
raise BatchError("fps must be greater than 0")
|
||
|
||
date_label, filename = rel.split("/", 1)
|
||
with db.cursor() as cur:
|
||
cur.execute(
|
||
"""INSERT INTO batches (project_id, video_path, date_label, batch_label,
|
||
start_sec, end_sec, fps, status, created_at)
|
||
VALUES (?, ?, ?, ?, ?, ?, ?, 'extracting', ?)""",
|
||
(project_id, video_path, date_label, library.batch_label(filename),
|
||
float(start_sec), float(end_sec), float(fps), time.time()),
|
||
)
|
||
batch_id = cur.lastrowid
|
||
|
||
jobs.create(
|
||
"extract",
|
||
params={"batch_id": batch_id},
|
||
project_id=project_id,
|
||
batch_id=batch_id,
|
||
message=f"{date_label}/{library.batch_label(filename)}",
|
||
)
|
||
return get(batch_id)
|
||
|
||
|
||
def get(batch_id: int) -> Optional[dict]:
|
||
with db.cursor() as cur:
|
||
cur.execute(
|
||
"""SELECT b.*, p.slug AS project_slug, p.name AS project_name
|
||
FROM batches b JOIN projects p ON p.id = b.project_id WHERE b.id = ?""",
|
||
(batch_id,),
|
||
)
|
||
row = cur.fetchone()
|
||
if row is None:
|
||
return None
|
||
return _row_to_dict(cur, row)
|
||
|
||
|
||
def listing(project_id: int) -> List[dict]:
|
||
with db.cursor() as cur:
|
||
cur.execute(
|
||
"""SELECT b.*, p.slug AS project_slug, p.name AS project_name
|
||
FROM batches b JOIN projects p ON p.id = b.project_id
|
||
WHERE b.project_id = ? ORDER BY b.created_at DESC""",
|
||
(project_id,),
|
||
)
|
||
return [_row_to_dict(cur, row) for row in cur.fetchall()]
|
||
|
||
|
||
def _row_to_dict(cur, row) -> dict:
|
||
cur.execute(
|
||
"""SELECT review_status, COUNT(*) FROM frames WHERE batch_id = ?
|
||
GROUP BY review_status""",
|
||
(row["id"],),
|
||
)
|
||
review = {"pending": 0, "approved": 0, "rejected": 0}
|
||
for status, count in cur.fetchall():
|
||
review[status] = count
|
||
|
||
cur.execute(
|
||
"""SELECT COUNT(*) FROM annotations a JOIN frames f ON f.id = a.frame_id
|
||
WHERE f.batch_id = ?""",
|
||
(row["id"],),
|
||
)
|
||
annotation_count = cur.fetchone()[0]
|
||
|
||
return {
|
||
"id": row["id"],
|
||
"project_id": row["project_id"],
|
||
"project_slug": row["project_slug"],
|
||
"project_name": row["project_name"],
|
||
"video_path": row["video_path"],
|
||
"date_label": row["date_label"],
|
||
"batch_label": row["batch_label"],
|
||
"start_sec": row["start_sec"],
|
||
"end_sec": row["end_sec"],
|
||
"fps": row["fps"],
|
||
"status": row["status"],
|
||
"frame_count": row["frame_count"],
|
||
"created_at": row["created_at"],
|
||
"merged_at": row["merged_at"],
|
||
"review": review,
|
||
"reviewed": review["approved"] + review["rejected"],
|
||
"annotation_count": annotation_count,
|
||
}
|
||
|
||
|
||
def frames(batch_id: int) -> List[dict]:
|
||
with db.cursor() as cur:
|
||
cur.execute(
|
||
"""SELECT f.*, (SELECT COUNT(*) FROM annotations a WHERE a.frame_id = f.id)
|
||
AS annotation_count
|
||
FROM frames f WHERE f.batch_id = ? ORDER BY f.idx""",
|
||
(batch_id,),
|
||
)
|
||
return [dict(row) for row in cur.fetchall()]
|
||
|
||
|
||
def frame_path(frame_id: int) -> Optional[str]:
|
||
with db.cursor() as cur:
|
||
cur.execute(
|
||
"""SELECT f.filename, b.id AS batch_id, p.slug
|
||
FROM frames f
|
||
JOIN batches b ON b.id = f.batch_id
|
||
JOIN projects p ON p.id = b.project_id
|
||
WHERE f.id = ?""",
|
||
(frame_id,),
|
||
)
|
||
row = cur.fetchone()
|
||
if row is None:
|
||
return None
|
||
return os.path.join(frames_dir(row["slug"], row["batch_id"]), row["filename"])
|
||
|
||
|
||
def set_status(batch_id: int, status: str) -> None:
|
||
with db.cursor() as cur:
|
||
cur.execute("UPDATE batches SET status = ? WHERE id = ?", (status, batch_id))
|
||
|
||
|
||
@jobs.handler("extract")
|
||
def _run_extract(job) -> None:
|
||
batch = get(job.params["batch_id"])
|
||
if batch is None:
|
||
raise BatchError("The batch disappeared before extraction started")
|
||
|
||
out_dir = frames_dir(batch["project_slug"], batch["id"])
|
||
expected = video.frame_count(batch["start_sec"], batch["end_sec"], batch["fps"])
|
||
job.log(f"Extracting {expected} frame(s) at {batch['fps']} fps from "
|
||
f"{batch['date_label']}/{batch['batch_label']} "
|
||
f"[{batch['start_sec']:.1f}s – {batch['end_sec']:.1f}s]")
|
||
job.progress(0, expected)
|
||
|
||
try:
|
||
names = video.extract_frames(
|
||
batch["video_path"], out_dir,
|
||
batch["start_sec"], batch["end_sec"], batch["fps"],
|
||
on_progress=lambda written: job.progress(written, expected),
|
||
should_stop=lambda: job.cancelled,
|
||
)
|
||
except video.VideoError as exc:
|
||
set_status(batch["id"], "failed")
|
||
raise BatchError(str(exc))
|
||
|
||
if not names:
|
||
set_status(batch["id"], "failed")
|
||
raise BatchError("ffmpeg produced no frames for that range")
|
||
|
||
with Image.open(os.path.join(out_dir, names[0])) as first:
|
||
width, height = first.size
|
||
|
||
with db.cursor() as cur:
|
||
cur.executemany(
|
||
"INSERT OR IGNORE INTO frames (batch_id, idx, filename, width, height) "
|
||
"VALUES (?, ?, ?, ?, ?)",
|
||
[(batch["id"], index, name, width, height) for index, name in enumerate(names)],
|
||
)
|
||
cur.execute("UPDATE batches SET frame_count = ?, status = 'extracted' WHERE id = ?",
|
||
(len(names), batch["id"]))
|
||
|
||
job.progress(len(names), len(names))
|
||
job.log(f"Extracted {len(names)} frame(s) at {width}×{height}")
|
||
|
||
|
||
def update(batch_id: int, patch: dict) -> dict:
|
||
batch = get(batch_id)
|
||
if batch is None:
|
||
raise BatchError("No such batch")
|
||
|
||
fields = []
|
||
args = []
|
||
if "batch_label" in patch and patch["batch_label"] is not None:
|
||
fields.append("batch_label = ?")
|
||
args.append(str(patch["batch_label"]).strip())
|
||
if "date_label" in patch and patch["date_label"] is not None:
|
||
fields.append("date_label = ?")
|
||
args.append(str(patch["date_label"]).strip())
|
||
if "status" in patch and patch["status"] is not None:
|
||
fields.append("status = ?")
|
||
args.append(str(patch["status"]).strip())
|
||
|
||
if fields:
|
||
args.append(batch_id)
|
||
with db.cursor() as cur:
|
||
cur.execute(f"UPDATE batches SET {', '.join(fields)} WHERE id = ?", args)
|
||
|
||
return get(batch_id)
|
||
|
||
|
||
def delete(batch_id: int) -> bool:
|
||
import shutil
|
||
batch = get(batch_id)
|
||
if batch is None:
|
||
return False
|
||
with db.cursor() as cur:
|
||
cur.execute("DELETE FROM batches WHERE id = ?", (batch_id,))
|
||
shutil.rmtree(batch_dir(batch["project_slug"], batch_id), ignore_errors=True)
|
||
return True
|
||
|
||
|
||
def approve_all_frames(batch_id: int) -> int:
|
||
with db.cursor() as cur:
|
||
cur.execute("UPDATE frames SET review_status = 'approved' WHERE batch_id = ?", (batch_id,))
|
||
return cur.rowcount
|
||
|
||
|