Gateway / src /api /space_client.py
Adidasa108's picture
Create space_client.py
bac4901 verified
Raw
History Blame Contribute Delete
9.37 kB
"""
space_client.py — HTTP client for distributed mode.
Calls the remote Face Space and Object Space for AI inference.
Handles:
- Persistent aiohttp session (created once at startup)
- Exponential-backoff retry for cold-start wakeups (HF free tier)
- Keep-alive background pinger to prevent spaces sleeping
- Structured exceptions that search.py / upload.py already handle
"""
import asyncio
import io
import logging
import os
import aiohttp
# ── Space URLs (set these as secrets in the Gateway HF Space) ────────────────
FACE_SPACE_URL = os.getenv("FACE_SPACE_URL", "").rstrip("/")
OBJECT_SPACE_URL = os.getenv("OBJECT_SPACE_URL", "").rstrip("/")
# ── Timeouts ─────────────────────────────────────────────────────────────────
# HF free-tier spaces can take 60-90 s to cold-start.
# Keep TOTAL high so we wait through the wake-up; CONNECT stays short.
_CONNECT_TIMEOUT = 10 # seconds — TCP handshake
_TOTAL_TIMEOUT = 120 # seconds — full request including model inference
# ── Retry config (exponential backoff for cold starts) ────────────────────────
_MAX_RETRIES = 3
_RETRY_BACKOFF = [5, 15, 30] # seconds to wait before each retry
logger = logging.getLogger("space_client")
# ── Singleton aiohttp session ─────────────────────────────────────────────────
_session: aiohttp.ClientSession | None = None
async def get_session() -> aiohttp.ClientSession:
"""Return (or lazily create) the shared aiohttp session."""
global _session
if _session is None or _session.closed:
connector = aiohttp.TCPConnector(
limit=20, # max simultaneous connections
limit_per_host=10,
keepalive_timeout=60,
)
timeout = aiohttp.ClientTimeout(
connect=_CONNECT_TIMEOUT,
total=_TOTAL_TIMEOUT,
)
_session = aiohttp.ClientSession(connector=connector, timeout=timeout)
return _session
async def close_session() -> None:
"""Close the session gracefully (called from lifespan shutdown)."""
global _session
if _session and not _session.closed:
await _session.close()
_session = None
# ── Core POST helper with retry ───────────────────────────────────────────────
async def _post_form_data(url: str, form: aiohttp.FormData, label: str) -> dict:
"""
POST multipart/form-data to `url` with retry + backoff.
Raises RuntimeError on final failure so callers can log a warning
and degrade gracefully (e.g. return [] instead of crashing).
"""
session = await get_session()
last_exc: Exception | None = None
for attempt in range(_MAX_RETRIES):
try:
async with session.post(url, data=form) as resp:
if resp.status == 200:
return await resp.json()
# 503 = space still waking up — retry
if resp.status == 503:
body = await resp.text()
logger.warning(
"[space_client] %s returned 503 (attempt %d/%d): %s",
label, attempt + 1, _MAX_RETRIES, body[:200],
)
last_exc = RuntimeError(f"{label} 503: {body[:200]}")
else:
body = await resp.text()
raise RuntimeError(f"{label} HTTP {resp.status}: {body[:300]}")
except asyncio.TimeoutError as e:
logger.warning(
"[space_client] %s timeout (attempt %d/%d)", label, attempt + 1, _MAX_RETRIES
)
last_exc = e
except aiohttp.ClientConnectionError as e:
logger.warning(
"[space_client] %s connection error (attempt %d/%d): %s",
label, attempt + 1, _MAX_RETRIES, e,
)
last_exc = e
# Wait before retrying (skip wait after last attempt)
if attempt < _MAX_RETRIES - 1:
wait = _RETRY_BACKOFF[attempt]
logger.info("[space_client] Retrying %s in %ds…", label, wait)
await asyncio.sleep(wait)
raise RuntimeError(f"{label} failed after {_MAX_RETRIES} attempts: {last_exc}")
# ── Public API: embed_face ────────────────────────────────────────────────────
async def embed_face(image_bytes: bytes, quality_gate: float = 0.35) -> list[dict]:
"""
Send image bytes to the Face Space and return a list of face dicts.
Returns [] on failure (caller logs search.face_space_unavailable).
"""
if not FACE_SPACE_URL:
logger.error("[space_client] FACE_SPACE_URL is not set.")
return []
url = f"{FACE_SPACE_URL}/embed/face"
try:
form = aiohttp.FormData()
form.add_field(
"file",
io.BytesIO(image_bytes),
filename="image.jpg",
content_type="image/jpeg",
)
form.add_field("quality_gate", str(quality_gate))
result = await _post_form_data(url, form, "face-space")
faces = result.get("faces", [])
logger.info("[space_client] embed_face → %d face(s)", len(faces))
return faces
except Exception as e:
logger.warning("[space_client] embed_face failed: %s", e)
return [] # Degrade gracefully — search.py handles empty list
# ── Public API: embed_object ──────────────────────────────────────────────────
async def embed_object(image_bytes: bytes) -> list[dict]:
"""
Send image bytes to the Object Space and return a list of object dicts.
Returns [] on failure (caller raises 503 to the user).
"""
if not OBJECT_SPACE_URL:
logger.error("[space_client] OBJECT_SPACE_URL is not set.")
return []
url = f"{OBJECT_SPACE_URL}/embed/object"
try:
form = aiohttp.FormData()
form.add_field(
"file",
io.BytesIO(image_bytes),
filename="image.jpg",
content_type="image/jpeg",
)
result = await _post_form_data(url, form, "object-space")
objects = result.get("objects", [])
logger.info("[space_client] embed_object → %d object vec(s)", len(objects))
return objects
except Exception as e:
logger.warning("[space_client] embed_object failed: %s", e)
return []
# ── Health check / ping ───────────────────────────────────────────────────────
async def ping_spaces() -> dict:
"""
Ping both spaces at startup (wakes them from sleep).
Returns {"face": bool, "object": bool}.
"""
session = await get_session()
results = {"face": False, "object": False}
async def _ping(url: str, key: str) -> None:
try:
async with session.get(
f"{url}/health",
timeout=aiohttp.ClientTimeout(total=90), # allow cold-start wake
) as resp:
results[key] = resp.status == 200
logger.info("[space_client] ping %s → %d", key, resp.status)
except Exception as e:
logger.warning("[space_client] ping %s failed: %s", key, e)
ping_tasks = []
if FACE_SPACE_URL:
ping_tasks.append(_ping(FACE_SPACE_URL, "face"))
if OBJECT_SPACE_URL:
ping_tasks.append(_ping(OBJECT_SPACE_URL, "object"))
await asyncio.gather(*ping_tasks)
return results
# ── Keep-warm background task ─────────────────────────────────────────────────
async def keep_spaces_warm(interval_seconds: int = 540) -> None:
"""
Ping both spaces every `interval_seconds` (default 9 min) to prevent
HF free-tier sleep (spaces sleep after ~15 min of inactivity).
Start this as a background task from your lifespan:
asyncio.create_task(keep_spaces_warm())
"""
session = await get_session()
while True:
await asyncio.sleep(interval_seconds)
urls = []
if FACE_SPACE_URL:
urls.append((FACE_SPACE_URL, "face"))
if OBJECT_SPACE_URL:
urls.append((OBJECT_SPACE_URL, "object"))
async def _warm(url: str, label: str) -> None:
try:
async with session.get(
f"{url}/health",
timeout=aiohttp.ClientTimeout(total=30),
) as resp:
logger.debug("[space_client] keep-warm %s → %d", label, resp.status)
except Exception as e:
logger.debug("[space_client] keep-warm %s error: %s", label, e)
await asyncio.gather(*[_warm(u, l) for u, l in urls], return_exceptions=True)