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