"""Helpers for lightweight, API-only image search across multiple sources. The app is a pure remote client: every source is queried over HTTP and only a small list of image URLs + metadata is kept in memory. No local indexes, dataset downloads, or persistent storage. """ from __future__ import annotations import base64 import csv import enum import html import json import os import tempfile import zipfile from pathlib import Path from typing import Any from urllib.parse import urlparse import requests # --- shared config --------------------------------------------------------- API_URL = "https://knn.laion.ai/knn-service" OPENVERSE_URL = "https://api.openverse.org/v1/images/" WIKIMEDIA_URL = "https://commons.wikimedia.org/w/api.php" DATASETS_SERVER = "https://datasets-server.huggingface.co" RELAION_DATASET = "laion/relaion-high-resolution" DEFAULT_INDEX = "laion5B-L-14" COMMON_INDICES = ["laion5B-L-14", "laion5B-H-14", "laion_400m"] MAX_RESULTS = 50 MAX_DOWNLOADS = 20 REQUEST_TIMEOUT = 20 HTTP_HEADERS = { "User-Agent": "laion-safe-search/1.0 (https://huggingface.co/spaces/zackdorsey/img)" } # --- source registry ------------------------------------------------------- # Each source declares which controls are meaningful for it so the UI can adapt. DEFAULT_SOURCE = "openverse" SOURCES: dict[str, dict[str, Any]] = { "openverse": { "label": "Openverse — CC / public-domain images (no key)", "supports_nsfw": True, # honours a mature/sensitive toggle "supports_punsafe": False, # no per-result CLIP safety score "supports_image": False, # text query only }, "wikimedia": { "label": "Wikimedia Commons — freely-licensed media (no key)", "supports_nsfw": False, "supports_punsafe": False, "supports_image": False, }, "relaion": { "label": "reLAION-HR — LAION safety re-release, high-res (needs HF token)", "supports_nsfw": True, # via the continuous punsafe range "supports_punsafe": True, # server-side punsafe filtering over 166M rows "supports_image": False, # keyword caption search, not CLIP-semantic }, "laion": { "label": "LAION-5B — clip-retrieval kNN (often offline)", "supports_nsfw": True, # via the p_unsafe range "supports_punsafe": True, "supports_image": True, }, } class SourceUnavailable(RuntimeError): """Raised when a remote source cannot be reached or refuses the request.""" class Modality(str, enum.Enum): """Query modality accepted by the clip-retrieval knn-service backend.""" TEXT = "text" IMAGE = "image" # --- LAION clip-retrieval client (vendored minimal HTTP surface) ----------- class ClipClient: """Minimal, dependency-free client for the public clip-retrieval knn-service. Vendors the small HTTP surface of ``clip_retrieval.clip_client.ClipClient`` so the app does not need the heavy ``clip-retrieval`` package. The backend computes CLIP embeddings server-side; we only send JSON. """ def __init__( self, url: str, indice_name: str, aesthetic_score: int = 9, aesthetic_weight: float = 0.5, modality: Modality = Modality.IMAGE, num_images: int = 40, deduplicate: bool = True, use_safety_model: bool = False, use_violence_detector: bool = False, ) -> None: self.url = url self.indice_name = indice_name self.aesthetic_score = int(aesthetic_score) self.aesthetic_weight = float(aesthetic_weight) self.modality = modality.value if isinstance(modality, Modality) else str(modality) self.num_images = int(num_images) self.deduplicate = deduplicate self.use_safety_model = use_safety_model self.use_violence_detector = use_violence_detector def query( self, text: str | None = None, image: str | None = None, embedding_input: list[float] | None = None, ) -> Any: """Search by text, image URL/path, or a raw embedding.""" if text and image: raise ValueError("Only one of text or image may be provided per query.") if text: return self.__search_knn_api__(text=text) if image: if image.startswith("http://") or image.startswith("https://"): return self.__search_knn_api__(image_url=image) return self.__search_knn_api__(image=self._encode_image(image)) if embedding_input is not None: return self.__search_knn_api__(embedding_input=embedding_input) raise ValueError("Provide a text prompt, an image, or an embedding.") @staticmethod def _encode_image(path: str) -> str: with open(path, "rb") as handle: return base64.b64encode(handle.read()).decode("utf-8") def __search_knn_api__( self, text: str | None = None, image: str | None = None, image_url: str | None = None, embedding_input: list[float] | None = None, ) -> Any: payload = { "text": text, "image": image, "image_url": image_url, "embedding_input": embedding_input, "deduplicate": self.deduplicate, "use_safety_model": self.use_safety_model, "use_violence_detector": self.use_violence_detector, "indice_name": self.indice_name, "use_mclip": False, "aesthetic_score": self.aesthetic_score, "aesthetic_weight": self.aesthetic_weight, "modality": self.modality, "num_images": self.num_images, "num_result_ids": self.num_images, } response = requests.post( self.url, data=json.dumps(payload), headers=HTTP_HEADERS, timeout=REQUEST_TIMEOUT ) response.raise_for_status() return response.json() def make_client( index_name: str = DEFAULT_INDEX, num_results: int = 20, aesthetic_score: int | None = None, aesthetic_weight: float | None = None, deduplicate: bool = True, ) -> ClipClient: """Create a lightweight remote-only clip-retrieval client.""" client_options: dict[str, Any] = { "url": API_URL, "indice_name": index_name, "modality": Modality.IMAGE, "num_images": max(1, min(int(num_results), MAX_RESULTS)), "deduplicate": deduplicate, } if aesthetic_score is not None: client_options["aesthetic_score"] = aesthetic_score if aesthetic_weight is not None: client_options["aesthetic_weight"] = aesthetic_weight return ClipClient(**client_options) def query_laion( text: str | None, image: str | None, index_name: str, num_results: int, aesthetic_score: int | None = None, aesthetic_weight: float | None = None, deduplicate: bool = True, ) -> list[dict[str, Any]]: """Query the public clip-retrieval API without local indexes or datasets.""" if not text and not image: raise ValueError("Enter a text prompt, image URL, or upload a small image.") client = make_client( index_name=index_name, num_results=num_results, aesthetic_score=aesthetic_score, aesthetic_weight=aesthetic_weight, deduplicate=deduplicate, ) try: raw = client.query(text=text or None, image=image or None) except requests.exceptions.RequestException as exc: # offline / reset / timeout raise SourceUnavailable( "LAION's public clip-retrieval backend (knn.laion.ai) is not responding — " "it is frequently offline. Switch the Source to Openverse or Wikimedia Commons." ) from exc if isinstance(raw, dict) and raw.get("message"): raise SourceUnavailable(f"clip-retrieval API returned: {raw['message']}") if not isinstance(raw, list): raise SourceUnavailable("Unexpected LAION API response format.") return [normalize_laion(item, rank) for rank, item in enumerate(raw, start=1)] # --- Openverse ------------------------------------------------------------- # Anonymous Openverse requests are capped at 20 results per page. OPENVERSE_ANON_PAGE_MAX = 20 # Openverse per-result sensitivity reasons (the ``unstable__sensitivity`` array). SENSITIVITY_TEXT = "sensitive_text" # automated text-content detection SENSITIVITY_USER = "user_reported_sensitivity" # flagged by Openverse user reports SENSITIVITY_PROVIDER = "provider_supplied_sensitivity" # upstream provider marked it mature SENSITIVITY_LABELS = { SENSITIVITY_TEXT: "auto-flagged text", SENSITIVITY_USER: "user-reported", SENSITIVITY_PROVIDER: "provider-marked", } # Graded content-sensitivity levels used across sources. SENSITIVITY_LEVELS = ("safe", "moderate", "unrestricted") def search_openverse( query: str, num_results: int, sensitivity: str = "safe", ) -> list[dict[str, Any]]: """Search Openverse (CC / public-domain images). No API key required. ``sensitivity`` grades how flagged content is handled: - ``"safe"`` (default): exclude everything Openverse flags as sensitive. - ``"moderate"``: include items flagged *only* by automated text detection (often false positives — art, medical, news), but still hide provider-marked mature and user-reported results. - ``"unrestricted"``: include all flagged results. """ if not query: raise ValueError("Enter a text prompt to search Openverse.") params: dict[str, Any] = { "q": query, "page_size": max(1, min(int(num_results), OPENVERSE_ANON_PAGE_MAX)), } if sensitivity in ("moderate", "unrestricted"): params["unstable__include_sensitive_results"] = "true" try: response = requests.get( OPENVERSE_URL, params=params, headers=HTTP_HEADERS, timeout=REQUEST_TIMEOUT ) response.raise_for_status() data = response.json() except requests.exceptions.RequestException as exc: raise SourceUnavailable(f"Openverse is not responding: {exc}") from exc results = [ normalize_openverse(item, rank) for rank, item in enumerate(data.get("results", []), start=1) ] if sensitivity == "moderate": # Keep unflagged items and those flagged *only* by automated text detection. results = [ item for item in results if not item.get("sensitivity") or set(item["sensitivity"]) <= {SENSITIVITY_TEXT} ] for new_rank, item in enumerate(results, start=1): item["rank"] = new_rank return results # --- Wikimedia Commons ----------------------------------------------------- def search_wikimedia(query: str, num_results: int) -> list[dict[str, Any]]: """Search Wikimedia Commons freely-licensed media. No API key required.""" if not query: raise ValueError("Enter a text prompt to search Wikimedia Commons.") params = { "action": "query", "generator": "search", "gsrnamespace": 6, # File: "gsrsearch": query, "gsrlimit": max(1, min(int(num_results), MAX_RESULTS)), "prop": "imageinfo", "iiprop": "url|size|extmetadata", "iiurlwidth": 320, "format": "json", } try: response = requests.get( WIKIMEDIA_URL, params=params, headers=HTTP_HEADERS, timeout=REQUEST_TIMEOUT ) response.raise_for_status() data = response.json() except requests.exceptions.RequestException as exc: raise SourceUnavailable(f"Wikimedia Commons is not responding: {exc}") from exc pages = ((data.get("query") or {}).get("pages") or {}) ordered = sorted(pages.values(), key=lambda page: page.get("index", 1_000_000)) results = [] for rank, page in enumerate(ordered, start=1): info = (page.get("imageinfo") or [{}])[0] results.append(normalize_wikimedia(page, info, rank)) return results # --- reLAION-HR (HF datasets-server, gated) -------------------------------- def _hf_token() -> str | None: return os.environ.get("HF_TOKEN") or os.environ.get("HUGGING_FACE_HUB_TOKEN") def search_relaion( query: str, num_results: int, min_punsafe: float = 0.0, min_pwatermark: float = 0.0, languages: list[str] | None = None, punsafe_dir: str = "above", ) -> list[dict[str, Any]]: """Keyword search over ReLAION-HR captions with threshold + language filters. Uses the Hugging Face datasets-server ``/filter`` endpoint against ``laion/relaion-high-resolution`` (the safety re-release of LAION-5B). Matches caption substrings (``TEXT ILIKE``) and filters server-side over all 166M rows: ``punsafe`` and ``pwatermark`` are lower-bound thresholds (show values at or above), and ``languages`` restricts the ``LANGUAGE`` column when non-empty. The dataset is gated, so an ``HF_TOKEN`` with the gate accepted is required. """ if not query: raise ValueError("Enter a text prompt to search reLAION-HR.") token = _hf_token() if not token: raise SourceUnavailable( "reLAION-HR is a gated dataset. Add your Hugging Face token as a Space " "secret named HF_TOKEN (Settings → Variables and secrets) and accept the " "dataset terms at huggingface.co/datasets/laion/relaion-high-resolution." ) safe_q = str(query).strip().replace("\\", "")[:100].replace("'", "''") punsafe_op = "<=" if punsafe_dir == "below" else ">=" clauses = [ f"\"TEXT\" ILIKE '%{safe_q}%'", f"\"punsafe\" {punsafe_op} {float(min_punsafe):.4f}", f"\"pwatermark\" >= {float(min_pwatermark):.4f}", ] langs = [str(code) for code in (languages or []) if str(code).replace("-", "").replace("_", "").isalnum()] if langs: # datasets-server /filter rejects IN (...); OR-joined equalities are supported. joined = " OR ".join(f"\"LANGUAGE\" = '{code}'" for code in langs) clauses.append(f"({joined})") where = " AND ".join(clauses) params = { "dataset": RELAION_DATASET, "config": "default", "split": "train", "where": where, "offset": 0, "length": max(1, min(int(num_results), 100)), } try: response = requests.get( f"{DATASETS_SERVER}/filter", params=params, headers={**HTTP_HEADERS, "Authorization": f"Bearer {token}"}, timeout=REQUEST_TIMEOUT + 15, ) if response.status_code in (401, 403): raise SourceUnavailable( "reLAION-HR access was denied (HTTP %d). Ensure the Space's HF_TOKEN " "belongs to an account that has accepted the dataset's gated terms." % response.status_code ) response.raise_for_status() data = response.json() except requests.exceptions.RequestException as exc: raise SourceUnavailable(f"reLAION-HR (datasets-server) is not responding: {exc}") from exc return [ normalize_relaion(entry.get("row", {}), rank) for rank, entry in enumerate(data.get("rows", []), start=1) ] def normalize_relaion(row: dict[str, Any], rank: int) -> dict[str, Any]: """Normalize a ReLAION-HR datasets-server row.""" url = str(row.get("URL") or "") return _normalized( rank=rank, source="relaion", url=url, thumb=url, caption=str(row.get("TEXT") or ""), similarity=_float_or_none(row.get("similarity")), p_unsafe=_float_or_none(row.get("punsafe")), pwatermark=_float_or_none(row.get("pwatermark")), language=row.get("LANGUAGE"), item_id=row.get("hash"), width=row.get("WIDTH"), height=row.get("HEIGHT"), ) # --- dispatch -------------------------------------------------------------- def run_search( source: str, text: str | None, image: str | None, num_results: int, sensitivity: str = "safe", index_name: str = DEFAULT_INDEX, aesthetic_score: int | None = None, aesthetic_weight: float | None = None, deduplicate: bool = True, min_punsafe: float = 0.0, min_pwatermark: float = 0.0, languages: list[str] | None = None, punsafe_dir: str = "above", ) -> list[dict[str, Any]]: """Route a query to the selected source and return normalized results.""" if source == "openverse": return search_openverse(text or "", num_results, sensitivity=sensitivity) if source == "wikimedia": return search_wikimedia(text or "", num_results) if source == "relaion": return search_relaion( text or "", num_results, min_punsafe=min_punsafe, min_pwatermark=min_pwatermark, languages=languages, punsafe_dir=punsafe_dir, ) if source == "laion": return query_laion( text=text, image=image, index_name=index_name, num_results=num_results, aesthetic_score=aesthetic_score, aesthetic_weight=aesthetic_weight, deduplicate=deduplicate, ) raise ValueError(f"Unknown source: {source}") # --- normalization --------------------------------------------------------- def _normalized( rank: int, source: str, url: str, thumb: str | None = None, caption: str = "", similarity: float | None = None, p_unsafe: float | None = None, pwatermark: float | None = None, mature: bool | None = None, sensitivity: list[str] | None = None, license: str | None = None, landing_url: str | None = None, language: str | None = None, item_id: Any = None, width: Any = None, height: Any = None, ) -> dict[str, Any]: """Unified result shape used by rendering, export, and download.""" return { "rank": rank, "source": source, "url": str(url or ""), "thumb": str(thumb or url or ""), "caption": str(caption or ""), "similarity": similarity, "p_unsafe": p_unsafe, "pwatermark": pwatermark, "mature": mature, "sensitivity": sensitivity or [], "license": license, "landing_url": landing_url or (url or ""), "language": language, "id": item_id, "width": width, "height": height, } def normalize_laion(item: dict[str, Any], rank: int) -> dict[str, Any]: """Normalize a clip-retrieval result.""" metadata = item.get("metadata") if isinstance(item.get("metadata"), dict) else {} merged = {**metadata, **item} url = str(merged.get("url") or merged.get("image_url") or "") return _normalized( rank=rank, source="laion", url=url, thumb=url, caption=str(merged.get("caption") or merged.get("text") or ""), similarity=_float_or_none(merged.get("similarity") or merged.get("score")), p_unsafe=_float_or_none( merged.get("p_unsafe", merged.get("NSFW", merged.get("unsafe", merged.get("punsafe")))) ), item_id=merged.get("id"), width=merged.get("width"), height=merged.get("height"), ) # Backwards-compatible alias (older name used elsewhere/tests). normalize_result = normalize_laion def normalize_openverse(item: dict[str, Any], rank: int) -> dict[str, Any]: """Normalize an Openverse image result.""" reasons = [str(r) for r in (item.get("unstable__sensitivity") or [])] mature = bool(item.get("mature")) or bool(reasons) lic = item.get("license") lic_version = item.get("license_version") license_str = " ".join(str(p) for p in (lic, lic_version) if p).upper() or None return _normalized( rank=rank, source="openverse", url=str(item.get("url") or ""), thumb=str(item.get("thumbnail") or item.get("url") or ""), caption=str(item.get("title") or item.get("creator") or ""), mature=mature, sensitivity=reasons, license=license_str, landing_url=item.get("foreign_landing_url") or item.get("url"), item_id=item.get("id"), width=item.get("width"), height=item.get("height"), ) def normalize_wikimedia(page: dict[str, Any], info: dict[str, Any], rank: int) -> dict[str, Any]: """Normalize a Wikimedia Commons imageinfo result.""" title = str(page.get("title") or "") caption = title[5:] if title.lower().startswith("file:") else title extmeta = info.get("extmetadata") or {} license_str = (extmeta.get("LicenseShortName") or {}).get("value") return _normalized( rank=rank, source="wikimedia", url=str(info.get("url") or info.get("thumburl") or ""), thumb=str(info.get("thumburl") or info.get("url") or ""), caption=caption, license=license_str, landing_url=info.get("descriptionurl") or info.get("url"), item_id=page.get("pageid"), width=info.get("width"), height=info.get("height"), ) # --- client-side filtering (LAION / reLAION metadata) ---------------------- def filter_results( results: list[dict[str, Any]], min_punsafe: float = 0.0, min_pwatermark: float = 0.0, languages: list[str] | None = None, include_missing: bool = False, punsafe_dir: str = "above", ) -> list[dict[str, Any]]: """Keep results matching the ``p_unsafe`` threshold (at or above it when ``punsafe_dir`` is ``"above"``, at or below it when ``"below"``), at or above the ``pwatermark`` lower bound, and — when a language set is given — in those languages. Results with a missing ``p_unsafe`` are only retained when ``include_missing`` is set (pwatermark/language filters are skipped for results that lack those fields). """ langs = {str(code).lower() for code in (languages or [])} thr_p = float(min_punsafe) lo_w = float(min_pwatermark) kept: list[dict[str, Any]] = [] for result in results: score = result.get("p_unsafe") if score is None: if not include_missing: continue else: score = float(score) if punsafe_dir == "below": if score > thr_p: continue elif score < thr_p: continue watermark = result.get("pwatermark") if watermark is not None and float(watermark) < lo_w: continue if langs: lang = str(result.get("language") or "").lower() if lang and lang not in langs: continue kept.append(result) return kept # --- rendering ------------------------------------------------------------- _SOURCE_BADGE = {"laion": "LAION", "openverse": "Openverse", "wikimedia": "Wikimedia"} def render_results(results: list[dict[str, Any]]) -> str: """Render a responsive HTML gallery using remote thumbnail URLs.""" if not results: return "
Similarity: {_format_float(result.get('similarity'))}
") if result.get("p_unsafe") is not None: rows.append(f"p_unsafe: {_format_float(result.get('p_unsafe'))}
") if result.get("pwatermark") is not None: rows.append(f"watermark: {_format_float(result.get('pwatermark'))}
") reasons = result.get("sensitivity") or [] if reasons: labels = ", ".join(SENSITIVITY_LABELS.get(r, r) for r in reasons) rows.append(f"⚠ sensitive: {html.escape(labels)}
") elif result.get("mature"): rows.append("⚠ mature/NSFW-flagged
") if result.get("license"): rows.append(f"License: {html.escape(str(result.get('license')))}
") if result.get("language"): rows.append(f"Language: {html.escape(str(result.get('language')))}
") cards.append( f"""