Compare commits
3
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
458d8862cc | ||
|
|
9f098d4b1f | ||
|
|
d9c0998400 |
+486
@@ -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
@@ -130,8 +130,22 @@ get '/api/v1/nexttask/:type' => sub {
|
||||
|
||||
post '/api/v1/task/:task/complete' => sub {
|
||||
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
|
||||
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);
|
||||
}
|
||||
|
||||
# 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=?");
|
||||
$update_sth->execute($results, $task);
|
||||
$update_sth->execute($results_json, $task);
|
||||
if($update_sth->rows == 0) {
|
||||
send_error("Failed to complete task", 500);
|
||||
}
|
||||
@@ -154,4 +169,18 @@ post '/api/v1/task/:task/complete' => sub {
|
||||
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();
|
||||
+3
-3
@@ -7,7 +7,7 @@ services:
|
||||
MYSQL_ROOT_PASSWORD: ${DB_ROOT_PASSWORD:-rootpass}
|
||||
MYSQL_DATABASE: ${DB_NAME:-videodetect}
|
||||
MYSQL_USER: ${DB_USER:-videodetect}
|
||||
MYSQL_PASSWORD: ${DB_PASSWORD:-videodetect123}
|
||||
MYSQL_PASSWORD: ${DB_PASSWORD:-changeme_videodetect}
|
||||
ports:
|
||||
- "3306:3306"
|
||||
volumes:
|
||||
@@ -38,7 +38,7 @@ services:
|
||||
- DB_PORT=3306
|
||||
- DB_NAME=videodetect
|
||||
- DB_USER=videodetect
|
||||
- DB_PASSWORD=${DB_PASSWORD:-videodetect123}
|
||||
- DB_PASSWORD=${DB_PASSWORD:-changeme_videodetect}
|
||||
- FLASK_ENV=production
|
||||
volumes:
|
||||
- ./ui:/app
|
||||
@@ -67,7 +67,7 @@ services:
|
||||
- DB_PORT=3306
|
||||
- DB_NAME=videodetect
|
||||
- DB_USER=videodetect
|
||||
- DB_PASSWORD=${DB_PASSWORD:-videodetect123}
|
||||
- DB_PASSWORD=${DB_PASSWORD:-changeme_videodetect}
|
||||
- DANCER_ENVIRONMENT=production
|
||||
volumes:
|
||||
- ./api:/app
|
||||
|
||||
@@ -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()
|
||||
@@ -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)
|
||||
@@ -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
@@ -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()
|
||||
@@ -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(),
|
||||
}
|
||||
@@ -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)
|
||||
@@ -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
@@ -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
|
||||
@@ -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
@@ -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,
|
||||
}
|
||||
@@ -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,
|
||||
}
|
||||
@@ -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())
|
||||
@@ -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()
|
||||
@@ -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()
|
||||
@@ -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()
|
||||
@@ -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()
|
||||
@@ -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()
|
||||
@@ -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()
|
||||
@@ -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
@@ -64,5 +64,5 @@ HEALTHCHECK --interval=30s --timeout=10s --retries=3 \
|
||||
# Switch to non-root user
|
||||
USER appuser
|
||||
|
||||
# Default command
|
||||
CMD ["python3", "-m", "src.main"]
|
||||
# Default command: process tasks from the API until drained
|
||||
CMD ["python3", "-m", "src.task_worker"]
|
||||
|
||||
@@ -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
|
||||
Reference in New Issue
Block a user