#!/usr/bin/env python3 """Smoke-test a deployed GeoText-1652 Inference Endpoint. Usage: export HF_TOKEN=hf_... python test_endpoint.py Optional: GEOTEXT_ENDPOINT_URL=https://...endpoints.huggingface.cloud \ python test_endpoint.py --patches --store-patches """ from __future__ import annotations import argparse import os import sys from typing import Any, Dict, Optional import requests DEFAULT_ENDPOINT = "https://6ac5162dc8b01bf6eda56cb3.endpoints.huggingface.cloud" DEFAULT_IMAGE = ( "https://huggingface.co/datasets/geobase/geoai-cogs/resolve/main/" "geoembeddings-demo/building-detection_sm.tif" ) def parse_args() -> argparse.Namespace: parser = argparse.ArgumentParser(description=__doc__) parser.add_argument( "--endpoint-url", default=os.environ.get("GEOTEXT_ENDPOINT_URL", DEFAULT_ENDPOINT), help="Endpoint base URL; defaults to GEOTEXT_ENDPOINT_URL or the deployed endpoint.", ) parser.add_argument("--image-url", default=DEFAULT_IMAGE, help="Public RGB GeoTIFF/COG URL.") parser.add_argument( "--patches", action="store_true", help="Also test tiled patch-embedding output." ) parser.add_argument( "--store-patches", action="store_true", help="Also test Parquet persistence; endpoint must have HF_TOKEN and HF_BUCKET configured.", ) return parser.parse_args() def call( session: requests.Session, method: str, url: str, payload: Optional[Dict[str, Any]] = None, ) -> Dict[str, Any]: response = session.request(method, url, json=payload, timeout=300) try: response.raise_for_status() except requests.HTTPError as error: raise RuntimeError(f"{method} {url} failed: {response.text}") from error return response.json() def check(name: str, result: Dict[str, Any], required: set[str]) -> None: missing = required.difference(result) if missing: raise RuntimeError(f"{name} response is missing: {sorted(missing)}\n{result}") print(f"PASS {name}") def main() -> None: args = parse_args() token = os.environ.get("HF_TOKEN") if not token: sys.exit("Set HF_TOKEN to a Hugging Face token with access to the private endpoint.") base_url = args.endpoint_url.rstrip("/") session = requests.Session() session.headers.update( {"Authorization": f"Bearer {token}", "Content-Type": "application/json"} ) health = call(session, "GET", f"{base_url}/health") check("health", health, {"status"}) text = call( session, "POST", f"{base_url}/embed/text", {"queries": ["a parking lot", "a building"]}, ) check("text embeddings", text, {"embedding_dimension", "text_embeddings"}) assert len(text["text_embeddings"]) == 2 assert len(text["text_embeddings"][0]["embedding"]) == 256 image = call( session, "POST", f"{base_url}/embed/image", {"image_url": args.image_url, "include_global_embedding": True}, ) check("global image embedding", image, {"image_embedding", "image_metadata"}) assert len(image["image_embedding"]) == 256 ranking = call( session, "POST", f"{base_url}/infer", {"image_url": args.image_url, "queries": ["a parking lot", "a building"]}, ) check("text-guided ranking", ranking, {"ranked_queries", "image_metadata"}) assert len(ranking["ranked_queries"]) == 2 if args.patches or args.store_patches: patch_request: Dict[str, Any] = { "image_url": args.image_url, "tile_size": 384, "tile_overlap": 64, "include_global_embedding": True, "include_patch_embeddings": args.patches, "store_patch_embeddings": args.store_patches, } patches = call(session, "POST", f"{base_url}/embed/image", patch_request) check("tiled patch embeddings", patches, {"tiling"}) if args.patches: check("patch-vector response", patches, {"patch_embeddings"}) assert patches["patch_embeddings"]["patches"] if args.store_patches: check("patch-vector storage", patches, {"patch_embeddings_storage"}) print("All requested GeoText endpoint checks passed.") if __name__ == "__main__": main()