import copy import io import json import os import tempfile import threading from datetime import datetime, timezone def utc_now_iso(): return datetime.now(timezone.utc).isoformat(timespec="seconds") class HubStateStore: """Small JSON state store backed by a private Hugging Face dataset.""" def __init__(self, data_dir): self.repo_id = os.getenv("STATE_HUB_REPO_ID") self.repo_type = os.getenv("STATE_HUB_REPO_TYPE", "dataset") self.filename = os.getenv("STATE_HUB_FILENAME", "monitor_state.json") token = os.getenv("HF_TOKEN") self.path = os.path.join(data_dir, self.filename) self._lock = threading.RLock() self._api = None self.last_error = None self.remote_synced = False if self.repo_id and token: try: from huggingface_hub import HfApi self._api = HfApi(token=token) except Exception as exc: self.last_error = str(exc) self._state = self._load() @property def enabled(self): return self._api is not None and bool(self.repo_id) def _empty(self): return {"updated_at": "1970-01-01T00:00:00+00:00", "paused_monitor_ids": [], "values": {}} def _decode(self, path): with open(path, "r", encoding="utf-8") as handle: state = json.load(handle) if not isinstance(state, dict): raise ValueError("state root must be an object") state.setdefault("paused_monitor_ids", []) state.setdefault("values", {}) if not isinstance(state["paused_monitor_ids"], list): state["paused_monitor_ids"] = [] if not isinstance(state["values"], dict): state["values"] = {} return state def _write_local(self, state): os.makedirs(os.path.dirname(self.path), exist_ok=True) fd, temp_path = tempfile.mkstemp(prefix=".monitor-state-", dir=os.path.dirname(self.path)) try: with os.fdopen(fd, "w", encoding="utf-8") as handle: json.dump(state, handle, ensure_ascii=False, indent=2, sort_keys=True) os.replace(temp_path, self.path) except Exception: try: os.unlink(temp_path) except OSError: pass raise def _load(self): local = self._empty() if os.path.exists(self.path): try: local = self._decode(self.path) except Exception as exc: self.last_error = str(exc) if not self.enabled: return local try: from huggingface_hub import hf_hub_download remote_path = hf_hub_download( repo_id=self.repo_id, repo_type=self.repo_type, filename=self.filename, force_download=True, ) remote = self._decode(remote_path) if str(remote.get("updated_at") or "") >= str(local.get("updated_at") or ""): self.remote_synced = True self._write_local(remote) return remote return local except Exception as exc: self.last_error = str(exc) return local def _upload(self, state): if not self.enabled: return False try: payload = json.dumps(state, ensure_ascii=False, indent=2, sort_keys=True).encode("utf-8") self._api.upload_file( path_or_fileobj=io.BytesIO(payload), path_in_repo=self.filename, repo_id=self.repo_id, repo_type=self.repo_type, commit_message="Update monitor state", ) self.last_error = None self.remote_synced = True return True except Exception as exc: self.last_error = str(exc) self.remote_synced = False return False def get(self, key, default=None): with self._lock: if key == "paused_monitor_ids": return copy.deepcopy(self._state.get(key, [])) return copy.deepcopy(self._state.get("values", {}).get(key, default)) def update(self, **updates): with self._lock: self._state.update(updates) self._state["updated_at"] = utc_now_iso() state = copy.deepcopy(self._state) self._write_local(state) self._upload(state) def set_kv(self, key, value): with self._lock: self._state.setdefault("values", {})[key] = copy.deepcopy(value) self._state["updated_at"] = utc_now_iso() state = copy.deepcopy(self._state) self._write_local(state) self._upload(state)