SatQuery-AI / satquery_agent /model_cache.py
AnirudhShashikumar's picture
Fix OpenCLIP safetensors loading on ZeroGPU
c306cf6
Raw
History Blame Contribute Delete
3.61 kB
"""Deterministic cache and retrieval helpers for public upstream models."""
from __future__ import annotations
import hashlib
import os
import shutil
import uuid
from pathlib import Path
from typing import Any, Callable, Mapping
APPLICATION_ROOT = Path(__file__).resolve().parents[1]
DEFAULT_MODEL_CACHE = APPLICATION_ROOT / ".cache" / "satquery-models"
class PublicModelRetrievalError(RuntimeError):
"""A pinned public upstream artifact could not be retrieved safely."""
def model_cache_dir(environment: Mapping[str, str] | None = None) -> Path:
"""Return the application-controlled model cache, honoring its sole override."""
values = os.environ if environment is None else environment
configured = values.get("SATQUERY_MODEL_CACHE", "").strip()
candidate = Path(configured).expanduser() if configured else DEFAULT_MODEL_CACHE
if not candidate.is_absolute():
candidate = APPLICATION_ROOT / candidate
return candidate.resolve()
def _materialized_public_model_path(*, repo_id: str, filename: str, revision: str) -> Path:
"""Return a deterministic, suffix-preserving path for one immutable artifact."""
if not filename or Path(filename).name != filename:
raise PublicModelRetrievalError("Public model filename must be a single path component")
identity = hashlib.sha256(f"{repo_id}\0{revision}\0{filename}".encode("utf-8")).hexdigest()
return model_cache_dir() / "materialized" / identity / filename
def _materialize_cached_file(source: Path, target: Path) -> Path:
"""Atomically hard-link or copy a cached artifact to its stable named path."""
if target.is_file():
return target.absolute()
target.parent.mkdir(parents=True, exist_ok=True)
temporary = target.with_name(f".{target.name}.{os.getpid()}.{uuid.uuid4().hex}.tmp")
try:
try:
os.link(source, temporary)
except OSError:
shutil.copyfile(source, temporary)
os.replace(temporary, target)
finally:
temporary.unlink(missing_ok=True)
return target.absolute()
def download_public_hf_file(
*,
repo_id: str,
filename: str,
revision: str,
downloader: Callable[..., Any] | None = None,
) -> Path:
"""Resolve one immutable public file to a stable path retaining its filename."""
target = _materialized_public_model_path(
repo_id=repo_id,
filename=filename,
revision=revision,
)
if target.is_file():
return target.absolute()
if downloader is None:
from huggingface_hub import hf_hub_download
downloader = hf_hub_download
cache_dir = model_cache_dir()
try:
cache_dir.mkdir(parents=True, exist_ok=True)
source = Path(
downloader(
repo_id=repo_id,
filename=filename,
revision=revision,
cache_dir=str(cache_dir),
token=False,
)
).resolve()
except Exception as error:
raise PublicModelRetrievalError(
f"Pinned public model artifact could not be retrieved: {repo_id}@{revision}/{filename}"
) from error
if not source.is_file():
raise PublicModelRetrievalError(
f"Pinned public model artifact was not materialized: {repo_id}@{revision}/{filename}"
)
try:
return _materialize_cached_file(source, target)
except OSError as error:
raise PublicModelRetrievalError(
f"Pinned public model artifact could not be named safely: {repo_id}@{revision}/{filename}"
) from error