From 2b54cc33538d5272b551f9a1d40101ea22c83937 Mon Sep 17 00:00:00 2001 From: jetson Date: Wed, 16 Sep 2026 15:14:44 +0700 Subject: [PATCH] =?UTF-8?q?fix:=20final=20review=20wave=20=E2=80=94=20CLI?= =?UTF-8?q?=20kwargs,=20class=20filter,=20job=20locks,=20upload=20order,?= =?UTF-8?q?=20tests?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- README.md | 2 +- app.py | 18 +++++------ cli.py | 14 ++++----- src/job.py | 71 ++++++++++++++++++++++++------------------ src/model_registry.py | 3 +- src/pipeline.py | 31 +++++++++++++----- tests/test_app.py | 12 +++++++ tests/test_job.py | 17 ++++++++++ tests/test_pipeline.py | 16 +++++++++- 9 files changed, 127 insertions(+), 57 deletions(-) diff --git a/README.md b/README.md index 9ec8fc7..a198bf5 100644 --- a/README.md +++ b/README.md @@ -38,7 +38,7 @@ recounter --video PATH Input video file --filter NAME Class filter (repeatable): sack, box, truck --sack-conf FLOAT Sack confidence threshold (default: 0.4) --truck-conf FLOAT Truck confidence threshold (default: 0.5) - --output PATH Output path (single model only) + --output PATH Output path (single model only) --output-dir DIR Output directory (default: ./output) --models-dir DIR Models directory (default: ./models) ``` diff --git a/app.py b/app.py index 2f14246..b8626eb 100644 --- a/app.py +++ b/app.py @@ -53,14 +53,6 @@ def upload(): if not safe_name or not safe_name.lower().endswith((".mp4", ".avi", ".mkv", ".mov", ".webm")): return "Invalid video file type", 400 - video_path = os.path.join(UPLOAD_DIR, safe_name) - base, ext = os.path.splitext(video_path) - n = 1 - while os.path.exists(video_path): - video_path = f"{base}_{n}{ext}" - n += 1 - video.save(video_path) - selected_models = request.form.getlist("models") models = scan_models(MODELS_DIR) by_name = {m.filename: m for m in models} @@ -81,6 +73,14 @@ def upload(): if not model_configs: return "No models selected", 400 + video_path = os.path.join(UPLOAD_DIR, safe_name) + base, ext = os.path.splitext(video_path) + n = 1 + while os.path.exists(video_path): + video_path = f"{base}_{n}{ext}" + n += 1 + video.save(video_path) + job = job_queue.add_job( video_path=video_path, model_configs=model_configs, @@ -138,7 +138,7 @@ def api_jobs(): "job_id": j.job_id, "status": j.status.name, "progress": j.progress, - "video_path": j.video_path, + "video_path": os.path.basename(j.video_path), "results": [ { "model": r.model_name, diff --git a/cli.py b/cli.py index 7d75d0a..20464bf 100644 --- a/cli.py +++ b/cli.py @@ -109,13 +109,13 @@ def main(argv: list[str] | None = None) -> None: print(f"\r [{_model}] frame {frame_idx}", end="", flush=True) result = run_pipeline( - args.video, - model_cfg, - out_path, - class_filter, - args.sack_conf, - args.truck_conf, - progress_callback, + video_path=args.video, + model_config=model_cfg, + output_path=out_path, + class_filter=class_filter, + sack_conf=args.sack_conf, + truck_conf=args.truck_conf, + progress_callback=progress_callback, ) print() print(f"Output: {result.output_path}") diff --git a/src/job.py b/src/job.py index 78c339c..731bb6d 100644 --- a/src/job.py +++ b/src/job.py @@ -92,6 +92,10 @@ class JobQueue: return job + # NOTE: get_job/list_jobs return the live Job object (Flask renders it + # directly), so readers must treat its fields as eventually consistent — + # the worker mutates them under self._lock while readers may observe + # a slightly stale snapshot. def get_job(self, job_id: str) -> Job | None: with self._lock: return self._jobs.get(job_id) @@ -120,33 +124,35 @@ class JobQueue: def _run_job(self, job_id: str) -> None: """Worker: process each model config sequentially.""" - job: Job | None = self._jobs.get(job_id) + with self._lock: + job: Job | None = self._jobs.get(job_id) if job is None: return - job.status = JobStatus.RUNNING - total_models = len(job.model_configs) + with self._lock: + job.status = JobStatus.RUNNING + total_models = len(job.model_configs) if total_models == 0: - job.status = JobStatus.COMPLETED - job.completed_at = time.time() + with self._lock: + job.status = JobStatus.COMPLETED + job.completed_at = time.time() return try: for i, model_cfg in enumerate(job.model_configs): - if job.status == JobStatus.CANCELLED: - break - - job.current_model = model_cfg.filename - job.progress = i / total_models + with self._lock: + if job.status == JobStatus.CANCELLED: + break + job.current_model = model_cfg.filename + job.progress = i / total_models + class_filter = job.class_filters.get(model_cfg.filename) output_path = os.path.join( job.output_dir, f"{model_cfg.stem}_annotated.mp4", ) - class_filter = job.class_filters.get(model_cfg.filename) - result: PipelineResult = run_pipeline( video_path=job.video_path, model_config=model_cfg, @@ -154,27 +160,32 @@ class JobQueue: class_filter=class_filter, ) - job.results.append( - JobResult( - model_name=model_cfg.filename, - output_path=result.output_path, - loading_count=result.loading_count, - unloading_count=result.unloading_count, - net_count=result.net_count, - batch_count=result.batch_count, - frame_count=result.frame_count, - duration_seconds=result.duration_seconds, + with self._lock: + job.results.append( + JobResult( + model_name=model_cfg.filename, + output_path=result.output_path, + loading_count=result.loading_count, + unloading_count=result.unloading_count, + net_count=result.net_count, + batch_count=result.batch_count, + frame_count=result.frame_count, + duration_seconds=result.duration_seconds, + ) ) - ) - if job.status != JobStatus.CANCELLED: - job.status = JobStatus.COMPLETED - job.progress = 1.0 + with self._lock: + if job.status != JobStatus.CANCELLED: + job.status = JobStatus.COMPLETED + job.progress = 1.0 except Exception as e: - job.status = JobStatus.FAILED - job.error = str(e) + with self._lock: + if job.status != JobStatus.CANCELLED: + job.status = JobStatus.FAILED + job.error = str(e) finally: - job.completed_at = time.time() - job.current_model = "" + with self._lock: + job.completed_at = time.time() + job.current_model = "" diff --git a/src/model_registry.py b/src/model_registry.py index 1b8cb30..0f6f60b 100644 --- a/src/model_registry.py +++ b/src/model_registry.py @@ -2,7 +2,6 @@ from __future__ import annotations -import os from dataclasses import dataclass, field from pathlib import Path @@ -40,7 +39,7 @@ def scan_models(models_dir: str) -> list[ModelConfig]: configs: list[ModelConfig] = [] for f in sorted(p.iterdir()): - if f.is_file() and f.suffix in MODEL_EXTENSIONS: + if f.is_file() and f.suffix.lower() in MODEL_EXTENSIONS: stem = f.stem known = KNOWN_MODEL_CLASSES.get(stem, []) configs.append( diff --git a/src/pipeline.py b/src/pipeline.py index 2caebf5..b6c476e 100644 --- a/src/pipeline.py +++ b/src/pipeline.py @@ -6,7 +6,8 @@ import time from dataclasses import dataclass import cv2 -import numpy as np + +from ultralytics import YOLO from src.batch import BatchLifecycleManager from src.counting import LineCrossCounter @@ -38,6 +39,18 @@ class PipelineResult: return self.loading_count - self.unloading_count +def apply_class_filter( + detections: list[Detection], class_filter: list[str] | None +) -> list[Detection]: + """Keep only detections whose class_name is in class_filter. + + None or an empty list means "keep all" (empty is treated as None). + """ + if not class_filter: + return detections + return [d for d in detections if d.class_name in class_filter] + + def run_pipeline( video_path: str, model_config: ModelConfig, @@ -72,12 +85,13 @@ def run_pipeline( w = int(cap.get(cv2.CAP_PROP_FRAME_WIDTH)) h = int(cap.get(cv2.CAP_PROP_FRAME_HEIGHT)) - # Build detector with class filtering + # Build detector with class filtering (single shared YOLO instance) effective_filter = class_filter or ( model_config.known_classes if model_config.known_classes else None ) + shared_model = YOLO(model_config.path) detector = BaseDetector( - model_config.path, conf=sack_conf, class_filter=effective_filter + shared_model, conf=sack_conf, class_filter=effective_filter ) # Truck detector: if model has "truck" class, use same model @@ -85,10 +99,10 @@ def run_pipeline( truck_detector = None if truck_has_truck: truck_detector = BaseDetector( - model_config.path, conf=truck_conf, class_filter=("truck",) + shared_model, conf=truck_conf, class_filter=("truck",) ) - tracker = ByteTrackTracker(model_config.path, conf=sack_conf) + tracker = ByteTrackTracker(shared_model, conf=sack_conf) stabilizer = BboxStabilizer() roi_tracker = TruckROITracker(frame_width=w, frame_height=h) counter = LineCrossCounter( @@ -100,7 +114,7 @@ def run_pipeline( batch_mgr = BatchLifecycleManager() dashboard = DashboardOverlay() - writer = AnnotatedVideoWriter(output_path, fps=fps, frame_size=(w, h)) + writer: AnnotatedVideoWriter | None = None start_time = time.time() frame_idx = 0 @@ -113,6 +127,7 @@ def run_pipeline( batch_mgr.on_batch_end(on_batch_end) try: + writer = AnnotatedVideoWriter(output_path, fps=fps, frame_size=(w, h)) while True: ret, frame = cap.read() if not ret: @@ -148,6 +163,7 @@ def run_pipeline( if batch_mgr.is_active: raw_tracked = tracker.update(frame, []) stable = stabilizer.update(raw_tracked) + stable = apply_class_filter(stable, effective_filter) if roi is not None: tracked_sacks = [ @@ -192,7 +208,8 @@ def run_pipeline( finally: cap.release() - writer.finish() + if writer is not None: + writer.finish() duration = time.time() - start_time return PipelineResult( diff --git a/tests/test_app.py b/tests/test_app.py index b1b7fe5..b24a05d 100644 --- a/tests/test_app.py +++ b/tests/test_app.py @@ -1,5 +1,7 @@ """Integration tests for Flask web app.""" +import io + import pytest from app import app @@ -45,6 +47,16 @@ def test_upload_no_video(client): assert resp.status_code == 400 +def test_upload_invalid_extension_rejected(client): + """POST /upload with a non-video file returns 400.""" + resp = client.post( + "/upload", + data={"video": (io.BytesIO(b"not a video"), "notes.txt")}, + content_type="multipart/form-data", + ) + assert resp.status_code == 400 + + def test_status_nonexistent(client): """GET /status/nonexistent returns 404.""" resp = client.get("/status/nonexistent") diff --git a/tests/test_job.py b/tests/test_job.py index 86f98b3..2d2f77a 100644 --- a/tests/test_job.py +++ b/tests/test_job.py @@ -60,6 +60,23 @@ def test_queue_cancel_pending(): assert q.cancel_job(job.job_id) in (True, False) # may have already started +def test_job_empty_config_completes(): + """add_job with [] model_configs completes with COMPLETED + empty results.""" + import time + q = JobQueue(output_dir="/tmp/output") + job = q.add_job(video_path="/tmp/test.mp4", model_configs=[]) + deadline = time.time() + 5.0 + while time.time() < deadline: + fetched = q.get_job(job.job_id) + if fetched is not None and fetched.status == JobStatus.COMPLETED: + break + time.sleep(0.05) + fetched = q.get_job(job.job_id) + assert fetched is not None + assert fetched.status == JobStatus.COMPLETED + assert fetched.results == [] + + def test_queue_status_counts(): """status_counts returns correct tally.""" q = JobQueue(output_dir="/tmp/output") diff --git a/tests/test_pipeline.py b/tests/test_pipeline.py index aaa3f7f..747aedd 100644 --- a/tests/test_pipeline.py +++ b/tests/test_pipeline.py @@ -5,7 +5,8 @@ import os import cv2 import numpy as np import pytest -from src.pipeline import run_pipeline, PipelineResult +from src.pipeline import apply_class_filter, run_pipeline, PipelineResult +from src.interfaces import Detection from src.model_registry import ModelConfig @@ -47,6 +48,19 @@ def test_run_pipeline_processes_video(tmp_path): pass # See integration test below for end-to-end with real models +def test_apply_class_filter_keeps_only_selected(): + """apply_class_filter keeps only classes in the filter; None/[] keep all.""" + dets = [ + Detection(bbox=(0, 0, 1, 1), confidence=0.9, class_id=0, class_name="sack"), + Detection(bbox=(0, 0, 1, 1), confidence=0.9, class_id=1, class_name="box"), + Detection(bbox=(0, 0, 1, 1), confidence=0.9, class_id=0, class_name="sack"), + ] + filtered = apply_class_filter(dets, ["sack"]) + assert [d.class_name for d in filtered] == ["sack", "sack"] + assert apply_class_filter(dets, None) == dets + assert apply_class_filter(dets, []) == dets + + def test_run_pipeline_no_model_raises(tmp_path): """run_pipeline raises RuntimeError if video can't be opened.""" with pytest.raises(RuntimeError, match="Cannot open video"):