Set request type explicitly (#66)

This commit is contained in:
Christian Byrne 2025-04-30 11:19:57 -07:00 committed by GitHub
parent f1eb74fee9
commit 7d3f693c00
No known key found for this signature in database
GPG Key ID: B5690EEEBB952194
2 changed files with 65 additions and 29 deletions

View File

@ -165,6 +165,45 @@ class ApiClient:
self.timeout = timeout self.timeout = timeout
self.verify_ssl = verify_ssl self.verify_ssl = verify_ssl
def _create_json_payload_args(
self,
data: Optional[Dict[str, Any]] = None,
headers: Optional[Dict[str, str]] = None,
) -> Dict[str, Any]:
return {
"json": data,
"headers": headers,
}
def _create_form_data_args(
self,
data: Dict[str, Any],
files: Dict[str, Any],
headers: Optional[Dict[str, str]] = None,
) -> Dict[str, Any]:
if headers:
del headers["Content-Type"]
return {
"data": data,
"files": files,
"headers": headers,
}
def _create_urlencoded_form_data_args(
self,
data: Dict[str, Any],
headers: Optional[Dict[str, str]] = None,
) -> Dict[str, Any]:
headers = headers or {}
headers["Content-Type"] = "application/x-www-form-urlencoded"
return {
"data": data,
"headers": headers,
}
def get_headers(self) -> Dict[str, str]: def get_headers(self) -> Dict[str, str]:
"""Get headers for API requests, including authentication if available""" """Get headers for API requests, including authentication if available"""
headers = {"Content-Type": "application/json", "Accept": "application/json"} headers = {"Content-Type": "application/json", "Accept": "application/json"}
@ -179,9 +218,10 @@ class ApiClient:
method: str, method: str,
path: str, path: str,
params: Optional[Dict[str, Any]] = None, params: Optional[Dict[str, Any]] = None,
json: Optional[Dict[str, Any]] = None, data: Optional[Dict[str, Any]] = None,
files: Optional[Dict[str, Any]] = None, files: Optional[Dict[str, Any]] = None,
headers: Optional[Dict[str, str]] = None, headers: Optional[Dict[str, str]] = None,
content_type: str = "application/json",
) -> Dict[str, Any]: ) -> Dict[str, Any]:
""" """
Make an HTTP request to the API Make an HTTP request to the API
@ -190,9 +230,10 @@ class ApiClient:
method: HTTP method (GET, POST, etc.) method: HTTP method (GET, POST, etc.)
path: API endpoint path (will be joined with base_url) path: API endpoint path (will be joined with base_url)
params: Query parameters params: Query parameters
json: JSON body data data: body data
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.
Returns: Returns:
Parsed JSON response Parsed JSON response
@ -214,34 +255,25 @@ class ApiClient:
logging.debug(f"[DEBUG] Request Headers: {request_headers}") logging.debug(f"[DEBUG] Request Headers: {request_headers}")
logging.debug(f"[DEBUG] Files: {files}") logging.debug(f"[DEBUG] Files: {files}")
logging.debug(f"[DEBUG] Params: {params}") logging.debug(f"[DEBUG] Params: {params}")
logging.debug(f"[DEBUG] Json: {json}") logging.debug(f"[DEBUG] Data: {data}")
match content_type:
case "application/x-www-form-urlencoded":
payload_args = self._create_urlencoded_form_data_args(data, request_headers)
case "multipart/form-data":
payload_args = self._create_form_data_args(data, files, request_headers)
case _:
payload_args = self._create_json_payload_args(data, request_headers)
try: try:
# If files are present, use data parameter instead of json response = requests.request(
if files: method=method,
form_data = {} url=url,
if json: params=params,
form_data.update(json) timeout=self.timeout,
response = requests.request( verify=self.verify_ssl,
method=method, **payload_args,
url=url, )
params=params,
data=form_data, # Use data instead of json
files=files,
headers=request_headers,
timeout=self.timeout,
verify=self.verify_ssl,
)
else:
response = requests.request(
method=method,
url=url,
params=params,
json=json,
headers=request_headers,
timeout=self.timeout,
verify=self.verify_ssl,
)
# Raise exception for error status codes # Raise exception for error status codes
response.raise_for_status() response.raise_for_status()
@ -367,6 +399,7 @@ class SynchronousOperation(Generic[T, R]):
auth_token: Optional[str] = None, auth_token: Optional[str] = None,
timeout: float = 604800.0, timeout: float = 604800.0,
verify_ssl: bool = True, verify_ssl: bool = True,
content_type: str = "application/json",
): ):
self.endpoint = endpoint self.endpoint = endpoint
self.request = request self.request = request
@ -377,6 +410,7 @@ class SynchronousOperation(Generic[T, R]):
self.timeout = timeout self.timeout = timeout
self.verify_ssl = verify_ssl self.verify_ssl = verify_ssl
self.files = files self.files = files
self.content_type = content_type
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"""
@ -408,9 +442,10 @@ class SynchronousOperation(Generic[T, R]):
resp = client.request( resp = client.request(
method=self.endpoint.method.value, method=self.endpoint.method.value,
path=self.endpoint.path, path=self.endpoint.path,
json=request_dict, data=request_dict,
params=self.endpoint.query_params, params=self.endpoint.query_params,
files=self.files, files=self.files,
content_type=self.content_type,
) )
# Debug log for response # Debug log for response

View File

@ -458,6 +458,7 @@ class OpenAIGPTImage1(ComfyNodeABC):
size=size, size=size,
), ),
files=files if files else None, files=files if files else None,
content_type="multipart/form-data",
auth_token=auth_token, auth_token=auth_token,
) )