import json import logging from fastapi import FastAPI, WebSocket, WebSocketDisconnect from fastapi.middleware.cors import CORSMiddleware import cv2 import numpy as np import config from pipeline import load_sign_pipeline, make_translator_for_connection logging.basicConfig(level=logging.INFO) logger = logging.getLogger("sign-to-text-server") app = FastAPI(title="Sign-to-Text Streaming API") app.add_middleware( CORSMiddleware, allow_origins=["*"], allow_credentials=False, allow_methods=["*"], allow_headers=["*"], ) pipeline_state = {"ready": False, "pipeline": None} @app.on_event("startup") def on_startup(): logger.info("Loading sign pipeline...") pipeline_state["pipeline"] = load_sign_pipeline() pipeline_state["ready"] = True logger.info("Sign pipeline ready. Server accepting connections.") @app.get("/health") def health(): return {"status": "ok" if pipeline_state["ready"] else "loading"} @app.get("/") def root(): return { "name": "Sign-to-Text Streaming API", "status": "ok" if pipeline_state["ready"] else "loading", "websocket_endpoint": "/ws", "protocol": { "client_sends": "binary WebSocket frames, each one JPEG-encoded RGB video frame, at config.TARGET_FPS", "server_sends": [ {"type": "segment", "start_seconds": "float", "end_seconds": "float", "num_frames": "int", "text": "string", "dino_batch_ms": "float", "translation_ms": "float"}, {"type": "error", "message": "string"}, ], }, } @app.websocket("/ws") async def websocket_endpoint(websocket: WebSocket): if not pipeline_state["ready"]: await websocket.close(code=1013) return await websocket.accept() logger.info("Client connected -- new translator created") translator = make_translator_for_connection(pipeline_state["pipeline"]) frame_index = 0 try: while True: raw_bytes = await websocket.receive_bytes() frame_bgr = cv2.imdecode(np.frombuffer(raw_bytes, dtype=np.uint8), cv2.IMREAD_COLOR) if frame_bgr is None: await websocket.send_text(json.dumps({"type": "error", "message": f"Could not decode frame {frame_index}"})) frame_index += 1 continue frame_rgb = cv2.cvtColor(frame_bgr, cv2.COLOR_BGR2RGB) timestamp_seconds = frame_index / config.TARGET_FPS try: completed_segments = translator.push_frame(frame_rgb, frame_index, timestamp_seconds) except Exception: logger.exception(f"Error while processing frame {frame_index}") await websocket.send_text(json.dumps({"type": "error", "message": f"Error while processing frame {frame_index}"})) completed_segments = [] for segment_result in completed_segments: await websocket.send_text(json.dumps({"type": "segment", **segment_result})) frame_index += 1 except WebSocketDisconnect: logger.info("Client disconnected") try: trailing_segments = translator.flush() for segment_result in trailing_segments: logger.info(f"Trailing segment on disconnect: {segment_result['text']}") except Exception: logger.exception("Error flushing trailing segment on disconnect") except Exception: logger.exception("Unexpected error in websocket loop") finally: translator.close()