Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
171 changes: 99 additions & 72 deletions dashscope/api_entities/api_request_factory.py
Original file line number Diff line number Diff line change
@@ -1,45 +1,123 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
from typing import Any, Dict, Optional, Union
from urllib.parse import urlencode

import aiohttp
import requests

import dashscope
from dashscope.api_entities.api_request_data import ApiRequestData
from dashscope.api_entities.encryption import Encryption
from dashscope.api_entities.http_request import HttpRequest
from dashscope.api_entities.websocket_request import WebSocketRequest
from dashscope.common.constants import (
REQUEST_TIMEOUT_KEYWORD,
SERVICE_API_PATH,
ApiProtocol,
HTTPMethod,
)
from dashscope.common.error import InputDataRequired, UnsupportedApiProtocol
from dashscope.common.logging import logger
from dashscope.protocol.websocket import WebsocketStreamingMode
from dashscope.api_entities.encryption import Encryption


def _get_protocol_params(kwargs):
api_protocol = kwargs.pop("api_protocol", ApiProtocol.HTTPS)
ws_stream_mode = kwargs.pop("ws_stream_mode", WebsocketStreamingMode.OUT)
is_binary_input = kwargs.pop("is_binary_input", False)
http_method = kwargs.pop("http_method", HTTPMethod.POST)
stream = kwargs.get("stream", False)
def _build_api_request( # pylint: disable=too-many-branches
# pylint: disable=too-many-arguments,too-many-locals
model: str,
input: object, # pylint: disable=redefined-builtin
task_group: str,
task: str,
function: str,
api_key: str,
is_service: bool = True,
# Protocol and connection configuration
api_protocol: ApiProtocol = ApiProtocol.HTTPS,
http_method: HTTPMethod = HTTPMethod.POST,
stream: bool = False,
async_request: bool = False,
request_timeout: Optional[int] = None,
# WebSocket specific
ws_stream_mode: WebsocketStreamingMode = WebsocketStreamingMode.OUT,
is_binary_input: bool = False,
# HTTP specific
query: bool = False,
headers: Optional[Dict[str, str]] = None,
form: Optional[Dict] = None,
resources: Optional[Dict] = None,
base_address: Optional[str] = None,
flattened_output: bool = False,
extra_url_parameters: Optional[Dict[str, Any]] = None,
user_agent: str = "",
session: Optional[Union[requests.Session, aiohttp.ClientSession]] = None,
task_id: Optional[str] = None,
enable_encryption: bool = False,
pre_task_id: Optional[str] = None,
# Additional parameters for API request data
**kwargs,
):
# pylint: disable=too-many-statements
"""Build API request object.

Args:
model (str): The model name.
input (object): The input data for the request.
task_group (str): The task group for the API path.
task (str): The task name for the API path.
function (str): The function name for the API path.
api_key (str): The API key for authentication.
is_service (bool, optional): Whether this is a service call.
Defaults to True.
api_protocol (ApiProtocol, optional): The protocol to use
(HTTP, HTTPS, WEBSOCKET). Defaults to ApiProtocol.HTTPS.
http_method (HTTPMethod, optional): The HTTP method (GET, POST).
Defaults to HTTPMethod.POST.
stream (bool, optional): Enable streaming output.
Defaults to False.
async_request (bool, optional): Enable async request.
Defaults to False.
request_timeout (int, optional): Request timeout in seconds.
Defaults to None.
ws_stream_mode (WebsocketStreamingMode, optional): WebSocket
streaming mode. Defaults to WebsocketStreamingMode.OUT.
is_binary_input (bool, optional): Whether input is binary data.
Defaults to False.
query (bool, optional): Whether this is a query request.
Defaults to False.
headers (Dict[str, str], optional): Additional HTTP headers.
Defaults to None.
form (Dict, optional): Form data for multipart requests.
Defaults to None.
resources (Dict, optional): Resource data. Defaults to None.
base_address (str, optional): Custom base URL for the API.
Defaults to None.
flattened_output (bool, optional): Whether to flatten output.
Defaults to False.
extra_url_parameters (Dict[str, Any], optional): Extra URL query
parameters. Defaults to None.
user_agent (str, optional): Custom user agent string.
Defaults to "".
session (Union[requests.Session, aiohttp.ClientSession], optional):
Custom session for connection reuse. Defaults to None.
task_id (str, optional): Task ID for the request.
Defaults to None.
enable_encryption (bool, optional): Enable request encryption.
Defaults to False.
pre_task_id (str, optional): Previous task ID for WebSocket.
Defaults to None.
**kwargs: Additional parameters passed to the API request data.

Returns:
HttpRequest or WebSocketRequest: The constructed request object.

Raises:
InputDataRequired: If input data is missing or invalid.
UnsupportedApiProtocol: If the API protocol is not supported.
"""
# Handle stream mode for WebSocket
if not stream and ws_stream_mode == WebsocketStreamingMode.OUT:
ws_stream_mode = WebsocketStreamingMode.NONE

async_request = kwargs.pop("async_request", False)
query = kwargs.pop("query", False)
headers = kwargs.pop("headers", None)
request_timeout = kwargs.pop(REQUEST_TIMEOUT_KEYWORD, None)
form = kwargs.pop("form", None)
resources = kwargs.pop("resources", None)
base_address = kwargs.pop("base_address", None)
flattened_output = kwargs.pop("flattened_output", False)
extra_url_parameters = kwargs.pop("extra_url_parameters", None)
session = kwargs.pop("session", None)

# Extract user_agent from kwargs (preferred) or from headers["user-agent"]
user_agent = kwargs.pop("user_agent", "")
# Handle user_agent from headers
if headers and "user-agent" in headers:
header_ua = headers.pop("user-agent")
if user_agent:
Expand All @@ -49,56 +127,6 @@ def _get_protocol_params(kwargs):
else:
user_agent = header_ua

return (
api_protocol,
ws_stream_mode,
is_binary_input,
http_method,
stream,
async_request,
query,
headers,
request_timeout,
form,
resources,
base_address,
flattened_output,
extra_url_parameters,
user_agent,
session,
)


def _build_api_request( # pylint: disable=too-many-branches
model: str,
input: object, # pylint: disable=redefined-builtin
task_group: str,
task: str,
function: str,
api_key: str,
is_service=True,
**kwargs,
):
(
api_protocol,
ws_stream_mode,
is_binary_input,
http_method,
stream,
async_request,
query,
headers,
request_timeout,
form,
resources,
base_address,
flattened_output,
extra_url_parameters,
user_agent,
session,
) = _get_protocol_params(kwargs)
task_id = kwargs.pop("task_id", None)
enable_encryption = kwargs.pop("enable_encryption", False)
encryption = None

if api_protocol in [ApiProtocol.HTTP, ApiProtocol.HTTPS]:
Expand Down Expand Up @@ -146,7 +174,6 @@ def _build_api_request( # pylint: disable=too-many-branches
websocket_url = base_address
else:
websocket_url = dashscope.base_websocket_api_url
pre_task_id = kwargs.pop("pre_task_id", None)
request = WebSocketRequest(
url=websocket_url,
api_key=api_key,
Expand Down
Loading
Loading