agAdvisor / src /accounts /hf_sync.py
tirtho149's picture
Deploy AgAdvisor
b30f068 verified
Raw
History Blame Contribute Delete
5.39 kB
"""Mirror the accounts SQLite DB to a private Hugging Face Dataset.
HF Space filesystems are ephemeral (wiped on rebuild/restart), so the DB is
pulled from a private Dataset repo on cold start and pushed back (debounced)
after writes. When no write token is configured — e.g. local development — every
method is a safe no-op and the local file is used as-is.
Concurrency/durability tradeoff: a single Streamlit process serializes writes;
pushes are coalesced (``min_push_interval``) to respect HF commit-rate limits.
A crash between a write and the next push loses at most the last exchange. For
true multi-writer concurrency, swap this for a hosted DB (Turso/Supabase) later.
Dependencies (``huggingface_hub``) are imported lazily and can be injected, so
this module is unit-testable without the package or network.
"""
from __future__ import annotations
import logging
import os
import shutil
import threading
import time
from pathlib import Path
from typing import Callable, Optional
logger = logging.getLogger("agadvisor.accounts.hf_sync")
DEFAULT_FILENAME = "agadvisor_users.db"
class HFSync:
def __init__(
self,
repo_id: Optional[str],
local_db_path: str | Path,
token: Optional[str] = None,
filename: str = DEFAULT_FILENAME,
min_push_interval: float = 10.0,
*,
api: Optional[object] = None,
downloader: Optional[Callable[..., str]] = None,
):
self.repo_id = repo_id or None
self.local_db_path = Path(local_db_path)
self.token = token or os.getenv("HF_DATA_TOKEN") or None
self.filename = filename
self.min_push_interval = float(min_push_interval)
self._api = api
self._downloader = downloader
self._lock = threading.Lock()
self._last_push = 0.0
self._pending = False
@property
def enabled(self) -> bool:
return bool(self.repo_id and self.token)
# -- lazy huggingface_hub handles ---------------------------------------
def _get_api(self):
if self._api is None:
from huggingface_hub import HfApi # lazy
self._api = HfApi(token=self.token)
return self._api
def _get_downloader(self) -> Callable[..., str]:
if self._downloader is None:
from huggingface_hub import hf_hub_download # lazy
self._downloader = hf_hub_download
return self._downloader
# -- operations ----------------------------------------------------------
def ensure_repo(self) -> None:
if not self.enabled:
return
try:
self._get_api().create_repo(
repo_id=self.repo_id,
repo_type="dataset",
private=True,
exist_ok=True,
token=self.token,
)
except Exception as e: # non-fatal
logger.warning("Could not ensure HF dataset repo exists: %s", e)
def pull(self) -> bool:
"""Download the DB from the dataset into ``local_db_path``. Returns True
if a file was fetched, False otherwise (no token / not present / error)."""
if not self.enabled:
return False
try:
path = self._get_downloader()(
repo_id=self.repo_id,
repo_type="dataset",
filename=self.filename,
token=self.token,
)
self.local_db_path.parent.mkdir(parents=True, exist_ok=True)
if Path(path) != self.local_db_path:
shutil.copyfile(path, self.local_db_path)
logger.info("Pulled accounts DB from HF dataset %s", self.repo_id)
return True
except Exception as e:
# Most commonly: file does not exist yet (fresh deploy). Not an error.
logger.info("No existing accounts DB pulled from HF (%s): %s", self.repo_id, e)
return False
def push(self, force: bool = False) -> bool:
"""Upload the local DB to the dataset. Debounced unless ``force``."""
if not self.enabled:
return False
with self._lock:
now = time.time()
if not force and (now - self._last_push) < self.min_push_interval:
self._pending = True
return False
if not self.local_db_path.exists():
return False
try:
self._get_api().upload_file(
path_or_fileobj=str(self.local_db_path),
path_in_repo=self.filename,
repo_id=self.repo_id,
repo_type="dataset",
token=self.token,
)
self._last_push = now
self._pending = False
logger.info("Pushed accounts DB to HF dataset %s", self.repo_id)
return True
except Exception as e:
logger.warning("Failed to push accounts DB to HF: %s", e)
return False
def maybe_push(self) -> bool:
"""Push if the debounce window has elapsed; otherwise mark pending."""
return self.push(force=False)
def flush(self) -> bool:
"""Force a push if there are pending unsynced writes (e.g. on logout)."""
if self._pending:
return self.push(force=True)
return False