Files
reTraining/backend/api/live_count.py
T
asus dee58e4ae5 fix: resolve model weight paths against DATA_DIR (REQ-187)
- config.py: resolve_data_path (legacy abs + rel) + rel_data_path
- all file-opening reads wrapped: preview, autolabel, training, model
  download, live count, projects.get; training_start_point hack replaced
- new writes store paths relative to data/
- legacy stale rows (/home/asus/reTraining/...) resolve without migration
- requirements: REQ-187 added; REQ-188 (per-class max box) + REQ-186
  copy-line amendment drafted for the next task
2026-10-02 16:59:25 +07:00

170 lines
6.2 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 config, 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):
path = config.resolve_data_path(version["weights_path"])
if version.get("weights_path") and os.path.isfile(path):
out.append({
"label": version.get("name") or f"v{version['version']}",
"path": 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"})