ARViewer / api /endpoints.py
wizardsmagic's picture
Add Wizara Vision API endpoints without changing UI or inference.
3820d5b
Raw
History Blame Contribute Delete
4.4 kB
"""Wizara Vision API endpoint handlers."""
from __future__ import annotations
import os
from typing import Any, Callable
from PIL import Image
from .responses import (
detections_to_objects,
error_response,
parse_advanced_settings,
success_response,
unsupported_response,
)
def resolve_image_path(image_file: Any, base_dir: str) -> str | None:
if isinstance(image_file, str):
candidate = image_file
if not os.path.isabs(candidate):
candidate = os.path.join(base_dir, candidate)
return candidate if os.path.exists(candidate) else None
if isinstance(image_file, dict):
path = image_file.get("path")
return path if path and os.path.exists(path) else None
path = getattr(image_file, "path", None)
return path if path and os.path.exists(path) else None
def merge_settings(
*,
model_mode: str,
temp: float,
top_p: float,
top_k: int,
short_size: int | None,
advanced_settings: str | None,
) -> dict[str, Any]:
settings = {
"model_mode": model_mode,
"temp": temp,
"top_p": top_p,
"top_k": top_k,
"short_size": short_size,
}
settings.update(parse_advanced_settings(advanced_settings))
return settings
def handle_detect(
*,
image_file: Any,
categories: str,
task_type: str,
model_mode: str,
temp: float,
top_p: float,
top_k: int,
short_size: int | None,
advanced_settings: str | None,
base_dir: str,
run_image_gpu: Callable[..., tuple],
generate_prompt: Callable[[str, str], str],
parse_results: Callable[[str, str], list[dict[str, Any]]],
) -> dict[str, Any]:
if not image_file:
return error_response("Image is required.", code="MISSING_IMAGE")
image_path = resolve_image_path(image_file, base_dir)
if not image_path:
return error_response("Invalid image file path.", code="INVALID_IMAGE")
settings = merge_settings(
model_mode=model_mode,
temp=temp,
top_p=top_p,
top_k=top_k,
short_size=short_size,
advanced_settings=advanced_settings,
)
category = categories.strip() or "objects"
task = task_type.strip() or "Detection"
question_override = settings.get("question_override")
if not question_override:
question_override = generate_prompt(task, category)
try:
image = Image.open(image_path).convert("RGB")
width, height = image.size
_, stats, raw_text, _, _ = run_image_gpu(
image_path,
category,
settings.get("model_mode", model_mode),
float(settings.get("temp", temp)),
float(settings.get("top_p", top_p)),
int(settings.get("top_k", top_k)),
settings.get("short_size", short_size),
question_override,
)
category_str = " ".join(part.strip() for part in category.split(",") if part.strip())
detections = parse_results(raw_text, category_str)
objects = detections_to_objects(detections)
return success_response(
image_width=width,
image_height=height,
objects=objects,
task=task.lower(),
extra={
"stats": stats,
"raw_text": raw_text,
"prompt": question_override,
},
)
except Exception as exc: # noqa: BLE001
return error_response(str(exc), code="INFERENCE_FAILED")
def handle_ocr(
*,
image_file: Any,
model_mode: str,
temp: float,
top_p: float,
top_k: int,
short_size: int | None,
advanced_settings: str | None,
base_dir: str,
run_image_gpu: Callable[..., tuple],
generate_prompt: Callable[[str, str], str],
parse_results: Callable[[str, str], list[dict[str, Any]]],
) -> dict[str, Any]:
return handle_detect(
image_file=image_file,
categories="text",
task_type="OCR",
model_mode=model_mode,
temp=temp,
top_p=top_p,
top_k=top_k,
short_size=short_size,
advanced_settings=advanced_settings,
base_dir=base_dir,
run_image_gpu=run_image_gpu,
generate_prompt=generate_prompt,
parse_results=parse_results,
)
def handle_unsupported(task: str) -> dict[str, Any]:
return unsupported_response(task)