169 lines
6.1 KiB
Python
169 lines
6.1 KiB
Python
"""Live counting routes: start/stop a session and watch it as MJPEG."""
|
|
|
|
import os
|
|
import time
|
|
|
|
from fastapi import APIRouter, HTTPException
|
|
from fastapi.responses import StreamingResponse
|
|
from pydantic import BaseModel
|
|
from typing import Optional
|
|
|
|
from backend import library, live_count, training
|
|
from backend.api.common import project_or_404
|
|
|
|
router = APIRouter(tags=["live-count"])
|
|
|
|
|
|
class StartRequest(BaseModel):
|
|
# Either a live stream — which must be a WebRTC (WHEP) URL, REQ-176 — or an
|
|
# archive-relative path like "2026-08-13/batch001.mp4", which the backend
|
|
# resolves so the frontend never needs to know where the archive is mounted.
|
|
source: str = ""
|
|
source_rel: Optional[str] = None
|
|
model_path: Optional[str] = None
|
|
model_version_id: Optional[int] = None
|
|
count_classes: list = ["sack"]
|
|
line_y: int = 266
|
|
line_x_start: int = 469
|
|
line_x_end: int = 910
|
|
conf: float = 0.35
|
|
dedup_radius: float = 60.0
|
|
margin: int = 5
|
|
imgsz: int = 640
|
|
# Ghost rejection and spatial dedup pull in opposite directions, so they are
|
|
# separate dials now (REQ-140).
|
|
entry_travel_min: float = 60.0
|
|
handoff_radius: float = 100.0
|
|
unload_confirm_frames: int = 3
|
|
min_area_scale: float = 1.0
|
|
spatial_dedup: bool = False
|
|
|
|
|
|
@router.get("/api/projects/{project_id}/live-count/models")
|
|
def available_models(project_id: int) -> dict:
|
|
"""Weights this project can count with: its trained versions, then its base."""
|
|
project = project_or_404(project_id)
|
|
out = []
|
|
for version in training.listing(project_id):
|
|
if version.get("weights_path") and os.path.isfile(version["weights_path"]):
|
|
out.append({
|
|
"label": version.get("name") or f"v{version['version']}",
|
|
"path": version["weights_path"],
|
|
"version_id": version["id"],
|
|
})
|
|
base = project.get("base_model_path")
|
|
if base and os.path.isfile(base):
|
|
out.append({"label": "base model", "path": base, "version_id": None})
|
|
return {"models": out}
|
|
|
|
|
|
@router.post("/api/projects/{project_id}/live-count/start")
|
|
def start(project_id: int, request: StartRequest) -> dict:
|
|
project = project_or_404(project_id)
|
|
|
|
source, whep = request.source, ""
|
|
if request.source_rel:
|
|
try:
|
|
source = library.resolve(project["video_root"], request.source_rel)
|
|
except library.LibraryError as exc:
|
|
raise HTTPException(400, str(exc))
|
|
elif source:
|
|
# A live source is a WebRTC URL and nothing else. The browser watches it
|
|
# over WebRTC; the counter decodes the RTSP leg of the same MediaMTX
|
|
# path, derived here (REQ-176).
|
|
try:
|
|
whep, source = live_count.whep_url(source), live_count.whep_to_rtsp(source)
|
|
except live_count.LiveCountError as exc:
|
|
raise HTTPException(400, str(exc))
|
|
if not source:
|
|
raise HTTPException(400, "Pick a video or enter a WebRTC stream URL")
|
|
|
|
path = request.model_path
|
|
if not path and request.model_version_id is not None:
|
|
for item in available_models(project_id)["models"]:
|
|
if item["version_id"] == request.model_version_id:
|
|
path = item["path"]
|
|
break
|
|
if not path:
|
|
raise HTTPException(400, "Pick a model to count with")
|
|
try:
|
|
return live_count.start(
|
|
source=source, model_path=path, line_y=request.line_y,
|
|
line_x_start=request.line_x_start, line_x_end=request.line_x_end,
|
|
conf=request.conf, dedup_radius=request.dedup_radius,
|
|
margin=request.margin, imgsz=request.imgsz,
|
|
entry_travel_min=request.entry_travel_min,
|
|
handoff_radius=request.handoff_radius,
|
|
unload_confirm_frames=request.unload_confirm_frames,
|
|
min_area_scale=request.min_area_scale,
|
|
spatial_dedup=request.spatial_dedup,
|
|
whep=whep,
|
|
count_classes=tuple(request.count_classes),
|
|
)
|
|
except live_count.LiveCountError as exc:
|
|
raise HTTPException(400, str(exc))
|
|
|
|
|
|
class LineRequest(BaseModel):
|
|
line_y: Optional[int] = None
|
|
line_x_start: Optional[int] = None
|
|
line_x_end: Optional[int] = None
|
|
|
|
|
|
@router.patch("/api/live-count/line")
|
|
def move_line(request: LineRequest) -> dict:
|
|
"""Reposition the counting line mid-session, without losing the counts."""
|
|
try:
|
|
return live_count.move_line(request.line_y, request.line_x_start, request.line_x_end)
|
|
except live_count.LiveCountError as exc:
|
|
raise HTTPException(400, str(exc))
|
|
|
|
|
|
@router.post("/api/live-count/stop")
|
|
def stop() -> dict:
|
|
return live_count.stop()
|
|
|
|
|
|
@router.get("/api/live-count/status")
|
|
def status() -> dict:
|
|
return live_count.status()
|
|
|
|
|
|
@router.get("/api/live-count/overlay")
|
|
def overlay() -> dict:
|
|
"""Geometry only — boxes, line and counts for the frame just processed.
|
|
|
|
What the WebRTC preview draws over the video, instead of the server
|
|
encoding a JPEG per frame for it (REQ-177).
|
|
"""
|
|
return live_count.overlay()
|
|
|
|
|
|
@router.get("/api/live-count/stream")
|
|
def stream():
|
|
"""MJPEG of the annotated frames — the archive-file preview. Ends with the session."""
|
|
if live_count.status().get("preview") == "webrtc":
|
|
# Nothing is encoding JPEGs for this session; without this the generator
|
|
# would sit on a worker thread for ten seconds producing nothing.
|
|
raise HTTPException(409, "This session is watched over WebRTC, not MJPEG")
|
|
|
|
def frames():
|
|
blank_streak = 0
|
|
while True:
|
|
jpeg = live_count.snapshot()
|
|
if jpeg is None:
|
|
blank_streak += 1
|
|
if blank_streak > 100 or not live_count.status().get("running"):
|
|
return
|
|
time.sleep(0.1)
|
|
continue
|
|
blank_streak = 0
|
|
yield (b"--frame\r\nContent-Type: image/jpeg\r\n"
|
|
b"Content-Length: " + str(len(jpeg)).encode() + b"\r\n\r\n"
|
|
+ jpeg + b"\r\n")
|
|
time.sleep(0.05)
|
|
|
|
return StreamingResponse(frames(),
|
|
media_type="multipart/x-mixed-replace; boundary=frame",
|
|
headers={"Cache-Control": "no-store"})
|