diff --git a/backend/open_webui/main.py b/backend/open_webui/main.py index 247146f7fb..4c231289cf 100644 --- a/backend/open_webui/main.py +++ b/backend/open_webui/main.py @@ -250,7 +250,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, @@ -1667,7 +1667,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) @@ -1716,11 +1720,10 @@ async def chat_completion( except Exception: pass - else: - # No chat_id/message_id → legacy/direct API path with no - # WebSocket error channel. We must surface the error as - # a proper HTTP response; without this the function would - # return None which FastAPI serializes as null. #23924 + # 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, @@ -1989,9 +1992,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 d354986d37..713ffe506d 100644 --- a/backend/open_webui/routers/ollama.py +++ b/backend/open_webui/routers/ollama.py @@ -36,7 +36,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_headers_and_cookies 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, @@ -133,7 +133,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: @@ -148,6 +152,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 b8a75c2bfb..8ef451ebdc 100644 --- a/backend/open_webui/routers/openai.py +++ b/backend/open_webui/routers/openai.py @@ -42,7 +42,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_headers_and_cookies, 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, @@ -1596,6 +1596,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( @@ -1608,7 +1609,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, @@ -1623,6 +1624,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 @@ -1649,10 +1651,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): @@ -1760,10 +1763,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: @@ -1888,10 +1892,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 @@ -2010,10 +2015,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 5c5dcce3a6..314bd8cff1 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()