coinpush / app /core /state.py
zt p
Persist monitor state across restarts
3be8894
Raw History Blame Contribute Delete
4.83 kB
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)