diff --git a/backend/open_webui/main.py b/backend/open_webui/main.py index dc1498b577..c29539c6ca 100644 --- a/backend/open_webui/main.py +++ b/backend/open_webui/main.py @@ -247,7 +247,7 @@ from open_webui.utils.middleware import ( process_chat_payload, process_chat_response, ) -from open_webui.utils.misc import get_response_error_detail, merge_model_params +from open_webui.utils.misc import get_response_error_detail, get_retry_after_headers, merge_model_params from open_webui.utils.model_ids import strip_provider_model_prefix from open_webui.utils.models import ( check_model_access, @@ -1660,7 +1660,11 @@ async def chat_completion( # raise so the except-block below emits a terminal # chat:message:error, unblocking the frontend. if isinstance(response, JSONResponse) and response.status_code >= 400: - raise Exception(get_response_error_detail(response)) + raise HTTPException( + status_code=response.status_code, + detail=get_response_error_detail(response), + headers=get_retry_after_headers(response.headers), + ) if ctx is None: ctx = await build_chat_response_context(request, form_data, user, model, metadata, tasks, events) @@ -1711,6 +1715,8 @@ async def chat_completion( pass # Legacy/direct callers await this response; returning None would send `null`. #23924 if not (metadata.get('session_id') and metadata.get('chat_id')): + if isinstance(e, HTTPException): + raise raise HTTPException( status_code=status.HTTP_400_BAD_REQUEST, detail=error_detail, @@ -1979,9 +1985,10 @@ async def passthrough_anthropic_messages(request: Request, form_data: dict, user requested_model=requested_model, upstream_error=response_data, ) + retry_after_headers = get_retry_after_headers(response.headers) if isinstance(response_data, (dict, list)): - return JSONResponse(status_code=response.status, content=response_data) - return Response(status_code=response.status, content=response_data) + return JSONResponse(status_code=response.status, content=response_data, headers=retry_after_headers) + return Response(status_code=response.status, content=response_data, headers=retry_after_headers) return response_data except HTTPException: diff --git a/backend/open_webui/routers/ollama.py b/backend/open_webui/routers/ollama.py index b283db9ae8..1a046c729e 100644 --- a/backend/open_webui/routers/ollama.py +++ b/backend/open_webui/routers/ollama.py @@ -38,7 +38,7 @@ from open_webui.utils.access_control import check_model_access from open_webui.utils.auth import get_admin_user, get_verified_user from open_webui.utils.headers import get_custom_headers, include_user_info_headers from open_webui.utils.json_codec import JSONCodec -from open_webui.utils.misc import calculate_sha256 +from open_webui.utils.misc import calculate_sha256, get_retry_after_headers from open_webui.utils.model_ids import strip_provider_model_prefix from open_webui.utils.payload import ( apply_model_params_to_body_ollama, @@ -149,7 +149,11 @@ async def send_request( upstream_error=res, ) if 'error' in res: - raise HTTPException(status_code=r.status, detail=res['error']) + raise HTTPException( + status_code=r.status, + detail=res['error'], + headers=get_retry_after_headers(r.headers), + ) except HTTPException: raise except Exception as e: @@ -164,6 +168,7 @@ async def send_request( raise HTTPException( status_code=r.status, detail=ERROR_MESSAGES.SERVER_CONNECTION_ERROR, + headers=get_retry_after_headers(r.headers), ) r.raise_for_status() diff --git a/backend/open_webui/routers/openai.py b/backend/open_webui/routers/openai.py index 63676d166e..ef65419ebc 100644 --- a/backend/open_webui/routers/openai.py +++ b/backend/open_webui/routers/openai.py @@ -44,7 +44,7 @@ from open_webui.utils.anthropic import ANTHROPIC_VERSION, get_anthropic_models, from open_webui.utils.auth import get_admin_user, get_verified_user from open_webui.utils.headers import get_custom_headers, include_user_info_headers from open_webui.utils.json_codec import JSONCodec -from open_webui.utils.misc import convert_logit_bias_input_to_json +from open_webui.utils.misc import convert_logit_bias_input_to_json, get_retry_after_headers from open_webui.utils.model_ids import strip_provider_model_prefix from open_webui.utils.payload import ( apply_model_params_to_body_openai, @@ -1679,6 +1679,7 @@ async def generate_chat_completion( r.status, error_body[:1000], ) + retry_after_headers = get_retry_after_headers(r.headers) try: error_json = JSONCodec.loads(error_body) await publish_model_provider_request_failed( @@ -1691,7 +1692,7 @@ async def generate_chat_completion( requested_model=requested_model, upstream_error=error_json, ) - return JSONResponse(status_code=r.status, content=error_json) + return JSONResponse(status_code=r.status, content=error_json, headers=retry_after_headers) except JSONCodec.JSONDecodeError: await publish_model_provider_request_failed( request, @@ -1706,6 +1707,7 @@ async def generate_chat_completion( return JSONResponse( status_code=r.status, content={'error': {'message': error_body, 'code': r.status}}, + headers=retry_after_headers, ) streaming = True @@ -1732,10 +1734,11 @@ async def generate_chat_completion( requested_model=requested_model, upstream_error=response, ) + retry_after_headers = get_retry_after_headers(r.headers) if isinstance(response, (dict, list)): - return JSONResponse(status_code=r.status, content=response) + return JSONResponse(status_code=r.status, content=response, headers=retry_after_headers) else: - return PlainTextResponse(status_code=r.status, content=response) + return PlainTextResponse(status_code=r.status, content=response, headers=retry_after_headers) # Convert Responses API result to simple format if is_responses and isinstance(response, dict): @@ -1843,10 +1846,11 @@ async def embeddings(request: Request, form_data: dict, user): requested_model=requested_model, upstream_error=response_data, ) + retry_after_headers = get_retry_after_headers(r.headers) if isinstance(response_data, (dict, list)): - return JSONResponse(status_code=r.status, content=response_data) + return JSONResponse(status_code=r.status, content=response_data, headers=retry_after_headers) else: - return PlainTextResponse(status_code=r.status, content=response_data) + return PlainTextResponse(status_code=r.status, content=response_data, headers=retry_after_headers) return response_data except Exception as e: @@ -1971,10 +1975,11 @@ async def responses( requested_model=payload.get('model'), upstream_error=response_data, ) + retry_after_headers = get_retry_after_headers(r.headers) if isinstance(response_data, (dict, list)): - return JSONResponse(status_code=r.status, content=response_data) + return JSONResponse(status_code=r.status, content=response_data, headers=retry_after_headers) else: - return PlainTextResponse(status_code=r.status, content=response_data) + return PlainTextResponse(status_code=r.status, content=response_data, headers=retry_after_headers) return response_data @@ -2093,10 +2098,11 @@ async def proxy(path: str, request: Request, user=Depends(get_verified_user)): requested_model=model_id, upstream_error=response_data, ) + retry_after_headers = get_retry_after_headers(r.headers) if isinstance(response_data, (dict, list)): - return JSONResponse(status_code=r.status, content=response_data) + return JSONResponse(status_code=r.status, content=response_data, headers=retry_after_headers) else: - return PlainTextResponse(status_code=r.status, content=response_data) + return PlainTextResponse(status_code=r.status, content=response_data, headers=retry_after_headers) return response_data diff --git a/backend/open_webui/utils/misc.py b/backend/open_webui/utils/misc.py index ad959d3bc7..b2ea2c2c5e 100644 --- a/backend/open_webui/utils/misc.py +++ b/backend/open_webui/utils/misc.py @@ -65,6 +65,10 @@ def get_response_error_detail(response: object) -> str: return detail if isinstance(detail, str) else str(detail) +def get_retry_after_headers(headers: collections.abc.Mapping) -> dict[str, str]: + return {name: headers[name] for name in ('Retry-After', 'retry-after-ms') if name in headers} + + def _strip_filter_entry(entry): # Compose list-form env syntax passes surrounding quotes through verbatim return (entry or '').strip().strip('"\'').strip()