mirror of
https://git.datalinker.icu/comfyanonymous/ComfyUI
synced 2026-09-04 19:07:05 +08:00
rework how errors are handled on the client side
This commit is contained in:
parent
9ad287ff20
commit
4cfedbc6c7
@ -94,11 +94,12 @@ from __future__ import annotations
|
|||||||
import logging
|
import logging
|
||||||
import time
|
import time
|
||||||
import io
|
import io
|
||||||
from typing import Dict, Type, Optional, Any, TypeVar, Generic, Callable
|
import socket
|
||||||
|
from typing import Dict, Type, Optional, Any, TypeVar, Generic, Callable, Tuple, List, Union
|
||||||
from enum import Enum
|
from enum import Enum
|
||||||
import json
|
import json
|
||||||
import requests
|
import requests
|
||||||
from urllib.parse import urljoin
|
from urllib.parse import urljoin, urlparse
|
||||||
from pydantic import BaseModel, Field
|
from pydantic import BaseModel, Field
|
||||||
|
|
||||||
from comfy.cli_args import args
|
from comfy.cli_args import args
|
||||||
@ -111,6 +112,21 @@ P = TypeVar("P", bound=BaseModel) # For poll response
|
|||||||
PROGRESS_BAR_MAX = 100
|
PROGRESS_BAR_MAX = 100
|
||||||
|
|
||||||
|
|
||||||
|
class NetworkError(Exception):
|
||||||
|
"""Base exception for network-related errors with diagnostic information."""
|
||||||
|
pass
|
||||||
|
|
||||||
|
|
||||||
|
class LocalNetworkError(NetworkError):
|
||||||
|
"""Exception raised when local network connectivity issues are detected."""
|
||||||
|
pass
|
||||||
|
|
||||||
|
|
||||||
|
class ApiServerError(NetworkError):
|
||||||
|
"""Exception raised when the API server is unreachable but internet is working."""
|
||||||
|
pass
|
||||||
|
|
||||||
|
|
||||||
class EmptyRequest(BaseModel):
|
class EmptyRequest(BaseModel):
|
||||||
"""Base class for empty request bodies.
|
"""Base class for empty request bodies.
|
||||||
For GET requests, fields will be sent as query parameters."""
|
For GET requests, fields will be sent as query parameters."""
|
||||||
@ -141,7 +157,7 @@ class HttpMethod(str, Enum):
|
|||||||
|
|
||||||
class ApiClient:
|
class ApiClient:
|
||||||
"""
|
"""
|
||||||
Client for making HTTP requests to an API with authentication and error handling.
|
Client for making HTTP requests to an API with authentication, error handling, and retry logic.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
def __init__(
|
def __init__(
|
||||||
@ -151,12 +167,22 @@ class ApiClient:
|
|||||||
comfy_api_key: Optional[str] = None,
|
comfy_api_key: Optional[str] = None,
|
||||||
timeout: float = 3600.0,
|
timeout: float = 3600.0,
|
||||||
verify_ssl: bool = True,
|
verify_ssl: bool = True,
|
||||||
|
max_retries: int = 3,
|
||||||
|
retry_delay: float = 1.0,
|
||||||
|
retry_backoff_factor: float = 2.0,
|
||||||
|
retry_status_codes: Optional[Tuple[int, ...]] = None,
|
||||||
):
|
):
|
||||||
self.base_url = base_url
|
self.base_url = base_url
|
||||||
self.auth_token = auth_token
|
self.auth_token = auth_token
|
||||||
self.comfy_api_key = comfy_api_key
|
self.comfy_api_key = comfy_api_key
|
||||||
self.timeout = timeout
|
self.timeout = timeout
|
||||||
self.verify_ssl = verify_ssl
|
self.verify_ssl = verify_ssl
|
||||||
|
self.max_retries = max_retries
|
||||||
|
self.retry_delay = retry_delay
|
||||||
|
self.retry_backoff_factor = retry_backoff_factor
|
||||||
|
# Default retry status codes: 408 (Request Timeout), 429 (Too Many Requests),
|
||||||
|
# 500, 502, 503, 504 (Server Errors)
|
||||||
|
self.retry_status_codes = retry_status_codes or (408, 429, 500, 502, 503, 504)
|
||||||
|
|
||||||
def _create_json_payload_args(
|
def _create_json_payload_args(
|
||||||
self,
|
self,
|
||||||
@ -211,6 +237,56 @@ class ApiClient:
|
|||||||
|
|
||||||
return headers
|
return headers
|
||||||
|
|
||||||
|
def _check_connectivity(self, target_url: str) -> Dict[str, bool]:
|
||||||
|
"""
|
||||||
|
Check connectivity to determine if network issues are local or server-related.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
target_url: URL to check connectivity to
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
Dictionary with connectivity status details
|
||||||
|
"""
|
||||||
|
results = {
|
||||||
|
"internet_accessible": False,
|
||||||
|
"api_accessible": False,
|
||||||
|
"is_local_issue": False,
|
||||||
|
"is_api_issue": False
|
||||||
|
}
|
||||||
|
|
||||||
|
# First check basic internet connectivity using a reliable external site
|
||||||
|
try:
|
||||||
|
# Use a reliable external domain for checking basic connectivity
|
||||||
|
check_response = requests.get("https://www.google.com",
|
||||||
|
timeout=5.0,
|
||||||
|
verify=self.verify_ssl)
|
||||||
|
if check_response.status_code < 500:
|
||||||
|
results["internet_accessible"] = True
|
||||||
|
except (requests.RequestException, socket.error):
|
||||||
|
results["internet_accessible"] = False
|
||||||
|
results["is_local_issue"] = True
|
||||||
|
return results
|
||||||
|
|
||||||
|
# Now check API server connectivity
|
||||||
|
try:
|
||||||
|
# Extract domain from the target URL to do a simpler health check
|
||||||
|
parsed_url = urlparse(target_url)
|
||||||
|
api_base = f"{parsed_url.scheme}://{parsed_url.netloc}"
|
||||||
|
|
||||||
|
# Try to reach the API domain
|
||||||
|
api_response = requests.get(f"{api_base}/health", timeout=5.0, verify=self.verify_ssl)
|
||||||
|
if api_response.status_code < 500:
|
||||||
|
results["api_accessible"] = True
|
||||||
|
else:
|
||||||
|
results["api_accessible"] = False
|
||||||
|
results["is_api_issue"] = True
|
||||||
|
except requests.RequestException:
|
||||||
|
results["api_accessible"] = False
|
||||||
|
# If we can reach the internet but not the API, it's an API issue
|
||||||
|
results["is_api_issue"] = True
|
||||||
|
|
||||||
|
return results
|
||||||
|
|
||||||
def request(
|
def request(
|
||||||
self,
|
self,
|
||||||
method: str,
|
method: str,
|
||||||
@ -221,9 +297,10 @@ class ApiClient:
|
|||||||
headers: Optional[Dict[str, str]] = None,
|
headers: Optional[Dict[str, str]] = None,
|
||||||
content_type: str = "application/json",
|
content_type: str = "application/json",
|
||||||
multipart_parser: Callable = None,
|
multipart_parser: Callable = None,
|
||||||
|
retry_count: int = 0, # Used internally for tracking retries
|
||||||
) -> Dict[str, Any]:
|
) -> Dict[str, Any]:
|
||||||
"""
|
"""
|
||||||
Make an HTTP request to the API
|
Make an HTTP request to the API with automatic retries for transient errors.
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
method: HTTP method (GET, POST, etc.)
|
method: HTTP method (GET, POST, etc.)
|
||||||
@ -233,12 +310,15 @@ class ApiClient:
|
|||||||
files: Files to upload
|
files: Files to upload
|
||||||
headers: Additional headers
|
headers: Additional headers
|
||||||
content_type: Content type of the request. Defaults to application/json.
|
content_type: Content type of the request. Defaults to application/json.
|
||||||
|
retry_count: Internal parameter for tracking retries, do not set manually
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
Parsed JSON response
|
Parsed JSON response
|
||||||
|
|
||||||
Raises:
|
Raises:
|
||||||
requests.RequestException: If the request fails
|
LocalNetworkError: If local network connectivity issues are detected
|
||||||
|
ApiServerError: If the API server is unreachable but internet is working
|
||||||
|
Exception: For other request failures
|
||||||
"""
|
"""
|
||||||
url = urljoin(self.base_url, path)
|
url = urljoin(self.base_url, path)
|
||||||
self.check_auth(self.auth_token, self.comfy_api_key)
|
self.check_auth(self.auth_token, self.comfy_api_key)
|
||||||
@ -275,17 +355,103 @@ class ApiClient:
|
|||||||
**payload_args,
|
**payload_args,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
# Check if we should retry based on status code
|
||||||
|
if (response.status_code in self.retry_status_codes and
|
||||||
|
retry_count < self.max_retries):
|
||||||
|
|
||||||
|
# Calculate delay with exponential backoff
|
||||||
|
delay = self.retry_delay * (self.retry_backoff_factor ** retry_count)
|
||||||
|
|
||||||
|
logging.warning(
|
||||||
|
f"Request failed with status {response.status_code}. "
|
||||||
|
f"Retrying in {delay:.2f}s ({retry_count + 1}/{self.max_retries})"
|
||||||
|
)
|
||||||
|
|
||||||
|
time.sleep(delay)
|
||||||
|
return self.request(
|
||||||
|
method=method,
|
||||||
|
path=path,
|
||||||
|
params=params,
|
||||||
|
data=data,
|
||||||
|
files=files,
|
||||||
|
headers=headers,
|
||||||
|
content_type=content_type,
|
||||||
|
multipart_parser=multipart_parser,
|
||||||
|
retry_count=retry_count + 1,
|
||||||
|
)
|
||||||
|
|
||||||
# Raise exception for error status codes
|
# Raise exception for error status codes
|
||||||
response.raise_for_status()
|
response.raise_for_status()
|
||||||
except requests.ConnectionError:
|
|
||||||
raise Exception(
|
|
||||||
f"Unable to connect to the API server at {self.base_url}. Please check your internet connection or verify the service is available."
|
|
||||||
)
|
|
||||||
|
|
||||||
except requests.Timeout:
|
except requests.ConnectionError as e:
|
||||||
|
# Only perform connectivity check if we've exhausted all retries
|
||||||
|
if retry_count >= self.max_retries:
|
||||||
|
# Check connectivity to determine if it's a local or API issue
|
||||||
|
connectivity = self._check_connectivity(self.base_url)
|
||||||
|
|
||||||
|
if connectivity["is_local_issue"]:
|
||||||
|
raise LocalNetworkError(
|
||||||
|
f"Unable to connect to the API server due to local network issues. "
|
||||||
|
f"Please check your internet connection and try again."
|
||||||
|
) from e
|
||||||
|
elif connectivity["is_api_issue"]:
|
||||||
|
raise ApiServerError(
|
||||||
|
f"The API server at {self.base_url} is currently unreachable. "
|
||||||
|
f"The service may be experiencing issues. Please try again later."
|
||||||
|
) from e
|
||||||
|
|
||||||
|
# If we haven't exhausted retries yet, retry the request
|
||||||
|
if retry_count < self.max_retries:
|
||||||
|
delay = self.retry_delay * (self.retry_backoff_factor ** retry_count)
|
||||||
|
logging.warning(
|
||||||
|
f"Connection error: {str(e)}. "
|
||||||
|
f"Retrying in {delay:.2f}s ({retry_count + 1}/{self.max_retries})"
|
||||||
|
)
|
||||||
|
time.sleep(delay)
|
||||||
|
return self.request(
|
||||||
|
method=method,
|
||||||
|
path=path,
|
||||||
|
params=params,
|
||||||
|
data=data,
|
||||||
|
files=files,
|
||||||
|
headers=headers,
|
||||||
|
content_type=content_type,
|
||||||
|
multipart_parser=multipart_parser,
|
||||||
|
retry_count=retry_count + 1,
|
||||||
|
)
|
||||||
|
|
||||||
|
# If we've exhausted retries and didn't identify the specific issue,
|
||||||
|
# raise a generic exception
|
||||||
raise Exception(
|
raise Exception(
|
||||||
f"Request timed out after {self.timeout} seconds. The server might be experiencing high load or the operation is taking longer than expected."
|
f"Unable to connect to the API server after {self.max_retries} attempts. "
|
||||||
)
|
f"Please check your internet connection or try again later."
|
||||||
|
) from e
|
||||||
|
|
||||||
|
except requests.Timeout as e:
|
||||||
|
# Retry timeouts if we haven't exhausted retries
|
||||||
|
if retry_count < self.max_retries:
|
||||||
|
delay = self.retry_delay * (self.retry_backoff_factor ** retry_count)
|
||||||
|
logging.warning(
|
||||||
|
f"Request timed out. "
|
||||||
|
f"Retrying in {delay:.2f}s ({retry_count + 1}/{self.max_retries})"
|
||||||
|
)
|
||||||
|
time.sleep(delay)
|
||||||
|
return self.request(
|
||||||
|
method=method,
|
||||||
|
path=path,
|
||||||
|
params=params,
|
||||||
|
data=data,
|
||||||
|
files=files,
|
||||||
|
headers=headers,
|
||||||
|
content_type=content_type,
|
||||||
|
multipart_parser=multipart_parser,
|
||||||
|
retry_count=retry_count + 1,
|
||||||
|
)
|
||||||
|
|
||||||
|
raise Exception(
|
||||||
|
f"Request timed out after {self.timeout} seconds and {self.max_retries} retry attempts. "
|
||||||
|
f"The server might be experiencing high load or the operation is taking longer than expected."
|
||||||
|
) from e
|
||||||
|
|
||||||
except requests.HTTPError as e:
|
except requests.HTTPError as e:
|
||||||
status_code = e.response.status_code if hasattr(e, "response") else None
|
status_code = e.response.status_code if hasattr(e, "response") else None
|
||||||
@ -310,14 +476,39 @@ class ApiClient:
|
|||||||
logging.debug(f"[DEBUG] API Error: {error_message} (Status: {status_code})")
|
logging.debug(f"[DEBUG] API Error: {error_message} (Status: {status_code})")
|
||||||
if hasattr(e, "response") and e.response.content:
|
if hasattr(e, "response") and e.response.content:
|
||||||
logging.debug(f"[DEBUG] Response content: {e.response.content}")
|
logging.debug(f"[DEBUG] Response content: {e.response.content}")
|
||||||
|
|
||||||
|
# Retry if the status code is in our retry list and we haven't exhausted retries
|
||||||
|
if (status_code in self.retry_status_codes and
|
||||||
|
retry_count < self.max_retries):
|
||||||
|
|
||||||
|
delay = self.retry_delay * (self.retry_backoff_factor ** retry_count)
|
||||||
|
logging.warning(
|
||||||
|
f"HTTP error {status_code}. "
|
||||||
|
f"Retrying in {delay:.2f}s ({retry_count + 1}/{self.max_retries})"
|
||||||
|
)
|
||||||
|
time.sleep(delay)
|
||||||
|
return self.request(
|
||||||
|
method=method,
|
||||||
|
path=path,
|
||||||
|
params=params,
|
||||||
|
data=data,
|
||||||
|
files=files,
|
||||||
|
headers=headers,
|
||||||
|
content_type=content_type,
|
||||||
|
multipart_parser=multipart_parser,
|
||||||
|
retry_count=retry_count + 1,
|
||||||
|
)
|
||||||
|
|
||||||
|
# Specific error messages for common status codes
|
||||||
if status_code == 401:
|
if status_code == 401:
|
||||||
error_message = "Unauthorized: Please login first to use this node."
|
error_message = "Unauthorized: Please login first to use this node."
|
||||||
if status_code == 402:
|
elif status_code == 402:
|
||||||
error_message = "Payment Required: Please add credits to your account to use this node."
|
error_message = "Payment Required: Please add credits to your account to use this node."
|
||||||
if status_code == 409:
|
elif status_code == 409:
|
||||||
error_message = "There is a problem with your account. Please contact support@comfy.org. "
|
error_message = "There is a problem with your account. Please contact support@comfy.org."
|
||||||
if status_code == 429:
|
elif status_code == 429:
|
||||||
error_message = "Rate Limit Exceeded: Please try again later."
|
error_message = "Rate Limit Exceeded: Please try again later."
|
||||||
|
|
||||||
raise Exception(error_message)
|
raise Exception(error_message)
|
||||||
|
|
||||||
# Parse and return JSON response
|
# Parse and return JSON response
|
||||||
@ -336,26 +527,76 @@ class ApiClient:
|
|||||||
upload_url: str,
|
upload_url: str,
|
||||||
file: io.BytesIO | str,
|
file: io.BytesIO | str,
|
||||||
content_type: str | None = None,
|
content_type: str | None = None,
|
||||||
|
max_retries: int = 3,
|
||||||
|
retry_delay: float = 1.0,
|
||||||
|
retry_backoff_factor: float = 2.0,
|
||||||
):
|
):
|
||||||
"""Upload a file to the API. Make sure the file has a filename equal to what the url expects.
|
"""Upload a file to the API with retry logic.
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
upload_url: The URL to upload to
|
upload_url: The URL to upload to
|
||||||
file: Either a file path string, BytesIO object, or tuple of (file_path, filename)
|
file: Either a file path string, BytesIO object, or tuple of (file_path, filename)
|
||||||
mime_type: Optional mime type to set for the upload
|
content_type: Optional mime type to set for the upload
|
||||||
|
max_retries: Maximum number of retry attempts
|
||||||
|
retry_delay: Initial delay between retries in seconds
|
||||||
|
retry_backoff_factor: Multiplier for the delay after each retry
|
||||||
"""
|
"""
|
||||||
headers = {}
|
headers = {}
|
||||||
if content_type:
|
if content_type:
|
||||||
headers["Content-Type"] = content_type
|
headers["Content-Type"] = content_type
|
||||||
|
|
||||||
|
# Prepare the file data
|
||||||
if isinstance(file, io.BytesIO):
|
if isinstance(file, io.BytesIO):
|
||||||
file.seek(0) # Ensure we're at the start of the file
|
file.seek(0) # Ensure we're at the start of the file
|
||||||
data = file.read()
|
data = file.read()
|
||||||
return requests.put(upload_url, data=data, headers=headers)
|
|
||||||
elif isinstance(file, str):
|
elif isinstance(file, str):
|
||||||
with open(file, "rb") as f:
|
with open(file, "rb") as f:
|
||||||
data = f.read()
|
data = f.read()
|
||||||
return requests.put(upload_url, data=data, headers=headers)
|
else:
|
||||||
|
raise ValueError("File must be either a BytesIO object or a file path string")
|
||||||
|
|
||||||
|
# Try the upload with retries
|
||||||
|
last_exception = None
|
||||||
|
for retry in range(max_retries + 1):
|
||||||
|
try:
|
||||||
|
response = requests.put(upload_url, data=data, headers=headers)
|
||||||
|
response.raise_for_status()
|
||||||
|
return response
|
||||||
|
|
||||||
|
except (requests.ConnectionError, requests.Timeout, requests.HTTPError) as e:
|
||||||
|
last_exception = e
|
||||||
|
if retry < max_retries:
|
||||||
|
delay = retry_delay * (retry_backoff_factor ** retry)
|
||||||
|
logging.warning(
|
||||||
|
f"File upload failed: {str(e)}. "
|
||||||
|
f"Retrying in {delay:.2f}s ({retry + 1}/{max_retries})"
|
||||||
|
)
|
||||||
|
time.sleep(delay)
|
||||||
|
else:
|
||||||
|
break
|
||||||
|
|
||||||
|
# If we've exhausted all retries, check if it's a network issue
|
||||||
|
try:
|
||||||
|
# Try to determine if it's a local network issue
|
||||||
|
check_response = requests.get("https://www.google.com", timeout=5.0)
|
||||||
|
if check_response.status_code < 500:
|
||||||
|
# Internet works but upload failed, likely a server or URL issue
|
||||||
|
raise Exception(
|
||||||
|
f"Failed to upload file after {max_retries} attempts. "
|
||||||
|
f"The upload service may be experiencing issues. Error: {str(last_exception)}"
|
||||||
|
) from last_exception
|
||||||
|
else:
|
||||||
|
# Even Google check failed, likely a local network issue
|
||||||
|
raise LocalNetworkError(
|
||||||
|
f"Failed to upload file due to network connectivity issues. "
|
||||||
|
f"Please check your internet connection and try again."
|
||||||
|
) from last_exception
|
||||||
|
except (requests.RequestException, socket.error):
|
||||||
|
# Could not reach Google, definitely a local network issue
|
||||||
|
raise LocalNetworkError(
|
||||||
|
f"Failed to upload file due to network connectivity issues. "
|
||||||
|
f"Please check your internet connection and try again."
|
||||||
|
) from last_exception
|
||||||
|
|
||||||
|
|
||||||
class ApiEndpoint(Generic[T, R]):
|
class ApiEndpoint(Generic[T, R]):
|
||||||
@ -403,6 +644,9 @@ class SynchronousOperation(Generic[T, R]):
|
|||||||
verify_ssl: bool = True,
|
verify_ssl: bool = True,
|
||||||
content_type: str = "application/json",
|
content_type: str = "application/json",
|
||||||
multipart_parser: Callable = None,
|
multipart_parser: Callable = None,
|
||||||
|
max_retries: int = 3,
|
||||||
|
retry_delay: float = 1.0,
|
||||||
|
retry_backoff_factor: float = 2.0,
|
||||||
):
|
):
|
||||||
self.endpoint = endpoint
|
self.endpoint = endpoint
|
||||||
self.request = request
|
self.request = request
|
||||||
@ -419,8 +663,12 @@ class SynchronousOperation(Generic[T, R]):
|
|||||||
self.files = files
|
self.files = files
|
||||||
self.content_type = content_type
|
self.content_type = content_type
|
||||||
self.multipart_parser = multipart_parser
|
self.multipart_parser = multipart_parser
|
||||||
|
self.max_retries = max_retries
|
||||||
|
self.retry_delay = retry_delay
|
||||||
|
self.retry_backoff_factor = retry_backoff_factor
|
||||||
|
|
||||||
def execute(self, client: Optional[ApiClient] = None) -> R:
|
def execute(self, client: Optional[ApiClient] = None) -> R:
|
||||||
"""Execute the API operation using the provided client or create one"""
|
"""Execute the API operation using the provided client or create one with retry support"""
|
||||||
try:
|
try:
|
||||||
# Create client if not provided
|
# Create client if not provided
|
||||||
if client is None:
|
if client is None:
|
||||||
@ -430,6 +678,9 @@ class SynchronousOperation(Generic[T, R]):
|
|||||||
comfy_api_key=self.comfy_api_key,
|
comfy_api_key=self.comfy_api_key,
|
||||||
timeout=self.timeout,
|
timeout=self.timeout,
|
||||||
verify_ssl=self.verify_ssl,
|
verify_ssl=self.verify_ssl,
|
||||||
|
max_retries=self.max_retries,
|
||||||
|
retry_delay=self.retry_delay,
|
||||||
|
retry_backoff_factor=self.retry_backoff_factor,
|
||||||
)
|
)
|
||||||
|
|
||||||
# Convert request model to dict, but use None for EmptyRequest
|
# Convert request model to dict, but use None for EmptyRequest
|
||||||
@ -443,11 +694,6 @@ class SynchronousOperation(Generic[T, R]):
|
|||||||
if isinstance(value, Enum):
|
if isinstance(value, Enum):
|
||||||
request_dict[key] = value.value
|
request_dict[key] = value.value
|
||||||
|
|
||||||
if request_dict:
|
|
||||||
for key, value in request_dict.items():
|
|
||||||
if isinstance(value, Enum):
|
|
||||||
request_dict[key] = value.value
|
|
||||||
|
|
||||||
# Debug log for request
|
# Debug log for request
|
||||||
logging.debug(
|
logging.debug(
|
||||||
f"[DEBUG] API Request: {self.endpoint.method.value} {self.endpoint.path}"
|
f"[DEBUG] API Request: {self.endpoint.method.value} {self.endpoint.path}"
|
||||||
@ -455,7 +701,7 @@ class SynchronousOperation(Generic[T, R]):
|
|||||||
logging.debug(f"[DEBUG] Request Data: {json.dumps(request_dict, indent=2)}")
|
logging.debug(f"[DEBUG] Request Data: {json.dumps(request_dict, indent=2)}")
|
||||||
logging.debug(f"[DEBUG] Query Params: {self.endpoint.query_params}")
|
logging.debug(f"[DEBUG] Query Params: {self.endpoint.query_params}")
|
||||||
|
|
||||||
# Make the request
|
# Make the request with built-in retry
|
||||||
resp = client.request(
|
resp = client.request(
|
||||||
method=self.endpoint.method.value,
|
method=self.endpoint.method.value,
|
||||||
path=self.endpoint.path,
|
path=self.endpoint.path,
|
||||||
@ -476,8 +722,18 @@ class SynchronousOperation(Generic[T, R]):
|
|||||||
# Parse and return the response
|
# Parse and return the response
|
||||||
return self._parse_response(resp)
|
return self._parse_response(resp)
|
||||||
|
|
||||||
|
except LocalNetworkError as e:
|
||||||
|
# Propagate specific network error types
|
||||||
|
logging.error(f"[ERROR] Local network error: {str(e)}")
|
||||||
|
raise
|
||||||
|
|
||||||
|
except ApiServerError as e:
|
||||||
|
# Propagate API server errors
|
||||||
|
logging.error(f"[ERROR] API server error: {str(e)}")
|
||||||
|
raise
|
||||||
|
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logging.error(f"[DEBUG] API Exception: {str(e)}")
|
logging.error(f"[ERROR] API Exception: {str(e)}")
|
||||||
raise Exception(str(e))
|
raise Exception(str(e))
|
||||||
|
|
||||||
def _parse_response(self, resp):
|
def _parse_response(self, resp):
|
||||||
@ -517,6 +773,10 @@ class PollingOperation(Generic[T, R]):
|
|||||||
comfy_api_key: Optional[str] = None,
|
comfy_api_key: Optional[str] = None,
|
||||||
auth_kwargs: Optional[Dict[str,str]] = None,
|
auth_kwargs: Optional[Dict[str,str]] = None,
|
||||||
poll_interval: float = 5.0,
|
poll_interval: float = 5.0,
|
||||||
|
max_poll_attempts: int = 120, # Default max polling attempts (10 minutes with 5s interval)
|
||||||
|
max_retries: int = 3, # Max retries per individual API call
|
||||||
|
retry_delay: float = 1.0,
|
||||||
|
retry_backoff_factor: float = 2.0,
|
||||||
):
|
):
|
||||||
self.poll_endpoint = poll_endpoint
|
self.poll_endpoint = poll_endpoint
|
||||||
self.request = request
|
self.request = request
|
||||||
@ -527,6 +787,10 @@ class PollingOperation(Generic[T, R]):
|
|||||||
self.auth_token = auth_kwargs.get("auth_token", self.auth_token)
|
self.auth_token = auth_kwargs.get("auth_token", self.auth_token)
|
||||||
self.comfy_api_key = auth_kwargs.get("comfy_api_key", self.comfy_api_key)
|
self.comfy_api_key = auth_kwargs.get("comfy_api_key", self.comfy_api_key)
|
||||||
self.poll_interval = poll_interval
|
self.poll_interval = poll_interval
|
||||||
|
self.max_poll_attempts = max_poll_attempts
|
||||||
|
self.max_retries = max_retries
|
||||||
|
self.retry_delay = retry_delay
|
||||||
|
self.retry_backoff_factor = retry_backoff_factor
|
||||||
|
|
||||||
# Polling configuration
|
# Polling configuration
|
||||||
self.status_extractor = status_extractor or (
|
self.status_extractor = status_extractor or (
|
||||||
@ -546,10 +810,29 @@ class PollingOperation(Generic[T, R]):
|
|||||||
if client is None:
|
if client is None:
|
||||||
client = ApiClient(
|
client = ApiClient(
|
||||||
base_url=self.api_base,
|
base_url=self.api_base,
|
||||||
|
<<<<<<< HEAD
|
||||||
auth_token=self.auth_token,
|
auth_token=self.auth_token,
|
||||||
comfy_api_key=self.comfy_api_key,
|
comfy_api_key=self.comfy_api_key,
|
||||||
|
=======
|
||||||
|
api_key=self.auth_token,
|
||||||
|
max_retries=self.max_retries,
|
||||||
|
retry_delay=self.retry_delay,
|
||||||
|
retry_backoff_factor=self.retry_backoff_factor,
|
||||||
|
>>>>>>> a998cdb6 (rework how errors are handled on the client side)
|
||||||
)
|
)
|
||||||
return self._poll_until_complete(client)
|
return self._poll_until_complete(client)
|
||||||
|
except LocalNetworkError as e:
|
||||||
|
# Provide clear message for local network issues
|
||||||
|
raise Exception(
|
||||||
|
f"Polling failed due to local network issues. Please check your internet connection. "
|
||||||
|
f"Details: {str(e)}"
|
||||||
|
) from e
|
||||||
|
except ApiServerError as e:
|
||||||
|
# Provide clear message for API server issues
|
||||||
|
raise Exception(
|
||||||
|
f"Polling failed due to API server issues. The service may be experiencing problems. "
|
||||||
|
f"Please try again later. Details: {str(e)}"
|
||||||
|
) from e
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
raise Exception(f"Error during polling: {str(e)}")
|
raise Exception(f"Error during polling: {str(e)}")
|
||||||
|
|
||||||
@ -569,10 +852,13 @@ class PollingOperation(Generic[T, R]):
|
|||||||
def _poll_until_complete(self, client: ApiClient) -> R:
|
def _poll_until_complete(self, client: ApiClient) -> R:
|
||||||
"""Poll until the task is complete"""
|
"""Poll until the task is complete"""
|
||||||
poll_count = 0
|
poll_count = 0
|
||||||
|
consecutive_errors = 0
|
||||||
|
max_consecutive_errors = min(5, self.max_retries * 2) # Limit consecutive errors
|
||||||
|
|
||||||
if self.progress_extractor:
|
if self.progress_extractor:
|
||||||
progress = utils.ProgressBar(PROGRESS_BAR_MAX)
|
progress = utils.ProgressBar(PROGRESS_BAR_MAX)
|
||||||
|
|
||||||
while True:
|
while poll_count < self.max_poll_attempts:
|
||||||
try:
|
try:
|
||||||
poll_count += 1
|
poll_count += 1
|
||||||
logging.debug(f"[DEBUG] Polling attempt #{poll_count}")
|
logging.debug(f"[DEBUG] Polling attempt #{poll_count}")
|
||||||
@ -599,8 +885,12 @@ class PollingOperation(Generic[T, R]):
|
|||||||
data=request_dict,
|
data=request_dict,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
# Successfully got a response, reset consecutive error count
|
||||||
|
consecutive_errors = 0
|
||||||
|
|
||||||
# Parse response
|
# Parse response
|
||||||
response_obj = self.poll_endpoint.response_model.model_validate(resp)
|
response_obj = self.poll_endpoint.response_model.model_validate(resp)
|
||||||
|
|
||||||
# Check if task is complete
|
# Check if task is complete
|
||||||
status = self._check_task_status(response_obj)
|
status = self._check_task_status(response_obj)
|
||||||
logging.debug(f"[DEBUG] Task Status: {status}")
|
logging.debug(f"[DEBUG] Task Status: {status}")
|
||||||
@ -630,6 +920,38 @@ class PollingOperation(Generic[T, R]):
|
|||||||
)
|
)
|
||||||
time.sleep(self.poll_interval)
|
time.sleep(self.poll_interval)
|
||||||
|
|
||||||
|
except (LocalNetworkError, ApiServerError) as e:
|
||||||
|
# For network-related errors, increment error count and potentially abort
|
||||||
|
consecutive_errors += 1
|
||||||
|
if consecutive_errors >= max_consecutive_errors:
|
||||||
|
raise Exception(
|
||||||
|
f"Polling aborted after {consecutive_errors} consecutive network errors: {str(e)}"
|
||||||
|
) from e
|
||||||
|
|
||||||
|
# Log the error but continue polling
|
||||||
|
logging.warning(
|
||||||
|
f"Network error during polling (attempt {poll_count}/{self.max_poll_attempts}): {str(e)}. "
|
||||||
|
f"Will retry in {self.poll_interval} seconds."
|
||||||
|
)
|
||||||
|
time.sleep(self.poll_interval)
|
||||||
|
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
|
# For other errors, increment count and potentially abort
|
||||||
|
consecutive_errors += 1
|
||||||
|
if consecutive_errors >= max_consecutive_errors:
|
||||||
|
raise Exception(
|
||||||
|
f"Polling aborted after {consecutive_errors} consecutive errors: {str(e)}"
|
||||||
|
) from e
|
||||||
|
|
||||||
logging.error(f"[DEBUG] Polling error: {str(e)}")
|
logging.error(f"[DEBUG] Polling error: {str(e)}")
|
||||||
raise Exception(f"Error while polling: {str(e)}")
|
logging.warning(
|
||||||
|
f"Error during polling (attempt {poll_count}/{self.max_poll_attempts}): {str(e)}. "
|
||||||
|
f"Will retry in {self.poll_interval} seconds."
|
||||||
|
)
|
||||||
|
time.sleep(self.poll_interval)
|
||||||
|
|
||||||
|
# If we've exhausted all polling attempts
|
||||||
|
raise Exception(
|
||||||
|
f"Polling timed out after {poll_count} attempts ({poll_count * self.poll_interval} seconds). "
|
||||||
|
f"The operation may still be running on the server but is taking longer than expected."
|
||||||
|
)
|
||||||
Loading…
x
Reference in New Issue
Block a user