diff --git a/system/athena/athenad.py b/system/athena/athenad.py index e78fc89bd..c31f1ef28 100755 --- a/system/athena/athenad.py +++ b/system/athena/athenad.py @@ -597,11 +597,7 @@ def startStream(sdp: str, video_enabled: bool | None = None) -> dict: resp = WEBRTCD_SESS.post(f"http://localhost:{WEBRTCD_PORT}/stream", json=asdict(body), timeout=10) t_end = time.monotonic() if not resp.ok: - try: - error_body = resp.json() - raise Exception(error_body.get("message", f"webrtcd returned {resp.status_code}")) - except ValueError: - resp.raise_for_status() + raise Exception(resp.json().get("message", f"webrtcd returned {resp.status_code}")) ret = resp.json() ret["time"] = (t_end - t_start) * 1000 return ret diff --git a/system/webrtc/webrtcd.py b/system/webrtc/webrtcd.py index f4f58c06d..8d29dafdd 100755 --- a/system/webrtc/webrtcd.py +++ b/system/webrtc/webrtcd.py @@ -400,12 +400,6 @@ async def get_stream(request: 'web.Request'): stream_dict[session.identifier] = session try: answer = await session.get_answer() - except ValueError as e: - await session.stop() - raise web.HTTPBadRequest( - text=json.dumps({"error": "invalid_sdp", "message": str(e)}), - content_type="application/json", - ) from e except Exception: await session.stop() stream_dict.pop(session.identifier, None) @@ -450,13 +444,24 @@ async def on_shutdown(app: 'web.Application'): del app['streams'] +@web.middleware +async def error_middleware(request: 'web.Request', handler): + try: + return await handler(request) + except web.HTTPException: + raise # intentional responses (400/404/etc.) pass through untouched + except Exception as e: + logging.getLogger("webrtcd").exception("Unhandled error handling %s", request.path) + return web.json_response({"error": "exception", "message": f"{type(e).__name__}: {e}"}, status=500) + + def webrtcd_thread(host: str, port: int, debug: bool): logging.basicConfig(level=logging.CRITICAL, handlers=[logging.StreamHandler()]) logging_level = logging.DEBUG if debug else logging.INFO logging.getLogger("WebRTCStream").setLevel(logging_level) logging.getLogger("webrtcd").setLevel(logging_level) - app = web.Application() + app = web.Application(middlewares=[error_middleware]) app['streams'] = dict() app['stream_lock'] = asyncio.Lock()