237 lines
8.5 KiB
Python
237 lines
8.5 KiB
Python
"""
|
|
In-memory index store with disk snapshot and WebSocket broadcast.
|
|
|
|
The IndexStore is the single source of truth for the media index.
|
|
All mutations happen synchronously in the asyncio event loop — no locks needed.
|
|
index.json on disk is a recovery snapshot only, written periodically via a
|
|
debounced background task.
|
|
"""
|
|
|
|
import asyncio
|
|
import logging
|
|
from datetime import datetime
|
|
from pathlib import Path
|
|
from typing import Optional
|
|
|
|
import msgspec
|
|
from fastapi import WebSocket
|
|
|
|
from hivescan.structs import (
|
|
IndexSnapshot,
|
|
MediaStats,
|
|
Movie,
|
|
Series,
|
|
TaskInfo,
|
|
WsInit,
|
|
WsInitData,
|
|
WsRemove,
|
|
WsTask,
|
|
WsUpsert,
|
|
)
|
|
|
|
logger = logging.getLogger("hivescan.index_store")
|
|
|
|
# Debounce interval for writing snapshots to disk (seconds)
|
|
SNAPSHOT_DEBOUNCE = 5.0
|
|
|
|
|
|
class IndexStore:
|
|
"""In-memory media index with WS broadcast and disk snapshots."""
|
|
|
|
def __init__(self, snapshot_path: Path, media_root: Optional[str] = None):
|
|
self.snapshot_path = snapshot_path
|
|
self.media_root = media_root
|
|
|
|
# The index: keyed by item id
|
|
self.movies: dict[str, Movie] = {}
|
|
self.series: dict[str, Series] = {}
|
|
|
|
# Connected WebSocket clients
|
|
self._clients: set[WebSocket] = set()
|
|
|
|
# Snapshot debounce state
|
|
self._snapshot_dirty = False
|
|
self._snapshot_task: Optional[asyncio.Task] = None
|
|
|
|
# ------------------------------------------------------------------
|
|
# Persistence
|
|
# ------------------------------------------------------------------
|
|
|
|
def load_snapshot(self) -> None:
|
|
"""Load index from disk snapshot (recovery on startup)."""
|
|
if not self.snapshot_path.exists():
|
|
logger.info("No snapshot found at %s, starting fresh", self.snapshot_path)
|
|
return
|
|
try:
|
|
data = msgspec.json.decode(
|
|
self.snapshot_path.read_bytes(), type=IndexSnapshot
|
|
)
|
|
for m in data.movies:
|
|
self.movies[m.id] = m
|
|
for s in data.series:
|
|
self.series[s.id] = s
|
|
logger.info(
|
|
"Loaded snapshot: %d movies, %d series",
|
|
len(self.movies),
|
|
len(self.series),
|
|
)
|
|
except Exception:
|
|
logger.exception("Failed to load snapshot from %s", self.snapshot_path)
|
|
|
|
def _write_snapshot(self) -> None:
|
|
"""Write current index to disk (synchronous, called from debounce task)."""
|
|
movies_list = sorted(
|
|
self.movies.values(), key=lambda x: (x.title.lower(), x.year or 0)
|
|
)
|
|
series_list = sorted(self.series.values(), key=lambda x: x.title.lower())
|
|
|
|
total_movie_versions = sum(len(m.versions) for m in movies_list)
|
|
total_series_episodes = sum(
|
|
sum(len(season.episodes) for season in s.seasons) for s in series_list
|
|
)
|
|
|
|
snapshot = IndexSnapshot(
|
|
generated_at=datetime.now().isoformat(),
|
|
media_root=self.media_root,
|
|
stats=MediaStats(
|
|
total_movies=len(movies_list),
|
|
total_movie_versions=total_movie_versions,
|
|
total_series=len(series_list),
|
|
total_series_episodes=total_series_episodes,
|
|
),
|
|
movies=movies_list,
|
|
series=series_list,
|
|
)
|
|
|
|
self.snapshot_path.parent.mkdir(parents=True, exist_ok=True)
|
|
tmp = self.snapshot_path.with_suffix(".tmp")
|
|
tmp.write_bytes(msgspec.json.format(msgspec.json.encode(snapshot), indent=2))
|
|
tmp.replace(self.snapshot_path)
|
|
logger.debug("Snapshot written to %s", self.snapshot_path)
|
|
|
|
def _schedule_snapshot(self) -> None:
|
|
"""Schedule a debounced snapshot write."""
|
|
self._snapshot_dirty = True
|
|
if self._snapshot_task is None or self._snapshot_task.done():
|
|
self._snapshot_task = asyncio.create_task(self._debounced_snapshot())
|
|
|
|
async def _debounced_snapshot(self) -> None:
|
|
"""Wait for debounce interval then write if still dirty."""
|
|
while self._snapshot_dirty:
|
|
self._snapshot_dirty = False
|
|
await asyncio.sleep(SNAPSHOT_DEBOUNCE)
|
|
# After the sleep, if no new mutations happened, write
|
|
self._write_snapshot()
|
|
|
|
async def flush_snapshot(self) -> None:
|
|
"""Force-write a snapshot immediately (e.g. on shutdown)."""
|
|
if self._snapshot_task and not self._snapshot_task.done():
|
|
self._snapshot_task.cancel()
|
|
try:
|
|
await self._snapshot_task
|
|
except asyncio.CancelledError:
|
|
pass
|
|
self._write_snapshot()
|
|
|
|
# ------------------------------------------------------------------
|
|
# Mutations
|
|
# ------------------------------------------------------------------
|
|
|
|
def upsert_movie(self, item: Movie) -> None:
|
|
"""Insert or update a movie in the index and broadcast."""
|
|
self.movies[item.id] = item
|
|
self._schedule_snapshot()
|
|
self._broadcast(WsUpsert(kind="movie", item=item))
|
|
|
|
def upsert_series(self, item: Series) -> None:
|
|
"""Insert or update a series in the index and broadcast."""
|
|
self.series[item.id] = item
|
|
self._schedule_snapshot()
|
|
self._broadcast(WsUpsert(kind="series", item=item))
|
|
|
|
def remove_movie(self, item_id: str) -> None:
|
|
"""Remove a movie from the index and broadcast."""
|
|
self.movies.pop(item_id, None)
|
|
self._schedule_snapshot()
|
|
self._broadcast(WsRemove(kind="movie", id=item_id))
|
|
|
|
def remove_series(self, item_id: str) -> None:
|
|
"""Remove a series from the index and broadcast."""
|
|
self.series.pop(item_id, None)
|
|
self._schedule_snapshot()
|
|
self._broadcast(WsRemove(kind="series", id=item_id))
|
|
|
|
# ------------------------------------------------------------------
|
|
# WebSocket management
|
|
# ------------------------------------------------------------------
|
|
|
|
async def connect(self, ws: WebSocket) -> None:
|
|
"""Accept a WS client and send the full index as init."""
|
|
await ws.accept()
|
|
self._clients.add(ws)
|
|
logger.info("WS client connected (%d total)", len(self._clients))
|
|
# Send full current state
|
|
msg = WsInit(
|
|
data=WsInitData(
|
|
movies=list(self.movies.values()),
|
|
series=list(self.series.values()),
|
|
)
|
|
)
|
|
await ws.send_bytes(msgspec.json.encode(msg))
|
|
|
|
def disconnect(self, ws: WebSocket) -> None:
|
|
"""Remove a WS client."""
|
|
self._clients.discard(ws)
|
|
logger.info("WS client disconnected (%d remaining)", len(self._clients))
|
|
|
|
def _broadcast(self, msg: object) -> None:
|
|
"""Broadcast a message to all connected WS clients (non-blocking)."""
|
|
data = msgspec.json.encode(msg)
|
|
dead: list[WebSocket] = []
|
|
for ws in self._clients:
|
|
asyncio.create_task(self._safe_send(ws, data, dead))
|
|
# Clean up dead connections after sends are scheduled
|
|
for ws in dead:
|
|
self._clients.discard(ws)
|
|
|
|
@staticmethod
|
|
async def _safe_send(ws: WebSocket, data: bytes, dead: list) -> None:
|
|
"""Send data to a WS client; mark as dead on failure."""
|
|
try:
|
|
await ws.send_bytes(data)
|
|
except Exception:
|
|
dead.append(ws)
|
|
|
|
def broadcast_task(self, task_info: TaskInfo) -> None:
|
|
"""Broadcast a task progress message to all WS clients."""
|
|
self._broadcast(WsTask(data=task_info))
|
|
|
|
# ------------------------------------------------------------------
|
|
# Read helpers
|
|
# ------------------------------------------------------------------
|
|
|
|
def get_full_index(self) -> IndexSnapshot:
|
|
"""Return the full index as an IndexSnapshot."""
|
|
movies_list = sorted(
|
|
self.movies.values(), key=lambda x: (x.title.lower(), x.year or 0)
|
|
)
|
|
series_list = sorted(self.series.values(), key=lambda x: x.title.lower())
|
|
|
|
total_movie_versions = sum(len(m.versions) for m in movies_list)
|
|
total_series_episodes = sum(
|
|
sum(len(season.episodes) for season in s.seasons) for s in series_list
|
|
)
|
|
|
|
return IndexSnapshot(
|
|
generated_at=datetime.now().isoformat(),
|
|
media_root=self.media_root,
|
|
stats=MediaStats(
|
|
total_movies=len(movies_list),
|
|
total_movie_versions=total_movie_versions,
|
|
total_series=len(series_list),
|
|
total_series_episodes=total_series_episodes,
|
|
),
|
|
movies=movies_list,
|
|
series=series_list,
|
|
)
|