140 lines
5.6 KiB
Python
140 lines
5.6 KiB
Python
|
|
"""实时 YOLO 预览会话(进程内、短轮询、只保留最新结果)。"""
|
||
|
|
import json
|
||
|
|
import threading
|
||
|
|
import time
|
||
|
|
import uuid
|
||
|
|
|
||
|
|
from app.analysis.event_bridge import get_event_bridge
|
||
|
|
from app.analysis.manager import AnalysisManager
|
||
|
|
|
||
|
|
|
||
|
|
class PreviewSessionRegistry(object):
|
||
|
|
TTL_SEC = 15.0
|
||
|
|
|
||
|
|
def __init__(self):
|
||
|
|
self._sessions = {}
|
||
|
|
self._lock = threading.RLock()
|
||
|
|
self._running = True
|
||
|
|
threading.Thread(target=self._cleanup_loop, name="preview-session-cleanup", daemon=True).start()
|
||
|
|
|
||
|
|
@staticmethod
|
||
|
|
def _zone_payload(zone):
|
||
|
|
try:
|
||
|
|
coords = json.loads(zone.coordinates or "[]")
|
||
|
|
except Exception:
|
||
|
|
coords = []
|
||
|
|
return {"id": zone.id, "name": zone.name, "color": zone.color,
|
||
|
|
"coords": coords if isinstance(coords, list) else []}
|
||
|
|
|
||
|
|
@staticmethod
|
||
|
|
def _point_in_polygon(point, polygon):
|
||
|
|
if len(polygon or []) < 3:
|
||
|
|
return False
|
||
|
|
x, y, inside, j = point[0], point[1], False, len(polygon) - 1
|
||
|
|
for i, current in enumerate(polygon):
|
||
|
|
previous = polygon[j]
|
||
|
|
if ((current[1] > y) != (previous[1] > y) and
|
||
|
|
x < (previous[0] - current[0]) * (y - current[1]) /
|
||
|
|
(previous[1] - current[1] + 1e-12) + current[0]):
|
||
|
|
inside = not inside
|
||
|
|
j = i
|
||
|
|
return inside
|
||
|
|
|
||
|
|
def start(self, owner, zone):
|
||
|
|
ok, mode = AnalysisManager().start_preview(zone)
|
||
|
|
if not ok:
|
||
|
|
raise RuntimeError(mode)
|
||
|
|
session_id = uuid.uuid4().hex
|
||
|
|
now = time.time()
|
||
|
|
algorithms = []
|
||
|
|
target_labels = set()
|
||
|
|
for ba in zone.algorithms.filter(state=1).select_related("small_model", "detector_model"):
|
||
|
|
model = ba.detector_model if int(ba.flow_type or 0) == 4 else ba.small_model
|
||
|
|
if model:
|
||
|
|
algorithms.append({"id": model.id, "name": model.name})
|
||
|
|
try:
|
||
|
|
labels = json.loads(ba.target_labels or "[]")
|
||
|
|
except Exception:
|
||
|
|
labels = []
|
||
|
|
target_labels.update(labels or [])
|
||
|
|
session = {
|
||
|
|
"session_id": session_id, "owner": owner, "stream_id": zone.stream_id,
|
||
|
|
"zone": self._zone_payload(zone), "last_heartbeat": now,
|
||
|
|
"created_at": now, "algorithms": algorithms,
|
||
|
|
"target_labels": sorted(target_labels),
|
||
|
|
}
|
||
|
|
with self._lock:
|
||
|
|
self._sessions[session_id] = session
|
||
|
|
stream = zone.stream
|
||
|
|
return {
|
||
|
|
"session_id": session_id, "mode": AnalysisManager().preview_mode(zone.stream_id),
|
||
|
|
"stream_id": zone.stream_id, "app": stream.app, "name": stream.name,
|
||
|
|
"models": algorithms, "zone": session["zone"],
|
||
|
|
}
|
||
|
|
|
||
|
|
def data(self, owner, session_id, since=None):
|
||
|
|
with self._lock:
|
||
|
|
session = self._sessions.get(session_id)
|
||
|
|
if not session or session["owner"] != owner:
|
||
|
|
raise PermissionError("预览会话不存在或不属于当前登录会话")
|
||
|
|
session["last_heartbeat"] = time.time()
|
||
|
|
stream_id = session["stream_id"]
|
||
|
|
zone = dict(session["zone"])
|
||
|
|
latest = get_event_bridge().latest_preview(stream_id)
|
||
|
|
mode = AnalysisManager().preview_mode(stream_id)
|
||
|
|
if not latest:
|
||
|
|
return {"changed": False, "sequence": 0, "mode": mode,
|
||
|
|
"zone": zone, "models": session["algorithms"]}
|
||
|
|
sequence = int(latest.get("sequence") or 0)
|
||
|
|
changed = since is None or sequence > int(since or 0)
|
||
|
|
if not changed:
|
||
|
|
return {"changed": False, "sequence": sequence, "mode": mode,
|
||
|
|
"timestamp": latest.get("timestamp")}
|
||
|
|
payload = dict(latest)
|
||
|
|
payload["tracks"] = [dict(track) for track in latest.get("tracks", [])]
|
||
|
|
payload.update({"changed": True, "mode": mode, "zone": zone})
|
||
|
|
# 叠框仅显示该布控目标类别;人员入侵不会再显示车辆和交通灯。
|
||
|
|
target_labels = set(session.get("target_labels") or [])
|
||
|
|
if target_labels:
|
||
|
|
payload["tracks"] = [t for t in payload.get("tracks", []) if t.get("label") in target_labels]
|
||
|
|
for track in payload.get("tracks", []):
|
||
|
|
box = track.get("box") or []
|
||
|
|
if len(box) != 4:
|
||
|
|
continue
|
||
|
|
foot = ((box[0] + box[2]) * .5,
|
||
|
|
box[3] - (box[3] - box[1]) * .05 if track.get("label") == "person"
|
||
|
|
else (box[1] + box[3]) * .5)
|
||
|
|
track["inside_zone_ids"] = ([zone["id"]] if self._point_in_polygon(foot, zone["coords"]) else [])
|
||
|
|
payload["zones"] = [zone]
|
||
|
|
return payload
|
||
|
|
|
||
|
|
def stop(self, owner, session_id):
|
||
|
|
with self._lock:
|
||
|
|
session = self._sessions.get(session_id)
|
||
|
|
if not session or session["owner"] != owner:
|
||
|
|
return False
|
||
|
|
stream_id = session["stream_id"]
|
||
|
|
self._sessions.pop(session_id, None)
|
||
|
|
has_viewers = any(s["stream_id"] == stream_id for s in self._sessions.values())
|
||
|
|
if not has_viewers:
|
||
|
|
AnalysisManager().stop_preview(stream_id)
|
||
|
|
return True
|
||
|
|
|
||
|
|
def _cleanup_loop(self):
|
||
|
|
while self._running:
|
||
|
|
now = time.time()
|
||
|
|
expired = []
|
||
|
|
with self._lock:
|
||
|
|
expired = [(sid, s["owner"]) for sid, s in self._sessions.items()
|
||
|
|
if now - s["last_heartbeat"] > self.TTL_SEC]
|
||
|
|
for sid, owner in expired:
|
||
|
|
self.stop(owner, sid)
|
||
|
|
threading.Event().wait(2.0)
|
||
|
|
|
||
|
|
|
||
|
|
_REGISTRY = PreviewSessionRegistry()
|
||
|
|
|
||
|
|
|
||
|
|
def get_preview_registry():
|
||
|
|
return _REGISTRY
|