import argparse
import json
from pathlib import Path
from google.cloud import vision
from google.api_core.exceptions import GoogleAPICallError
def classify_image(path: Path, threshold: float) -> list[dict]:
if not 0.0 <= threshold <= 1.0:
raise ValueError("Threshold must be between 0 and 1")
content = path.read_bytes()
if not content:
raise ValueError("Image file is empty")
client = vision.ImageAnnotatorClient()
response = client.label_detection(
image=vision.Image(content=content), max_results=20, timeout=30.0
)
if response.error.message:
raise RuntimeError(
f"Vision error {response.error.code}: {response.error.message}"
)
return sorted(
[{"label": label.description, "score": float(label.score)}
for label in response.label_annotations
if label.score >= threshold],
key=lambda item: item["score"], reverse=True
)
if __name__ == "__main__":
parser = argparse.ArgumentParser(description="Cloud Vision label detection")
parser.add_argument("image", type=Path)
parser.add_argument("--threshold", type=float, default=0.7)
args = parser.parse_args()
try:
labels = classify_image(args.image, args.threshold)
print(json.dumps({"labels": labels}, indent=2))
except (OSError, ValueError, RuntimeError, GoogleAPICallError) as error:
parser.exit(1, f"Classification failed: {error}\n")