208 lines
7.0 KiB
Python
208 lines
7.0 KiB
Python
"""Project, library and video-streaming routes (REQ-001…013)."""
|
|
|
|
import os
|
|
import shutil
|
|
import tempfile
|
|
from typing import Dict, List, Optional
|
|
|
|
from fastapi import APIRouter, File, HTTPException, Request, UploadFile
|
|
from fastapi.responses import FileResponse, StreamingResponse
|
|
from pydantic import BaseModel, Field
|
|
|
|
from backend import config, library, video
|
|
from backend import projects as project_store
|
|
from backend.api.common import project_or_404
|
|
|
|
router = APIRouter(prefix="/api/projects", tags=["projects"])
|
|
|
|
VIDEO_MEDIA = {".mp4": "video/mp4", ".m4v": "video/mp4", ".mkv": "video/x-matroska",
|
|
".mov": "video/quicktime", ".webm": "video/webm", ".avi": "video/x-msvideo"}
|
|
STREAM_CHUNK = 1024 * 512
|
|
|
|
|
|
class ClassSpec(BaseModel):
|
|
name: str
|
|
prompt: Optional[str] = None
|
|
|
|
|
|
class ProjectRequest(BaseModel):
|
|
name: str
|
|
label_type: str = "bbox"
|
|
video_root: str
|
|
classes: List[ClassSpec] = Field(default_factory=list)
|
|
val_every: int = 5
|
|
|
|
|
|
class ProjectPatch(BaseModel):
|
|
prompts: Optional[Dict[int, str]] = None
|
|
val_every: Optional[int] = None
|
|
video_root: Optional[str] = None
|
|
|
|
|
|
@router.get("")
|
|
def list_projects() -> dict:
|
|
return {"projects": project_store.listing(), "video_root_default": config.VIDEO_ROOT}
|
|
|
|
|
|
@router.post("")
|
|
def create_project(request: ProjectRequest) -> dict:
|
|
try:
|
|
return project_store.create(
|
|
name=request.name,
|
|
label_type=request.label_type,
|
|
video_root=request.video_root,
|
|
classes=[item.model_dump() for item in request.classes],
|
|
val_every=request.val_every,
|
|
)
|
|
except project_store.ProjectError as exc:
|
|
raise HTTPException(400, str(exc))
|
|
|
|
|
|
@router.get("/{project_id}")
|
|
def read_project(project_id: int) -> dict:
|
|
return project_or_404(project_id)
|
|
|
|
|
|
@router.patch("/{project_id}")
|
|
def patch_project(project_id: int, request: ProjectPatch) -> dict:
|
|
try:
|
|
return project_store.update(project_id, prompts=request.prompts,
|
|
val_every=request.val_every,
|
|
video_root=request.video_root)
|
|
except project_store.ProjectError as exc:
|
|
raise HTTPException(400, str(exc))
|
|
|
|
|
|
@router.post("/{project_id}/base-model")
|
|
async def upload_base_model(project_id: int, file: UploadFile = File(...)) -> dict:
|
|
if not (file.filename or "").endswith(".pt"):
|
|
raise HTTPException(400, "Base model must be a .pt checkpoint")
|
|
# Staged to a temp file first: reading the classes can fail, and a rejected
|
|
# upload must not leave a broken model.pt in the project.
|
|
with tempfile.NamedTemporaryFile(suffix=".pt", delete=False) as staged:
|
|
shutil.copyfileobj(file.file, staged)
|
|
staged_path = staged.name
|
|
await file.close()
|
|
try:
|
|
return project_store.set_base_model(project_id, staged_path)
|
|
except project_store.ProjectError as exc:
|
|
raise HTTPException(400, str(exc))
|
|
finally:
|
|
os.unlink(staged_path)
|
|
|
|
|
|
@router.post("/{project_id}/secondary-model")
|
|
async def upload_secondary_model(project_id: int, file: UploadFile = File(...)) -> dict:
|
|
project_or_404(project_id)
|
|
if not file.filename.endswith(".pt"):
|
|
raise HTTPException(400, "The secondary model must be a .pt file")
|
|
|
|
with tempfile.NamedTemporaryFile(suffix=".pt", delete=False) as staged:
|
|
shutil.copyfileobj(file.file, staged)
|
|
staged_path = staged.name
|
|
await file.close()
|
|
try:
|
|
return project_store.set_secondary_model(project_id, staged_path, name=file.filename)
|
|
except project_store.ProjectError as exc:
|
|
raise HTTPException(400, str(exc))
|
|
finally:
|
|
os.unlink(staged_path)
|
|
|
|
|
|
|
|
@router.post("/{project_id}/classes")
|
|
def add_class(project_id: int, request: ClassSpec) -> dict:
|
|
try:
|
|
return project_store.add_class(project_id, request.name, request.prompt)
|
|
except project_store.ProjectError as exc:
|
|
raise HTTPException(400, str(exc))
|
|
|
|
|
|
@router.delete("/{project_id}/classes/{class_id}")
|
|
def delete_class(project_id: int, class_id: int) -> dict:
|
|
try:
|
|
return project_store.delete_class(project_id, class_id)
|
|
except project_store.ProjectError as exc:
|
|
raise HTTPException(400, str(exc))
|
|
|
|
|
|
@router.delete("/{project_id}")
|
|
def delete_project(project_id: int) -> dict:
|
|
return {"deleted": project_store.delete(project_id)}
|
|
|
|
|
|
@router.get("/{project_id}/library")
|
|
def list_library(project_id: int) -> dict:
|
|
project = project_or_404(project_id)
|
|
try:
|
|
return {"video_root": project["video_root"],
|
|
"dates": library.list_dates(project["video_root"])}
|
|
except library.LibraryError as exc:
|
|
raise HTTPException(400, str(exc))
|
|
|
|
|
|
@router.get("/{project_id}/library/{date}")
|
|
def list_library_date(project_id: int, date: str) -> dict:
|
|
project = project_or_404(project_id)
|
|
try:
|
|
return {"date": date,
|
|
"videos": library.list_videos(project["video_root"], date, project_id)}
|
|
except library.LibraryError as exc:
|
|
raise HTTPException(400, str(exc))
|
|
|
|
|
|
@router.get("/{project_id}/video/info")
|
|
def video_info(project_id: int, rel: str) -> dict:
|
|
project = project_or_404(project_id)
|
|
try:
|
|
path = library.resolve(project["video_root"], rel)
|
|
info = dict(video.probe(path))
|
|
except (library.LibraryError, video.VideoError) as exc:
|
|
raise HTTPException(400, str(exc))
|
|
info.update({"rel": rel, "batch_label": library.batch_label(os.path.basename(path)),
|
|
"date_label": rel.split("/", 1)[0]})
|
|
return info
|
|
|
|
|
|
@router.get("/{project_id}/video")
|
|
def stream_video(project_id: int, rel: str, request: Request):
|
|
"""Serve a video with Range support so the player can seek (REQ-013)."""
|
|
project = project_or_404(project_id)
|
|
try:
|
|
path = library.resolve(project["video_root"], rel)
|
|
except library.LibraryError as exc:
|
|
raise HTTPException(404, str(exc))
|
|
|
|
media = VIDEO_MEDIA.get(os.path.splitext(path)[1].lower(), "application/octet-stream")
|
|
size = os.path.getsize(path)
|
|
header = request.headers.get("range")
|
|
if not header or not header.startswith("bytes="):
|
|
return FileResponse(path, media_type=media, headers={"Accept-Ranges": "bytes"})
|
|
|
|
start_text, _, end_text = header[6:].partition("-")
|
|
start = int(start_text) if start_text else 0
|
|
end = int(end_text) if end_text else size - 1
|
|
start, end = max(0, start), min(end, size - 1)
|
|
if start > end:
|
|
raise HTTPException(416, "Requested range is not satisfiable")
|
|
|
|
def chunks():
|
|
with open(path, "rb") as handle:
|
|
handle.seek(start)
|
|
remaining = end - start + 1
|
|
while remaining > 0:
|
|
block = handle.read(min(STREAM_CHUNK, remaining))
|
|
if not block:
|
|
break
|
|
remaining -= len(block)
|
|
yield block
|
|
|
|
return StreamingResponse(
|
|
chunks(), status_code=206, media_type=media,
|
|
headers={
|
|
"Content-Range": f"bytes {start}-{end}/{size}",
|
|
"Accept-Ranges": "bytes",
|
|
"Content-Length": str(end - start + 1),
|
|
},
|
|
)
|