import io
import time
from fastapi import APIRouter, File, Request, UploadFile
from PIL import Image
from tools import (
ocr,
text_detection,
layout_detection,
table_recognition,
extract_text_from_image,
)
from vllm_batcher import vllm_ocr_batcher
from vllm_tools import vllm_backend_info
from surya.endpoint.schemas import ApiResponse, Info
from surya.endpoint.service import (
_dump_model_or_list,
_error_response,
_load_base64_image,
_log_exception,
_ocr_response_data,
_success_response,
logger,
)
router = APIRouter()
@router.post("/home")
def home():
return "
Welcome to SURYA OCR API!
"
@router.post("/v1/api/ai/suya_ocr", include_in_schema=False, response_model=ApiResponse)
@router.post("/v1/api/ai/suya_ocr/", response_model=ApiResponse)
def run_ocr(p: Info, request: Request):
try:
pil_image, pil_image_highres = _load_base64_image(p)
rec_img, pred, box_img = ocr(
pil_image,
pil_image_highres,
p.skip_text_detection,
p.recognize_math,
with_bboxes=p.ocr_with_boxes,
)
return _success_response(_ocr_response_data(pred))
except Exception as e:
results = _error_response(e)
_log_exception("suya_ocr", request, e, results)
return results
@router.get("/v1/api/ai/suya_ocr_vllm/health", response_model=ApiResponse)
async def suya_ocr_vllm_health():
return _success_response(vllm_backend_info())
@router.post("/v1/api/ai/suya_ocr_vllm", include_in_schema=False, response_model=ApiResponse)
@router.post("/v1/api/ai/suya_ocr_vllm/", response_model=ApiResponse)
def run_ocr_vllm(p: Info, request: Request):
try:
pil_image, pil_image_highres = _load_base64_image(p)
start = time.perf_counter()
_, pred, _ = vllm_ocr_batcher.submit(
pil_image,
pil_image_highres,
skip_text_detection=p.skip_text_detection,
recognize_math=p.recognize_math,
with_bboxes=False,
request_id=getattr(request.state, "request_id", "-"),
)
logger.info(
"api_vllm_submit_wait_complete request_id=%s duration_ms=%.2f",
getattr(request.state, "request_id", "-"),
(time.perf_counter() - start) * 1000,
)
return _success_response(_ocr_response_data(pred))
except Exception as e:
results = _error_response(e)
_log_exception("suya_ocr_vllm", request, e, results)
return results
@router.post("/v1/api/ai/suya_text_det", include_in_schema=False, response_model=ApiResponse)
@router.post("/v1/api/ai/suya_text_det/", response_model=ApiResponse)
def run_text_det(p: Info, request: Request):
try:
pil_image, _ = _load_base64_image(p)
det_img, text_pred = text_detection(pil_image)
text_lines = text_pred.model_dump(exclude=["heatmap", "affinity_map"])
return _success_response({"text_lines": text_lines})
except Exception as e:
results = _error_response(e)
_log_exception("suya_text_det", request, e, results)
return results
@router.post("/v1/api/ai/suya_layout_det", include_in_schema=False, response_model=ApiResponse)
@router.post("/v1/api/ai/suya_layout_det/", response_model=ApiResponse)
def run_layout_det(p: Info, request: Request):
try:
pil_image, _ = _load_base64_image(p)
layout_img, pred = layout_detection(pil_image)
text_lines = pred.model_dump(exclude=["segmentation_map"])
return _success_response({"text_lines": text_lines})
except Exception as e:
results = _error_response(e)
_log_exception("suya_layout_det", request, e, results)
return results
@router.post("/v1/api/ai/suya_table_rec", include_in_schema=False, response_model=ApiResponse)
@router.post("/v1/api/ai/suya_table_rec/", response_model=ApiResponse)
def run_table_rec(p: Info, request: Request):
try:
pil_image, pil_image_highres = _load_base64_image(p)
table_img, pred = table_recognition(
pil_image, pil_image_highres, p.skip_table_detection
)
text_json = _dump_model_or_list(pred)
text_lines = ""
if isinstance(pred, list):
text_lines = "\n".join(
[item.html for item in pred if getattr(item, "html", None)]
)
elif hasattr(pred, "text_lines"):
text_lines = "\n".join([p.text for p in pred.text_lines])
return _success_response({"ocr_text_json": text_json, "text_lines": text_lines})
except Exception as e:
results = _error_response(e)
_log_exception("suya_table_rec", request, e, results)
return results
@router.post("/image2text", response_model=ApiResponse)
async def image_to_text(request: Request, file: UploadFile = File(...)):
try:
contents = await file.read()
image = Image.open(io.BytesIO(contents))
logger.info(
"image2text_upload request_id=%s filename=%s content_type=%s size_bytes=%s image_size=%s",
getattr(request.state, "request_id", "-"),
file.filename,
file.content_type,
len(contents),
image.size,
)
full_text = extract_text_from_image(image)
return _success_response(full_text)
except Exception as e:
results = _error_response(e)
_log_exception("image2text", request, e, results)
return results