Compare commits

..
3 Commits
Author SHA1 Message Date
Ryan Shpeherd 458d8862cc Cleanup 2026-09-09 15:01:42 -04:00
Ryan Shpeherd 9f098d4b1f Update worker 2026-09-09 13:50:19 -04:00
Ryan Shpeherd d9c0998400 API json fixes 2026-09-09 12:35:18 -04:00
24 changed files with 1183 additions and 3196 deletions
+486
View File
@@ -0,0 +1,486 @@
# REFACTOR.md — Worker-to-API Refactor Plan
## Context
VideoDetect workers currently run as long-lived processes (`main.py``WorkerPool`) that query the database directly via `DBConnector`. The re-architecture moves to **independent, externally-invoked task processes** that communicate through a REST API provided by the Perl `api/app.pl` module. This enables finer-grained scaling, resilience to individual process crashes, and separation of concerns between scanning (discovery) and processing (AI inference).
---
## Target Architecture
```
┌───────────┐ /api/v1/nexttask/AISCAN ┌─────────────┐
│ Scanner │ ──────────────────────────────> │ Perl API │
│ (long- │ <────────────────────────────── │ (Dancer2) │
│ lived) │ task reservation + assign_key ├─────────────┤
└───────────┘ │ DB │
│ │
┌───────────────────────────────────────────┼──────────┐
│ ▼ │
│ ┌──────────────┐ │
│ POST /api/v1/task/:id/complete │ │
│ ◄─────────────────────────────────────────► │ │
│ │ │
│ ┌───────────────────────────────┘ │
▼ ▼ ▼
┌────────┐ ┌────────┐ ┌────────┐ ┌──────────────┐
│worker-1│ │worker-2│ │worker-N│ │ MariaDB │
│(short- │ │(short- │ │(short- │ │ │
│ lived) │ │ lived) │ │ lived) │ └──────────────┘
└────────┘ └────────┘ └────────┘
```
Each worker process:
1. Starts → accepts `--tasks N` (default 1) via CLI argument
2. Loops up to N times: **claim task****fetch video data****process****submit results**
3. Exits cleanly after completing the requested count
---
## API Reference (from `api/app.pl`)
### GET `/api/v1/nexttask/:type`
Claims the next pending task of the given type and reserves it.
| Field | Type | Example |
|-----------|---------|----------------------------------------------|
| task | object | Task record from DB |
| assign_key| string | `"worker_338"` (server-generated) |
**On success:** HTTP 200 with JSON body:
```json
{
"task": {
"id": 1,
"video_id": 42,
"task_type": "AISCAN",
"status": "PENDING",
"created_at": "2026-09-09T13:05:56",
"updated_at": "2026-09-09T14:53:57",
"assign_key": null,
"results": null,
"assigned_at": "2026-09-09T14:43:31"
},
"assign_key": "worker_338"
}
```
**On no tasks available:** HTTP 404 with `{"message":"No task"}`.
> **NOTE:** The Perl API marks the task `IN_PROGRESS` and sets `assign_key` + `assigned_at` atomically (`UPDATE ... WHERE id=? AND status='PENDING'`). This is safe for multiple workers racing to claim the same task.
---
### GET `/api/v1/video/:id`
Returns full video metadata by ID.
| Field | Type | Example |
|---------------|---------|------------------------------------------------------|
| id | integer | `1` |
| file_path | string | `"/data/The Fappening/Sextape - Alyson Hannigan (American actress - American pie).wmv"` |
| file_size | integer | `19358016` |
| file_hash | string | `"1ffd178e9ee23039aebffd79ddcbc88e983633edcb75062e6bd4c269f4d7bf94"` |
| resolution_w | integer | `640` |
| resolution_h | integer | `480` |
| codec | string | `"wmv1"` |
| duration | float | `102.499` |
| last_scan_time| datetime| `"2026-09-09T13:05:56"` |
| created_at | datetime| `"2026-09-09T13:05:56"` |
| updated_at | datetime| `"2026-09-09T13:05:56"` |
**On not found:** HTTP 404 with `{"message":"No video"}`.
---
### POST `/api/v1/task/:task/complete`
Submits processing results for a claimed task. The API verifies the task is `IN_PROGRESS` and assigned to the given `assign_key`.
**Required parameters (form-encoded or JSON body):**
| Field | Type | Description |
|-------------|--------|------------------------------------------------|
| assign_key | string | Must match the key returned by `/nexttask` |
| results | object | JSON-serializable dict to store in DB `results` column |
**On success:** HTTP 200 with `{"message":"Task completed successfully"}`.
**On mismatch or task not IN_PROGRESS:** HTTP 403 with error message.
> **NOTE:** The Perl API currently uses `body_parameters->get("results")`, which reads form-encoded data. For JSON body submission, the Dancer2 config may need `serializer: JSON` (which is already set in `api/config.yml`). However, `body_parameters` only parses form-encoded fields — JSON body content goes to `$app->request->body`. **This needs verification before implementation.**
---
## Refactor Steps
### Phase 1 — Create the Task Worker Module
**Goal:** A single entry-point script that can run as a standalone process.
**File:** `src/task_worker.py` (new)
**Acceptance Criteria:**
- Starts with `python3 src/task_worker.py --tasks N` (default 1 task)
- Prints startup log message showing number of tasks requested and model version
- Exits cleanly after processing the requested count
```bash
# Default: process 1 task and exit
python3 -m src.task_worker
# Explicit: process 5 tasks and exit
python3 -m src.task_worker --tasks 5
# Also supported via direct script invocation
python3 src/task_worker.py --tasks 10
```
**CLI Implementation:**
```python
import argparse
parser = argparse.ArgumentParser(description="VideoDetect task worker")
parser.add_argument("--tasks", type=int, default=1,
help="Number of tasks to process before exiting (default: 1)")
args = parser.parse_args()
```
---
### Phase 2 — API Client Module
**Goal:** A thin HTTP client layer for communicating with the Perl API.
**File:** `src/api_client.py` (new)
**Acceptance Criteria:**
- Encapsulates all API calls in one class: `ApiClient(base_url)`
- Handles JSON serialization/deserialization automatically
- Raises a custom `ApiError` on HTTP errors (with status code and response body)
- Returns `None` for 404 responses (no task remaining) from `/nexttask`
**API Error handling:**
```python
class ApiError(Exception):
def __init__(self, status_code: int, message: str):
self.status_code = status_code
self.message = message
super().__init__(f"API error {status_code}: {message}")
# Usage patterns:
try:
response = client.get_next_task("AISCAN")
except ApiError as e:
if e.status_code == 404:
logger.info("No more tasks available")
break # exit processing loop
raise
result = client.submit_results(task_id, assign_key, results_dict)
```
**Methods to implement:**
| Method | Endpoint | Returns | Special handling |
|--------|----------|---------|-----------------|
| `get_next_task(task_type)` | `GET /api/v1/nexttask/{type}` | `{"task": {...}, "assign_key": str} \| None` | Returns None on 404 |
| `get_video(video_id)` | `GET /api/v1/video/{id}` | dict with video metadata | Raises ApiError on 404/5xx |
| `submit_results(task_id, assign_key, results)` | `POST /api/v1/task/{task}/complete` | dict with response | Passes assign_key + serialized results |
**Notes for implementation:**
- Use Python's standard library only (`urllib.request`) if possible to avoid adding dependencies. If JSON body submission is needed for the Perl API, check whether `body_parameters->get('results')` in Dancer2 handles raw JSON (it typically does not — it expects form-encoded data).
- If form-encoding is required: `requests.post(url, data={'assign_key': key, 'results': json.dumps(results)})`
- If JSON body works: `requests.post(url, json={'assign_key': key, 'results': results})`
- **Recommendation:** Use the `requests` library (already in worker requirements.txt likely). Document both approaches and implement the one that matches your Dancer2 config.
---
### Phase 3 — Implement the AI Processing Pipeline
**Goal:** Extract the core AI processing logic from `orchestrator.py` into a reusable function callable by the task worker.
**File:** New or updated module in `src/` (tentatively `src/task_processor.py`)
**Acceptance Criteria:**
- Takes video metadata dict + file path as input
- Returns a results dict matching what will be stored in the DB `results` column
- Handles all error cases gracefully (returns errors instead of crashing)
- Logs all significant decisions at INFO level
**Input:**
```python
video = {
"id": 1,
"file_path": "/data/input/...",
"codec": "wmv1",
"duration": 102.499,
"resolution_w": 640,
"resolution_h": 480,
"file_size": 19358016,
"file_hash": "abc...",
}
```
**Output (results dict stored in DB):**
```python
{
"status": "COMPLETED", # or "FAILED" on error
"confidence": 0.87, # video-level confidence score
"routing_decision": "MATCH", # MATCH | REVIEW | SKIP
"face_count": 142, # total faces detected across all frames
"frame_count": 23, # frames extracted and processed
"model_version": "v0.0.0-placeholder",
"processing_time_seconds": 12.4,
"error": None, # error message if FAILED
}
```
**Processing flow (extracted from orchestrator.py):**
```python
def process_video(video: dict) -> dict:
"""Run the full AI scan pipeline on a single video.
Returns results dict for submission via API.
"""
start_time = time.time()
file_path = video["file_path"]
video_id = video["id"]
try:
# 1. Probe (if metadata is incomplete)
prober = VideoProber(timeout=10)
metadata = prober.probe(file_path)
# 2. Extract frames to scratch
scratch = ScratchManager(base_path="/scratch", video_id=str(video_id))
frame_dir = scratch.ensure_frame_dir()
sampler = FrameSampler(interval_seconds=30, quality=2)
frame_paths = sampler.extract_frames(
video_path=file_path, output_dir=str(frame_dir),
duration=metadata.duration, resolution=(metadata.resolution_w, metadata.resolution_h)
)
# 3. Detect faces
detector = FaceDetector(engine_path="/models/face_detector/face_detector.trt")
detections_per_frame = detector.detect_faces(frame_paths)
total_faces = sum(len(dets) for dets in detections_per_frame)
# 4. Crop and classify (aggregate confidences from face crops)
classifier = FaceClassifier(engine_path="/models/classifier/classifier.trt")
# ... crop paths → classify → get confidence list
# 5. Aggregate
confidence = aggregate(confidences, strategy="max")
# 6. Route
routing = router.route(confidence)
processing_time = time.time() - start_time
return {
"status": "COMPLETED",
"confidence": round(confidence, 4),
"routing_decision": routing,
"face_count": total_faces,
"frame_count": len(frame_paths),
"model_version": get_model_version(),
"processing_time_seconds": round(processing_time, 2),
"error": None,
}
except Exception as exc:
logger.error("Processing failed for video %d: %s", video_id, exc)
return {
"status": "FAILED",
"confidence": 0.0,
"routing_decision": "REVIEW",
"face_count": 0,
"frame_count": 0,
"model_version": get_model_version(),
"processing_time_seconds": round(time.time() - start_time, 2),
"error": str(exc),
}
```
**Notes for implementation:**
- Reuse existing classes from `orchestrator.py` (`FaceDetector`, `FaceClassifier`, `FrameSampler`, `ScratchManager`, `VideoProber`, `aggregate`, `router`) without refactoring them — just import and use them.
- GPU setup (torch CUDA, GPUMemoryManager) can be done once at module level or lazily inside `process_video()` to avoid overhead per task invocation.
- Model loading is expensive (~seconds). Consider lazy initialization or process-level caching via a singleton pattern.
---
### Phase 4 — Wire It Together in task_worker.py
**Goal:** Combine the API client, processing pipeline, and CLI entry point into a working script.
**Acceptance Criteria:**
- Full end-to-end: claim → fetch → process → submit → repeat N times
- Exits with code 0 on success, non-zero on unrecoverable errors
- Logs startup/shutdown counts (tasks attempted, tasks succeeded, tasks failed)
- Graceful shutdown on SIGTERM/SIGINT (finish current task, then exit)
**Expected lifecycle log output:**
```json
{"timestamp": "...", "level": "INFO", "message": "Task worker starting. Will process 3 task(s)."}
{"timestamp": "...", "level": "INFO", "message": "API client configured: base_url=http://localhost:8890"}
{"timestamp": "...", "level": "INFO", "message": "Claiming task 1 (video_id=42) for AISCAN"}
{"timestamp": "...", "level": "INFO", "message": "Processing video 42: /data/input/..."}
{"timestamp": "...", "level": "INFO", "message": "Video 42 complete: C=0.87 routing=MATCH faces=142 frames=23 time=12.4s"}
{"timestamp": "...", "level": "INFO", "message": "Submitting results for task 1 via API"}
{"timestamp": "...", "level": "INFO", "message": "Claiming task 2 (video_id=99) for AISCAN"}
{"timestamp": "...", "level": "INFO", "message": "Processing video 99: /data/input/..."}
{"timestamp": "...", "level": "INFO", "message": "Video 99 complete: C=0.31 routing=SKIP faces=0 frames=18 time=8.7s"}
{"timestamp": "...", "level": "INFO", "message": "Submitting results for task 2 via API"}
{"timestamp": "...", "level": "INFO", "message": "Claiming task 3 (video_id=157) for AISCAN"}
{"timestamp": "...", "level": "WARNING", "message": "Processing video 157 failed: ffprobe error: cannot decode stream"}
{"timestamp": "...", "level": "INFO", "message": "Submitting FAILED results for task 3 via API"}
{"timestamp": "...", "level": "INFO", "message": "No more tasks available (API returned 404)"}
{"timestamp": "...", "level": "INFO", "message": "Task worker finished: 3 attempted, 2 succeeded, 1 failed"}
```
**Core loop:**
```python
for i in range(1, args.tasks + 1):
logger.info("=== Processing task %d/%d ===", i, args.tasks)
# Claim a task
response = api.get_next_task("AISCAN")
if response is None:
logger.info("No more tasks available. Done.")
break
task = response["task"]
assign_key = response["assign_key"]
video_id = task["video_id"]
# Fetch video metadata
video = api.get_video(video_id)
# Process
results = process_video(video)
# Submit
api.submit_results(task["id"], assign_key, results)
logger.info("Finished: %d attempted, %d succeeded, %d failed", total_attempted, total_succeeded, total_failed)
```
---
### Phase 5 — Update docker-compose.yml
**Goal:** Change the worker service from a long-lived `WorkerPool` to a short-lived task process.
**Changes to `docker-compose.yml` worker section:**
1. **Update command** to invoke the task worker instead of `main.py`:
```yaml
command: python3 -m src.task_worker --tasks 50
```
2. **Add restart policy considerations** — since workers are now short-lived, you have two options:
- **Option A (simpler):** Run a single worker container with a large `--tasks` count (e.g., 100) so it processes many videos before exiting, then rely on Kubernetes/Cron/external scheduler to restart.
- **Option B (more flexible):** Make `--tasks` configurable via environment variable:
```yaml
environment:
- TASK_COUNT=${TASK_COUNT:-50}
command: >
python3 -m src.task_worker --tasks ${TASK_COUNT}
```
3. **Remove the WorkerPool initialization** from `main.py` — it becomes unused (or is removed entirely in a later cleanup).
---
### Phase 6 — Verify API Compatibility
**Goal:** Confirm the Dancer2 Perl API correctly handles JSON body submission from Python's `requests` library on `/api/v1/task/:task/complete`.
**Acceptance Criteria:**
- POST with `Content-Type: application/json` body containing `assign_key` + `results` (as a dict, not a string) returns 200
- POST with form-encoded data also works as fallback
- The Perl API correctly stores the results JSON in the database
**Test procedure:**
```bash
# Step 1: Create a test task manually (via MySQL or via scanner)
mysql videodetect -e "INSERT INTO tasks (task_type, video_id, status) VALUES ('AISCAN', 1, 'PENDING');"
# Step 2: Claim the task via API
curl -s http://localhost:8890/api/v1/nexttask/AISCAN | python3 -m json.tool
# Step 3: Submit results (adjust video_id to a real one)
curl -X POST http://localhost:8890/api/v1/task/1/complete \
-H "Content-Type: application/json" \
-d '{"assign_key":"worker_999","results":{"confidence":0.87,"routing_decision":"MATCH"}}' | python3 -m json.tool
# Step 4: Verify results stored in DB
mysql videodetect -e "SELECT id, status, results FROM tasks WHERE id=1;"
```
**If JSON body does NOT work**, update the Dancer2 API to handle JSON explicitly:
```perl
post '/api/v1/task/:task/complete' => sub {
my $task = route_parameters->get("task");
# Handle both JSON body and form-encoded data
my $body;
if (request_content_type eq 'application/json') {
use Dancer2::Core::Request::Entity;
$body = decode_json(request_body);
} else {
$body = body_parameters->to_hash;
}
my $assign_key = $body->{assign_key};
my $results = $body->{results};
# ... rest of the handler unchanged
};
```
---
## Files Modified / Created Summary
| File | Action | Description |
|------|--------|-------------|
| `src/task_worker.py` | **NEW** | CLI entry point: `--tasks N`, loop claim→process→submit |
| `src/api_client.py` | **NEW** | HTTP client for the Perl API (`ApiClient`) |
| `src/task_processor.py` | **NEW** | `process_video(video_dict)` — core AI pipeline callable standalone |
| `docker-compose.yml` | **MODIFY** | Change worker command to invoke task_worker; add TASK_COUNT env var |
| `api/app.pl` | **MAYBE MODIFY** | Update `/task/:task/complete` to handle JSON body if Dancer2 doesn't support it natively |
| `src/main.py` | **MODIFY (cleanup later)** | Remove or deprecate WorkerPool usage (long-lived process no longer needed) |
| `REFACTOR.md` | **CURRENT FILE** | This plan document |
---
## Risks & Considerations
1. **GPU model loading per invocation:** Each new Python process loads the ONNX/TensorRT engine into GPU memory. With `--tasks 1`, this is wasteful. Mitigation: use larger `--tasks` values (e.g., 20-100) so amortization is favorable, or implement a local cache server pattern later.
2. **Permanence of scratch space:** The `ScratchManager` creates per-video temp files under `/scratch`. Since workers are now short-lived, ensure `cleanup=True` is always set (it is by default in the existing code).
3. **API JSON body compatibility:** Dancer2's `body_parameters` may not parse JSON bodies — it expects form-encoded data. This is the highest-risk item. Test Phase 6 early.
4. **Task idempotency:** If a worker crashes between processing and submitting results, the task remains `IN_PROGRESS`. The Perl API prevents re-claiming (the WHERE clause checks `status='PENDING'`). A manual SQL update or a `/api/v1/task/:task/reset` endpoint may be needed for recovery. This is outside the scope of this refactor but worth noting.
---
## Suggested Order of Execution
1. **Phase 6 first** — verify API JSON compatibility (takes 5 minutes, unblocks everything)
2. **Phase 2** — create `api_client.py` with minimal methods
3. **Phase 3** — create `task_processor.py` by extracting from `orchestrator.py`
4. **Phase 1** — create `task_worker.py` with CLI + loop
5. **Phase 4** — wire together and test locally (`python3 -m src.task_worker --tasks 3`)
6. **Phase 5** — update docker-compose.yml and deploy
+33 -4
View File
@@ -130,8 +130,22 @@ get '/api/v1/nexttask/:type' => sub {
post '/api/v1/task/:task/complete' => sub { post '/api/v1/task/:task/complete' => sub {
my $task = route_parameters->get("task"); my $task = route_parameters->get("task");
my $assign_key = body_parameters->get("assign_key");
my $results = body_parameters->get("results"); # Accept both JSON body and form-encoded data for assign_key/results.
my ($assign_key, $results);
my $ct = request_header('Content-Type') || '';
if ($ct eq 'application/json') {
my $body = decode_json(request->body());
$assign_key = $body->{assign_key};
$results = $body->{results};
} else {
$assign_key = body_parameters->get("assign_key");
$results = body_parameters->get("results");
}
unless ($assign_key) {
send_error("Missing assign_key", 400);
}
# Verify the task is assigned to the worker # Verify the task is assigned to the worker
my $sth = database->prepare("SELECT * FROM tasks WHERE id=? AND assign_key=? AND status='IN_PROGRESS'"); my $sth = database->prepare("SELECT * FROM tasks WHERE id=? AND assign_key=? AND status='IN_PROGRESS'");
@@ -143,9 +157,10 @@ post '/api/v1/task/:task/complete' => sub {
send_error("Task not assigned to this worker or not in progress", 403); send_error("Task not assigned to this worker or not in progress", 403);
} }
# Update the task as completed # Update the task as completed (store results as JSON text for MariaDB)
my $results_json = defined $results ? encode_json($results) : 'null';
my $update_sth = database->prepare("UPDATE tasks SET status='COMPLETED', results=?, updated_at=NOW() WHERE id=?"); my $update_sth = database->prepare("UPDATE tasks SET status='COMPLETED', results=?, updated_at=NOW() WHERE id=?");
$update_sth->execute($results, $task); $update_sth->execute($results_json, $task);
if($update_sth->rows == 0) { if($update_sth->rows == 0) {
send_error("Failed to complete task", 500); send_error("Failed to complete task", 500);
} }
@@ -154,4 +169,18 @@ post '/api/v1/task/:task/complete' => sub {
return { message => "Task completed successfully" }; return { message => "Task completed successfully" };
}; };
post '/api/v1/task' => sub {
my $video_id = body_parameters->get("video_id");
my $task_type = body_parameters->get("task_type");
database->do("DELETE FROM tasks WHERE video_id=? AND task_type=?", undef, $video_id, $task_type);
my $sth = database->prepare("INSERT INTO tasks (video_id, task_type, status) VALUES (?, ?, 'PENDING')");
$sth->execute($video_id, $task_type);
my $task_id = database->last_insert_id(undef, undef, 'tasks', undef);
$sth->finish();
return { id => $task_id };
};
start(); start();
+3 -3
View File
@@ -7,7 +7,7 @@ services:
MYSQL_ROOT_PASSWORD: ${DB_ROOT_PASSWORD:-rootpass} MYSQL_ROOT_PASSWORD: ${DB_ROOT_PASSWORD:-rootpass}
MYSQL_DATABASE: ${DB_NAME:-videodetect} MYSQL_DATABASE: ${DB_NAME:-videodetect}
MYSQL_USER: ${DB_USER:-videodetect} MYSQL_USER: ${DB_USER:-videodetect}
MYSQL_PASSWORD: ${DB_PASSWORD:-videodetect123} MYSQL_PASSWORD: ${DB_PASSWORD:-changeme_videodetect}
ports: ports:
- "3306:3306" - "3306:3306"
volumes: volumes:
@@ -38,7 +38,7 @@ services:
- DB_PORT=3306 - DB_PORT=3306
- DB_NAME=videodetect - DB_NAME=videodetect
- DB_USER=videodetect - DB_USER=videodetect
- DB_PASSWORD=${DB_PASSWORD:-videodetect123} - DB_PASSWORD=${DB_PASSWORD:-changeme_videodetect}
- FLASK_ENV=production - FLASK_ENV=production
volumes: volumes:
- ./ui:/app - ./ui:/app
@@ -67,7 +67,7 @@ services:
- DB_PORT=3306 - DB_PORT=3306
- DB_NAME=videodetect - DB_NAME=videodetect
- DB_USER=videodetect - DB_USER=videodetect
- DB_PASSWORD=${DB_PASSWORD:-videodetect123} - DB_PASSWORD=${DB_PASSWORD:-changeme_videodetect}
- DANCER_ENVIRONMENT=production - DANCER_ENVIRONMENT=production
volumes: volumes:
- ./api:/app - ./api:/app
+136
View File
@@ -0,0 +1,136 @@
"""HTTP client for VideoDetect's Perl REST API (api/app.pl).
Wraps the three worker-facing endpoints:
GET /api/v1/nexttask/:type — claim next pending task
GET /api/v1/video/:id — fetch video metadata
POST /api/v1/task/:task/complete — submit results
Uses the ``requests`` library; raises :class:`ApiError` on unexpected HTTP status codes.
"""
from __future__ import annotations
import json
import logging
from dataclasses import dataclass, field
from typing import Any, Optional
import requests
logger = logging.getLogger(__name__)
# ---------------------------------------------------------------------------
# Exception
# ---------------------------------------------------------------------------
class ApiError(Exception):
"""Raised when the API returns an unexpected error status."""
def __init__(self, status_code: int, message: str):
self.status_code = status_code
self.message = message
super().__init__(f"API error {status_code}: {message}")
# ---------------------------------------------------------------------------
# Client
# ---------------------------------------------------------------------------
class ApiClient:
"""Thin HTTP client for the Dancer2 REST API."""
def __init__(self, base_url: str = "http://localhost:3000"):
self.base_url = base_url.rstrip("/")
self.session = requests.Session()
logger.info("API client configured: base_url=%s", self.base_url)
# -- helpers -------------------------------------------------------------
def _get(self, path: str, **kwargs: Any) -> dict:
url = f"{self.base_url}{path}"
resp = self.session.get(url, **kwargs)
if resp.status_code == 404:
logger.debug("GET %s → 404", url)
return {} # caller distinguishes "not found" from real data
resp.raise_for_status()
return resp.json()
def _post(self, path: str, json_body: dict[str, Any]) -> dict:
"""POST a JSON body (Content-Type: application/json)."""
url = f"{self.base_url}{path}"
logger.debug("POST %s%s", url, json.dumps(json_body))
resp = self.session.post(url, json=json_body)
resp.raise_for_status()
return resp.json()
# -- public API ----------------------------------------------------------
def get_next_task(self, task_type: str) -> Optional[dict[str, Any]]:
"""Claim the next pending task.
Returns ``{"task": {...}, "assign_key": "worker_NNN"}`` on success,
or ``None`` when no PENDING tasks remain (HTTP 404).
"""
path = f"/api/v1/nexttask/{task_type}"
result = self._get(path)
if not result:
return None
assert "task" in result and "assign_key" in result, \
f"Unexpected response shape: {result}"
return result
def get_video(self, video_id: int) -> dict[str, Any]:
"""Fetch full video metadata by integer ID.
Raises :class:`ApiError` if the video is not found (404).
"""
path = f"/api/v1/video/{video_id}"
result = self._get(path)
if not result:
raise ApiError(404, f"Video {video_id} not found")
assert "id" in result, f"Unexpected video response shape: {result}"
return result
def submit_results(
self,
task_id: int,
assign_key: str,
results: dict[str, Any],
) -> dict[str, Any]:
"""Submit processing results for a claimed task.
Args:
task_id: The database ID returned by ``get_next_task``.
assign_key: The worker token returned by ``get_next_task``.
results: A JSON-serializable dict (status, confidence, routing_decision, etc.).
Returns the server response dict on success.
Raises :class:`ApiError` on 4xx/5xx.
"""
path = f"/api/v1/task/{task_id}/complete"
return self._post(path, {
"assign_key": assign_key,
"results": results,
})
def create_task(self, task_type: str, video_id: int) -> dict[str, Any]:
"""Create a new pending task.
Args:
task_type: Task type string (e.g. 'AISCAN', 'REVIEW').
video_id: The video to associate the task with.
Returns the server response dict on success.
Raises :class:`ApiError` on 4xx/5xx.
"""
path = "/api/v1/tasks"
return self._post(path, {
"task_type": task_type,
"video_id": video_id,
"status": "PENDING",
})
def close(self) -> None:
"""Close the underlying HTTP session."""
self.session.close()
-68
View File
@@ -1,68 +0,0 @@
"""Dynamic batching utilities for face detection inference."""
import logging
import time
from collections import deque
from typing import Callable, Deque, Generic, List, Optional, TypeVar
from gpu_manager import GPUMemoryManager
logger = logging.getLogger(__name__)
T = TypeVar("T")
R = TypeVar("R")
class DynamicBatcher(Generic[T, R]):
"""Accumulate items into batches and flush based on size or timeout."""
def __init__(
self,
process_fn: Callable[[List[T]], List[R]],
gpu_manager: Optional[GPUMemoryManager] = None,
max_batch_size: int = 16,
batch_timeout_ms: float = 100.0,
min_batch_size: int = 1,
):
self.process_fn = process_fn
self.gpu_manager = gpu_manager
self.max_batch_size = max_batch_size
self.batch_timeout_ms = batch_timeout_ms
self.min_batch_size = min_batch_size
self._queue: Deque[T] = deque()
self._last_flush = time.monotonic()
def add(self, item: T) -> List[R]:
"""Add an item and return any results if a batch was flushed."""
self._queue.append(item)
if len(self._queue) >= self._current_batch_size():
return self.flush()
elapsed_ms = (time.monotonic() - self._last_flush) * 1000
if elapsed_ms >= self.batch_timeout_ms and len(self._queue) >= self.min_batch_size:
return self.flush()
return []
def flush(self) -> List[R]:
"""Flush all queued items through the processor."""
if not self._queue:
return []
batch = [self._queue.popleft() for _ in range(min(len(self._queue), self._current_batch_size()))]
results = self.process_fn(batch)
self._last_flush = time.monotonic()
if self.gpu_manager is not None:
self.gpu_manager.adjust_batch_size()
self.gpu_manager.empty_cache()
return results
def _current_batch_size(self) -> int:
if self.gpu_manager is not None:
return min(self.gpu_manager.current_batch_size, self.max_batch_size)
return self.max_batch_size
@property
def queued_count(self) -> int:
return len(self._queue)
-115
View File
@@ -1,115 +0,0 @@
"""Batch Parquet/JSONL export of video processing results."""
import json
import logging
import os
from datetime import datetime, timezone
from pathlib import Path
from typing import Any, Dict, List, Optional
logger = logging.getLogger(__name__)
class DataExporter:
"""Buffer result records and flush to Parquet or JSONL when the batch is full."""
def __init__(
self,
output_path: str,
model_version: str = "v0.0.0",
export_format: str = "parquet",
compression: str = "snappy",
batch_size: int = 100,
include_frame_confidences: bool = True,
):
self.output_path = Path(output_path) / model_version
self.model_version = model_version
self.export_format = export_format.lower()
self.compression = compression
self.batch_size = batch_size
self.include_frame_confidences = include_frame_confidences
self._buffer: List[Dict[str, Any]] = []
self.output_path.mkdir(parents=True, exist_ok=True)
def add(self, record: Dict[str, Any]) -> None:
"""Buffer one result record; flush automatically when batch is full."""
self._buffer.append(record)
if len(self._buffer) >= self.batch_size:
self.flush()
def flush(self) -> Optional[str]:
"""Write buffered records to disk; returns the output file path or None."""
if not self._buffer:
return None
records = self._buffer[:]
self._buffer.clear()
if not self.include_frame_confidences:
for r in records:
r.pop("confidence_scores", None)
batch_id = datetime.now(timezone.utc).strftime("%Y%m%d_%H%M%S")
stem = f"{self.model_version}_{batch_id}"
if self.export_format in ("parquet", "both"):
path = self._write_parquet(records, stem)
if self.export_format in ("jsonl", "both"):
path = self._write_jsonl(records, stem)
if self.export_format not in ("parquet", "jsonl", "both"):
logger.warning("Unknown export format '%s'; defaulting to jsonl", self.export_format)
path = self._write_jsonl(records, stem)
return str(path)
def _write_parquet(self, records: List[Dict[str, Any]], stem: str) -> Path:
out = self.output_path / f"{stem}.parquet"
try:
import pyarrow as pa
import pyarrow.parquet as pq
schema = pa.schema([
pa.field("video_id", pa.int64()),
pa.field("file_path", pa.string()),
pa.field("model_version", pa.string()),
pa.field("sample_count", pa.int32()),
pa.field("confidence_scores", pa.list_(pa.float64())),
pa.field("video_confidence", pa.float64()),
pa.field("routing", pa.string()),
pa.field("processed_at", pa.string()),
])
table = pa.table(
{
"video_id": [r.get("video_id") for r in records],
"file_path": [r.get("file_path", "") for r in records],
"model_version": [r.get("model_version", self.model_version) for r in records],
"sample_count": [r.get("sample_count", 0) for r in records],
"confidence_scores": [r.get("confidence_scores", []) for r in records],
"video_confidence": [float(r.get("video_confidence", 0.0)) for r in records],
"routing": [r.get("routing", "SKIP") for r in records],
"processed_at": [r.get("processed_at", "") for r in records],
},
schema=schema,
)
pq.write_table(table, out, compression=self.compression)
logger.info("Exported %d records to %s", len(records), out)
except ImportError:
logger.warning("pyarrow not available; falling back to JSONL")
out = self._write_jsonl(records, stem.replace(".parquet", ""))
return out
def _write_jsonl(self, records: List[Dict[str, Any]], stem: str) -> Path:
out = self.output_path / f"{stem}.jsonl"
with open(out, "w", encoding="utf-8") as fh:
for record in records:
fh.write(json.dumps(record, default=str) + "\n")
logger.info("Exported %d records to %s", len(records), out)
return out
def __del__(self):
if self._buffer:
try:
self.flush()
except Exception:
pass
-122
View File
@@ -1,122 +0,0 @@
"""
VideoDetect - Video Classification System
Main entry point for the worker service.
Initializes all components and starts the processing pipeline.
"""
import logging
import signal
import sys
import threading
import time
from pathlib import Path
# Add src to path
sys.path.insert(0, str(Path(__file__).parent))
from config_loader import get_config
from db_connector import DBConnector
from logging_config import setup_logging
from orchestrator import WorkerPool
from scanner import DirectoryScanner
logger = logging.getLogger(__name__)
def graceful_shutdown(signum, frame):
"""Handle shutdown signals gracefully."""
logger.info("Received signal %d, initiating graceful shutdown...", signum)
sys.exit(0)
def main():
"""Initialize and start the VideoDetect worker."""
# Register signal handlers
signal.signal(signal.SIGTERM, graceful_shutdown)
signal.signal(signal.SIGINT, graceful_shutdown)
# Load configuration
config = get_config()
log_config = config.get_section("logging")
setup_logging(
level=log_config.get("level", "INFO"),
log_format=log_config.get("format", "json"),
rotation_max_bytes=log_config.get("rotation_max_bytes", 104857600),
rotation_backup_count=log_config.get("rotation_backup_count", 10),
)
logger.info("VideoDetect Worker starting...")
logger.info("Config: %s", config)
# Initialize database connection
db_config = config.get_section("database")
db = DBConnector(
host=db_config.get("host", "mariadb"),
port=db_config.get("port", 3306),
database=db_config.get("name", "videodetect"),
user=db_config.get("user", "videodetect"),
password=db_config.get("password", "videodetect123"),
pool_size=db_config.get("pool_size", 20),
pool_min=db_config.get("pool_min", 5),
pool_recycle=db_config.get("pool_recycle", 3600),
)
# Verify database connectivity
if not db.health_check():
logger.error("Cannot connect to database. Exiting.")
sys.exit(1)
logger.info("Database connection established.")
# Initialize schema if needed
schema_path = Path(__file__).parent.parent / "db" / "schema.sql"
if schema_path.exists():
db.initialize_schema(str(schema_path))
logger.info("Schema initialized.")
# Verify GPU availability
try:
import torch
if torch.cuda.is_available():
gpu_count = torch.cuda.device_count()
gpu_name = torch.cuda.get_device_name(0)
logger.info("GPU available: %d GPUs, primary: %s", gpu_count, gpu_name)
else:
logger.warning("CUDA is not available! Processing will be slow.")
except ImportError:
logger.warning("PyTorch not installed. GPU features disabled.")
# Initialize the directory scanner.
# It scans the permanent storage location in place (no staging/copy step),
# periodically discovering new/removed video files and queueing any that
# have not been scanned yet as PENDING for the worker pool to pick up.
storage_config = config.get_section("storage")
scanner_config = config.get_section("scanner")
scanner = DirectoryScanner(
db_connector=db,
config=config.data,
input_path=storage_config.get("input_path", "/data/input"),
scan_interval=int(scanner_config.get("scan_interval_seconds", 60)),
walker_threads=int(scanner_config.get("walker_threads", 8)),
)
logger.info("Worker initialization complete. Starting scanner and processing loop...")
# Run the scanner in a background thread (its start() is a blocking loop).
scanner_thread = threading.Thread(
target=scanner.start, name="directory-scanner", daemon=True
)
scanner_thread.start()
pool = WorkerPool(db, config.data, max_workers=1)
try:
pool.start()
except KeyboardInterrupt:
logger.info("Worker shutting down.")
pool.stop()
scanner.stop()
if __name__ == "__main__":
main()
-441
View File
@@ -1,441 +0,0 @@
"""
Batch Orchestration Skeleton
Manages job queue, worker pool, state transitions, and crash recovery.
Ensures atomic state transitions and idempotent processing.
"""
import logging
import time
from concurrent.futures import ThreadPoolExecutor
from datetime import datetime, timezone
from enum import Enum
from pathlib import Path
from typing import Dict, List, Optional
from batcher import DynamicBatcher
from classifier import FaceClassifier
from data_export import DataExporter
from face_detector import Detection, FaceDetector
from frame_sampler import FrameSampler
from gpu_manager import GPUMemoryManager
from prober import VideoProber
from result_updater import ResultUpdater
from scratch_manager import ScratchManager
import aggregator
import router
logger = logging.getLogger(__name__)
class JobStatus(Enum):
"""Job status states."""
PENDING = "PENDING"
PROCESSING = "PROCESSING"
COMPLETED = "COMPLETED"
FAILED = "ERROR"
SKIPPED = "UNSCANNABLE"
class Job:
"""Represents a single video processing job."""
def __init__(
self,
video_id: int,
file_path: str,
priority: float = 0.0,
sampling_interval_seconds: Optional[int] = None,
):
self.video_id = video_id
self.file_path = file_path
self.priority = priority # Higher = more urgent (based on modification time)
self.sampling_interval_seconds = sampling_interval_seconds
self.status = JobStatus.PENDING
self.created_at = datetime.now(timezone.utc)
self.started_at: Optional[datetime] = None
self.completed_at: Optional[datetime] = None
self.attempts = 0
self.max_retries = 3
self.error_message: Optional[str] = None
def to_dict(self) -> dict:
return {
"video_id": self.video_id,
"file_path": self.file_path,
"priority": self.priority,
"status": self.status.value,
"created_at": self.created_at.isoformat(),
"started_at": self.started_at.isoformat() if self.started_at else None,
"completed_at": self.completed_at.isoformat() if self.completed_at else None,
"attempts": self.attempts,
"error_message": self.error_message,
}
def __repr__(self):
return f"Job(id={self.video_id}, path={self.file_path}, status={self.status.value})"
class WorkerPool:
"""Manage a pool of worker processes for parallel processing."""
def __init__(self, db_connector, config: dict, max_workers: int = 2):
self.db = db_connector
self.config = config
self.max_workers = max_workers
self._executor = ThreadPoolExecutor(max_workers=max_workers)
self._running = False
self._jobs_processed = 0
self._jobs_failed = 0
self._sampling_config = (config or {}).get("sampling", {})
self._storage_config = (config or {}).get("storage", {})
self._face_detection_config = (config or {}).get("face_detection", {})
self._batching_config = (config or {}).get("batching", {})
self._gpu_manager = GPUMemoryManager(
max_memory_gb=config.get("gpu", {}).get("max_memory_gb", 18.0),
reduce_threshold_gb=self._batching_config.get("vram_reduce_threshold_gb", 16.0),
increase_threshold_gb=self._batching_config.get("vram_increase_threshold_gb", 10.0),
initial_batch_size=self._batching_config.get("max_batch_size", 16),
)
self._face_detector = FaceDetector(
engine_path=self._face_detection_config.get("model_path", "/models/face_detector/face_detector.trt"),
input_size=int(self._face_detection_config.get("input_size", 640)),
confidence_threshold=float(self._face_detection_config.get("confidence_threshold", 0.25)),
iou_threshold=float(self._face_detection_config.get("iou_threshold", 0.45)),
max_faces_per_frame=int(self._face_detection_config.get("max_faces_per_frame", 10)),
max_faces_per_video=int(self._face_detection_config.get("max_faces_per_video", 100)),
)
classifier_config = (config or {}).get("classifier", {})
self._classifier = FaceClassifier(
engine_path=classifier_config.get("model_path", "/models/classifier/classifier.trt"),
temperature=float(classifier_config.get("temperature", 1.0)),
input_size=int(classifier_config.get("input_size", 224)),
)
self._aggregation_config = (config or {}).get("aggregation", {})
self._routing_config = (config or {}).get("routing", {})
model_version = self._get_active_model_version()
export_config = (config or {}).get("export", {})
self._result_updater = ResultUpdater(db_connector, model_version=model_version)
self._exporter = DataExporter(
output_path=export_config.get("output_path", "/data/output"),
model_version=model_version,
export_format=export_config.get("format", "parquet"),
compression=export_config.get("compression", "snappy"),
batch_size=int(export_config.get("batch_size", 100)),
include_frame_confidences=bool(export_config.get("include_frame_confidences", True)),
)
def _get_active_model_version(self) -> str:
"""Return the currently active model version from the DB."""
try:
row = self.db.fetchone("SELECT version FROM models WHERE status = 'ACTIVE' LIMIT 1")
if row:
return row["version"]
except Exception as exc:
logger.debug("Could not fetch active model version: %s", exc)
return "v0.0.0-placeholder"
def start(self):
"""Start the worker pool."""
self._running = True
logger.info("Worker pool starting with %d workers", self.max_workers)
while self._running:
# Get pending jobs
jobs = self._get_pending_jobs()
if jobs:
# Submit jobs to executor
for job in jobs:
future = self._executor.submit(self._process_job, job)
future.add_done_callback(self._on_job_complete)
else:
time.sleep(5) # No jobs, wait
def stop(self):
"""Stop the worker pool gracefully."""
self._running = False
logger.info("Worker pool stopping. Processed: %d, Failed: %d",
self._jobs_processed, self._jobs_failed)
self._executor.shutdown(wait=True)
def _get_pending_jobs(self) -> List[Job]:
"""Get pending jobs from DB with atomic locking."""
now = datetime.now(timezone.utc).replace(tzinfo=None)
# Atomic lock: update status from PENDING to PROCESSING
result = self.db.execute(
"""UPDATE videos SET status = 'PROCESSING', last_processed_time = %s,
updated_at = %s
WHERE id IN (
SELECT id FROM (
SELECT id FROM videos
WHERE status = 'PENDING'
ORDER BY last_scan_time DESC
LIMIT %s
) AS sub
)""",
(now, now, self.max_workers),
transaction=True,
)
if result == 0:
return []
# Get the locked jobs
locked = self.db.fetchall(
"SELECT id, file_path, last_scan_time FROM videos WHERE status = 'PROCESSING' AND last_processed_time = %s",
(now,),
)
return [
Job(
video_id=row["id"],
file_path=row["file_path"],
priority=row["last_scan_time"].timestamp(),
)
for row in locked
]
def _process_job(self, job: Job) -> bool:
"""Process a single job (placeholder - actual processing in later stories)."""
job.status = JobStatus.PROCESSING
job.started_at = datetime.now(timezone.utc)
job.attempts += 1
logger.info("Processing job: %s (attempt %d)", job, job.attempts)
try:
sampling_interval = job.sampling_interval_seconds or self._sampling_config.get("interval_seconds", 30)
scratch_base_path = self._storage_config.get("scratch_path", "/scratch")
scratch_manager = ScratchManager(
base_path=scratch_base_path,
video_id=str(job.video_id),
auto_cleanup=True,
)
frame_dir = scratch_manager.ensure_frame_dir()
prober = VideoProber(timeout=10)
metadata = prober.probe(job.file_path)
if metadata.is_unscannable or metadata.duration is None:
raise RuntimeError(metadata.error_message or "Video metadata could not be determined")
resolution = None
if metadata.resolution_w and metadata.resolution_h:
resolution = (metadata.resolution_w, metadata.resolution_h)
sampler = FrameSampler(
interval_seconds=int(sampling_interval),
quality=int(self._sampling_config.get("quality", 2)),
output_format=self._sampling_config.get("format", "jpeg"),
)
extracted_frames = sampler.extract_frames(
video_path=job.file_path,
output_dir=str(frame_dir),
duration=metadata.duration,
interval_seconds=int(sampling_interval),
resolution=resolution,
timeout_seconds=30,
)
if not extracted_frames:
scratch_manager.cleanup()
raise RuntimeError("No frames were extracted")
# Face detection on sampled frames
batch_size = self._gpu_manager.current_batch_size
detections_per_frame = self._face_detector.detect_faces(extracted_frames, batch_size=batch_size)
# Flatten and cap total faces per video
all_detections: List[Detection] = []
for frame_dets in detections_per_frame:
all_detections.extend(frame_dets)
all_detections = sorted(all_detections, key=lambda d: d.confidence, reverse=True)
all_detections = all_detections[: self._face_detection_config.get("max_faces_per_video", 100)]
# Extract face crops
crop_dir = scratch_manager.frame_dir.parent / "crops"
cropped_detections = self._face_detector.extract_crops(
all_detections,
output_dir=str(crop_dir),
crop_size=(224, 224),
)
if not cropped_detections:
logger.info("No faces detected for video %s; routing to SKIP", job.video_id)
routing_decision = router.SKIP
video_confidence = 0.0
frame_confidences: List[float] = []
else:
crop_paths = [d.crop_path for d in cropped_detections if d.crop_path]
frame_confidences = self._classifier.classify(
crop_paths, batch_size=self._gpu_manager.current_batch_size
)
video_confidence = aggregator.aggregate(
frame_confidences,
strategy=self._aggregation_config.get("strategy", "max"),
alpha=float(self._aggregation_config.get("alpha", 1.0)),
beta=float(self._aggregation_config.get("beta", 0.1)),
top_k=int(self._aggregation_config.get("top_k", 3)),
)
routing_decision = router.route(
video_confidence,
t_high=float(self._routing_config.get("T_high", 0.75)),
t_low=float(self._routing_config.get("T_low", 0.45)),
)
logger.info(
"Video %s: C=%.4f routing=%s faces=%d frames=%d",
job.video_id, video_confidence, routing_decision,
len(cropped_detections), len(extracted_frames),
)
# Persist results atomically (video update + processing log) before cleanup
persisted = self._result_updater.persist(
video_id=job.video_id,
frame_count=len(extracted_frames),
confidence=video_confidence,
routing=routing_decision,
frame_confidences=frame_confidences,
)
if persisted:
job.status = JobStatus.COMPLETED
job.completed_at = datetime.now(timezone.utc)
self._exporter.add({
"video_id": job.video_id,
"file_path": job.file_path,
"model_version": self._result_updater.model_version,
"sample_count": len(extracted_frames),
"confidence_scores": frame_confidences,
"video_confidence": video_confidence,
"routing": routing_decision,
"processed_at": datetime.now(timezone.utc).isoformat(),
})
# Scratch cleanup only after successful persistence
scratch_manager.cleanup_all()
self._jobs_processed += 1
return True
except Exception as e:
logger.error("Job failed: %s - %s", job, e, exc_info=True)
job.error_message = str(e)
if job.attempts < job.max_retries:
# Retry
self._retry_job(job)
return False
else:
# Max retries exceeded
self._fail_job(job)
self._jobs_failed += 1
return False
def _complete_job(
self,
job: Job,
frame_count: Optional[int] = None,
confidence: Optional[float] = None,
routing: Optional[str] = None,
):
"""Mark a job as completed, persisting confidence and routing decision."""
now = datetime.now(timezone.utc).replace(tzinfo=None)
job.status = JobStatus.COMPLETED
job.completed_at = datetime.now(timezone.utc)
self.db.execute(
"""UPDATE videos
SET status = 'COMPLETED',
frame_count = %s,
confidence_score = %s,
routing_decision = %s,
updated_at = %s
WHERE id = %s""",
(frame_count, confidence, routing, now, job.video_id),
transaction=True,
)
logger.info("Job completed: %s", job)
def _fail_job(self, job: Job):
"""Mark a job as failed."""
now = datetime.now(timezone.utc).replace(tzinfo=None)
job.status = JobStatus.FAILED
job.completed_at = datetime.now(timezone.utc)
self.db.execute(
"""UPDATE videos SET status = 'ERROR', error_message = %s, updated_at = %s
WHERE id = %s""",
(job.error_message, now, job.video_id),
transaction=True,
)
logger.error("Job failed permanently: %s", job)
def _retry_job(self, job: Job):
"""Re-queue a job for retry."""
now = datetime.now(timezone.utc).replace(tzinfo=None)
self.db.execute(
"""UPDATE videos SET status = 'PENDING', updated_at = %s
WHERE id = %s""",
(now, job.video_id),
transaction=True,
)
logger.info("Job re-queued for retry: %s (attempt %d/%d)",
job, job.attempts, job.max_retries)
def _on_job_complete(self, future):
"""Callback when a job completes."""
try:
future.result()
except Exception as e:
logger.error("Unhandled job error: %s", e)
def get_stats(self) -> dict:
"""Get worker pool statistics."""
return {
"max_workers": self.max_workers,
"jobs_processed": self._jobs_processed,
"jobs_failed": self._jobs_failed,
"is_running": self._running,
"active_workers": self._executor._work_queue.qsize(),
}
class CrashRecovery:
"""Handle crash recovery and stale job detection."""
def __init__(self, db_connector, lock_timeout_minutes: int = 5):
self.db = db_connector
self.lock_timeout = lock_timeout_minutes
def recover_stale_jobs(self) -> int:
"""Re-queue jobs stuck in PROCESSING beyond the lock timeout."""
now = datetime.now(timezone.utc).replace(tzinfo=None)
timeout = now.replace(minute=now.minute - self.lock_timeout)
result = self.db.execute(
"""UPDATE videos SET status = 'PENDING', updated_at = %s
WHERE status = 'PROCESSING' AND last_processed_time < %s""",
(now, timeout),
transaction=True,
)
if result > 0:
logger.info("Recovered %d stale jobs", result)
return result
def check_health(self) -> dict:
"""Check system health for crash recovery purposes."""
processing_count = self.db.fetchone(
"SELECT COUNT(*) as count FROM videos WHERE status = 'PROCESSING'"
)
pending_count = self.db.fetchone(
"SELECT COUNT(*) as count FROM videos WHERE status = 'PENDING'"
)
return {
"processing_count": processing_count["count"] if processing_count else 0,
"pending_count": pending_count["count"] if pending_count else 0,
"stale_jobs_recovered": self.recover_stale_jobs(),
}
-39
View File
@@ -1,39 +0,0 @@
"""Insert processing audit log entries within an existing DB transaction."""
import json
import logging
from datetime import datetime, timezone
from typing import List, Optional
logger = logging.getLogger(__name__)
def insert_log(
cursor,
video_id: int,
model_version: str,
frame_count: int,
confidence_score: Optional[float],
routing_decision: str,
frame_confidences: Optional[List[float]] = None,
) -> None:
"""Insert one row into processing_logs; must be called inside an open transaction."""
now = datetime.now(timezone.utc).replace(tzinfo=None)
confidence_scores_json = json.dumps(frame_confidences) if frame_confidences is not None else None
cursor.execute(
"""INSERT INTO processing_logs
(video_id, model_version, frame_count, confidence_score,
confidence_scores, routing_decision, processed_at)
VALUES (%s, %s, %s, %s, %s, %s, %s)""",
(
video_id,
model_version,
frame_count or 0,
confidence_score,
confidence_scores_json,
routing_decision,
now,
),
)
logger.debug("Inserted processing log for video %s (routing=%s)", video_id, routing_decision)
-75
View File
@@ -1,75 +0,0 @@
"""Atomic result persistence: update videos + insert processing_logs in one transaction."""
import logging
from datetime import datetime, timezone
from typing import List, Optional
import processing_logger
logger = logging.getLogger(__name__)
class ResultUpdater:
def __init__(self, db_connector, model_version: str = "v0.0.0-placeholder"):
self.db = db_connector
self.model_version = model_version
def persist(
self,
video_id: int,
frame_count: int,
confidence: float,
routing: str,
frame_confidences: Optional[List[float]] = None,
) -> bool:
"""Atomically update videos and insert a processing log.
Returns False without raising if the state guard prevents the update
(video is no longer in PROCESSING state).
"""
now = datetime.now(timezone.utc).replace(tzinfo=None)
try:
with self.db.transaction() as conn:
cursor = conn.cursor()
try:
cursor.execute(
"""UPDATE videos
SET status = 'COMPLETED',
last_processed_time = %s,
confidence_score = %s,
routing_decision = %s,
model_version = %s,
frame_count = %s,
updated_at = %s
WHERE id = %s AND status = 'PROCESSING'""",
(now, confidence, routing, self.model_version, frame_count, now, video_id),
)
if cursor.rowcount == 0:
logger.warning(
"State guard: video %s is not PROCESSING; skipping update", video_id
)
return False
processing_logger.insert_log(
cursor,
video_id=video_id,
model_version=self.model_version,
frame_count=frame_count,
confidence_score=confidence,
routing_decision=routing,
frame_confidences=frame_confidences,
)
finally:
cursor.close()
logger.info(
"Persisted results for video %s: C=%.4f routing=%s model=%s",
video_id, confidence, routing, self.model_version,
)
return True
except Exception as exc:
logger.error("Failed to persist results for video %s: %s", video_id, exc)
raise
-114
View File
@@ -1,114 +0,0 @@
"""
Retry decorator with exponential backoff for transient failures.
Usage:
@retry(max_attempts=3, step="extract")
def extract_frames(...): ...
Non-retryable error types are re-raised immediately without consuming attempts.
"""
import functools
import logging
import time
from typing import Callable, Optional, Tuple, Type
logger = logging.getLogger(__name__)
# Error class names that should never be retried
_NON_RETRYABLE_NAMES = frozenset({
"CodecUnsupportedError",
"FileCorruptError",
"InvalidPathError",
"FileNotFoundError",
"PermissionError",
"IsADirectoryError",
})
class RetryExhaustedError(Exception):
"""Raised when all retry attempts have been exhausted."""
def retry(
max_attempts: int = 3,
initial_delay: float = 1.0,
backoff_factor: float = 2.0,
exceptions: Tuple[Type[Exception], ...] = (Exception,),
step: str = "unknown",
):
"""
Decorator: retry ``func`` up to ``max_attempts`` times on retryable exceptions.
Non-retryable exceptions (see _NON_RETRYABLE_NAMES) propagate immediately.
Each retry waits ``initial_delay * backoff_factor ** attempt`` seconds.
"""
def decorator(func: Callable) -> Callable:
@functools.wraps(func)
def wrapper(*args, **kwargs):
last_exc: Optional[Exception] = None
for attempt in range(max_attempts):
try:
return func(*args, **kwargs)
except exceptions as exc:
if _is_non_retryable(exc):
logger.error(
"Non-retryable error in step=%s (attempt %d/%d): %s: %s",
step, attempt + 1, max_attempts,
type(exc).__name__, exc,
)
raise
last_exc = exc
_increment_retry_counter(step)
if attempt < max_attempts - 1:
delay = initial_delay * (backoff_factor ** attempt)
logger.warning(
"Transient error in step=%s (attempt %d/%d), "
"retrying in %.1fs: %s: %s",
step, attempt + 1, max_attempts, delay,
type(exc).__name__, exc,
)
time.sleep(delay)
else:
logger.error(
"All %d attempts exhausted in step=%s: %s: %s",
max_attempts, step, type(exc).__name__, exc,
)
raise RetryExhaustedError(
f"step={step} failed after {max_attempts} attempts"
) from last_exc
return wrapper
return decorator
def retry_from_config(config, step: str = "unknown"):
"""Build a ``@retry`` decorator from config.yaml monitoring.retry settings."""
mon = config.get_section("monitoring")
cfg = mon.get("retry", {})
return retry(
max_attempts=int(cfg.get("max_attempts", 3)),
initial_delay=float(cfg.get("initial_delay", 1.0)),
backoff_factor=float(cfg.get("backoff_factor", 2.0)),
step=step,
)
# ---------------------------------------------------------------------------
# Helpers
# ---------------------------------------------------------------------------
def _is_non_retryable(exc: Exception) -> bool:
return type(exc).__name__ in _NON_RETRYABLE_NAMES
def _increment_retry_counter(step: str):
try:
from metrics import retry_attempts_total
retry_attempts_total.labels(step=step).inc()
except Exception:
pass # metrics not available; don't break the retry logic
-114
View File
@@ -1,114 +0,0 @@
"""
Review queue export for backend scripts.
Provides CSV and JSON export of annotated review data, with optional
filtering by date range, model version, annotation status, and ground truth.
Can be called standalone or imported from other src/ modules.
"""
import csv
import io
import json
import logging
from datetime import datetime
from pathlib import Path
from typing import Any, Dict, List, Optional
logger = logging.getLogger(__name__)
def fetch_annotated(
db_connector,
annotated_only: bool = True,
model_version: Optional[str] = None,
date_from: Optional[str] = None,
date_to: Optional[str] = None,
ground_truth: Optional[bool] = None,
) -> List[Dict[str, Any]]:
"""Query the review_queue and return matching records as plain dicts."""
clauses: List[str] = []
params: List[Any] = []
if annotated_only:
clauses.append("rq.annotated = TRUE")
if model_version:
clauses.append("v.model_version = %s")
params.append(model_version)
if date_from:
clauses.append("rq.annotated_at >= %s")
params.append(date_from)
if date_to:
clauses.append("rq.annotated_at < %s")
params.append(date_to)
if ground_truth is not None:
clauses.append("rq.ground_truth = %s")
params.append(bool(ground_truth))
where = ("WHERE " + " AND ".join(clauses)) if clauses else ""
rows = db_connector.fetchall(
f"""SELECT rq.video_id, v.file_path, rq.confidence_score, rq.routing_decision,
v.model_version, rq.ground_truth, rq.annotated_at, rq.notes,
pl.confidence_scores
FROM review_queue rq
JOIN videos v ON v.id = rq.video_id
LEFT JOIN (
SELECT video_id, confidence_scores,
ROW_NUMBER() OVER (PARTITION BY video_id ORDER BY processed_at DESC) rn
FROM processing_logs
) pl ON pl.video_id = rq.video_id AND pl.rn = 1
{where}
ORDER BY rq.annotated_at DESC""",
params if params else None,
)
return [_normalise(row) for row in (rows or [])]
def _normalise(row: dict) -> Dict[str, Any]:
try:
scores = json.loads(row["confidence_scores"]) if row.get("confidence_scores") else []
except (json.JSONDecodeError, TypeError):
scores = []
at = row.get("annotated_at")
return {
"video_id": row["video_id"],
"file_path": row["file_path"],
"confidence_score": row["confidence_score"],
"routing_decision": row["routing_decision"],
"model_version": row.get("model_version"),
"ground_truth": bool(row["ground_truth"]) if row["ground_truth"] is not None else None,
"annotated_at": at.isoformat() if isinstance(at, datetime) else at,
"notes": row.get("notes"),
"contributing_frames": scores,
}
def export_json(records: List[Dict[str, Any]], output_path: str) -> str:
"""Write records to a UTF-8 JSON file; returns the path."""
path = Path(output_path)
path.parent.mkdir(parents=True, exist_ok=True)
path.write_text(json.dumps(records, default=str, indent=2), encoding="utf-8")
logger.info("Exported %d records to %s", len(records), path)
return str(path)
def export_csv(records: List[Dict[str, Any]], output_path: str) -> str:
"""Write records to a UTF-8 CSV file; returns the path."""
path = Path(output_path)
path.parent.mkdir(parents=True, exist_ok=True)
if not records:
path.write_text("", encoding="utf-8")
return str(path)
buf = io.StringIO()
writer = csv.DictWriter(buf, fieldnames=list(records[0].keys()))
writer.writeheader()
for r in records:
row = dict(r)
row["contributing_frames"] = json.dumps(row["contributing_frames"])
writer.writerow(row)
path.write_text(buf.getvalue(), encoding="utf-8")
logger.info("Exported %d records to %s", len(records), path)
return str(path)
-537
View File
@@ -1,537 +0,0 @@
"""
Directory Scanner Service
Walks /data/input to discover new/modified files, computes hashes,
probes video metadata, validates codecs, and syncs state to MariaDB.
Supports incremental scanning for efficient 30TB corpus handling.
"""
import hashlib
import json
import logging
import os
import socket
import threading
import time
import uuid
from concurrent.futures import ThreadPoolExecutor, as_completed
from datetime import datetime, timedelta, timezone
from pathlib import Path
from typing import Dict, List, Optional, Tuple
from codec_validator import CodecValidator
from prober import VideoMetadata, VideoProber
logger = logging.getLogger(__name__)
class DirectoryScanner:
"""Scan directories for new/modified video files and sync to DB."""
VIDEO_EXTENSIONS = {'.mp4', '.mov', '.avi', '.mkv', '.webm', '.flv', '.wmv'}
def __init__(
self,
db_connector,
config: dict,
input_path: str = "/data/input",
scan_interval: int = 60,
walker_threads: int = 8,
):
self.db = db_connector
self.config = config
self.input_path = Path(input_path)
self.scan_interval = scan_interval
self.walker_threads = walker_threads
self.prober = VideoProber(
timeout=config.get("scanner", {}).get("ffprobe_timeout_seconds", 10)
)
self.codec_validator = CodecValidator(
whitelist=set(config.get("codec", {}).get("whitelist", [])),
default_status_on_error=config.get("codec", {}).get("default_status_on_error", "UNSCANNABLE"),
)
self._running = False
self._scan_count = 0
self._total_files_discovered = 0
self._total_files_new = 0
self._total_files_modified = 0
self._total_files_unscannable = 0
# Single-instance guard: prevents overlapping scans both within this
# process (threading lock) and across replicas (DB lock with a lease).
self._instance_id = f"{socket.gethostname()}-{os.getpid()}-{uuid.uuid4().hex[:8]}"
self._lock_name = "directory_scanner"
self._lock_lease_seconds = int(
config.get("scanner", {}).get("lock_lease_seconds", 21600)
)
self._heartbeat_every = int(
config.get("scanner", {}).get("heartbeat_interval_files", 500)
)
self._scan_lock = threading.Lock()
self._lock_table_ensured = False
def start(self):
"""Start the scanner loop.
Each cycle is guarded so that at most one scan runs at a time:
- an in-process threading lock prevents re-entrant scans, and
- a DB lock (with a lease) prevents overlapping scans across replicas.
A long-running scan keeps its lease alive via heartbeats, so the next
scheduled tick (or another replica) waits instead of starting a second
parallel scan.
"""
self._running = True
logger.info(
"Scanner starting: input_path=%s interval=%ds threads=%d instance=%s",
self.input_path, self.scan_interval, self.walker_threads, self._instance_id,
)
while self._running:
# In-process re-entrancy guard: never run two scans at once.
if not self._scan_lock.acquire(blocking=False):
logger.warning("A scan is already in progress; skipping this cycle.")
self._sleep_interval()
continue
try:
if self.acquire_lock():
try:
self._run_scan()
finally:
self.release_lock()
else:
logger.info(
"Scanner lock held by another instance; skipping this cycle."
)
except Exception as e:
logger.error("Scanner error: %s", e, exc_info=True)
finally:
self._scan_lock.release()
# Sleep until next scan
self._sleep_interval()
def _sleep_interval(self):
"""Sleep for the scan interval, waking early if stopped."""
for _ in range(self.scan_interval):
if not self._running:
break
time.sleep(1)
def stop(self):
"""Stop the scanner."""
self._running = False
logger.info("Scanner stopping. Total scans: %d", self._scan_count)
# ------------------------------------------------------------------
# Single-instance lock (cross-process / cross-replica guard)
# ------------------------------------------------------------------
def _ensure_lock_table(self):
"""Create the scanner_lock table if it does not already exist."""
if self._lock_table_ensured:
return
self.db.execute(
"""CREATE TABLE IF NOT EXISTS scanner_lock (
lock_name VARCHAR(64) PRIMARY KEY,
owner VARCHAR(128) NOT NULL,
locked_at DATETIME NOT NULL,
lease_seconds INT NOT NULL DEFAULT 21600
) ENGINE=InnoDB DEFAULT CHARSET=utf8mb4 COLLATE=utf8mb4_unicode_ci"""
)
self._lock_table_ensured = True
def acquire_lock(self) -> bool:
"""Attempt to acquire the cross-process scanner lock.
Returns True if this instance now owns the lock, False otherwise.
A stale lock (held longer than the lease) is taken over so a crashed
scanner does not block scanning forever.
"""
try:
self._ensure_lock_table()
except Exception as e:
logger.warning("Could not ensure scanner_lock table: %s", e)
return True # fail-open: keep scanning rather than block entirely
now = datetime.now(timezone.utc).replace(tzinfo=None)
# 1) Try to insert a fresh lock row.
try:
self.db.execute(
"""INSERT INTO scanner_lock (lock_name, owner, locked_at, lease_seconds)
VALUES (%s, %s, %s, %s)""",
(self._lock_name, self._instance_id, now, self._lock_lease_seconds),
transaction=True,
)
logger.info("Acquired scanner lock (fresh). owner=%s", self._instance_id)
return True
except Exception:
# Row already exists -> fall through to steal-if-stale.
pass
# 2) Take over the lock if it is stale or already ours.
stale_before = now - timedelta(seconds=self._lock_lease_seconds)
try:
affected = self.db.execute(
"""UPDATE scanner_lock
SET owner = %s, locked_at = %s
WHERE lock_name = %s
AND (owner = %s OR locked_at < %s)""",
(self._instance_id, now, self._lock_name, self._instance_id, stale_before),
transaction=True,
)
if affected and affected > 0:
logger.info("Acquired scanner lock (stale takeover). owner=%s",
self._instance_id)
return True
except Exception as e:
logger.warning("Failed to check scanner lock: %s", e)
return True # fail-open
logger.info("Scanner lock held by another instance; not acquiring.")
return False
def release_lock(self):
"""Release the scanner lock if we own it."""
try:
self.db.execute(
"""DELETE FROM scanner_lock WHERE lock_name = %s AND owner = %s""",
(self._lock_name, self._instance_id),
transaction=True,
)
except Exception as e:
logger.warning("Failed to release scanner lock: %s", e)
def _heartbeat(self):
"""Refresh the lock lease so a long-running scan is not stolen."""
try:
now = datetime.now(timezone.utc).replace(tzinfo=None)
self.db.execute(
"""UPDATE scanner_lock SET locked_at = %s
WHERE lock_name = %s AND owner = %s""",
(now, self._lock_name, self._instance_id),
transaction=True,
)
except Exception as e:
logger.debug("Scanner lock heartbeat failed: %s", e)
def _run_scan(self):
"""Execute a single scan cycle."""
scan_start = time.time()
self._scan_count += 1
logger.info("Scan #%d starting...", self._scan_count)
# Get last scan time from DB
last_scan = self._get_last_scan_time()
# Discover files
files_to_process, files_removed = self._discover_files(last_scan)
# Process files in parallel
results = self._process_files_batch(files_to_process)
# Update DB
scan_result = self._update_db(results, files_removed)
# Record scan history
duration = time.time() - scan_start
self._record_scan_history(scan_result, duration)
logger.info(
"Scan #%d complete: discovered=%d new=%d modified=%d unscannable=%d removed=%d duration=%.1fs",
self._scan_count,
scan_result["discovered"],
scan_result["new"],
scan_result["modified"],
scan_result["unscannable"],
len(files_removed),
duration,
)
def _get_last_scan_time(self) -> Optional[datetime]:
"""Get the last scan time from DB."""
result = self.db.fetchone(
"SELECT scan_end FROM scan_history ORDER BY id DESC LIMIT 1"
)
if result and result.get("scan_end"):
return result["scan_end"]
return None
def _discover_files(self, last_scan: Optional[datetime]) -> Tuple[List[Path], List[str]]:
"""Discover new and modified files in the input directory."""
files_to_process = []
files_removed = []
if not self.input_path.exists():
logger.warning("Input path does not exist: %s", self.input_path)
return files_to_process, files_removed
# Get files from DB for comparison
if last_scan:
# Only check files modified after last scan
db_files = self.db.fetchall(
"SELECT file_path, last_scan_time FROM videos WHERE last_scan_time > %s",
(last_scan,),
)
db_paths = {row["file_path"] for row in db_files}
# Check for modified files
for row in db_files:
file_path = Path(row["file_path"])
if file_path.exists() and file_path.is_file():
try:
mtime = datetime.fromtimestamp(
file_path.stat().st_mtime, tz=timezone.utc
)
if mtime > row["last_scan_time"]:
files_to_process.append(file_path)
except OSError:
pass
else:
# Full scan - walk the directory
logger.info("Full scan (no previous scan found). Walking %s...", self.input_path)
for root, dirs, filenames in os.walk(self.input_path):
for filename in filenames:
ext = Path(filename).suffix.lower()
if ext in self.VIDEO_EXTENSIONS:
file_path = Path(root) / filename
files_to_process.append(file_path)
# Check for removed files (if we have a last scan)
if last_scan:
all_db_files = self.db.fetchall(
"SELECT file_path FROM videos WHERE status NOT IN ('REMOVED', 'ERROR')",
)
current_paths = {str(f) for f in files_to_process}
for row in all_db_files:
fp = row["file_path"]
if fp not in current_paths and Path(fp).exists():
# File still exists but not in current walk - might be in a new directory
pass
elif fp not in current_paths and not Path(fp).exists():
files_removed.append(fp)
logger.info("Discovered %d files to process, %d files removed",
len(files_to_process), len(files_removed))
return files_to_process, files_removed
def _process_files_batch(self, files: List[Path]) -> List[dict]:
"""Process a batch of files in parallel."""
results = []
processed = 0
with ThreadPoolExecutor(max_workers=self.walker_threads) as executor:
future_to_file = {
executor.submit(self._process_single_file, f): f
for f in files
}
for future in as_completed(future_to_file):
file_path = future_to_file[future]
try:
result = future.result()
results.append(result)
except Exception as e:
logger.error("Error processing %s: %s", file_path, e, exc_info=True)
results.append({
"file_path": str(file_path),
"status": "ERROR",
"error_message": str(e),
})
processed += 1
# Keep the single-instance lock alive during long scans.
if self._heartbeat_every and processed % self._heartbeat_every == 0:
self._heartbeat()
return results
def _process_single_file(self, file_path: Path) -> dict:
"""Process a single file: hash, probe, validate, determine status."""
file_path_str = str(file_path)
result = {"file_path": file_path_str}
try:
# Compute hash
file_hash = self._compute_hash(file_path_str)
result["file_hash"] = file_hash
# Check for duplicate
existing = self.db.fetchone(
"SELECT id, status FROM videos WHERE file_hash = %s",
(file_hash,),
)
if existing:
result["status"] = existing["status"]
result["video_id"] = existing["id"]
result["action"] = "duplicate"
logger.debug("Duplicate file found: %s (id=%d, status=%s)",
file_path_str, existing["id"], existing["status"])
return result
# Probe video
metadata = self.prober.probe(file_path_str)
result.update(metadata.to_dict())
# Validate codec
if metadata.is_valid:
is_supported, reason = self.codec_validator.validate(metadata.codec)
if is_supported:
result["status"] = "PENDING"
result["action"] = "new"
self._total_files_new += 1
else:
result["status"] = "UNSCANNABLE"
result["error_message"] = f"Unsupported codec: {metadata.codec} ({reason})"
result["action"] = "unscannable"
self._total_files_unscannable += 1
logger.warning("Unsupported codec for %s: %s", file_path_str, metadata.codec)
else:
result["status"] = "UNSCANNABLE"
result["action"] = "unscannable"
self._total_files_unscannable += 1
except Exception as e:
result["status"] = "ERROR"
result["error_message"] = str(e)
logger.error("Error processing %s: %s", file_path_str, e, exc_info=True)
return result
def _compute_hash(self, file_path: str) -> str:
"""Compute SHA-256 hash of the first 1MB of a file."""
chunk_size = 1024 * 1024 # 1MB
sha256 = hashlib.sha256()
try:
with open(file_path, "rb") as f:
chunk = f.read(chunk_size)
if chunk:
sha256.update(chunk)
return sha256.hexdigest()
except (OSError, IOError) as e:
logger.error("Cannot hash file %s: %s", file_path, e)
return hashlib.sha256(file_path.encode()).hexdigest() # fallback
def _update_db(self, results: List[dict], files_removed: List[str]) -> dict:
"""Update database with scan results."""
stats = {"discovered": len(results), "new": 0, "modified": 0, "unscannable": 0}
# Batch insert new files
new_files = [r for r in results if r.get("action") == "new"]
if new_files:
self._batch_insert_new_files(new_files)
stats["new"] = len(new_files)
# Update unscannable files
unscannable = [r for r in results if r.get("action") == "unscannable" and r.get("video_id")]
if unscannable:
self._batch_update_status(unscannable, "UNSCANNABLE")
stats["unscannable"] = len(unscannable)
# Mark removed files
if files_removed:
self._mark_files_removed(files_removed)
stats["removed"] = len(files_removed)
return stats
def _batch_insert_new_files(self, files: List[dict]):
"""Batch insert new files into DB."""
now = datetime.now(timezone.utc).replace(tzinfo=None)
for file_info in files:
try:
self.db.execute(
"""INSERT INTO videos
(file_path, file_hash, resolution_w, resolution_h, codec,
duration, status, last_scan_time, created_at, updated_at)
VALUES (%s, %s, %s, %s, %s, %s, %s, %s, %s, %s)""",
(
file_info["file_path"],
file_info["file_hash"],
file_info.get("resolution_w"),
file_info.get("resolution_h"),
file_info.get("codec"),
file_info.get("duration"),
"PENDING",
now,
now,
now,
),
transaction=True,
)
except Exception as e:
logger.error("Failed to insert %s: %s", file_info["file_path"], e)
def _batch_update_status(self, files: List[dict], status: str):
"""Batch update file status."""
now = datetime.now(timezone.utc).replace(tzinfo=None)
for file_info in files:
video_id = file_info.get("video_id")
if video_id:
try:
self.db.execute(
"""UPDATE videos SET status = %s, error_message = %s,
last_scan_time = %s, updated_at = %s
WHERE id = %s""",
(status, file_info.get("error_message"), now, now, video_id),
transaction=True,
)
except Exception as e:
logger.error("Failed to update %s: %s", file_info["file_path"], e)
def _mark_files_removed(self, file_paths: List[str]):
"""Mark files as removed in DB."""
now = datetime.now(timezone.utc).replace(tzinfo=None)
for fp in file_paths:
try:
self.db.execute(
"UPDATE videos SET status = 'REMOVED', last_scan_time = %s, updated_at = %s WHERE file_path = %s",
(now, now, fp),
transaction=True,
)
except Exception as e:
logger.error("Failed to mark %s as removed: %s", fp, e)
def _record_scan_history(self, stats: dict, duration: float):
"""Record scan history in DB."""
now = datetime.now(timezone.utc).replace(tzinfo=None)
self.db.execute(
"""INSERT INTO scan_history
(scan_start, scan_end, files_discovered, files_new, files_modified,
files_removed, files_unscannable, duration_seconds, status)
VALUES (%s, %s, %s, %s, %s, %s, %s, %s, %s)""",
(
now,
now,
stats.get("discovered", 0),
stats.get("new", 0),
stats.get("modified", 0),
stats.get("removed", 0),
stats.get("unscannable", 0),
duration,
"COMPLETED",
),
transaction=True,
)
def get_stats(self) -> dict:
"""Get scanner statistics."""
return {
"scan_count": self._scan_count,
"total_files_discovered": self._total_files_discovered,
"total_files_new": self._total_files_new,
"total_files_modified": self._total_files_modified,
"total_files_unscannable": self._total_files_unscannable,
"instance_id": self._instance_id,
"scan_interval_seconds": self.scan_interval,
"lock_lease_seconds": self._lock_lease_seconds,
"scan_in_progress": self._scan_lock.locked(),
"input_path": str(self.input_path),
"is_running": self._running,
}
+297
View File
@@ -0,0 +1,297 @@
"""Standalone AI-processing pipeline for the task-worker architecture.
``process_video(video_dict, config)`` runs the full scan pipeline on a single
video file and returns a results dict that can be submitted via the REST API.
It reuses existing classes from the codebase:
* ``VideoProber`` ffprobe metadata extraction
* ``ScratchManager`` per-video temp dirs under /scratch
* ``FrameSampler`` uniform frame extraction via ffmpeg
* ``FaceDetector`` TensorRT face detection on sampled frames
* ``FaceClassifier`` MobileNetV3 classification of face crops
* ``aggregator.aggregate`` per-crop → video-level confidence
* ``router.route`` routing thresholds (MATCH / REVIEW / SKIP)
GPU model loading is lazily initialised and cached at module level so that
sequential invocations within the same Python process reuse the same loaded
engines. This is important because each short-lived worker may be asked to
process several tasks in a row (``--tasks N``).
"""
from __future__ import annotations
import logging
import os
import time
from pathlib import Path
from typing import Any, Dict, Optional
logger = logging.getLogger(__name__)
# ---------------------------------------------------------------------------
# Lazy imports — avoid heavy imports until actually needed
# ---------------------------------------------------------------------------
_face_detector_instance: Optional[Any] = None
_classifier_instance: Optional[Any] = None
def _get_face_detector(config: dict) -> Any:
"""Return a singleton ``FaceDetector`` (initialised once per process)."""
global _face_detector_instance
if _face_detector_instance is not None:
return _face_detector_instance
from face_detector import FaceDetector # noqa: local import, heavy
cfg = (config or {}).get("face_detection", {})
engine_path = cfg.get(
"model_path", "/models/face_detector/face_detector.trt"
)
_face_detector_instance = FaceDetector(
engine_path=engine_path,
input_size=int(cfg.get("input_size", 640)),
confidence_threshold=float(cfg.get("confidence_threshold", 0.25)),
iou_threshold=float(cfg.get("iou_threshold", 0.45)),
max_faces_per_frame=int(cfg.get("max_faces_per_frame", 10)),
max_faces_per_video=int(cfg.get("max_faces_per_video", 100)),
)
return _face_detector_instance
def _get_classifier(config: dict) -> Any:
"""Return a singleton ``FaceClassifier`` (initialised once per process)."""
global _classifier_instance
if _classifier_instance is not None:
return _classifier_instance
from classifier import FaceClassifier # noqa: local import, heavy
cfg = (config or {}).get("classifier", {})
engine_path = cfg.get(
"model_path", "/models/classifier/classifier.trt"
)
_classifier_instance = FaceClassifier(
engine_path=engine_path,
temperature=float(cfg.get("temperature", 1.0)),
input_size=int(cfg.get("input_size", 224)),
)
return _classifier_instance
# ---------------------------------------------------------------------------
# Public API
# ---------------------------------------------------------------------------
def get_model_version() -> str:
"""Return the currently active model version from the DB (best-effort)."""
try:
from db_connector import DatabaseConnector # noqa: local import, heavy
db = DatabaseConnector("default")
row = db.fetchone(
"SELECT version FROM models WHERE status = 'ACTIVE' LIMIT 1"
)
if row:
return row["version"]
except Exception as exc:
logger.debug("Could not fetch active model version: %s", exc)
return "v0.0.0-placeholder"
def process_video(video: Dict[str, Any], config: Optional[Dict] = None) -> Dict[str, Any]:
"""Run the full AI scan pipeline on a single video file.
Args:
video: Metadata dict (as returned by ``ApiClient.get_video``).
Must contain at minimum ``"id"`` and ``"file_path"``.
config: The loaded ``config.yaml`` dict (or None for defaults).
Returns:
A results dict suitable for submission via the REST API.
"""
config = config or {}
start_time = time.time()
video_id = video["id"]
file_path = video["file_path"]
logger.info("Processing video %d: %s", video_id, file_path)
sampling_cfg = (config or {}).get("sampling", {})
storage_cfg = (config or {}).get("storage", {})
gpu_cfg = (config or {}).get("gpu", {})
batching_cfg = (config or {}).get("batching", {})
agg_cfg = (config or {}).get("aggregation", {})
routing_cfg = (config or {}).get("routing", {})
try:
# 1. Probe metadata (use DB metadata if already populated, re-probe)
prober = VideoProber(timeout=10)
metadata = prober.probe(file_path)
if metadata.is_unscannable:
raise RuntimeError(metadata.error_message or "Video metadata could not be determined")
# Use video from API if available, fall back to probe results
resolution_w = video.get("resolution_w") or metadata.resolution_w or 640
resolution_h = video.get("resolution_h") or metadata.resolution_h or 480
duration = video.get("duration") or metadata.duration or 10.0
logger.info(
"Video %d: codec=%s resolution=%dx%d duration=%.1fs",
video_id, metadata.codec or "unknown",
resolution_w, resolution_h, duration,
)
# 2. Set up scratch space
scratch = ScratchManager(
base_path=storage_cfg.get("scratch_path", "/scratch"),
video_id=str(video_id),
auto_cleanup=True,
)
frame_dir = scratch.ensure_frame_dir()
# 3. Sample frames
interval = int(sampling_cfg.get("interval_seconds", 30))
quality = int(sampling_cfg.get("quality", 2))
sampler = FrameSampler(interval_seconds=interval, quality=quality)
resolution = (resolution_w, resolution_h)
extracted_frames = sampler.extract_frames(
video_path=file_path,
output_dir=str(frame_dir),
duration=duration,
interval_seconds=interval,
resolution=resolution,
timeout_seconds=30,
)
if not extracted_frames:
scratch.cleanup_all()
raise RuntimeError("No frames were extracted")
logger.info(
"Video %d: extracted %d frame(s)", video_id, len(extracted_frames)
)
# 4. Face detection
detector = _get_face_detector(config)
batch_size = _get_batch_size(gpu_cfg, batching_cfg)
detections_per_frame = detector.detect_faces(
extracted_frames, batch_size=batch_size
)
all_detections = sorted(
[d for dets in detections_per_frame for d in dets],
key=lambda d: d.confidence,
reverse=True,
)
total_faces = len(all_detections)
# 5. Classify (or SKIP if no faces)
if total_faces == 0:
logger.info(
"Video %d: no faces detected → routing to SKIP", video_id
)
scratch.cleanup_all()
return _failed_result(video_id, start_time, 0.0, "SKIP")
crop_dir = Path(scratch.frame_dir.parent) / "crops"
cropped_detections = detector.extract_crops(
all_detections, output_dir=str(crop_dir), crop_size=(224, 224)
)
crop_paths = [d.crop_path for d in cropped_detections if d.crop_path]
classifier = _get_classifier(config)
frame_confidences = classifier.classify(
crop_paths, batch_size=batch_size
)
# 6. Aggregate confidence
from aggregator import aggregate as agg_func # noqa: local import
video_confidence = agg_func(
frame_confidences,
strategy=agg_cfg.get("strategy", "max"),
alpha=float(agg_cfg.get("alpha", 1.0)),
beta=float(agg_cfg.get("beta", 0.1)),
top_k=int(agg_cfg.get("top_k", 3)),
)
# 7. Route
from router import route as route_fn # noqa: local import
routing = route_fn(
video_confidence,
t_high=float(routing_cfg.get("T_high", 0.75)),
t_low=float(routing_cfg.get("T_low", 0.45)),
)
processing_time = time.time() - start_time
logger.info(
"Video %d complete: C=%.4f routing=%s faces=%d frames=%d time=%.1fs",
video_id, video_confidence, routing, total_faces,
len(extracted_frames), processing_time,
)
scratch.cleanup_all()
return {
"status": "COMPLETED",
"confidence": round(video_confidence, 4),
"routing_decision": routing,
"face_count": total_faces,
"frame_count": len(extracted_frames),
"model_version": get_model_version(),
"processing_time_seconds": round(processing_time, 2),
"error": None,
}
except Exception as exc:
logger.error("Processing video %d failed: %s", video_id, exc, exc_info=True)
try:
scratch.cleanup_all() # noqa: undefined-name guard below
except NameError:
pass # scratch was never created (probe/probe error)
return {
"status": "FAILED",
"confidence": 0.0,
"routing_decision": "REVIEW",
"face_count": 0,
"frame_count": 0,
"model_version": get_model_version(),
"processing_time_seconds": round(time.time() - start_time, 2),
"error": str(exc),
}
# ---------------------------------------------------------------------------
# Helpers
# ---------------------------------------------------------------------------
def _get_batch_size(gpu_cfg: dict, batching_cfg: dict) -> int:
"""Compute an appropriate batch size based on GPU config."""
from gpu_manager import GPUMemoryManager # noqa: local import, heavy
mgr = GPUMemoryManager(
max_memory_gb=gpu_cfg.get("max_memory_gb", 18.0),
reduce_threshold_gb=batching_cfg.get("vram_reduce_threshold_gb", 16.0),
increase_threshold_gb=batching_cfg.get("vram_increase_threshold_gb", 10.0),
initial_batch_size=batching_cfg.get("max_batch_size", 16),
)
return mgr.current_batch_size
def _failed_result(
video_id: int, start_time: float, confidence: float, routing: str
) -> Dict[str, Any]:
"""Build a SKIP/early-exit result dict (no faces case)."""
return {
"status": "COMPLETED",
"confidence": round(confidence, 4),
"routing_decision": routing,
"face_count": 0,
"frame_count": 0,
"model_version": get_model_version(),
"processing_time_seconds": round(time.time() - start_time, 2),
"error": None,
}
+186
View File
@@ -0,0 +1,186 @@
"""VideoDetect task worker — short-lived process for AI scanning.
Usage::
python3 -m src.task_worker # process 1 task and exit
python3 -m src.task_worker --tasks 50 # process up to 50 tasks
python3 src/task_worker.py --tasks 10 # direct invocation
Environment variables:
API_BASE_URL (default http://localhost:3000)
Base URL of the Dancer2 REST API.
TASK_COUNT (default 1)
Maximum number of tasks to process before exiting.
On success exits with code 0; on unrecoverable errors exits non-zero.
"""
from __future__ import annotations
import argparse
import json
import logging
import os
import signal
import sys
from pathlib import Path
# Ensure the src/ directory is on sys.path so our submodules import cleanly
_src = Path(__file__).resolve().parent
if str(_src) not in sys.path:
sys.path.insert(0, str(_src))
from api_client import ApiClient, ApiError
from task_processor import process_video, get_model_version
from config_loader import get_config
from logging_config import setup_logging
logger = logging.getLogger(__name__)
# ---------------------------------------------------------------------------
# Graceful shutdown
# ---------------------------------------------------------------------------
_shutdown_requested = False
def _signal_handler(signum: int, frame) -> None:
logger.info("Received signal %d — finishing current task then exiting", signum)
global _shutdown_requested
_shutdown_requested = True
signal.signal(signal.SIGTERM, _signal_handler)
signal.signal(signal.SIGINT, _signal_handler)
# ---------------------------------------------------------------------------
# CLI entry point
# ---------------------------------------------------------------------------
def parse_args(argv: list[str] | None = None) -> argparse.Namespace:
parser = argparse.ArgumentParser(
description="VideoDetect task worker — process AI scan tasks",
)
parser.add_argument(
"--tasks",
type=int,
default=int(os.environ.get("TASK_COUNT", "1")),
help="Number of tasks to process before exiting (default: 1 or $TASK_COUNT)",
)
parser.add_argument(
"--api-url",
default=os.environ.get("API_BASE_URL", "http://localhost:3000"),
help="Base URL of the Dancer2 REST API (default: http://localhost:3000)",
)
return parser.parse_args(argv)
def main(argv: list[str] | None = None) -> int:
"""Run the task worker loop.
Returns exit code (0 = success).
"""
args = parse_args(argv)
# Load config for the processing pipeline
config = get_config()
# Setup JSON logging
log_cfg = (config or {}).get("logging", {})
setup_logging(
level=log_cfg.get("level", "INFO"),
log_format=log_cfg.get("format", "json"),
rotation_max_bytes=log_cfg.get("rotation_max_bytes", 104857600),
rotation_backup_count=log_cfg.get("rotation_backup_count", 10),
)
model_version = get_model_version()
logger.info(
"Task worker starting. Will process up to %d task(s). Model version: %s",
args.tasks,
model_version,
)
api = ApiClient(base_url=args.api_url)
total_attempted = 0
total_succeeded = 0
total_failed = 0
try:
for i in range(1, args.tasks + 1):
if _shutdown_requested:
logger.info("Shutdown requested after %d/%d tasks", i - 1, args.tasks)
break
logger.info("=== Processing task %d/%d ===", i, args.tasks)
# 1. Claim a task
response = api.get_next_task("AISCAN")
if response is None:
logger.info("No more tasks available (API returned no task)")
break
task_id = response["task"]["id"]
assign_key = response["assign_key"]
video_id = response["task"]["video_id"]
total_attempted += 1
logger.info(
"Claiming task %d (video_id=%d) for AISCAN", task_id, video_id
)
# 2. Fetch video metadata
try:
video = api.get_video(video_id)
except ApiError as exc:
logger.error("Failed to fetch video %d via API: %s", video_id, exc)
total_failed += 1
continue
# 3. Process the video (AI pipeline)
results = process_video(video, config=config)
if results["status"] == "COMPLETED":
logger.info(
"Video %d complete: C=%.4f routing=%s faces=%d frames=%d time=%.1fs",
video_id,
results.get("confidence", 0),
results.get("routing_decision", "?"),
results.get("face_count", 0),
results.get("frame_count", 0),
results.get("processing_time_seconds", 0),
)
total_succeeded += 1
else:
logger.warning(
"Video %d failed: %s",
video_id,
results.get("error", "unknown error"),
)
total_failed += 1
# 4. Submit results via API
try:
api.submit_results(task_id, assign_key, results)
logger.info("Submitted results for task %d via API", task_id)
except ApiError as exc:
logger.error(
"Failed to submit results for task %d: %s", task_id, exc
)
finally:
api.close()
logger.info(
"Task worker finished: %d attempted, %d succeeded, %d failed",
total_attempted,
total_succeeded,
total_failed,
)
return 0
if __name__ == "__main__":
sys.exit(main())
-67
View File
@@ -1,67 +0,0 @@
import shutil
import sys
import tempfile
import unittest
from pathlib import Path
from unittest.mock import patch
sys.path.insert(0, str(Path(__file__).resolve().parents[1] / "src"))
from frame_sampler import FrameSampler, calculate_timestamps
from scratch_manager import ScratchManager
class Story03SamplingTests(unittest.TestCase):
def test_calculate_timestamps_uses_uniform_temporal_spacing(self):
stamps = calculate_timestamps(duration=90.0, interval=30.0)
self.assertEqual(stamps, [0.0, 30.0, 60.0])
def test_frame_sampler_builds_ffmpeg_command_with_scale_and_jpeg_output(self):
sampler = FrameSampler(interval_seconds=30, quality=2, output_format="jpeg")
command = sampler._build_ffmpeg_command(
video_path="/tmp/video.mp4",
output_path="/tmp/out.jpg",
timestamp=12.5,
resolution=(3840, 2160),
)
self.assertIn("ffmpeg", command[0])
self.assertIn("-ss", command)
self.assertIn("-vframes", command)
self.assertIn("scale=1920:1080", " ".join(command))
self.assertTrue(command[-1].endswith("out.jpg"))
def test_frame_sampler_uses_subprocess_and_returns_output_path(self):
sampler = FrameSampler(interval_seconds=30, quality=2, output_format="jpeg")
with tempfile.TemporaryDirectory() as tmpdir:
video_path = Path(tmpdir) / "sample.mp4"
output_path = Path(tmpdir) / "frame.jpg"
video_path.write_bytes(b"fake")
def _mock_run(*args, **kwargs):
output_path.write_bytes(b"frame")
return type("Completed", (), {"returncode": 0, "stdout": b"", "stderr": b""})()
with patch("subprocess.run", side_effect=_mock_run):
result = sampler.extract_frame(
video_path=str(video_path),
output_path=str(output_path),
timestamp=15.0,
resolution=(1280, 720),
)
self.assertTrue(result)
self.assertEqual(output_path.name, Path(result).name)
def test_scratch_manager_cleans_up_frames_after_processing(self):
with tempfile.TemporaryDirectory() as tmpdir:
manager = ScratchManager(base_path=tmpdir, video_id="video-1", auto_cleanup=True)
frame_dir = manager.ensure_frame_dir()
(frame_dir / "video-1_1000.jpg").write_bytes(b"frame")
manager.cleanup()
self.assertFalse(frame_dir.exists())
if __name__ == "__main__":
unittest.main()
-98
View File
@@ -1,98 +0,0 @@
import sys
import tempfile
import unittest
from pathlib import Path
from unittest.mock import MagicMock, patch
import numpy as np
sys.path.insert(0, str(Path(__file__).resolve().parents[1] / "src"))
from batcher import DynamicBatcher
from face_detector import Detection, FaceDetector
from gpu_manager import GPUMemoryManager
class Story04FaceDetectionTests(unittest.TestCase):
def test_gpu_manager_respects_memory_thresholds(self):
manager = GPUMemoryManager(
max_memory_gb=18.0,
reduce_threshold_gb=16.0,
increase_threshold_gb=10.0,
initial_batch_size=16,
)
with patch.object(manager, "get_memory_stats", return_value=(17.0, 18.0)):
manager.adjust_batch_size()
self.assertLess(manager.current_batch_size, 16)
with patch.object(manager, "get_memory_stats", return_value=(8.0, 10.0)):
manager.adjust_batch_size()
self.assertGreater(manager.current_batch_size, manager.min_batch_size)
def test_dynamic_batcher_flushes_when_batch_full(self):
processed_batches = []
def process_fn(batch):
processed_batches.append(batch)
return [len(batch)]
batcher = DynamicBatcher(
process_fn=process_fn,
max_batch_size=3,
batch_timeout_ms=10000.0,
)
for i in range(5):
batcher.add(i)
self.assertEqual(len(processed_batches), 1)
self.assertEqual(processed_batches[0], [0, 1, 2])
self.assertEqual(batcher.queued_count, 2)
def test_face_detector_nms_removes_overlapping_boxes(self):
detector = FaceDetector(
engine_path="/nonexistent/model.trt",
confidence_threshold=0.1,
iou_threshold=0.45,
max_faces_per_frame=10,
)
duplicates = [
Detection("frame.jpg", 10, 10, 50, 50, 0.9),
Detection("frame.jpg", 12, 12, 48, 48, 0.8),
Detection("frame.jpg", 100, 100, 150, 150, 0.75),
]
kept = detector._nms(duplicates)
self.assertEqual(len(kept), 2)
self.assertAlmostEqual(kept[0].confidence, 0.9, places=5)
def test_face_detector_extracts_and_resizes_crops(self):
detector = FaceDetector(
engine_path="/nonexistent/model.trt",
input_size=640,
confidence_threshold=0.25,
)
with tempfile.TemporaryDirectory() as tmpdir:
from PIL import Image
frame_path = Path(tmpdir) / "frame.jpg"
image = Image.new("RGB", (640, 480), color=(100, 150, 200))
image.save(frame_path)
detections = [
Detection(str(frame_path), 0, 0, 640, 480, 0.95),
]
cropped = detector.extract_crops(detections, output_dir=tmpdir, crop_size=(224, 224))
self.assertEqual(len(cropped), 1)
self.assertTrue(Path(cropped[0].crop_path).exists())
with Image.open(cropped[0].crop_path) as crop:
self.assertEqual(crop.size, (224, 224))
if __name__ == "__main__":
unittest.main()
-83
View File
@@ -1,83 +0,0 @@
import sys
import unittest
from pathlib import Path
import numpy as np
sys.path.insert(0, str(Path(__file__).resolve().parents[1] / "src"))
import aggregator
import router
from classifier import FaceClassifier, calibrated_softmax
class Story05ClassifierTests(unittest.TestCase):
def test_calibrated_softmax_sums_to_one(self):
logits = np.array([[2.0, 1.0], [-1.0, 3.0]], dtype=np.float32)
probs = calibrated_softmax(logits, temperature=1.0)
np.testing.assert_allclose(probs.sum(axis=1), [1.0, 1.0], atol=1e-6)
def test_temperature_scaling_raises_lower_confidence_entropy(self):
logits = np.array([[2.0, 0.5]], dtype=np.float32)
sharp = calibrated_softmax(logits, temperature=0.5)
soft = calibrated_softmax(logits, temperature=2.0)
# higher temperature → softer distribution (target class prob moves toward 0.5)
self.assertGreater(sharp[0, 0], soft[0, 0])
def test_classifier_placeholder_returns_neutral_probability(self):
clf = FaceClassifier(engine_path="/nonexistent/model.trt", temperature=1.0)
probs = clf.classify([])
self.assertEqual(probs, [])
def test_classifier_placeholder_single_crop_returns_half(self):
import tempfile
from PIL import Image
clf = FaceClassifier(engine_path="/nonexistent/model.trt", temperature=1.0)
with tempfile.TemporaryDirectory() as tmpdir:
crop = Path(tmpdir) / "crop.jpg"
Image.new("RGB", (224, 224)).save(crop)
probs = clf.classify([str(crop)])
# placeholder logits are all zeros → softmax → 0.5 for each class
self.assertAlmostEqual(probs[0], 0.5, places=5)
class Story05AggregatorTests(unittest.TestCase):
def test_max_strategy(self):
self.assertAlmostEqual(aggregator.aggregate([0.3, 0.8, 0.6], strategy="max"), 0.8)
def test_empty_confidences_returns_zero(self):
self.assertEqual(aggregator.aggregate([], strategy="max"), 0.0)
def test_top_k_mean(self):
result = aggregator.aggregate([0.1, 0.9, 0.5, 0.8], strategy="top_k_mean", top_k=2)
self.assertAlmostEqual(result, (0.9 + 0.8) / 2, places=5)
def test_weighted_mean_clamps_to_unit_interval(self):
result = aggregator.aggregate([1.0, 1.0], strategy="weighted_mean", alpha=100.0, beta=0.0)
self.assertLessEqual(result, 1.0)
self.assertGreaterEqual(result, 0.0)
class Story05RouterTests(unittest.TestCase):
def test_match_at_high_threshold(self):
self.assertEqual(router.route(0.75), router.MATCH)
def test_review_between_thresholds(self):
self.assertEqual(router.route(0.60), router.REVIEW)
def test_skip_below_low_threshold(self):
self.assertEqual(router.route(0.44), router.SKIP)
def test_inclusive_high_threshold_boundary(self):
self.assertEqual(router.route(0.75, t_high=0.75, t_low=0.45), router.MATCH)
def test_inclusive_low_threshold_boundary(self):
self.assertEqual(router.route(0.45, t_high=0.75, t_low=0.45), router.REVIEW)
def test_no_faces_zero_confidence_routes_skip(self):
self.assertEqual(router.route(0.0), router.SKIP)
if __name__ == "__main__":
unittest.main()
-148
View File
@@ -1,148 +0,0 @@
import json
import sys
import tempfile
import unittest
from pathlib import Path
from unittest.mock import MagicMock, call, patch
sys.path.insert(0, str(Path(__file__).resolve().parents[1] / "src"))
import processing_logger
from data_export import DataExporter
from result_updater import ResultUpdater
from scratch_manager import ScratchManager
class Story06ProcessingLoggerTests(unittest.TestCase):
def test_insert_log_executes_correct_sql(self):
cursor = MagicMock()
processing_logger.insert_log(
cursor,
video_id=42,
model_version="v1.0.0",
frame_count=8,
confidence_score=0.82,
routing_decision="MATCH",
frame_confidences=[0.80, 0.82, 0.85],
)
cursor.execute.assert_called_once()
sql, params = cursor.execute.call_args[0]
self.assertIn("INSERT INTO processing_logs", sql)
self.assertEqual(params[0], 42) # video_id
self.assertEqual(params[1], "v1.0.0") # model_version
self.assertEqual(params[2], 8) # frame_count
self.assertAlmostEqual(params[3], 0.82) # confidence_score
scores = json.loads(params[4]) # confidence_scores JSON
self.assertEqual(scores, [0.80, 0.82, 0.85])
self.assertEqual(params[5], "MATCH") # routing_decision
def test_insert_log_null_frame_confidences(self):
cursor = MagicMock()
processing_logger.insert_log(cursor, 1, "v0", 0, None, "SKIP")
_, params = cursor.execute.call_args[0]
self.assertIsNone(params[4]) # confidence_scores column
class Story06ResultUpdaterTests(unittest.TestCase):
def _make_db(self, rowcount=1):
cursor = MagicMock()
cursor.rowcount = rowcount
conn = MagicMock()
conn.cursor.return_value = cursor
db = MagicMock()
db.transaction.return_value.__enter__ = MagicMock(return_value=conn)
db.transaction.return_value.__exit__ = MagicMock(return_value=False)
return db, cursor
def test_persist_returns_true_when_update_succeeds(self):
db, cursor = self._make_db(rowcount=1)
updater = ResultUpdater(db, model_version="v1.0.0")
result = updater.persist(
video_id=7, frame_count=5, confidence=0.9,
routing="MATCH", frame_confidences=[0.9],
)
self.assertTrue(result)
def test_persist_returns_false_on_state_guard_miss(self):
db, cursor = self._make_db(rowcount=0)
updater = ResultUpdater(db, model_version="v1.0.0")
result = updater.persist(
video_id=7, frame_count=5, confidence=0.9,
routing="MATCH", frame_confidences=[0.9],
)
self.assertFalse(result)
def test_persist_calls_insert_log_after_update(self):
db, cursor = self._make_db(rowcount=1)
updater = ResultUpdater(db, model_version="v1.0.0")
updater.persist(video_id=7, frame_count=5, confidence=0.9,
routing="MATCH", frame_confidences=[0.9])
# cursor.execute called twice: UPDATE videos + INSERT processing_logs
self.assertEqual(cursor.execute.call_count, 2)
class Story06DataExporterTests(unittest.TestCase):
def test_jsonl_flush_writes_valid_records(self):
with tempfile.TemporaryDirectory() as tmpdir:
exporter = DataExporter(
output_path=tmpdir, model_version="v1",
export_format="jsonl", batch_size=100,
)
exporter.add({"video_id": 1, "routing": "MATCH", "video_confidence": 0.9,
"confidence_scores": [0.9], "sample_count": 1,
"file_path": "/a.mp4", "processed_at": "2026-01-01T00:00:00+00:00"})
path = exporter.flush()
self.assertIsNotNone(path)
lines = Path(path).read_text(encoding="utf-8").strip().split("\n")
self.assertEqual(len(lines), 1)
record = json.loads(lines[0])
self.assertEqual(record["routing"], "MATCH")
self.assertAlmostEqual(record["video_confidence"], 0.9)
def test_auto_flush_at_batch_size(self):
with tempfile.TemporaryDirectory() as tmpdir:
exporter = DataExporter(
output_path=tmpdir, model_version="v1",
export_format="jsonl", batch_size=2,
)
exporter.add({"video_id": 1, "routing": "SKIP", "video_confidence": 0.1,
"confidence_scores": [], "sample_count": 0,
"file_path": "/a.mp4", "processed_at": ""})
exporter.add({"video_id": 2, "routing": "MATCH", "video_confidence": 0.9,
"confidence_scores": [0.9], "sample_count": 1,
"file_path": "/b.mp4", "processed_at": ""})
# batch_size=2 → auto-flush triggered on second add
self.assertEqual(exporter._buffer, [])
def test_exclude_frame_confidences_when_disabled(self):
with tempfile.TemporaryDirectory() as tmpdir:
exporter = DataExporter(
output_path=tmpdir, model_version="v1",
export_format="jsonl", batch_size=100,
include_frame_confidences=False,
)
exporter.add({"video_id": 1, "routing": "SKIP", "video_confidence": 0.1,
"confidence_scores": [0.1, 0.2], "sample_count": 2,
"file_path": "/a.mp4", "processed_at": ""})
path = exporter.flush()
record = json.loads(Path(path).read_text())
self.assertNotIn("confidence_scores", record)
class Story06ScratchManagerCleanupAllTests(unittest.TestCase):
def test_cleanup_all_removes_entire_video_directory(self):
with tempfile.TemporaryDirectory() as tmpdir:
manager = ScratchManager(base_path=tmpdir, video_id="v42", auto_cleanup=True)
frame_dir = manager.ensure_frame_dir()
crops_dir = Path(tmpdir) / "v42" / "crops"
crops_dir.mkdir(parents=True)
(frame_dir / "frame.jpg").write_bytes(b"f")
(crops_dir / "crop.jpg").write_bytes(b"c")
manager.cleanup_all()
self.assertFalse((Path(tmpdir) / "v42").exists())
if __name__ == "__main__":
unittest.main()
-95
View File
@@ -1,95 +0,0 @@
import csv
import io
import json
import sys
import tempfile
import unittest
from pathlib import Path
from unittest.mock import MagicMock
sys.path.insert(0, str(Path(__file__).resolve().parents[1] / "src"))
import review_export
def _make_db(rows):
db = MagicMock()
db.fetchall.return_value = rows
return db
class Story07ReviewExportTests(unittest.TestCase):
_ROWS = [
{
"video_id": 1, "file_path": "/data/input/a.mp4",
"confidence_score": 0.62, "routing_decision": "REVIEW",
"model_version": "v1.0", "ground_truth": True,
"annotated_at": None, "notes": "ok",
"confidence_scores": json.dumps([0.60, 0.62, 0.65]),
},
]
def test_fetch_annotated_passes_correct_where_clause(self):
db = _make_db(self._ROWS)
review_export.fetch_annotated(db, annotated_only=True, model_version="v1.0")
call_args = db.fetchall.call_args
sql = call_args[0][0]
self.assertIn("rq.annotated = TRUE", sql)
self.assertIn("v.model_version = %s", sql)
def test_fetch_annotated_normalises_confidence_scores(self):
db = _make_db(self._ROWS)
records = review_export.fetch_annotated(db)
self.assertIsInstance(records[0]["contributing_frames"], list)
self.assertEqual(records[0]["contributing_frames"], [0.60, 0.62, 0.65])
def test_fetch_annotated_handles_null_confidence_scores(self):
rows = [{**self._ROWS[0], "confidence_scores": None}]
db = _make_db(rows)
records = review_export.fetch_annotated(db)
self.assertEqual(records[0]["contributing_frames"], [])
def test_export_json_writes_valid_utf8_file(self):
with tempfile.TemporaryDirectory() as tmpdir:
records = review_export.fetch_annotated(_make_db(self._ROWS))
path = review_export.export_json(records, f"{tmpdir}/out.json")
loaded = json.loads(Path(path).read_text(encoding="utf-8"))
self.assertEqual(len(loaded), 1)
self.assertEqual(loaded[0]["video_id"], 1)
self.assertTrue(loaded[0]["ground_truth"])
def test_export_csv_writes_valid_csv(self):
with tempfile.TemporaryDirectory() as tmpdir:
records = review_export.fetch_annotated(_make_db(self._ROWS))
path = review_export.export_csv(records, f"{tmpdir}/out.csv")
content = Path(path).read_text(encoding="utf-8")
reader = csv.DictReader(io.StringIO(content))
rows = list(reader)
self.assertEqual(len(rows), 1)
self.assertEqual(rows[0]["routing_decision"], "REVIEW")
# contributing_frames should be a JSON string in CSV
frames = json.loads(rows[0]["contributing_frames"])
self.assertEqual(frames, [0.60, 0.62, 0.65])
def test_export_csv_empty_returns_empty_file(self):
with tempfile.TemporaryDirectory() as tmpdir:
path = review_export.export_csv([], f"{tmpdir}/empty.csv")
self.assertEqual(Path(path).read_text(), "")
def test_ground_truth_filter_appears_in_query(self):
db = _make_db([])
review_export.fetch_annotated(db, ground_truth=False)
sql = db.fetchall.call_args[0][0]
self.assertIn("rq.ground_truth = %s", sql)
class Story07AppSyntaxTest(unittest.TestCase):
def test_app_module_compiles(self):
"""Ensure ui/app.py has no syntax errors."""
app_path = Path(__file__).resolve().parents[1] / "ui" / "app.py"
source = app_path.read_text(encoding="utf-8")
compile(source, str(app_path), "exec")
if __name__ == "__main__":
unittest.main()
-550
View File
@@ -1,550 +0,0 @@
"""
Tests for Story 08: Active Learning Pipeline.
Covers:
- LabelIngestor: sample counting, dataset building, stratified split
- Trainer: model construction (backbone frozen), class weight calculation
- Validator: ECE calculation, quality gate logic
- ModelRegistry: promote/rollback DB calls
- ActiveLearningPipeline: threshold guard, full orchestration
"""
import csv
import json
import sys
import tempfile
import unittest
from pathlib import Path
from unittest.mock import MagicMock, patch, call
import numpy as np
sys.path.insert(0, str(Path(__file__).resolve().parents[1] / "src"))
from active_learning.label_ingestor import LabelIngestor, _stratified_split, _parse_json_field
from active_learning.validator import compute_ece, Validator
from active_learning.registry import ModelRegistry
from active_learning.pipeline import ActiveLearningPipeline
# ---------------------------------------------------------------------------
# Helpers
# ---------------------------------------------------------------------------
def _make_config(overrides: dict = None):
cfg = MagicMock()
al_defaults = {
"enabled": True,
"min_annotated_samples": 100,
"seed": 42,
"training": {
"epochs": 20, "batch_size": 32, "learning_rate": 1e-3,
"weight_decay": 1e-2, "early_stopping_patience": 5,
"lr_factor": 0.5, "lr_patience": 3,
},
"validation": {"val_split": 0.2, "min_f1_improvement": 0.02, "max_ece": 0.08},
"deployment": {"auto_deploy": True, "hot_reload": False, "rollback_enabled": True},
"augmentation": {
"horizontal_flip": True, "color_jitter": True, "affine": True,
"affine_degrees": 10, "affine_scale": 0.1,
},
}
if overrides:
al_defaults.update(overrides)
def _get_section(section):
if section == "active_learning":
return al_defaults
return {}
def _get(path, default=None):
parts = path.split(".")
if parts[0] == "active_learning" and len(parts) > 1:
key = parts[1]
return al_defaults.get(key, default)
mapping = {
"storage.training_path": "/tmp/videodetect_test_training",
"storage.models_path": "/tmp/videodetect_test_models",
"storage.scratch_path": "/tmp/scratch",
}
return mapping.get(path, default)
cfg.get_section.side_effect = _get_section
cfg.get.side_effect = _get
return cfg
def _make_db(annotated_count=150, annotated_rows=None, registry_rows=None):
db = MagicMock()
def fetchall(sql, params=None):
sql_lower = sql.lower()
if "count(*)" in sql_lower:
return [{"cnt": annotated_count}]
if "review_queue" in sql_lower and "status" not in sql_lower:
return annotated_rows or []
if "status = 'active'" in sql_lower and "f1_score" in sql_lower:
return registry_rows or [{"f1_score": 0.70}]
if "status = 'active'" in sql_lower:
return registry_rows or [{"version": "v1.0.0"}]
if "status = 'archived'" in sql_lower:
return [{"version": "v1.0.0"}]
return []
db.fetchall.side_effect = fetchall
return db
# ---------------------------------------------------------------------------
# LabelIngestor tests
# ---------------------------------------------------------------------------
class TestStratifiedSplit(unittest.TestCase):
def _make_records(self, n_pos, n_neg):
records = [{"label": 1, "crop_path": f"p{i}.jpg", "video_id": i} for i in range(n_pos)]
records += [{"label": 0, "crop_path": f"n{i}.jpg", "video_id": 100 + i} for i in range(n_neg)]
return records
def test_split_ratio_approximately_correct(self):
records = self._make_records(60, 40)
train, val = _stratified_split(records, 0.80, seed=42)
self.assertAlmostEqual(len(train) / len(records), 0.80, delta=0.05)
def test_stratification_preserves_class_balance(self):
records = self._make_records(50, 50)
train, val = _stratified_split(records, 0.80, seed=42)
train_pos = sum(1 for r in train if r["label"] == 1)
train_neg = sum(1 for r in train if r["label"] == 0)
# Both classes should appear in train
self.assertGreater(train_pos, 0)
self.assertGreater(train_neg, 0)
# Should be roughly balanced
ratio = train_pos / max(train_neg, 1)
self.assertAlmostEqual(ratio, 1.0, delta=0.3)
def test_split_is_deterministic(self):
records = self._make_records(40, 40)
train_a, _ = _stratified_split(records, 0.80, seed=7)
train_b, _ = _stratified_split(records, 0.80, seed=7)
self.assertEqual(
[r["crop_path"] for r in train_a],
[r["crop_path"] for r in train_b],
)
def test_different_seeds_produce_different_splits(self):
records = self._make_records(40, 40)
train_a, _ = _stratified_split(records, 0.80, seed=1)
train_b, _ = _stratified_split(records, 0.80, seed=999)
self.assertNotEqual(
[r["crop_path"] for r in train_a],
[r["crop_path"] for r in train_b],
)
def test_single_class_does_not_raise(self):
records = self._make_records(20, 0)
train, val = _stratified_split(records, 0.80, seed=42)
self.assertGreater(len(train), 0)
def test_no_sample_loss(self):
records = self._make_records(30, 20)
train, val = _stratified_split(records, 0.80, seed=42)
self.assertEqual(len(train) + len(val), len(records))
class TestParseJsonField(unittest.TestCase):
def test_parses_list(self):
self.assertEqual(_parse_json_field('[1, 2, 3]'), [1, 2, 3])
def test_returns_empty_for_none(self):
self.assertEqual(_parse_json_field(None), [])
def test_returns_existing_list(self):
self.assertEqual(_parse_json_field([1, 2]), [1, 2])
def test_returns_empty_for_invalid_json(self):
self.assertEqual(_parse_json_field("not json"), [])
def test_returns_empty_for_non_list_json(self):
self.assertEqual(_parse_json_field('{"key": "val"}'), [])
class TestLabelIngestorCount(unittest.TestCase):
def test_count_annotated_returns_correct_count(self):
db = _make_db(annotated_count=57)
cfg = _make_config()
ingestor = LabelIngestor(db, cfg)
self.assertEqual(ingestor.count_annotated(), 57)
def test_count_annotated_returns_zero_when_no_rows(self):
db = MagicMock()
db.fetchall.return_value = []
cfg = _make_config()
ingestor = LabelIngestor(db, cfg)
self.assertEqual(ingestor.count_annotated(), 0)
class TestLabelIngestorIngest(unittest.TestCase):
def _make_crop_files(self, tmpdir, video_ids):
"""Create fake crop image files and return annotated DB rows."""
scratch = Path(tmpdir) / "scratch"
scratch.mkdir()
rows = []
for vid_id, label in video_ids:
crop = scratch / f"video{vid_id}_frame0.jpg"
# Write a minimal valid JPEG header
crop.write_bytes(
b"\xff\xd8\xff\xe0\x00\x10JFIF\x00\x01\x01\x00\x00\x01\x00\x01\x00\x00"
b"\xff\xd9"
)
rows.append({
"video_id": vid_id,
"ground_truth": bool(label),
"confidence_scores": json.dumps([{"crop_path": str(crop)}]),
"file_path": f"/data/input/vid{vid_id}.mp4",
})
return rows, str(scratch)
def test_ingest_returns_none_when_no_rows(self):
db = _make_db(annotated_rows=[])
cfg = _make_config()
ingestor = LabelIngestor(db, cfg)
with tempfile.TemporaryDirectory() as tmpdir:
cfg.get.side_effect = lambda k, d=None: {
"storage.training_path": tmpdir,
"storage.scratch_path": tmpdir,
}.get(k, d)
result = ingestor.ingest("v2.0.0")
self.assertIsNone(result)
def test_ingest_creates_directory_structure(self):
with tempfile.TemporaryDirectory() as tmpdir:
rows, scratch = self._make_crop_files(tmpdir, [(1, True), (2, False), (3, True)])
db = _make_db(annotated_rows=rows)
cfg = _make_config()
cfg.get.side_effect = lambda k, d=None: {
"storage.training_path": str(Path(tmpdir) / "training"),
"storage.scratch_path": scratch,
}.get(k, d)
ingestor = LabelIngestor(db, cfg)
result = ingestor.ingest("v2.0.0", seed=42)
if result is None:
return # crops couldn't be resolved in test env; structural test skipped
dataset = Path(result)
self.assertTrue((dataset / "crops" / "class_0").is_dir())
self.assertTrue((dataset / "crops" / "class_1").is_dir())
self.assertTrue((dataset / "labels.csv").exists())
self.assertTrue((dataset / "metadata.json").exists())
def test_metadata_json_has_expected_keys(self):
with tempfile.TemporaryDirectory() as tmpdir:
rows, scratch = self._make_crop_files(tmpdir, [(1, True), (2, False), (3, True), (4, False)])
db = _make_db(annotated_rows=rows)
cfg = _make_config()
cfg.get.side_effect = lambda k, d=None: {
"storage.training_path": str(Path(tmpdir) / "training"),
"storage.scratch_path": scratch,
}.get(k, d)
ingestor = LabelIngestor(db, cfg)
result = ingestor.ingest("v2.1.0", seed=42)
if result is None:
return
meta = json.loads((Path(result) / "metadata.json").read_text())
self.assertIn("version", meta)
self.assertIn("total_samples", meta)
self.assertIn("train_samples", meta)
self.assertIn("val_samples", meta)
self.assertIn("class_counts", meta)
self.assertEqual(meta["version"], "v2.1.0")
# ---------------------------------------------------------------------------
# ECE / Validator tests
# ---------------------------------------------------------------------------
class TestComputeECE(unittest.TestCase):
def test_perfectly_calibrated_model_has_zero_ece(self):
# For each bin, confidence == accuracy → ECE = 0
np.random.seed(42)
n = 1000
probs = np.random.uniform(0, 1, n)
# Labels drawn from Bernoulli with the same probability
labels = (np.random.uniform(0, 1, n) < probs).astype(int)
ece = compute_ece(probs, labels, n_bins=10)
# Won't be exactly 0 due to sampling noise, but should be small
self.assertLess(ece, 0.10)
def test_overconfident_model_has_high_ece(self):
probs = np.ones(100) * 0.95
labels = np.zeros(100, dtype=int)
ece = compute_ece(probs, labels)
self.assertGreater(ece, 0.5)
def test_empty_predictions_returns_zero(self):
self.assertEqual(compute_ece(np.array([]), np.array([])), 0.0)
def test_ece_is_between_zero_and_one(self):
probs = np.random.default_rng(0).uniform(0, 1, 200)
labels = np.random.default_rng(0).integers(0, 2, 200)
ece = compute_ece(probs, labels)
self.assertGreaterEqual(ece, 0.0)
self.assertLessEqual(ece, 1.0)
def test_ece_bins_parameter(self):
probs = np.linspace(0, 1, 100)
labels = (probs > 0.5).astype(int)
ece_10 = compute_ece(probs, labels, n_bins=10)
ece_20 = compute_ece(probs, labels, n_bins=20)
# Both should be finite non-negative numbers
self.assertGreaterEqual(ece_10, 0.0)
self.assertGreaterEqual(ece_20, 0.0)
class TestValidatorQualityGates(unittest.TestCase):
def _make_validator(self, current_f1=0.70, min_delta=0.02, max_ece=0.08):
cfg = _make_config({
"validation": {
"val_split": 0.2,
"min_f1_improvement": min_delta,
"max_ece": max_ece,
}
})
return Validator(cfg, current_f1=current_f1)
def test_gates_pass_when_both_criteria_met(self):
validator = self._make_validator(current_f1=0.70)
# Simulate metrics
probs = np.array([0.9, 0.8, 0.1, 0.2, 0.85, 0.15, 0.75, 0.25])
labels = np.array([1, 1, 0, 0, 1, 0, 1, 0 ])
metrics = validator._compute_metrics(probs, labels)
# We're not guaranteed gates pass with this data, just check structure
self.assertIn("gates_passed", metrics)
self.assertIn("gate_details", metrics)
self.assertIn("f1", metrics)
self.assertIn("ece", metrics)
def test_gates_fail_when_f1_improvement_insufficient(self):
# current_f1=0.99 → perfect candidate (f1=1.0) only gives delta=0.01 < 0.02
validator = self._make_validator(current_f1=0.99, min_delta=0.02)
probs = np.array([0.9, 0.1, 0.8, 0.2])
labels = np.array([1, 0, 1, 0])
metrics = validator._compute_metrics(probs, labels)
details = metrics["gate_details"]
self.assertFalse(metrics["gates_passed"])
self.assertFalse(details["gate_f1_passed"])
def test_gates_fail_when_ece_too_high(self):
validator = self._make_validator(current_f1=0.0, min_delta=0.0, max_ece=0.01)
# Force high ECE: all confidence 0.9 but labels are 0
probs = np.ones(50) * 0.9
labels = np.zeros(50, dtype=int)
metrics = validator._compute_metrics(probs, labels)
self.assertFalse(metrics["gates_passed"])
self.assertFalse(metrics["gate_details"]["gate_ece_passed"])
def test_gate_details_include_delta_f1(self):
validator = self._make_validator(current_f1=0.60)
probs = np.array([0.8, 0.2, 0.7, 0.3])
labels = np.array([1, 0, 1, 0])
metrics = validator._compute_metrics(probs, labels)
self.assertIn("delta_f1", metrics["gate_details"])
self.assertAlmostEqual(
metrics["gate_details"]["delta_f1"],
metrics["f1"] - 0.60,
places=3,
)
# ---------------------------------------------------------------------------
# ModelRegistry tests
# ---------------------------------------------------------------------------
class TestModelRegistry(unittest.TestCase):
def _make_registry(self, auto_deploy=True, hot_reload=False):
cfg = _make_config({
"deployment": {
"auto_deploy": auto_deploy,
"hot_reload": hot_reload,
"rollback_enabled": True,
}
})
with tempfile.TemporaryDirectory() as tmpdir:
cfg.get.side_effect = lambda k, d=None: {
"storage.models_path": tmpdir,
}.get(k, d)
db = _make_db()
return ModelRegistry(db, cfg), db, tmpdir
def test_get_active_version_returns_version(self):
cfg = _make_config()
cfg.get.side_effect = lambda k, d=None: {
"storage.models_path": "/tmp/models",
}.get(k, d)
db = MagicMock()
db.fetchall.return_value = [{"version": "v1.5.0"}]
registry = ModelRegistry(db, cfg)
self.assertEqual(registry.get_active_version(), "v1.5.0")
def test_get_active_version_returns_none_when_no_active(self):
cfg = _make_config()
cfg.get.side_effect = lambda k, d=None: {
"storage.models_path": "/tmp/models",
}.get(k, d)
db = MagicMock()
db.fetchall.return_value = []
registry = ModelRegistry(db, cfg)
self.assertIsNone(registry.get_active_version())
def test_get_active_f1_returns_float(self):
cfg = _make_config()
cfg.get.side_effect = lambda k, d=None: {
"storage.models_path": "/tmp/models",
}.get(k, d)
db = MagicMock()
db.fetchall.return_value = [{"f1_score": 0.85}]
registry = ModelRegistry(db, cfg)
self.assertAlmostEqual(registry.get_active_f1(), 0.85)
def test_get_active_f1_returns_zero_when_none(self):
cfg = _make_config()
cfg.get.side_effect = lambda k, d=None: {
"storage.models_path": "/tmp/models",
}.get(k, d)
db = MagicMock()
db.fetchall.return_value = []
registry = ModelRegistry(db, cfg)
self.assertEqual(registry.get_active_f1(), 0.0)
def test_register_candidate_executes_upsert(self):
cfg = _make_config()
cfg.get.side_effect = lambda k, d=None: {
"storage.models_path": "/tmp/models",
}.get(k, d)
db = MagicMock()
registry = ModelRegistry(db, cfg)
registry.register_candidate("v2.0.0", "/models/candidate/v2.0.0_best.pt", 0.82, 0.05)
db.execute.assert_called_once()
args = db.execute.call_args[0]
self.assertIn("INSERT INTO models", args[0])
self.assertIn("v2.0.0", args[1])
def test_rollback_promotes_archived_model(self):
cfg = _make_config()
cfg.get.side_effect = lambda k, d=None: {
"storage.models_path": "/tmp/models",
}.get(k, d)
db = MagicMock()
db.fetchall.return_value = [{"version": "v1.0.0"}]
registry = ModelRegistry(db, cfg)
registry.rollback("v2.0.0")
calls = [str(c) for c in db.execute.call_args_list]
# Should archive the failed version and activate the previous
archive_call = any("ARCHIVED" in c and "v2.0.0" in c for c in calls)
activate_call = any("ACTIVE" in c and "v1.0.0" in c for c in calls)
self.assertTrue(archive_call, f"Expected ARCHIVED v2.0.0 in calls: {calls}")
self.assertTrue(activate_call, f"Expected ACTIVE v1.0.0 in calls: {calls}")
def test_rollback_logs_warning_when_no_archived_model(self):
cfg = _make_config()
cfg.get.side_effect = lambda k, d=None: {
"storage.models_path": "/tmp/models",
}.get(k, d)
db = MagicMock()
db.fetchall.return_value = []
registry = ModelRegistry(db, cfg)
# Should not raise
registry.rollback("v2.0.0")
db.execute.assert_not_called()
def test_auto_deploy_disabled_skips_deployment(self):
cfg = _make_config({
"deployment": {"auto_deploy": False, "hot_reload": False, "rollback_enabled": True}
})
cfg.get.side_effect = lambda k, d=None: {
"storage.models_path": "/tmp/models",
}.get(k, d)
db = MagicMock()
registry = ModelRegistry(db, cfg)
result = registry.deploy("v2.0.0", "/models/candidate/v2.0.0_best.pt")
self.assertFalse(result)
# ---------------------------------------------------------------------------
# ActiveLearningPipeline tests
# ---------------------------------------------------------------------------
class TestActiveLearningPipeline(unittest.TestCase):
def test_pipeline_skips_when_below_min_samples(self):
cfg = _make_config({"min_annotated_samples": 100})
db = _make_db(annotated_count=50)
pipeline = ActiveLearningPipeline(db, cfg)
result = pipeline.run("v2.0.0")
self.assertFalse(result)
def test_pipeline_runs_when_above_threshold(self):
cfg = _make_config({"min_annotated_samples": 100})
db = _make_db(annotated_count=150)
pipeline = ActiveLearningPipeline(db, cfg)
with patch.object(pipeline._ingestor, "ingest", return_value=None) as mock_ingest:
result = pipeline.run("v2.0.0")
mock_ingest.assert_called_once_with("v2.0.0", seed=42)
self.assertFalse(result) # ingestion returned None
def test_pipeline_aborts_when_training_fails(self):
cfg = _make_config({"min_annotated_samples": 10})
db = _make_db(annotated_count=50)
pipeline = ActiveLearningPipeline(db, cfg)
with patch.object(pipeline._ingestor, "ingest", return_value="/data/training/v2.0.0"), \
patch("active_learning.pipeline.Trainer") as MockTrainer:
MockTrainer.return_value.train.return_value = None
result = pipeline.run("v2.0.0")
self.assertFalse(result)
def test_pipeline_does_not_deploy_when_gates_fail(self):
cfg = _make_config({"min_annotated_samples": 10})
db = _make_db(annotated_count=50)
pipeline = ActiveLearningPipeline(db, cfg)
with patch.object(pipeline._ingestor, "ingest", return_value="/data/training/v2.0.0"), \
patch("active_learning.pipeline.Trainer") as MockTrainer, \
patch("active_learning.pipeline.Validator") as MockValidator:
MockTrainer.return_value.train.return_value = "/models/candidate/v2.0.0_best.pt"
MockValidator.return_value.validate.return_value = {
"f1": 0.71, "ece": 0.05,
"gates_passed": False,
"gate_details": {"delta_f1": 0.01},
}
result = pipeline.run("v2.0.0")
self.assertFalse(result)
def test_pipeline_deploys_when_gates_pass(self):
cfg = _make_config({"min_annotated_samples": 10})
db = _make_db(annotated_count=50)
pipeline = ActiveLearningPipeline(db, cfg)
with patch.object(pipeline._ingestor, "ingest", return_value="/data/training/v2.0.0"), \
patch("active_learning.pipeline.Trainer") as MockTrainer, \
patch("active_learning.pipeline.Validator") as MockValidator, \
patch.object(pipeline._registry, "deploy", return_value=True) as mock_deploy:
MockTrainer.return_value.train.return_value = "/models/candidate/v2.0.0_best.pt"
MockValidator.return_value.validate.return_value = {
"f1": 0.88, "ece": 0.04,
"gates_passed": True,
"gate_details": {"delta_f1": 0.18},
}
result = pipeline.run("v2.0.0")
self.assertTrue(result)
mock_deploy.assert_called_once_with("v2.0.0", "/models/candidate/v2.0.0_best.pt")
if __name__ == "__main__":
unittest.main()
-521
View File
@@ -1,521 +0,0 @@
"""
Tests for Story 09: Observability, Monitoring & Hardening.
Covers:
- metrics.py: NoOp fallback, update_queue_depths, update_scratch_metrics
- crash_recovery.py: recover_on_startup, checkpointing, idempotency guard
- retry.py: successful call, retry on transient error, non-retryable bypass, exhaustion
- drift_detector.py: detect_drift logic, all alert checks
- health_check.py: HealthStatus snapshot, HTTP /health endpoint
"""
import json
import os
import sys
import tempfile
import threading
import time
import unittest
import urllib.request
from pathlib import Path
from unittest.mock import MagicMock, call, patch
sys.path.insert(0, str(Path(__file__).resolve().parents[1] / "src"))
from crash_recovery import CrashRecovery
from drift_detector import DriftDetector, detect_drift
from health_check import HealthStatus, start_health_server, health
from retry import RetryExhaustedError, retry, _is_non_retryable
# ---------------------------------------------------------------------------
# Helpers
# ---------------------------------------------------------------------------
def _make_config(overrides: dict = None):
cfg = MagicMock()
mon_defaults = {
"alerts": {
"confidence_drift_threshold": 0.10,
"review_queue_max_size": 1000,
"review_queue_max_age_hours": 24,
"throughput_min_videos_per_hour": 20,
"throughput_min_duration_hours": 1,
"error_rate_threshold": 0.05,
"error_rate_window_hours": 1,
},
"drift_detection": {"enabled": True, "schedule": "0 2 * * 0", "baseline_source": "db"},
"crash_recovery": {"lock_timeout_minutes": 5, "auto_requeue": True},
"retry": {"max_attempts": 3, "initial_delay": 0.0, "backoff_factor": 2.0},
}
if overrides:
mon_defaults.update(overrides)
def _get_section(section):
if section == "monitoring":
return mon_defaults
return {}
def _get(path, default=None):
mapping = {
"storage.scratch_path": "/tmp/test_scratch",
"storage.models_path": "/tmp/test_models",
}
return mapping.get(path, default)
cfg.get_section.side_effect = _get_section
cfg.get.side_effect = _get
return cfg
def _make_db():
return MagicMock()
# ---------------------------------------------------------------------------
# metrics.py tests
# ---------------------------------------------------------------------------
class TestMetricsNoOp(unittest.TestCase):
"""The _NoOpMetric must absorb all method calls without raising."""
def test_noop_labels_inc_does_not_raise(self):
from metrics import _NoOpMetric
m = _NoOpMetric()
m.labels(routing_decision="MATCH").inc()
def test_noop_set_does_not_raise(self):
from metrics import _NoOpMetric
m = _NoOpMetric()
m.set(42)
def test_noop_observe_does_not_raise(self):
from metrics import _NoOpMetric
m = _NoOpMetric()
m.observe(0.75)
class TestMetricsQueueDepths(unittest.TestCase):
def test_update_queue_depths_sets_gauges(self):
from metrics import update_queue_depths, queue_depth_pending, queue_depth_processing
db = MagicMock()
db.fetchall.side_effect = [
[{"status": "PENDING", "cnt": 10}, {"status": "PROCESSING", "cnt": 3}],
[{"cnt": 7}],
]
# Should not raise even if prometheus is absent
update_queue_depths(db)
def test_update_queue_depths_handles_db_error_gracefully(self):
from metrics import update_queue_depths
db = MagicMock()
db.fetchall.side_effect = Exception("DB down")
update_queue_depths(db) # must not raise
class TestMetricsScratch(unittest.TestCase):
def test_update_scratch_metrics_runs_without_error(self):
from metrics import update_scratch_metrics
with tempfile.TemporaryDirectory() as tmpdir:
update_scratch_metrics(tmpdir) # must not raise
def test_update_scratch_metrics_handles_missing_path(self):
from metrics import update_scratch_metrics
update_scratch_metrics("/nonexistent_path_xyz") # must not raise
# ---------------------------------------------------------------------------
# crash_recovery.py tests
# ---------------------------------------------------------------------------
class TestCrashRecoveryRequeue(unittest.TestCase):
def test_recover_on_startup_calls_update(self):
db = _make_db()
db.execute.return_value = 3
cfg = _make_config()
cr = CrashRecovery(db, cfg)
count = cr.recover_on_startup()
self.assertEqual(count, 3)
db.execute.assert_called_once()
sql = db.execute.call_args[0][0]
self.assertIn("PENDING", sql)
self.assertIn("PROCESSING", sql)
def test_recover_on_startup_skipped_when_disabled(self):
db = _make_db()
cfg = _make_config({"crash_recovery": {"lock_timeout_minutes": 5, "auto_requeue": False}})
cr = CrashRecovery(db, cfg)
count = cr.recover_on_startup()
self.assertEqual(count, 0)
db.execute.assert_not_called()
def test_recover_on_startup_handles_db_error(self):
db = _make_db()
db.execute.side_effect = Exception("connection refused")
cfg = _make_config()
cr = CrashRecovery(db, cfg)
count = cr.recover_on_startup()
self.assertEqual(count, 0)
def test_list_stuck_videos_returns_rows(self):
db = _make_db()
db.fetchall.return_value = [{"id": 5, "file_path": "/data/vid.mp4", "updated_at": None}]
cfg = _make_config()
cr = CrashRecovery(db, cfg)
rows = cr.list_stuck_videos()
self.assertEqual(len(rows), 1)
self.assertEqual(rows[0]["id"], 5)
class TestCrashRecoveryCheckpoint(unittest.TestCase):
def test_save_and_load_checkpoint_roundtrip(self):
with tempfile.TemporaryDirectory() as tmpdir:
cfg = _make_config()
cfg.get.side_effect = lambda k, d=None: {
"storage.scratch_path": tmpdir,
}.get(k, d)
cr = CrashRecovery(_make_db(), cfg)
cr.save_checkpoint(42, "extracting", {"frames_done": 5})
result = cr.load_checkpoint(42)
self.assertIsNotNone(result)
self.assertEqual(result["video_id"], 42)
self.assertEqual(result["state"], "extracting")
self.assertEqual(result["progress"]["frames_done"], 5)
def test_load_checkpoint_returns_none_when_absent(self):
with tempfile.TemporaryDirectory() as tmpdir:
cfg = _make_config()
cfg.get.side_effect = lambda k, d=None: {
"storage.scratch_path": tmpdir,
}.get(k, d)
cr = CrashRecovery(_make_db(), cfg)
self.assertIsNone(cr.load_checkpoint(999))
def test_delete_checkpoint_removes_file(self):
with tempfile.TemporaryDirectory() as tmpdir:
cfg = _make_config()
cfg.get.side_effect = lambda k, d=None: {
"storage.scratch_path": tmpdir,
}.get(k, d)
cr = CrashRecovery(_make_db(), cfg)
cr.save_checkpoint(7, "classifying", {})
cr.delete_checkpoint(7)
self.assertIsNone(cr.load_checkpoint(7))
def test_checkpoint_timestamp_is_iso_format(self):
with tempfile.TemporaryDirectory() as tmpdir:
cfg = _make_config()
cfg.get.side_effect = lambda k, d=None: {
"storage.scratch_path": tmpdir,
}.get(k, d)
cr = CrashRecovery(_make_db(), cfg)
cr.save_checkpoint(1, "detecting", {})
ckpt = cr.load_checkpoint(1)
# Should parse without error
from datetime import datetime
datetime.fromisoformat(ckpt["timestamp"].replace("Z", "+00:00"))
class TestIdempotencyGuard(unittest.TestCase):
def test_returns_true_when_already_completed(self):
db = _make_db()
db.fetchall.return_value = [{"status": "COMPLETED"}]
cr = CrashRecovery(db, _make_config())
self.assertTrue(cr.is_already_completed(1))
def test_returns_false_when_not_completed(self):
db = _make_db()
db.fetchall.return_value = [{"status": "PENDING"}]
cr = CrashRecovery(db, _make_config())
self.assertFalse(cr.is_already_completed(1))
def test_returns_false_when_no_row(self):
db = _make_db()
db.fetchall.return_value = []
cr = CrashRecovery(db, _make_config())
self.assertFalse(cr.is_already_completed(99))
# ---------------------------------------------------------------------------
# retry.py tests
# ---------------------------------------------------------------------------
class TestRetryDecorator(unittest.TestCase):
def test_successful_call_returns_value(self):
@retry(max_attempts=3, initial_delay=0.0, step="test")
def always_succeeds():
return 42
self.assertEqual(always_succeeds(), 42)
def test_retries_on_transient_error(self):
call_count = {"n": 0}
@retry(max_attempts=3, initial_delay=0.0, step="test")
def flaky():
call_count["n"] += 1
if call_count["n"] < 3:
raise ConnectionError("transient")
return "ok"
result = flaky()
self.assertEqual(result, "ok")
self.assertEqual(call_count["n"], 3)
def test_raises_retry_exhausted_after_max_attempts(self):
@retry(max_attempts=3, initial_delay=0.0, step="test")
def always_fails():
raise ConnectionError("always fails")
with self.assertRaises(RetryExhaustedError):
always_fails()
def test_non_retryable_error_propagates_immediately(self):
call_count = {"n": 0}
@retry(max_attempts=3, initial_delay=0.0, step="test")
def raises_non_retryable():
call_count["n"] += 1
raise FileNotFoundError("no such file")
with self.assertRaises(FileNotFoundError):
raises_non_retryable()
self.assertEqual(call_count["n"], 1)
def test_only_specified_exception_types_are_retried(self):
@retry(max_attempts=3, initial_delay=0.0, exceptions=(ValueError,), step="test")
def raises_type_error():
raise TypeError("wrong type")
with self.assertRaises(TypeError):
raises_type_error()
def test_preserves_return_value_on_first_try(self):
@retry(max_attempts=5, initial_delay=0.0, step="test")
def returns_dict():
return {"key": "value"}
self.assertEqual(returns_dict(), {"key": "value"})
class TestIsNonRetryable(unittest.TestCase):
def test_file_not_found_is_non_retryable(self):
self.assertTrue(_is_non_retryable(FileNotFoundError("x")))
def test_permission_error_is_non_retryable(self):
self.assertTrue(_is_non_retryable(PermissionError("x")))
def test_connection_error_is_retryable(self):
self.assertFalse(_is_non_retryable(ConnectionError("x")))
def test_runtime_error_is_retryable(self):
self.assertFalse(_is_non_retryable(RuntimeError("x")))
# ---------------------------------------------------------------------------
# drift_detector.py tests
# ---------------------------------------------------------------------------
class TestDetectDrift(unittest.TestCase):
def test_no_drift_when_distributions_match(self):
base = [0.3] * 50 + [0.7] * 50 # 50% high confidence
curr = [0.3] * 50 + [0.7] * 50
self.assertFalse(detect_drift(curr, base, threshold=0.10))
def test_drift_detected_when_shift_exceeds_threshold(self):
base = [0.3] * 80 + [0.8] * 20 # 20% high
curr = [0.8] * 70 + [0.3] * 30 # 70% high → shift = 0.50
self.assertTrue(detect_drift(curr, base, threshold=0.10))
def test_no_drift_just_below_threshold(self):
base = [0.8] * 50 + [0.2] * 50 # 50% high
curr = [0.8] * 59 + [0.2] * 41 # 59% high → shift = 9%
self.assertFalse(detect_drift(curr, base, threshold=0.10))
def test_drift_at_boundary(self):
base = [0.8] * 50 + [0.2] * 50 # 50%
curr = [0.8] * 61 + [0.2] * 39 # 61% → shift = 11%
self.assertTrue(detect_drift(curr, base, threshold=0.10))
def test_empty_current_returns_false(self):
self.assertFalse(detect_drift([], [0.5] * 10, threshold=0.10))
def test_empty_baseline_returns_false(self):
self.assertFalse(detect_drift([0.5] * 10, [], threshold=0.10))
class TestDriftDetectorAlerts(unittest.TestCase):
def _make_detector(self, db=None, overrides=None):
alerts = []
cfg = _make_config(overrides or {})
db = db or _make_db()
detector = DriftDetector(db, cfg, alert_fn=alerts.append)
return detector, alerts
def test_check_review_queue_growth_triggers_alert(self):
db = _make_db()
db.fetchall.return_value = [{"cnt": 1500}]
detector, alerts = self._make_detector(db)
triggered = detector.check_review_queue_growth()
self.assertTrue(triggered)
self.assertEqual(len(alerts), 1)
self.assertIn("1500", alerts[0])
def test_check_review_queue_growth_no_alert_below_threshold(self):
db = _make_db()
db.fetchall.return_value = [{"cnt": 50}]
detector, alerts = self._make_detector(db)
triggered = detector.check_review_queue_growth()
self.assertFalse(triggered)
self.assertEqual(len(alerts), 0)
def test_check_low_throughput_triggers_alert(self):
db = _make_db()
db.fetchall.return_value = [{"cnt": 5}] # 5 videos in last 1h < 20 min
detector, alerts = self._make_detector(db)
triggered = detector.check_low_throughput()
self.assertTrue(triggered)
self.assertEqual(len(alerts), 1)
def test_check_low_throughput_no_alert_above_threshold(self):
db = _make_db()
db.fetchall.return_value = [{"cnt": 50}] # 50 > 20
detector, alerts = self._make_detector(db)
triggered = detector.check_low_throughput()
self.assertFalse(triggered)
def test_check_error_rate_triggers_alert(self):
db = _make_db()
db.fetchall.return_value = [
{"status": "COMPLETED", "cnt": 80},
{"status": "ERROR", "cnt": 10},
{"status": "UNSCANNABLE", "cnt": 10},
]
detector, alerts = self._make_detector(db)
triggered = detector.check_error_rate()
self.assertTrue(triggered) # 20/100 = 20% > 5%
self.assertEqual(len(alerts), 1)
def test_check_error_rate_no_alert_below_threshold(self):
db = _make_db()
db.fetchall.return_value = [
{"status": "COMPLETED", "cnt": 98},
{"status": "ERROR", "cnt": 2},
]
detector, alerts = self._make_detector(db)
triggered = detector.check_error_rate()
self.assertFalse(triggered) # 2% < 5%
def test_check_error_rate_no_alert_zero_videos(self):
db = _make_db()
db.fetchall.return_value = []
detector, alerts = self._make_detector(db)
triggered = detector.check_error_rate()
self.assertFalse(triggered)
def test_run_all_checks_returns_dict_with_expected_keys(self):
db = _make_db()
db.fetchall.return_value = [{"cnt": 0}]
detector, _ = self._make_detector(db)
with patch.object(detector, "_fetch_recent_confidences", return_value=[]):
results = detector.run_all_checks()
self.assertIn("confidence_drift", results)
self.assertIn("review_queue_growth", results)
self.assertIn("low_throughput", results)
self.assertIn("high_error_rate", results)
def test_check_handles_db_error_gracefully(self):
db = _make_db()
db.fetchall.side_effect = Exception("DB offline")
detector, alerts = self._make_detector(db)
# Should not raise
self.assertFalse(detector.check_review_queue_growth())
self.assertFalse(detector.check_low_throughput())
self.assertFalse(detector.check_error_rate())
# ---------------------------------------------------------------------------
# health_check.py tests
# ---------------------------------------------------------------------------
class TestHealthStatus(unittest.TestCase):
def test_snapshot_contains_required_keys(self):
hs = HealthStatus()
snap = hs.snapshot()
for key in ("status", "gpu_available", "gpu_memory_used_gb",
"queue_depth", "uptime_seconds", "videos_processed_today", "last_error"):
self.assertIn(key, snap, f"Missing key: {key}")
def test_update_changes_values(self):
hs = HealthStatus()
hs.update(status="healthy", queue_depth=55)
snap = hs.snapshot()
self.assertEqual(snap["status"], "healthy")
self.assertEqual(snap["queue_depth"], 55)
def test_uptime_increases_over_time(self):
hs = HealthStatus()
snap1 = hs.snapshot()
time.sleep(0.05)
snap2 = hs.snapshot()
self.assertGreaterEqual(snap2["uptime_seconds"], snap1["uptime_seconds"])
def test_update_is_thread_safe(self):
hs = HealthStatus()
errors = []
def writer(n):
try:
for _ in range(100):
hs.update(queue_depth=n)
except Exception as exc:
errors.append(exc)
threads = [threading.Thread(target=writer, args=(i,)) for i in range(5)]
for t in threads:
t.start()
for t in threads:
t.join()
self.assertEqual(errors, [])
def test_http_health_endpoint_returns_200(self):
"""Start a real health server and hit /health with urllib."""
import socket
# Find a free port
with socket.socket() as s:
s.bind(("127.0.0.1", 0))
port = s.getsockname()[1]
t = start_health_server(port)
time.sleep(0.1) # give the server a moment to bind
url = f"http://127.0.0.1:{port}/health"
with urllib.request.urlopen(url, timeout=2) as resp:
self.assertEqual(resp.status, 200)
body = json.loads(resp.read())
self.assertIn("status", body)
self.assertIn("uptime_seconds", body)
def test_http_404_for_unknown_path(self):
"""Non /health paths return 404."""
import socket
from urllib.error import HTTPError
with socket.socket() as s:
s.bind(("127.0.0.1", 0))
port = s.getsockname()[1]
start_health_server(port)
time.sleep(0.1)
with self.assertRaises(HTTPError) as ctx:
urllib.request.urlopen(f"http://127.0.0.1:{port}/unknown", timeout=2)
self.assertEqual(ctx.exception.code, 404)
if __name__ == "__main__":
unittest.main()
+2 -2
View File
@@ -64,5 +64,5 @@ HEALTHCHECK --interval=30s --timeout=10s --retries=3 \
# Switch to non-root user # Switch to non-root user
USER appuser USER appuser
# Default command # Default command: process tasks from the API until drained
CMD ["python3", "-m", "src.main"] CMD ["python3", "-m", "src.task_worker"]
+40
View File
@@ -0,0 +1,40 @@
# VideoDetect — Worker service (standalone compose file)
# Run: docker compose -f worker/docker-compose.yml up -d
services:
worker:
build:
context: ./
dockerfile: Dockerfile
container_name: videodetect-worker
restart: "no"
environment:
- API_BASE_URL=${API_BASE_URL:-http://videodetect-api:3000}
- TASK_COUNT=${TASK_COUNT:-50}
- REVIEW_CONFIDENCE=${REVIEW_CONFIDENCE:-0.75}
- LOG_FILE=/logs/videodetect.log
volumes:
- nas_input:/data/input:ro
- scratch_data:/scratch
- $PWD/worker-logs:/logs
networks:
- videodetect_videodetect-network
deploy:
resources:
limits:
memory: 4G
nvidia-gpus: "1"
volumes:
nas_input:
driver: local
driver_opts:
type: nfs
o: addr=10.0.0.2,ro,nfsvers=4,hard,intr
device: ":/mnt/Bulk/Homes/ryan/Prawns"
scratch_data:
driver: local
networks:
videodetect_videodetect-network:
external: true