2239 lines
78 KiB
Python
2239 lines
78 KiB
Python
# Copyright 2026 Google LLC
|
|
#
|
|
# Licensed under the Apache License, Version 2.0 (the "License");
|
|
# you may not use this file except in compliance with the License.
|
|
# You may obtain a copy of the License at
|
|
#
|
|
# http://www.apache.org/licenses/LICENSE-2.0
|
|
#
|
|
# Unless required by applicable law or agreed to in writing, software
|
|
# distributed under the License is distributed on an "AS IS" BASIS,
|
|
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
|
# See the License for the specific language governing permissions and
|
|
# limitations under the License.
|
|
#
|
|
|
|
|
|
"""Base client for calling HTTP APIs sending and receiving JSON.
|
|
|
|
The BaseApiClient is intended to be a private module and is subject to change.
|
|
"""
|
|
|
|
import asyncio
|
|
from collections.abc import Generator
|
|
import copy
|
|
from dataclasses import dataclass
|
|
import inspect
|
|
import io
|
|
import json
|
|
import logging
|
|
import math
|
|
import os
|
|
import random
|
|
import ssl
|
|
import sys
|
|
import threading
|
|
import time
|
|
from typing import Any, AsyncIterator, Iterator, Optional, TYPE_CHECKING, Tuple, Union
|
|
from urllib.parse import urlparse
|
|
from urllib.parse import urlunparse
|
|
import warnings
|
|
|
|
import anyio
|
|
import certifi
|
|
import google.auth
|
|
import google.auth.credentials
|
|
from google.auth.credentials import Credentials
|
|
from google.auth.transport import mtls
|
|
from google.auth.transport.requests import AuthorizedSession
|
|
from google.auth import exceptions as auth_exceptions
|
|
import httpx
|
|
from pydantic import BaseModel
|
|
from pydantic import ValidationError
|
|
import requests
|
|
from requests.structures import CaseInsensitiveDict
|
|
import tenacity
|
|
|
|
from . import _common
|
|
from . import errors
|
|
from . import version
|
|
from .types import HttpOptions
|
|
from .types import HttpOptionsOrDict
|
|
from .types import HttpResponse as SdkHttpResponse
|
|
from .types import HttpRetryOptions
|
|
from .types import ResourceScope
|
|
|
|
|
|
try:
|
|
from websockets.asyncio.client import connect as ws_connect
|
|
except ModuleNotFoundError:
|
|
# This try/except is for TAP, mypy complains about it which is why we have the type: ignore
|
|
from websockets.client import connect as ws_connect # type: ignore
|
|
|
|
|
|
has_aiohttp = False
|
|
try:
|
|
import aiohttp
|
|
|
|
has_aiohttp = True
|
|
except ImportError:
|
|
pass
|
|
|
|
|
|
if TYPE_CHECKING:
|
|
from multidict import CIMultiDictProxy
|
|
from google.auth.aio.transport.sessions import AsyncAuthorizedSession
|
|
|
|
|
|
logger = logging.getLogger('google_genai._api_client')
|
|
CHUNK_SIZE = 8 * 1024 * 1024 # 8 MB chunk size
|
|
READ_BUFFER_SIZE = 2**22
|
|
MAX_RETRY_COUNT = 3
|
|
INITIAL_RETRY_DELAY = 1 # second
|
|
DELAY_MULTIPLIER = 2
|
|
|
|
_MULTI_REGIONAL_LOCATIONS = {'us', 'eu'}
|
|
|
|
|
|
class EphemeralTokenAPIKeyError(ValueError):
|
|
"""Error raised when the API key is invalid."""
|
|
|
|
|
|
# This method checks for the API key in the environment variables. Google API
|
|
# key is precedenced over Gemini API key.
|
|
def get_env_api_key() -> Optional[str]:
|
|
"""Gets the API key from environment variables, prioritizing GOOGLE_API_KEY.
|
|
|
|
Returns:
|
|
The API key string if found, otherwise None. Empty string is considered
|
|
invalid.
|
|
"""
|
|
env_google_api_key = os.environ.get('GOOGLE_API_KEY', None)
|
|
env_gemini_api_key = os.environ.get('GEMINI_API_KEY', None)
|
|
if env_google_api_key and env_gemini_api_key:
|
|
logger.warning(
|
|
'Both GOOGLE_API_KEY and GEMINI_API_KEY are set. Using GOOGLE_API_KEY.'
|
|
)
|
|
|
|
return env_google_api_key or env_gemini_api_key or None
|
|
|
|
|
|
def append_library_version_headers(headers: dict[str, str]) -> None:
|
|
"""Appends the telemetry header to the headers dict."""
|
|
library_label = f'google-genai-sdk/{version.__version__}'
|
|
language_label = 'gl-python/' + sys.version.split()[0]
|
|
version_header_value = f'{library_label} {language_label}'
|
|
if (
|
|
'user-agent' in headers
|
|
and library_label not in headers['user-agent']
|
|
):
|
|
headers['user-agent'] = f'{version_header_value} ' + headers['user-agent']
|
|
elif 'user-agent' in headers and language_label not in headers['user-agent']:
|
|
headers['user-agent'] = f'{headers["user-agent"]} {language_label}'
|
|
elif 'user-agent' not in headers:
|
|
headers['user-agent'] = version_header_value
|
|
if (
|
|
'x-goog-api-client' in headers
|
|
and library_label not in headers['x-goog-api-client']
|
|
):
|
|
headers['x-goog-api-client'] = (
|
|
f'{version_header_value} ' + headers['x-goog-api-client']
|
|
)
|
|
elif (
|
|
'x-goog-api-client' in headers
|
|
and language_label not in headers['x-goog-api-client']
|
|
):
|
|
headers['x-goog-api-client'] = (
|
|
f"{headers['x-goog-api-client']} {language_label}"
|
|
)
|
|
elif 'x-goog-api-client' not in headers:
|
|
headers['x-goog-api-client'] = version_header_value
|
|
|
|
|
|
def patch_http_options(
|
|
options: HttpOptions, patch_options: HttpOptions
|
|
) -> HttpOptions:
|
|
copy_option = options.model_copy()
|
|
|
|
options_headers = copy_option.headers or {}
|
|
patch_options_headers = patch_options.headers or {}
|
|
copy_option.headers = {
|
|
**options_headers,
|
|
**patch_options_headers,
|
|
}
|
|
|
|
http_options_keys = HttpOptions.model_fields.keys()
|
|
|
|
for key in http_options_keys:
|
|
if key == 'headers':
|
|
continue
|
|
patch_value = getattr(patch_options, key, None)
|
|
if patch_value is not None:
|
|
setattr(copy_option, key, patch_value)
|
|
else:
|
|
setattr(copy_option, key, getattr(options, key))
|
|
|
|
if copy_option.headers is not None:
|
|
append_library_version_headers(copy_option.headers)
|
|
return copy_option
|
|
|
|
|
|
def populate_server_timeout_header(
|
|
headers: dict[str, str], timeout_in_seconds: Optional[Union[float, int]]
|
|
) -> None:
|
|
"""Populates the server timeout header in the headers dict."""
|
|
if timeout_in_seconds and 'X-Server-Timeout' not in headers:
|
|
headers['X-Server-Timeout'] = str(math.ceil(timeout_in_seconds))
|
|
|
|
|
|
def join_url_path(base_url: str, path: str) -> str:
|
|
parsed_base = urlparse(base_url)
|
|
base_path = (
|
|
parsed_base.path[:-1]
|
|
if parsed_base.path.endswith('/')
|
|
else parsed_base.path
|
|
)
|
|
path = path[1:] if path.startswith('/') else path
|
|
return urlunparse(parsed_base._replace(path=base_path + '/' + path))
|
|
|
|
|
|
def load_auth(*, project: Union[str, None]) -> Tuple[Credentials, str]:
|
|
"""Loads google auth credentials and project id."""
|
|
credentials, loaded_project_id = google.auth.default( # type: ignore[no-untyped-call]
|
|
scopes=['https://www.googleapis.com/auth/cloud-platform'],
|
|
)
|
|
|
|
if not project:
|
|
project = loaded_project_id
|
|
|
|
if not project:
|
|
raise ValueError(
|
|
'Could not resolve project using application default credentials.'
|
|
)
|
|
|
|
return credentials, project
|
|
|
|
|
|
def refresh_auth(credentials: Credentials) -> Credentials:
|
|
from google.auth.transport.requests import Request
|
|
credentials.refresh(Request()) # type: ignore[no-untyped-call]
|
|
return credentials
|
|
|
|
|
|
def get_timeout_in_seconds(
|
|
timeout: Optional[Union[float, int]],
|
|
) -> Optional[float]:
|
|
"""Converts the timeout to seconds."""
|
|
if timeout:
|
|
# HttpOptions.timeout is in milliseconds. But httpx.Client.request()
|
|
# expects seconds.
|
|
timeout_in_seconds = timeout / 1000.0
|
|
else:
|
|
timeout_in_seconds = None
|
|
return timeout_in_seconds
|
|
|
|
|
|
@dataclass
|
|
class HttpRequest:
|
|
headers: dict[str, str]
|
|
url: str
|
|
method: str
|
|
data: Union[dict[str, object], bytes]
|
|
timeout: Optional[float] = None
|
|
|
|
|
|
class HttpResponse:
|
|
|
|
def __init__(
|
|
self,
|
|
headers: Union[
|
|
dict[str, str],
|
|
httpx.Headers,
|
|
'CIMultiDictProxy[str]',
|
|
CaseInsensitiveDict,
|
|
],
|
|
response_stream: Union[Any, str] = None,
|
|
byte_stream: Union[Any, bytes] = None,
|
|
):
|
|
if isinstance(headers, dict):
|
|
self.headers = headers
|
|
elif isinstance(headers, httpx.Headers):
|
|
self.headers = {
|
|
key: ', '.join(headers.get_list(key)) for key in headers.keys()
|
|
}
|
|
elif isinstance(headers, CaseInsensitiveDict):
|
|
self.headers = {key: value for key, value in headers.items()}
|
|
elif type(headers).__name__ == 'CIMultiDictProxy':
|
|
self.headers = {
|
|
key: ', '.join(headers.getall(key)) for key in headers.keys()
|
|
}
|
|
|
|
self.status_code: int = 200
|
|
self.response_stream = response_stream
|
|
self.byte_stream = byte_stream
|
|
|
|
# Async iterator for async streaming.
|
|
def __aiter__(self) -> 'HttpResponse':
|
|
self.segment_iterator = self.async_segments()
|
|
return self
|
|
|
|
async def __anext__(self) -> Any:
|
|
try:
|
|
return await self.segment_iterator.__anext__()
|
|
except StopIteration:
|
|
raise StopAsyncIteration
|
|
|
|
@property
|
|
def json(self) -> Any:
|
|
# Handle case where response_stream is not a list (e.g., aiohttp.ClientResponse)
|
|
# This can happen when the API returns an error and the response object
|
|
# is passed directly instead of being wrapped in a list.
|
|
# See: https://github.com/googleapis/python-genai/issues/1897
|
|
if not isinstance(self.response_stream, list):
|
|
return None
|
|
if not self.response_stream or not self.response_stream[0]: # Empty response
|
|
return ''
|
|
return self._load_json_from_response(self.response_stream[0])
|
|
|
|
def segments(self) -> Generator[Any, None, None]:
|
|
if isinstance(self.response_stream, list):
|
|
# list of objects retrieved from replay or from non-streaming API.
|
|
for chunk in self.response_stream:
|
|
yield self._load_json_from_response(chunk) if chunk else {}
|
|
elif self.response_stream is None:
|
|
yield from []
|
|
else:
|
|
# Iterator of objects retrieved from the API.
|
|
for chunk in self._iter_response_stream():
|
|
yield self._load_json_from_response(chunk)
|
|
|
|
async def async_segments(self) -> AsyncIterator[Any]:
|
|
if isinstance(self.response_stream, list):
|
|
# list of objects retrieved from replay or from non-streaming API.
|
|
for chunk in self.response_stream:
|
|
yield self._load_json_from_response(chunk) if chunk else {}
|
|
elif self.response_stream is None:
|
|
async for c in []: # type: ignore[attr-defined]
|
|
yield c
|
|
else:
|
|
# Iterator of objects retrieved from the API.
|
|
async for chunk in self._aiter_response_stream():
|
|
yield self._load_json_from_response(chunk)
|
|
|
|
def byte_segments(self) -> Generator[Union[bytes, Any], None, None]:
|
|
if isinstance(self.byte_stream, list):
|
|
# list of objects retrieved from replay or from non-streaming API.
|
|
yield from self.byte_stream
|
|
elif self.byte_stream is None:
|
|
yield from []
|
|
else:
|
|
raise ValueError(
|
|
'Byte segments are not supported for streaming responses.'
|
|
)
|
|
|
|
def _copy_to_dict(self, response_payload: dict[str, object]) -> None:
|
|
# Cannot pickle 'generator' object.
|
|
delattr(self, 'segment_iterator')
|
|
for attribute in dir(self):
|
|
response_payload[attribute] = copy.deepcopy(getattr(self, attribute))
|
|
|
|
def _iter_response_stream(self) -> Iterator[str]:
|
|
"""Iterates over chunks retrieved from the API."""
|
|
if not (
|
|
isinstance(self.response_stream, httpx.Response)
|
|
or isinstance(self.response_stream, requests.Response)
|
|
):
|
|
raise TypeError(
|
|
'Expected self.response_stream to be an httpx.Response object, '
|
|
f'but got {type(self.response_stream).__name__}.'
|
|
)
|
|
|
|
chunk = ''
|
|
balance = 0
|
|
data_buffer: list[str] = []
|
|
if isinstance(self.response_stream, httpx.Response):
|
|
response_stream = self.response_stream.iter_lines()
|
|
else:
|
|
response_stream = self.response_stream.iter_lines(decode_unicode=True)
|
|
for line in response_stream:
|
|
if not line:
|
|
if data_buffer:
|
|
yield '\n'.join(data_buffer)
|
|
data_buffer = []
|
|
continue
|
|
|
|
# In streaming mode, the response of JSON is prefixed with "data: " which
|
|
# we must strip before parsing.
|
|
if line.startswith('data: '):
|
|
data_buffer.append(line[len('data: '):])
|
|
continue
|
|
|
|
# When API returns an error message, it comes line by line. So we buffer
|
|
# the lines until a complete JSON string is read. A complete JSON string
|
|
# is found when the balance is 0.
|
|
for c in line:
|
|
if c == '{':
|
|
balance += 1
|
|
elif c == '}':
|
|
balance -= 1
|
|
|
|
chunk += line
|
|
if balance == 0:
|
|
yield chunk
|
|
chunk = ''
|
|
|
|
# If there is any remaining chunk, yield it.
|
|
if chunk:
|
|
yield chunk
|
|
if data_buffer:
|
|
yield '\n'.join(data_buffer)
|
|
|
|
async def _aiter_response_stream(self) -> AsyncIterator[str]:
|
|
"""Asynchronously iterates over chunks retrieved from the API."""
|
|
is_valid_response = isinstance(self.response_stream, httpx.Response) or (
|
|
has_aiohttp and isinstance(self.response_stream, aiohttp.ClientResponse)
|
|
)
|
|
if not is_valid_response:
|
|
raise TypeError(
|
|
'Expected self.response_stream to be an httpx.Response or'
|
|
' aiohttp.ClientResponse object, but got'
|
|
f' {type(self.response_stream).__name__}.'
|
|
)
|
|
|
|
chunk = ''
|
|
balance = 0
|
|
data_buffer: list[str] = []
|
|
# httpx.Response has a dedicated async line iterator.
|
|
if isinstance(self.response_stream, httpx.Response):
|
|
try:
|
|
async for line in self.response_stream.aiter_lines():
|
|
if not line:
|
|
if data_buffer:
|
|
yield '\n'.join(data_buffer)
|
|
data_buffer = []
|
|
continue
|
|
# In streaming mode, the response of JSON is prefixed with "data: "
|
|
# which we must strip before parsing.
|
|
if line.startswith('data: '):
|
|
data_buffer.append(line[len('data: '):])
|
|
continue
|
|
|
|
# When API returns an error message, it comes line by line. So we buffer
|
|
# the lines until a complete JSON string is read. A complete JSON string
|
|
# is found when the balance is 0.
|
|
for c in line:
|
|
if c == '{':
|
|
balance += 1
|
|
elif c == '}':
|
|
balance -= 1
|
|
|
|
chunk += line
|
|
if balance == 0:
|
|
yield chunk
|
|
chunk = ''
|
|
# If there is any remaining chunk, yield it.
|
|
if chunk:
|
|
yield chunk
|
|
if data_buffer:
|
|
yield '\n'.join(data_buffer)
|
|
finally:
|
|
# Close the response and release the connection.
|
|
await self.response_stream.aclose()
|
|
|
|
# aiohttp.ClientResponse uses a content stream that we read line by line.
|
|
elif has_aiohttp and isinstance(
|
|
self.response_stream, aiohttp.ClientResponse
|
|
):
|
|
try:
|
|
while True:
|
|
# Read a line from the stream. This returns bytes.
|
|
try:
|
|
line_bytes = await self.response_stream.content.readline(
|
|
max_line_length=READ_BUFFER_SIZE
|
|
)
|
|
except TypeError:
|
|
# Ensure backwards compatibility with older versions of
|
|
# aiohttp that do not support max_line_length.
|
|
line_bytes = await self.response_stream.content.readline()
|
|
if not line_bytes:
|
|
break
|
|
# Decode the bytes and remove trailing whitespace and newlines.
|
|
line = line_bytes.decode('utf-8').rstrip()
|
|
if not line:
|
|
if data_buffer:
|
|
yield '\n'.join(data_buffer)
|
|
data_buffer = []
|
|
continue
|
|
|
|
# In streaming mode, the response of JSON is prefixed with "data: "
|
|
# which we must strip before parsing.
|
|
if line.startswith('data: '):
|
|
data_buffer.append(line[len('data: '):])
|
|
continue
|
|
|
|
# When API returns an error message, it comes line by line. So we
|
|
# buffer the lines until a complete JSON string is read. A complete
|
|
# JSON strings found when the balance is 0.
|
|
for c in line:
|
|
if c == '{':
|
|
balance += 1
|
|
elif c == '}':
|
|
balance -= 1
|
|
|
|
chunk += line
|
|
if balance == 0:
|
|
yield chunk
|
|
chunk = ''
|
|
# If there is any remaining chunk, yield it.
|
|
if chunk:
|
|
yield chunk
|
|
if data_buffer:
|
|
yield '\n'.join(data_buffer)
|
|
finally:
|
|
# Release the connection back to the pool for potential reuse.
|
|
self.response_stream.release()
|
|
|
|
@classmethod
|
|
def _load_json_from_response(cls, response: Any) -> Any:
|
|
"""Loads JSON from the response, or raises an error if the parsing fails."""
|
|
try:
|
|
return json.loads(response)
|
|
except json.JSONDecodeError as e:
|
|
raise errors.UnknownApiResponseError(
|
|
f'Failed to parse response as JSON. Raw response: {response}'
|
|
) from e
|
|
|
|
|
|
# Default retry options.
|
|
# The config is based on https://cloud.google.com/storage/docs/retry-strategy.
|
|
# By default, the client will retry 4 times with approximately 1.0, 2.0, 4.0,
|
|
# 8.0 seconds between each attempt.
|
|
_RETRY_ATTEMPTS = 5 # including the initial call.
|
|
_RETRY_INITIAL_DELAY = 1.0 # seconds
|
|
_RETRY_MAX_DELAY = 60.0 # seconds
|
|
_RETRY_EXP_BASE = 2
|
|
_RETRY_JITTER = 1
|
|
_RETRY_HTTP_STATUS_CODES = (
|
|
408, # Request timeout.
|
|
429, # Too many requests.
|
|
500, # Internal server error.
|
|
502, # Bad gateway.
|
|
503, # Service unavailable.
|
|
504, # Gateway timeout
|
|
)
|
|
|
|
|
|
def retry_args(options: Optional[HttpRetryOptions]) -> _common.StringDict:
|
|
"""Returns the retry args for the given http retry options.
|
|
|
|
Args:
|
|
options: The http retry options to use for the retry configuration. If None,
|
|
the 'never retry' stop strategy will be used.
|
|
|
|
Returns:
|
|
The arguments passed to the tenacity.(Async)Retrying constructor.
|
|
"""
|
|
if options is None:
|
|
return {'stop': tenacity.stop_after_attempt(1), 'reraise': True}
|
|
if options.attempts == 0:
|
|
options.attempts = 1
|
|
stop = tenacity.stop_after_attempt(options.attempts or _RETRY_ATTEMPTS)
|
|
retriable_codes = options.http_status_codes or _RETRY_HTTP_STATUS_CODES
|
|
retry = tenacity.retry_if_exception(
|
|
lambda e: (isinstance(e, errors.APIError) and e.code in retriable_codes)
|
|
or isinstance(e, (httpx.TimeoutException, httpx.ConnectError)),
|
|
)
|
|
wait = tenacity.wait_exponential_jitter(
|
|
initial=options.initial_delay or _RETRY_INITIAL_DELAY,
|
|
max=options.max_delay or _RETRY_MAX_DELAY,
|
|
exp_base=options.exp_base or _RETRY_EXP_BASE,
|
|
jitter=options.jitter or _RETRY_JITTER,
|
|
)
|
|
return {
|
|
'stop': stop,
|
|
'retry': retry,
|
|
'reraise': True,
|
|
'wait': wait,
|
|
'before_sleep': tenacity.before_sleep_log(logger, logging.INFO),
|
|
}
|
|
|
|
|
|
class SyncHttpxClient(httpx.Client):
|
|
"""Sync httpx client."""
|
|
|
|
def __init__(self, **kwargs: Any) -> None:
|
|
"""Initializes the httpx client."""
|
|
kwargs.setdefault('follow_redirects', True)
|
|
super().__init__(**kwargs)
|
|
|
|
def __del__(self) -> None:
|
|
"""Closes the httpx client."""
|
|
try:
|
|
if self.is_closed:
|
|
return
|
|
except Exception:
|
|
pass
|
|
try:
|
|
self.close()
|
|
except Exception:
|
|
pass
|
|
|
|
|
|
class AsyncHttpxClient(httpx.AsyncClient):
|
|
"""Async httpx client."""
|
|
|
|
def __init__(self, **kwargs: Any) -> None:
|
|
"""Initializes the httpx client."""
|
|
kwargs.setdefault('follow_redirects', True)
|
|
super().__init__(**kwargs)
|
|
|
|
def __del__(self) -> None:
|
|
try:
|
|
if self.is_closed:
|
|
return
|
|
except Exception:
|
|
pass
|
|
try:
|
|
asyncio.get_running_loop().create_task(self.aclose())
|
|
except Exception:
|
|
pass
|
|
|
|
|
|
class BaseApiClient:
|
|
"""Client for calling HTTP APIs sending and receiving JSON."""
|
|
|
|
def __init__(
|
|
self,
|
|
vertexai: Optional[bool] = None,
|
|
api_key: Optional[str] = None,
|
|
credentials: Optional[google.auth.credentials.Credentials] = None,
|
|
project: Optional[str] = None,
|
|
location: Optional[str] = None,
|
|
http_options: Optional[HttpOptionsOrDict] = None,
|
|
):
|
|
self.vertexai = vertexai
|
|
self.custom_base_url = None
|
|
if self.vertexai is None:
|
|
env_enterprise_str = os.environ.get('GOOGLE_GENAI_USE_ENTERPRISE', None)
|
|
env_vertexai_str = os.environ.get('GOOGLE_GENAI_USE_VERTEXAI', None)
|
|
|
|
env_enterprise = None
|
|
if env_enterprise_str is not None:
|
|
env_enterprise = env_enterprise_str.lower() in ['true', '1']
|
|
|
|
env_vertexai = None
|
|
if env_vertexai_str is not None:
|
|
env_vertexai = env_vertexai_str.lower() in ['true', '1']
|
|
|
|
if (
|
|
env_enterprise is not None
|
|
and env_vertexai is not None
|
|
and env_enterprise != env_vertexai
|
|
):
|
|
warnings.warn(
|
|
'Warning: Both GOOGLE_GENAI_USE_ENTERPRISE and'
|
|
' GOOGLE_GENAI_USE_VERTEXAI are set with conflicting values. The'
|
|
' value of GOOGLE_GENAI_USE_ENTERPRISE will be used.'
|
|
)
|
|
|
|
if env_enterprise is not None:
|
|
self.vertexai = env_enterprise
|
|
elif env_vertexai is not None:
|
|
self.vertexai = env_vertexai
|
|
|
|
# Validate explicitly set initializer values.
|
|
if (project or location) and api_key:
|
|
# API cannot consume both project/location and api_key.
|
|
raise ValueError(
|
|
'Project/location and API key are mutually exclusive in the client'
|
|
' initializer.'
|
|
)
|
|
elif credentials and api_key:
|
|
# API cannot consume both credentials and api_key.
|
|
raise ValueError(
|
|
'Credentials and API key are mutually exclusive in the client'
|
|
' initializer.'
|
|
)
|
|
|
|
# Validate http_options if it is provided.
|
|
validated_http_options = HttpOptions()
|
|
if isinstance(http_options, dict):
|
|
try:
|
|
validated_http_options = HttpOptions.model_validate(http_options)
|
|
except ValidationError as e:
|
|
raise ValueError('Invalid http_options') from e
|
|
elif http_options and _common.is_duck_type_of(http_options, HttpOptions):
|
|
validated_http_options = http_options
|
|
|
|
if (
|
|
validated_http_options.base_url_resource_scope
|
|
and not validated_http_options.base_url
|
|
):
|
|
# base_url_resource_scope is only valid when base_url is set.
|
|
raise ValueError(
|
|
'base_url must be set when base_url_resource_scope is set.'
|
|
)
|
|
|
|
# Retrieve implicitly set values from the environment.
|
|
env_project = os.environ.get('GOOGLE_CLOUD_PROJECT', None)
|
|
env_location = os.environ.get('GOOGLE_CLOUD_LOCATION', None)
|
|
env_api_key = get_env_api_key()
|
|
self.project = project or env_project
|
|
self.location = location or env_location
|
|
self.api_key = api_key or env_api_key
|
|
|
|
self._credentials = credentials
|
|
self._http_options = HttpOptions()
|
|
# Initialize the lock. This lock will be used to protect access to the
|
|
# credentials. This is crucial for thread safety when multiple coroutines
|
|
# might be accessing the credentials at the same time.
|
|
self._sync_auth_lock = threading.Lock()
|
|
self._async_auth_lock: Optional[asyncio.Lock] = None
|
|
self._async_auth_lock_creation_lock: Optional[asyncio.Lock] = None
|
|
|
|
# Handle when to use Vertex AI in express mode (api key).
|
|
# Explicit initializer arguments are already validated above.
|
|
if self.vertexai:
|
|
if credentials and env_api_key:
|
|
# Explicit credentials take precedence over implicit api_key.
|
|
logger.info(
|
|
'The user provided Google Cloud credentials will take precedence'
|
|
+ ' over the API key from the environment variable.'
|
|
)
|
|
self.api_key = None
|
|
elif (env_location or env_project) and api_key:
|
|
# Explicit api_key takes precedence over implicit project/location.
|
|
logger.info(
|
|
'The user provided Vertex AI API key will take precedence over the'
|
|
+ ' project/location from the environment variables.'
|
|
)
|
|
self.project = None
|
|
self.location = None
|
|
elif (project or location) and env_api_key:
|
|
# Explicit project/location takes precedence over implicit api_key.
|
|
logger.info(
|
|
'The user provided project/location will take precedence over the'
|
|
+ ' Vertex AI API key from the environment variable.'
|
|
)
|
|
self.api_key = None
|
|
elif (env_location or env_project) and env_api_key:
|
|
# Implicit project/location takes precedence over implicit api_key.
|
|
logger.info(
|
|
'The project/location from the environment variables will take'
|
|
+ ' precedence over the API key from the environment variables.'
|
|
)
|
|
self.api_key = None
|
|
|
|
self.custom_base_url = (
|
|
validated_http_options.base_url
|
|
if validated_http_options.base_url
|
|
else None
|
|
)
|
|
|
|
if (
|
|
not self.location
|
|
and not self.api_key
|
|
):
|
|
if not self.custom_base_url:
|
|
self.location = 'global'
|
|
elif self.custom_base_url.endswith('.googleapis.com'):
|
|
self.location = 'global'
|
|
|
|
# Skip fetching project from ADC if base url is provided in http options.
|
|
if (
|
|
not self.project
|
|
and not self.api_key
|
|
and not self.custom_base_url
|
|
):
|
|
credentials, self.project = load_auth(project=None)
|
|
if not self._credentials:
|
|
self._credentials = credentials
|
|
|
|
has_sufficient_auth = (self.project and self.location) or self.api_key
|
|
|
|
if not has_sufficient_auth and not self.custom_base_url:
|
|
# Skip sufficient auth check if base url is provided in http options.
|
|
raise ValueError(
|
|
'Project or API key must be set when using the Vertex AI API.'
|
|
)
|
|
if (
|
|
self.api_key or self.location == 'global'
|
|
) and not self.custom_base_url:
|
|
self._http_options.base_url = f'https://aiplatform.googleapis.com/'
|
|
elif (
|
|
self.location in _MULTI_REGIONAL_LOCATIONS
|
|
and not self.custom_base_url
|
|
):
|
|
self._http_options.base_url = (
|
|
f'https://aiplatform.{self.location}.rep.googleapis.com/'
|
|
)
|
|
elif (
|
|
self.custom_base_url
|
|
and not self.custom_base_url.endswith('.googleapis.com')
|
|
) and not ((project and location) or api_key):
|
|
# Avoid setting default base url and api version if base_url provided.
|
|
# API gateway proxy can use the auth in custom headers, not url.
|
|
# Enable custom url if auth is not sufficient.
|
|
self._http_options.base_url = self.custom_base_url
|
|
# Clear project and location if base_url is provided.
|
|
self.project = None
|
|
self.location = None
|
|
else:
|
|
self._http_options.base_url = (
|
|
f'https://{self.location}-aiplatform.googleapis.com/'
|
|
)
|
|
self._http_options.api_version = 'v1beta1'
|
|
else: # Implicit initialization or missing arguments.
|
|
if not self.api_key:
|
|
raise ValueError(
|
|
'No API key was provided. Please pass a valid API key. Learn how to'
|
|
' create an API key at'
|
|
' https://ai.google.dev/gemini-api/docs/api-key.'
|
|
)
|
|
self._http_options.base_url = 'https://generativelanguage.googleapis.com/'
|
|
self._http_options.api_version = 'v1beta'
|
|
# Default options for both clients.
|
|
self._http_options.headers = {'Content-Type': 'application/json'}
|
|
if self.api_key:
|
|
self.api_key = self.api_key.strip()
|
|
if self._http_options.headers is not None:
|
|
self._http_options.headers['x-goog-api-key'] = self.api_key
|
|
# Update the http options with the user provided http options.
|
|
if http_options:
|
|
self._http_options = patch_http_options(
|
|
self._http_options, validated_http_options
|
|
)
|
|
else:
|
|
if self._http_options.headers is not None:
|
|
append_library_version_headers(self._http_options.headers)
|
|
|
|
client_args, async_client_args = self._ensure_httpx_ssl_ctx(
|
|
self._http_options
|
|
)
|
|
self._async_httpx_client_args = async_client_args
|
|
self._authorized_session: Optional[AuthorizedSession] = None
|
|
|
|
if self._use_google_auth_sync():
|
|
self._httpx_client = None
|
|
elif self._http_options.httpx_client:
|
|
self._httpx_client = self._http_options.httpx_client
|
|
else:
|
|
self._httpx_client = SyncHttpxClient(**client_args)
|
|
|
|
if self._use_google_auth_async():
|
|
self._async_httpx_client = None
|
|
elif self._http_options.httpx_async_client:
|
|
self._async_httpx_client = self._http_options.httpx_async_client
|
|
else:
|
|
self._async_httpx_client = AsyncHttpxClient(**async_client_args)
|
|
|
|
if self._http_options.httpx_async_client:
|
|
self._async_httpx_client = self._http_options.httpx_async_client
|
|
else:
|
|
self._async_httpx_client = AsyncHttpxClient(**async_client_args)
|
|
|
|
# Initialize the aiohttp client session.
|
|
self._aiohttp_session: Optional[Union['aiohttp.ClientSession', 'AsyncAuthorizedSession']] = None
|
|
if self._use_aiohttp():
|
|
try:
|
|
import aiohttp # pylint: disable=g-import-not-at-top
|
|
|
|
if self._http_options.aiohttp_client:
|
|
self._aiohttp_session = self._http_options.aiohttp_client
|
|
# Do it once at the genai.Client level. Share among all requests.
|
|
self._async_client_session_request_args = (
|
|
self._ensure_aiohttp_ssl_ctx(self._http_options)
|
|
)
|
|
if self._use_google_auth_async():
|
|
self._async_client_session_request_args['ssl'] = True # type: ignore[no-untyped-call]
|
|
self._async_client_session_request_args['max_allowed_time'] = float(
|
|
'inf'
|
|
) if self._http_options.timeout is None else float(
|
|
self._http_options.timeout
|
|
)
|
|
self._async_client_session_request_args['total_attempts'] = 1
|
|
except ImportError:
|
|
pass
|
|
|
|
retry_kwargs = retry_args(self._http_options.retry_options)
|
|
self._websocket_ssl_ctx = self._ensure_websocket_ssl_ctx(self._http_options)
|
|
self._retry = tenacity.Retrying(**retry_kwargs)
|
|
self._async_retry = tenacity.AsyncRetrying(**retry_kwargs)
|
|
|
|
def _use_google_auth_sync(self) -> bool:
|
|
if not hasattr(mtls, 'should_use_client_cert'):
|
|
return False
|
|
return bool(
|
|
self.vertexai
|
|
and mtls.should_use_client_cert() # type: ignore[no-untyped-call]
|
|
and mtls.has_default_client_cert_source() # type: ignore[no-untyped-call]
|
|
and not (
|
|
self._http_options.httpx_client or self._http_options.client_args
|
|
)
|
|
)
|
|
|
|
def _use_google_auth_async(self) -> bool:
|
|
return bool(
|
|
has_aiohttp
|
|
and self.vertexai
|
|
and hasattr(mtls, 'should_use_client_cert')
|
|
and mtls.should_use_client_cert() # type: ignore[no-untyped-call]
|
|
and mtls.has_default_client_cert_source() # type: ignore[no-untyped-call]
|
|
and not self._http_options.httpx_async_client
|
|
)
|
|
|
|
async def _get_aiohttp_session(
|
|
self,
|
|
) -> Union['aiohttp.ClientSession', 'AsyncAuthorizedSession']:
|
|
"""Returns the aiohttp client session."""
|
|
|
|
if self._aiohttp_session is None and self._use_google_auth_async():
|
|
try:
|
|
from google.auth.aio.credentials import Credentials as AsyncCredentials
|
|
from google.auth.aio.transport.sessions import AsyncAuthorizedSession
|
|
|
|
class _RefreshableAsyncCredentials(AsyncCredentials): # type: ignore[misc, valid-type]
|
|
"""Adapter to use the client's sync credentials in an AsyncAuthorizedSession."""
|
|
|
|
def __init__(self, client: 'BaseApiClient'):
|
|
super().__init__() # type: ignore[no-untyped-call]
|
|
self._client = client
|
|
|
|
async def before_request(
|
|
self, request: Any, method: str, url: str, headers: dict[str, str]
|
|
) -> None:
|
|
token = await self._client._async_access_token()
|
|
headers['Authorization'] = f'Bearer {token}'
|
|
if (
|
|
self._client._credentials
|
|
and self._client._credentials.quota_project_id
|
|
):
|
|
headers['x-goog-user-project'] = (
|
|
self._client._credentials.quota_project_id
|
|
)
|
|
|
|
@property
|
|
def valid(self) -> bool:
|
|
if not self._client._credentials:
|
|
return False
|
|
return not self._client._credentials.expired
|
|
|
|
self._aiohttp_session = AsyncAuthorizedSession(_RefreshableAsyncCredentials(self)) # type: ignore[no-untyped-call,assignment]
|
|
return self._aiohttp_session # type: ignore[return-value]
|
|
except ImportError:
|
|
pass
|
|
|
|
if not self._use_google_auth_async() and (
|
|
self._aiohttp_session is None
|
|
or self._aiohttp_session.closed # type: ignore[union-attr]
|
|
or self._aiohttp_session._loop.is_closed() # type: ignore[union-attr]
|
|
): # pylint: disable=protected-access
|
|
# Initialize the aiohttp client session if it's not set up or closed.
|
|
class AiohttpClientSession(aiohttp.ClientSession): # type: ignore[misc]
|
|
|
|
def __del__(self, _warnings: Any = warnings) -> None:
|
|
if not self.closed:
|
|
context = {
|
|
'client_session': self,
|
|
'message': 'Unclosed client session',
|
|
}
|
|
if self._source_traceback is not None:
|
|
context['source_traceback'] = self._source_traceback
|
|
# Remove this self._loop.call_exception_handler(context)
|
|
|
|
class AiohttpTCPConnector(aiohttp.TCPConnector): # type: ignore[misc]
|
|
|
|
def __del__(self, _warnings: Any = warnings) -> None:
|
|
if self._closed:
|
|
return
|
|
if not self._conns:
|
|
return
|
|
conns = [repr(c) for c in self._conns.values()]
|
|
# After v3.13.2, it may change to self._close_immediately()
|
|
self._close()
|
|
context = {
|
|
'connector': self,
|
|
'connections': conns,
|
|
'message': 'Unclosed connector',
|
|
}
|
|
if self._source_traceback is not None:
|
|
context['source_traceback'] = self._source_traceback
|
|
# Remove this self._loop.call_exception_handler(context)
|
|
self._aiohttp_session = AiohttpClientSession(
|
|
connector=AiohttpTCPConnector(limit=0),
|
|
trust_env=True,
|
|
read_bufsize=READ_BUFFER_SIZE,
|
|
)
|
|
|
|
return self._aiohttp_session # type: ignore[return-value]
|
|
|
|
@staticmethod
|
|
def _ensure_httpx_ssl_ctx(
|
|
options: HttpOptions,
|
|
) -> Tuple[_common.StringDict, _common.StringDict]:
|
|
"""Ensures the SSL context is present in the HTTPX client args.
|
|
|
|
Creates a default SSL context if one is not provided.
|
|
|
|
Args:
|
|
options: The http options to check for SSL context.
|
|
|
|
Returns:
|
|
A tuple of sync/async httpx client args.
|
|
"""
|
|
|
|
verify = 'verify'
|
|
args = options.client_args
|
|
async_args = options.async_client_args
|
|
ctx = (
|
|
args.get(verify)
|
|
if args
|
|
else None or async_args.get(verify)
|
|
if async_args
|
|
else None
|
|
)
|
|
|
|
if not ctx:
|
|
# Initialize the SSL context for the httpx client.
|
|
# Unlike requests, the httpx package does not automatically pull in the
|
|
# environment variables SSL_CERT_FILE or SSL_CERT_DIR. They need to be
|
|
# enabled explicitly.
|
|
ctx = ssl.create_default_context(
|
|
cafile=os.environ.get('SSL_CERT_FILE', certifi.where()),
|
|
capath=os.environ.get('SSL_CERT_DIR'),
|
|
)
|
|
|
|
def _maybe_set(
|
|
args: Optional[_common.StringDict],
|
|
ctx: ssl.SSLContext,
|
|
) -> _common.StringDict:
|
|
"""Sets the SSL context in the client args if not set.
|
|
|
|
Does not override the SSL context if it is already set.
|
|
|
|
Args:
|
|
args: The client args to to check for SSL context.
|
|
ctx: The SSL context to set.
|
|
|
|
Returns:
|
|
The client args with the SSL context included.
|
|
"""
|
|
args = (args or {}).copy()
|
|
if not args.get(verify):
|
|
args[verify] = ctx
|
|
if 'timeout' not in args:
|
|
args['timeout'] = None
|
|
# Drop the args that isn't used by the httpx client.
|
|
copied_args = args.copy()
|
|
for key in copied_args.copy():
|
|
if key not in inspect.signature(httpx.Client.__init__).parameters:
|
|
del copied_args[key]
|
|
return copied_args
|
|
|
|
return (
|
|
_maybe_set(args, ctx),
|
|
_maybe_set(async_args, ctx),
|
|
)
|
|
|
|
@staticmethod
|
|
def _ensure_aiohttp_ssl_ctx(options: HttpOptions) -> _common.StringDict:
|
|
"""Ensures the SSL context is present in the async client args.
|
|
|
|
Creates a default SSL context if one is not provided.
|
|
|
|
Args:
|
|
options: The http options to check for SSL context.
|
|
|
|
Returns:
|
|
An async aiohttp ClientSession._request args.
|
|
"""
|
|
verify = 'ssl' # keep it consistent with aiohttp.
|
|
async_args = options.async_client_args
|
|
ctx = async_args.get(verify) if async_args else None
|
|
|
|
if not ctx:
|
|
ctx = ssl.create_default_context(
|
|
cafile=os.environ.get('SSL_CERT_FILE', certifi.where()),
|
|
capath=os.environ.get('SSL_CERT_DIR'),
|
|
)
|
|
|
|
def _maybe_set(
|
|
args: Optional[_common.StringDict],
|
|
ctx: ssl.SSLContext,
|
|
) -> _common.StringDict:
|
|
"""Sets the SSL context in the client args if not set.
|
|
|
|
Does not override the SSL context if it is already set.
|
|
|
|
Args:
|
|
args: The client args to to check for SSL context.
|
|
ctx: The SSL context to set.
|
|
|
|
Returns:
|
|
The client args with the SSL context included.
|
|
"""
|
|
if not args or not args.get(verify):
|
|
args = (args or {}).copy()
|
|
args[verify] = ctx
|
|
# Drop the args that isn't in the aiohttp RequestOptions.
|
|
copied_args = args.copy()
|
|
for key in copied_args.copy():
|
|
if (
|
|
key
|
|
not in inspect.signature(aiohttp.ClientSession._request).parameters
|
|
):
|
|
del copied_args[key]
|
|
return copied_args
|
|
|
|
return _maybe_set(async_args, ctx)
|
|
|
|
@staticmethod
|
|
def _ensure_websocket_ssl_ctx(options: HttpOptions) -> _common.StringDict:
|
|
"""Ensures the SSL context is present in the async client args.
|
|
|
|
Creates a default SSL context if one is not provided.
|
|
|
|
Args:
|
|
options: The http options to check for SSL context.
|
|
|
|
Returns:
|
|
An async aiohttp ClientSession._request args.
|
|
"""
|
|
|
|
verify = 'ssl' # keep it consistent with httpx.
|
|
async_args = options.async_client_args
|
|
ctx = async_args.get(verify) if async_args else None
|
|
|
|
if not ctx:
|
|
# Initialize the SSL context for the httpx client.
|
|
# Unlike requests, the aiohttp package does not automatically pull in the
|
|
# environment variables SSL_CERT_FILE or SSL_CERT_DIR. They need to be
|
|
# enabled explicitly. Instead of 'verify' at client level in httpx,
|
|
# aiohttp uses 'ssl' at request level.
|
|
ctx = ssl.create_default_context(
|
|
cafile=os.environ.get('SSL_CERT_FILE', certifi.where()),
|
|
capath=os.environ.get('SSL_CERT_DIR'),
|
|
)
|
|
|
|
def _maybe_set(
|
|
args: Optional[_common.StringDict],
|
|
ctx: ssl.SSLContext,
|
|
) -> _common.StringDict:
|
|
"""Sets the SSL context in the client args if not set.
|
|
|
|
Does not override the SSL context if it is already set.
|
|
|
|
Args:
|
|
args: The client args to to check for SSL context.
|
|
ctx: The SSL context to set.
|
|
|
|
Returns:
|
|
The client args with the SSL context included.
|
|
"""
|
|
if not args or not args.get(verify):
|
|
args = (args or {}).copy()
|
|
args[verify] = ctx
|
|
# Drop the args that isn't in the aiohttp RequestOptions.
|
|
copied_args = args.copy()
|
|
for key in copied_args.copy():
|
|
if key not in inspect.signature(ws_connect).parameters and key != 'ssl':
|
|
del copied_args[key]
|
|
return copied_args
|
|
|
|
return _maybe_set(async_args, ctx)
|
|
|
|
def _use_aiohttp(self) -> bool:
|
|
# If the instantiator has passed a custom transport, they want httpx not
|
|
# aiohttp.
|
|
return (
|
|
has_aiohttp
|
|
and (self._http_options.async_client_args or {}).get('transport')
|
|
is None
|
|
and (self._http_options.httpx_async_client is None)
|
|
)
|
|
|
|
def _websocket_base_url(self) -> str:
|
|
has_sufficient_auth = (self.project and self.location) or self.api_key
|
|
if self.custom_base_url and not has_sufficient_auth:
|
|
# API gateway proxy can use the auth in custom headers, not url.
|
|
# Enable custom url if auth is not sufficient.
|
|
return self.custom_base_url
|
|
url_parts = urlparse(self._http_options.base_url)
|
|
return url_parts._replace(scheme='wss').geturl() # type: ignore[arg-type, return-value]
|
|
|
|
def _access_token(self) -> str:
|
|
"""Retrieves the access token for the credentials."""
|
|
with self._sync_auth_lock:
|
|
if not self._credentials:
|
|
self._credentials, project = load_auth(project=self.project)
|
|
if not self.project:
|
|
self.project = project
|
|
|
|
if self._credentials:
|
|
return get_token_from_credentials(self, self._credentials) # type: ignore[no-any-return]
|
|
else:
|
|
raise RuntimeError('Could not resolve API token from the environment')
|
|
|
|
async def _get_async_auth_lock(self) -> asyncio.Lock:
|
|
"""Lazily initializes and returns an asyncio.Lock for async authentication.
|
|
|
|
This method ensures that a single `asyncio.Lock` instance is created and
|
|
shared among all asynchronous operations that require authentication,
|
|
preventing race conditions when accessing or refreshing credentials.
|
|
|
|
The lock is created on the first call to this method. An internal async lock
|
|
is used to protect the creation of the main authentication lock to ensure
|
|
it's a singleton within the client instance.
|
|
|
|
Returns:
|
|
The asyncio.Lock instance for asynchronous authentication operations.
|
|
"""
|
|
if self._async_auth_lock is None:
|
|
# Create async creation lock if needed
|
|
if self._async_auth_lock_creation_lock is None:
|
|
self._async_auth_lock_creation_lock = asyncio.Lock()
|
|
|
|
async with self._async_auth_lock_creation_lock:
|
|
if self._async_auth_lock is None:
|
|
self._async_auth_lock = asyncio.Lock()
|
|
|
|
return self._async_auth_lock
|
|
|
|
async def _async_access_token(self) -> Union[str, Any]:
|
|
"""Retrieves the access token for the credentials asynchronously."""
|
|
if not self._credentials:
|
|
async_auth_lock = await self._get_async_auth_lock()
|
|
async with async_auth_lock:
|
|
# This ensures that only one coroutine can execute the auth logic at a
|
|
# time for thread safety.
|
|
if not self._credentials:
|
|
# Double check that the credentials are not set before loading them.
|
|
self._credentials, project = await asyncio.to_thread(
|
|
load_auth, project=self.project
|
|
)
|
|
if not self.project:
|
|
self.project = project
|
|
|
|
if self._credentials:
|
|
return await async_get_token_from_credentials(
|
|
self,
|
|
self._credentials
|
|
) # type: ignore[no-any-return]
|
|
else:
|
|
raise RuntimeError('Could not resolve API token from the environment')
|
|
|
|
def _build_request(
|
|
self,
|
|
http_method: str,
|
|
path: str,
|
|
request_dict: dict[str, object],
|
|
http_options: Optional[HttpOptionsOrDict] = None,
|
|
) -> HttpRequest:
|
|
# Remove all special dict keys such as _url and _query.
|
|
keys_to_delete = [key for key in request_dict.keys() if key.startswith('_')]
|
|
for key in keys_to_delete:
|
|
del request_dict[key]
|
|
# patch the http options with the user provided settings.
|
|
if http_options:
|
|
if isinstance(http_options, HttpOptions):
|
|
patched_http_options = patch_http_options(
|
|
self._http_options,
|
|
http_options,
|
|
)
|
|
else:
|
|
patched_http_options = patch_http_options(
|
|
self._http_options, HttpOptions.model_validate(http_options)
|
|
)
|
|
else:
|
|
patched_http_options = self._http_options
|
|
# Skip adding project and locations when getting Vertex AI base models.
|
|
query_vertex_base_models = False
|
|
if (
|
|
self.vertexai
|
|
and http_method == 'get'
|
|
and path.startswith('publishers/')
|
|
):
|
|
query_vertex_base_models = True
|
|
if (
|
|
self.vertexai
|
|
and not path.startswith('projects/')
|
|
and not query_vertex_base_models
|
|
and (self.project or self.location)
|
|
and not (
|
|
self.custom_base_url
|
|
and patched_http_options.base_url_resource_scope
|
|
== ResourceScope.COLLECTION
|
|
)
|
|
):
|
|
path = f'projects/{self.project}/locations/{self.location}/' + path
|
|
|
|
if patched_http_options.api_version is None:
|
|
versioned_path = f'/{path}'
|
|
else:
|
|
versioned_path = f'{patched_http_options.api_version}/{path}'
|
|
|
|
if (
|
|
patched_http_options.base_url is None
|
|
or not patched_http_options.base_url
|
|
):
|
|
raise ValueError('Base URL must be set.')
|
|
else:
|
|
base_url = patched_http_options.base_url
|
|
|
|
if (
|
|
hasattr(patched_http_options, 'extra_body')
|
|
and patched_http_options.extra_body
|
|
):
|
|
_common.recursive_dict_update(
|
|
request_dict, patched_http_options.extra_body
|
|
)
|
|
url = base_url
|
|
if (
|
|
not self.custom_base_url
|
|
or (self.project and self.location)
|
|
or self.api_key
|
|
):
|
|
if (
|
|
patched_http_options.base_url_resource_scope
|
|
== ResourceScope.COLLECTION
|
|
):
|
|
url = join_url_path(base_url, path)
|
|
else:
|
|
url = join_url_path(
|
|
base_url,
|
|
versioned_path,
|
|
)
|
|
elif(
|
|
self.custom_base_url
|
|
and patched_http_options.base_url_resource_scope == ResourceScope.COLLECTION
|
|
):
|
|
url = join_url_path(base_url, path)
|
|
|
|
if self.api_key and self.api_key.startswith('auth_tokens/'):
|
|
raise EphemeralTokenAPIKeyError(
|
|
'Ephemeral tokens can only be used with the live API.'
|
|
)
|
|
|
|
timeout_in_seconds = get_timeout_in_seconds(patched_http_options.timeout)
|
|
|
|
if patched_http_options.headers is None:
|
|
raise ValueError('Request headers must be set.')
|
|
populate_server_timeout_header(
|
|
patched_http_options.headers, timeout_in_seconds
|
|
)
|
|
return HttpRequest(
|
|
method=http_method,
|
|
url=url,
|
|
headers=patched_http_options.headers,
|
|
data=request_dict,
|
|
timeout=timeout_in_seconds,
|
|
)
|
|
|
|
def _request_once(
|
|
self,
|
|
http_request: HttpRequest,
|
|
stream: bool = False,
|
|
) -> HttpResponse:
|
|
data: Optional[Union[str, bytes]] = None
|
|
# If using proj/location, fetch ADC
|
|
if self.vertexai and (self.project or self.location):
|
|
http_request.headers['Authorization'] = f'Bearer {self._access_token()}'
|
|
if self._credentials and self._credentials.quota_project_id:
|
|
http_request.headers['x-goog-user-project'] = (
|
|
self._credentials.quota_project_id
|
|
)
|
|
data = json.dumps(http_request.data) if http_request.data else None
|
|
else:
|
|
if http_request.data:
|
|
if not isinstance(http_request.data, bytes):
|
|
data = json.dumps(http_request.data) if http_request.data else None
|
|
else:
|
|
data = http_request.data
|
|
|
|
if self._use_google_auth_sync():
|
|
url = str(http_request.url)
|
|
if self._authorized_session is None:
|
|
self._authorized_session = AuthorizedSession( # type: ignore[no-untyped-call]
|
|
self._credentials,
|
|
max_refresh_attempts=1,
|
|
)
|
|
client_cert_source = mtls.default_client_cert_source() # type: ignore[no-untyped-call]
|
|
self._authorized_session.configure_mtls_channel(
|
|
client_cert_source
|
|
) # type: ignore[no-untyped-call]
|
|
if self._authorized_session._is_mtls and 'googleapis.com' in url:
|
|
if 'sandbox' in url:
|
|
url = url.replace(
|
|
'sandbox.googleapis.com', 'mtls.sandbox.googleapis.com'
|
|
)
|
|
else:
|
|
url = url.replace('googleapis.com', 'mtls.googleapis.com')
|
|
response = self._authorized_session.request( # type: ignore[no-untyped-call]
|
|
method=http_request.method.upper(),
|
|
url=url,
|
|
data=data,
|
|
headers=http_request.headers,
|
|
timeout=http_request.timeout,
|
|
stream=stream,
|
|
)
|
|
else:
|
|
httpx_request = self._httpx_client.build_request( # type: ignore[union-attr]
|
|
method=http_request.method,
|
|
url=http_request.url,
|
|
content=data,
|
|
headers=http_request.headers,
|
|
timeout=http_request.timeout,
|
|
)
|
|
response = self._httpx_client.send(httpx_request, stream=stream) # type: ignore[union-attr]
|
|
errors.APIError.raise_for_response(response)
|
|
return HttpResponse(
|
|
response.headers, response if stream else [response.text]
|
|
)
|
|
|
|
def _request(
|
|
self,
|
|
http_request: HttpRequest,
|
|
http_options: Optional[HttpOptionsOrDict] = None,
|
|
stream: bool = False,
|
|
) -> HttpResponse:
|
|
if http_options:
|
|
parameter_model = (
|
|
HttpOptions(**http_options)
|
|
if isinstance(http_options, dict)
|
|
else http_options
|
|
)
|
|
# Support per request retry options.
|
|
if parameter_model.retry_options:
|
|
retry_kwargs = retry_args(parameter_model.retry_options)
|
|
retry = tenacity.Retrying(**retry_kwargs)
|
|
return retry(self._request_once, http_request, stream) # type: ignore[no-any-return]
|
|
|
|
return self._retry(self._request_once, http_request, stream) # type: ignore[no-any-return]
|
|
|
|
async def _async_request_once(
|
|
self, http_request: HttpRequest, stream: bool = False
|
|
) -> HttpResponse:
|
|
data: Optional[bytes] = None
|
|
|
|
# If using proj/location, fetch ADC
|
|
if self.vertexai and (self.project or self.location):
|
|
http_request.headers['Authorization'] = (
|
|
f'Bearer {await self._async_access_token()}'
|
|
)
|
|
if self._credentials and self._credentials.quota_project_id:
|
|
http_request.headers['x-goog-user-project'] = (
|
|
self._credentials.quota_project_id
|
|
)
|
|
if http_request.data:
|
|
if not isinstance(http_request.data, bytes):
|
|
data = json.dumps(http_request.data).encode('utf-8')
|
|
else:
|
|
data = http_request.data
|
|
|
|
if stream:
|
|
if self._use_aiohttp():
|
|
self._aiohttp_session = await self._get_aiohttp_session() # type: ignore[assignment]
|
|
url = http_request.url
|
|
if self._use_google_auth_async():
|
|
client_cert_source = mtls.default_client_cert_source() # type: ignore[no-untyped-call]
|
|
await self._aiohttp_session.configure_mtls_channel( # type: ignore[union-attr]
|
|
client_cert_source
|
|
)
|
|
if self._aiohttp_session._is_mtls and 'googleapis.com' in url: # type: ignore[union-attr]
|
|
if 'sandbox' in url:
|
|
url = url.replace(
|
|
'sandbox.googleapis.com', 'mtls.sandbox.googleapis.com'
|
|
)
|
|
else:
|
|
url = url.replace('googleapis.com', 'mtls.googleapis.com')
|
|
try:
|
|
response = await self._aiohttp_session.request( # type: ignore[union-attr]
|
|
method=http_request.method,
|
|
url=url,
|
|
headers=http_request.headers,
|
|
data=data,
|
|
timeout=aiohttp.ClientTimeout(total=http_request.timeout),
|
|
**self._async_client_session_request_args,
|
|
)
|
|
except (
|
|
aiohttp.ClientConnectorError,
|
|
aiohttp.ClientConnectorDNSError,
|
|
aiohttp.ClientOSError,
|
|
aiohttp.ServerDisconnectedError,
|
|
auth_exceptions.TransportError,
|
|
) as e:
|
|
await asyncio.sleep(1 + random.randint(0, 9))
|
|
logger.info('Retrying due to aiohttp error: %s' % e)
|
|
# Retrieve the SSL context from the session.
|
|
self._async_client_session_request_args = (
|
|
self._ensure_aiohttp_ssl_ctx(self._http_options)
|
|
)
|
|
# Instantiate a new session with the updated SSL context.
|
|
self._aiohttp_session = await self._get_aiohttp_session() # type: ignore[assignment]
|
|
response = await self._aiohttp_session.request( # type: ignore[union-attr]
|
|
method=http_request.method,
|
|
url=url,
|
|
headers=http_request.headers,
|
|
data=data,
|
|
timeout=aiohttp.ClientTimeout(total=http_request.timeout),
|
|
**self._async_client_session_request_args,
|
|
)
|
|
|
|
await errors.APIError.raise_for_async_response(response)
|
|
if hasattr(response, '_response'):
|
|
# Extract the underlying aiohttp.ClientResponse from the
|
|
# AsyncAuthorizedSession Response.
|
|
response = response._response
|
|
return HttpResponse(response.headers, response)
|
|
else:
|
|
# aiohttp is not available. Fall back to httpx.
|
|
httpx_request = self._async_httpx_client.build_request( # type: ignore[union-attr]
|
|
method=http_request.method,
|
|
url=http_request.url,
|
|
content=data,
|
|
headers=http_request.headers,
|
|
timeout=http_request.timeout,
|
|
)
|
|
client_response = await self._async_httpx_client.send( # type: ignore[union-attr]
|
|
httpx_request,
|
|
stream=stream,
|
|
)
|
|
await errors.APIError.raise_for_async_response(client_response)
|
|
return HttpResponse(client_response.headers, client_response)
|
|
else:
|
|
if self._use_aiohttp():
|
|
self._aiohttp_session = await self._get_aiohttp_session() # type: ignore[assignment]
|
|
url = http_request.url
|
|
if self._use_google_auth_async():
|
|
client_cert_source = mtls.default_client_cert_source() # type: ignore[no-untyped-call]
|
|
await self._aiohttp_session.configure_mtls_channel( # type: ignore[union-attr]
|
|
client_cert_source
|
|
)
|
|
if self._aiohttp_session._is_mtls and 'googleapis.com' in url: # type: ignore[union-attr]
|
|
if 'sandbox' in url:
|
|
url = url.replace(
|
|
'sandbox.googleapis.com', 'mtls.sandbox.googleapis.com'
|
|
)
|
|
else:
|
|
url = url.replace('googleapis.com', 'mtls.googleapis.com')
|
|
try:
|
|
response = await self._aiohttp_session.request( # type: ignore[union-attr]
|
|
method=http_request.method,
|
|
url=url,
|
|
headers=http_request.headers,
|
|
data=data,
|
|
timeout=aiohttp.ClientTimeout(total=http_request.timeout),
|
|
**self._async_client_session_request_args,
|
|
)
|
|
await errors.APIError.raise_for_async_response(response)
|
|
unwrapped_response: Any = response
|
|
if hasattr(unwrapped_response, '_response'):
|
|
unwrapped_response = unwrapped_response._response
|
|
|
|
return HttpResponse(
|
|
unwrapped_response.headers, [await unwrapped_response.text()]
|
|
)
|
|
except (
|
|
aiohttp.ClientConnectorError,
|
|
aiohttp.ClientConnectorDNSError,
|
|
aiohttp.ClientOSError,
|
|
aiohttp.ServerDisconnectedError,
|
|
auth_exceptions.TransportError,
|
|
) as e:
|
|
await asyncio.sleep(1 + random.randint(0, 9))
|
|
logger.info('Retrying due to aiohttp error: %s' % e)
|
|
# Retrieve the SSL context from the session.
|
|
self._async_client_session_request_args = (
|
|
self._ensure_aiohttp_ssl_ctx(self._http_options)
|
|
)
|
|
# Instantiate a new session with the updated SSL context.
|
|
self._aiohttp_session = await self._get_aiohttp_session() # type: ignore[assignment]
|
|
response = await self._aiohttp_session.request( # type: ignore[union-attr]
|
|
method=http_request.method,
|
|
url=url,
|
|
headers=http_request.headers,
|
|
data=data,
|
|
timeout=aiohttp.ClientTimeout(total=http_request.timeout),
|
|
**self._async_client_session_request_args,
|
|
)
|
|
await errors.APIError.raise_for_async_response(response)
|
|
unwrapped_retry_response: Any = response
|
|
if hasattr(unwrapped_retry_response, '_response'):
|
|
unwrapped_retry_response = unwrapped_retry_response._response
|
|
|
|
return HttpResponse(
|
|
unwrapped_retry_response.headers,
|
|
[await unwrapped_retry_response.text()],
|
|
)
|
|
else:
|
|
# aiohttp is not available. Fall back to httpx.
|
|
client_response = await self._async_httpx_client.request( # type: ignore[union-attr]
|
|
method=http_request.method,
|
|
url=http_request.url,
|
|
headers=http_request.headers,
|
|
content=data,
|
|
timeout=http_request.timeout,
|
|
)
|
|
await errors.APIError.raise_for_async_response(client_response)
|
|
return HttpResponse(client_response.headers, [client_response.text])
|
|
|
|
async def _async_request(
|
|
self,
|
|
http_request: HttpRequest,
|
|
http_options: Optional[HttpOptionsOrDict] = None,
|
|
stream: bool = False,
|
|
) -> HttpResponse:
|
|
if http_options:
|
|
parameter_model = (
|
|
HttpOptions(**http_options)
|
|
if isinstance(http_options, dict)
|
|
else http_options
|
|
)
|
|
# Support per request retry options.
|
|
if parameter_model.retry_options:
|
|
retry_kwargs = retry_args(parameter_model.retry_options)
|
|
retry = tenacity.AsyncRetrying(**retry_kwargs)
|
|
return await retry(self._async_request_once, http_request, stream) # type: ignore[no-any-return]
|
|
return await self._async_retry( # type: ignore[no-any-return]
|
|
self._async_request_once, http_request, stream
|
|
)
|
|
|
|
def get_read_only_http_options(self) -> _common.StringDict:
|
|
if isinstance(self._http_options, BaseModel):
|
|
copied = self._http_options.model_dump()
|
|
else:
|
|
copied = self._http_options
|
|
return copied
|
|
|
|
def request(
|
|
self,
|
|
http_method: str,
|
|
path: str,
|
|
request_dict: dict[str, object],
|
|
http_options: Optional[HttpOptionsOrDict] = None,
|
|
) -> SdkHttpResponse:
|
|
http_request = self._build_request(
|
|
http_method, path, request_dict, http_options
|
|
)
|
|
response = self._request(http_request, http_options, stream=False)
|
|
response_body = (
|
|
response.response_stream[0] if response.response_stream else ''
|
|
)
|
|
return SdkHttpResponse(headers=response.headers, body=response_body)
|
|
|
|
def request_streamed(
|
|
self,
|
|
http_method: str,
|
|
path: str,
|
|
request_dict: dict[str, object],
|
|
http_options: Optional[HttpOptionsOrDict] = None,
|
|
) -> Generator[SdkHttpResponse, None, None]:
|
|
http_request = self._build_request(
|
|
http_method, path, request_dict, http_options
|
|
)
|
|
|
|
session_response = self._request(http_request, http_options, stream=True)
|
|
for chunk in session_response.segments():
|
|
chunk_dump = json.dumps(chunk)
|
|
try:
|
|
if chunk_dump.startswith('{"error":'):
|
|
chunk_json = json.loads(chunk_dump)
|
|
errors.APIError.raise_error(
|
|
chunk_json.get('error', {}).get('code'),
|
|
chunk_json,
|
|
session_response,
|
|
)
|
|
except json.decoder.JSONDecodeError:
|
|
logger.debug(
|
|
'Failed to decode chunk that contains an error: %s' % chunk_dump
|
|
)
|
|
pass
|
|
yield SdkHttpResponse(headers=session_response.headers, body=chunk_dump)
|
|
|
|
async def async_request(
|
|
self,
|
|
http_method: str,
|
|
path: str,
|
|
request_dict: dict[str, object],
|
|
http_options: Optional[HttpOptionsOrDict] = None,
|
|
) -> SdkHttpResponse:
|
|
http_request = self._build_request(
|
|
http_method, path, request_dict, http_options
|
|
)
|
|
|
|
result = await self._async_request(
|
|
http_request=http_request, http_options=http_options, stream=False
|
|
)
|
|
response_body = result.response_stream[0] if result.response_stream else ''
|
|
return SdkHttpResponse(headers=result.headers, body=response_body)
|
|
|
|
async def async_request_streamed(
|
|
self,
|
|
http_method: str,
|
|
path: str,
|
|
request_dict: dict[str, object],
|
|
http_options: Optional[HttpOptionsOrDict] = None,
|
|
) -> Any:
|
|
http_request = self._build_request(
|
|
http_method, path, request_dict, http_options
|
|
)
|
|
|
|
response = await self._async_request(
|
|
http_request=http_request, http_options=http_options, stream=True
|
|
)
|
|
|
|
async def async_generator(): # type: ignore[no-untyped-def]
|
|
async for chunk in response:
|
|
chunk_dump = json.dumps(chunk)
|
|
try:
|
|
if chunk_dump.startswith('{"error":'):
|
|
chunk_json = json.loads(chunk_dump)
|
|
await errors.APIError.raise_error_async(
|
|
chunk_json.get('error', {}).get('code'),
|
|
chunk_json,
|
|
response,
|
|
)
|
|
except json.decoder.JSONDecodeError:
|
|
logger.debug(
|
|
'Failed to decode chunk that contains an error: %s' % chunk_dump
|
|
)
|
|
pass
|
|
yield SdkHttpResponse(headers=response.headers, body=chunk_dump)
|
|
|
|
return async_generator() # type: ignore[no-untyped-call]
|
|
|
|
def upload_file(
|
|
self,
|
|
file_path: Union[str, io.IOBase],
|
|
upload_url: str,
|
|
upload_size: int,
|
|
*,
|
|
http_options: Optional[HttpOptionsOrDict] = None,
|
|
) -> HttpResponse:
|
|
"""Transfers a file to the given URL.
|
|
|
|
Args:
|
|
file_path: The full path to the file or a file like object inherited from
|
|
io.BytesIO. If the local file path is not found, an error will be
|
|
raised.
|
|
upload_url: The URL to upload the file to.
|
|
upload_size: The size of file content to be uploaded, this will have to
|
|
match the size requested in the resumable upload request.
|
|
http_options: The http options to use for the request.
|
|
|
|
returns:
|
|
The HttpResponse object from the finalize request.
|
|
"""
|
|
if isinstance(file_path, io.IOBase):
|
|
return self._upload_fd(
|
|
file_path, upload_url, upload_size, http_options=http_options
|
|
)
|
|
else:
|
|
with open(file_path, 'rb') as file:
|
|
return self._upload_fd(
|
|
file, upload_url, upload_size, http_options=http_options
|
|
)
|
|
|
|
def _upload_fd(
|
|
self,
|
|
file: io.IOBase,
|
|
upload_url: str,
|
|
upload_size: int,
|
|
*,
|
|
http_options: Optional[HttpOptionsOrDict] = None,
|
|
) -> HttpResponse:
|
|
"""Transfers a file to the given URL.
|
|
|
|
Args:
|
|
file: A file like object inherited from io.BytesIO.
|
|
upload_url: The URL to upload the file to.
|
|
upload_size: The size of file content to be uploaded, this will have to
|
|
match the size requested in the resumable upload request.
|
|
http_options: The http options to use for the request.
|
|
|
|
returns:
|
|
The HttpResponse object from the finalize request.
|
|
"""
|
|
offset = 0
|
|
http_options = http_options if http_options else self._http_options
|
|
base_url = (
|
|
http_options.get('base_url')
|
|
if isinstance(http_options, dict)
|
|
else getattr(http_options, 'base_url', None)
|
|
)
|
|
if base_url:
|
|
parsed_base = urlparse(base_url)
|
|
parsed_upload = urlparse(upload_url)
|
|
upload_url = urlunparse(
|
|
parsed_upload._replace(
|
|
scheme=parsed_base.scheme, netloc=parsed_base.netloc
|
|
)
|
|
)
|
|
|
|
# Upload the file in chunks
|
|
while True:
|
|
file_chunk = file.read(CHUNK_SIZE)
|
|
chunk_size = 0
|
|
if file_chunk:
|
|
chunk_size = len(file_chunk)
|
|
upload_command = 'upload'
|
|
# If last chunk, finalize the upload.
|
|
if chunk_size + offset >= upload_size:
|
|
upload_command += ', finalize'
|
|
timeout = (
|
|
http_options.get('timeout')
|
|
if isinstance(http_options, dict)
|
|
else http_options.timeout
|
|
)
|
|
if timeout is None:
|
|
# Per request timeout is not configured. Check the global timeout.
|
|
timeout = (
|
|
self._http_options.timeout
|
|
if isinstance(self._http_options, dict)
|
|
else self._http_options.timeout
|
|
)
|
|
timeout_in_seconds = get_timeout_in_seconds(timeout)
|
|
user_headers = (
|
|
http_options.get('headers', {})
|
|
if isinstance(http_options, dict)
|
|
else (getattr(http_options, 'headers', {}) or {})
|
|
)
|
|
upload_headers = dict(user_headers) if user_headers else {}
|
|
upload_headers.update({
|
|
'X-Goog-Upload-Command': upload_command,
|
|
'X-Goog-Upload-Offset': str(offset),
|
|
'Content-Length': str(chunk_size),
|
|
})
|
|
populate_server_timeout_header(upload_headers, timeout_in_seconds)
|
|
retry_count = 0
|
|
while retry_count < MAX_RETRY_COUNT:
|
|
response = self._httpx_client.request( # type: ignore[union-attr]
|
|
method='POST',
|
|
url=upload_url,
|
|
headers=upload_headers,
|
|
content=file_chunk,
|
|
timeout=timeout_in_seconds,
|
|
)
|
|
if response.headers.get('x-goog-upload-status'):
|
|
break
|
|
delay_seconds = INITIAL_RETRY_DELAY * (DELAY_MULTIPLIER**retry_count)
|
|
retry_count += 1
|
|
time.sleep(delay_seconds)
|
|
|
|
offset += chunk_size
|
|
if response.headers.get('x-goog-upload-status') != 'active':
|
|
break # upload is complete or it has been interrupted.
|
|
if upload_size <= offset: # Status is not finalized.
|
|
raise ValueError(
|
|
f'All content has been uploaded, but the upload status is not'
|
|
f' finalized.'
|
|
)
|
|
errors.APIError.raise_for_response(response)
|
|
if response.headers.get('x-goog-upload-status') != 'final':
|
|
raise ValueError('Failed to upload file: Upload status is not finalized.')
|
|
return HttpResponse(response.headers, response_stream=[response.text])
|
|
|
|
def download_file(
|
|
self,
|
|
path: str,
|
|
*,
|
|
http_options: Optional[HttpOptionsOrDict] = None,
|
|
) -> Union[Any, bytes]:
|
|
"""Downloads the file data.
|
|
|
|
Args:
|
|
path: The request path with query params.
|
|
http_options: The http options to use for the request.
|
|
|
|
returns:
|
|
The file bytes
|
|
"""
|
|
http_request = self._build_request(
|
|
'get', path=path, request_dict={}, http_options=http_options
|
|
)
|
|
|
|
data: Optional[Union[str, bytes]] = None
|
|
if http_request.data:
|
|
if not isinstance(http_request.data, bytes):
|
|
data = json.dumps(http_request.data)
|
|
else:
|
|
data = http_request.data
|
|
|
|
response = self._httpx_client.request( # type: ignore[union-attr]
|
|
method=http_request.method,
|
|
url=http_request.url,
|
|
headers=http_request.headers,
|
|
content=data,
|
|
timeout=http_request.timeout,
|
|
)
|
|
|
|
errors.APIError.raise_for_response(response)
|
|
return HttpResponse(
|
|
response.headers, byte_stream=[response.read()]
|
|
).byte_stream[0]
|
|
|
|
async def async_upload_file(
|
|
self,
|
|
file_path: Union[str, io.IOBase],
|
|
upload_url: str,
|
|
upload_size: int,
|
|
*,
|
|
http_options: Optional[HttpOptionsOrDict] = None,
|
|
) -> HttpResponse:
|
|
"""Transfers a file asynchronously to the given URL.
|
|
|
|
Args:
|
|
file_path: The full path to the file. If the local file path is not found,
|
|
an error will be raised.
|
|
upload_url: The URL to upload the file to.
|
|
upload_size: The size of file content to be uploaded, this will have to
|
|
match the size requested in the resumable upload request.
|
|
http_options: The http options to use for the request.
|
|
|
|
returns:
|
|
The HttpResponse object from the finalize request.
|
|
"""
|
|
if isinstance(file_path, io.IOBase):
|
|
return await self._async_upload_fd(
|
|
file_path, upload_url, upload_size, http_options=http_options
|
|
)
|
|
else:
|
|
file = anyio.Path(file_path)
|
|
fd = await file.open('rb')
|
|
async with fd:
|
|
return await self._async_upload_fd(
|
|
fd, upload_url, upload_size, http_options=http_options
|
|
)
|
|
|
|
async def _async_upload_fd(
|
|
self,
|
|
file: Union[io.IOBase, anyio.AsyncFile[Any]],
|
|
upload_url: str,
|
|
upload_size: int,
|
|
*,
|
|
http_options: Optional[HttpOptionsOrDict] = None,
|
|
) -> HttpResponse:
|
|
"""Transfers a file asynchronously to the given URL.
|
|
|
|
Args:
|
|
file: A file like object inherited from io.BytesIO.
|
|
upload_url: The URL to upload the file to.
|
|
upload_size: The size of file content to be uploaded, this will have to
|
|
match the size requested in the resumable upload request.
|
|
http_options: The http options to use for the request.
|
|
|
|
returns:
|
|
The HttpResponse object from the finalized request.
|
|
"""
|
|
offset = 0
|
|
http_options = http_options if http_options else self._http_options
|
|
base_url = (
|
|
http_options.get('base_url')
|
|
if isinstance(http_options, dict)
|
|
else getattr(http_options, 'base_url', None)
|
|
)
|
|
if base_url:
|
|
parsed_base = urlparse(base_url)
|
|
parsed_upload = urlparse(upload_url)
|
|
upload_url = urlunparse(
|
|
parsed_upload._replace(
|
|
scheme=parsed_base.scheme, netloc=parsed_base.netloc
|
|
)
|
|
)
|
|
|
|
# Upload the file in chunks
|
|
if self._use_aiohttp(): # pylint: disable=g-import-not-at-top
|
|
self._aiohttp_session = await self._get_aiohttp_session() # type: ignore[assignment]
|
|
while True:
|
|
if isinstance(file, io.IOBase):
|
|
file_chunk = file.read(CHUNK_SIZE)
|
|
else:
|
|
file_chunk = await file.read(CHUNK_SIZE)
|
|
chunk_size = 0
|
|
if file_chunk:
|
|
chunk_size = len(file_chunk)
|
|
upload_command = 'upload'
|
|
# If last chunk, finalize the upload.
|
|
if chunk_size + offset >= upload_size:
|
|
upload_command += ', finalize'
|
|
timeout = (
|
|
http_options.get('timeout')
|
|
if isinstance(http_options, dict)
|
|
else http_options.timeout
|
|
)
|
|
if timeout is None:
|
|
# Per request timeout is not configured. Check the global timeout.
|
|
timeout = (
|
|
self._http_options.timeout
|
|
if isinstance(self._http_options, dict)
|
|
else self._http_options.timeout
|
|
)
|
|
timeout_in_seconds = get_timeout_in_seconds(timeout)
|
|
user_headers = (
|
|
http_options.get('headers', {})
|
|
if isinstance(http_options, dict)
|
|
else (getattr(http_options, 'headers', {}) or {})
|
|
)
|
|
upload_headers = dict(user_headers) if user_headers else {}
|
|
upload_headers.update({
|
|
'X-Goog-Upload-Command': upload_command,
|
|
'X-Goog-Upload-Offset': str(offset),
|
|
'Content-Length': str(chunk_size),
|
|
})
|
|
populate_server_timeout_header(upload_headers, timeout_in_seconds)
|
|
|
|
retry_count = 0
|
|
response = None
|
|
while retry_count < MAX_RETRY_COUNT:
|
|
response = await self._aiohttp_session.request( # type: ignore[union-attr]
|
|
method='POST',
|
|
url=upload_url,
|
|
data=file_chunk,
|
|
headers=upload_headers,
|
|
timeout=aiohttp.ClientTimeout(total=timeout_in_seconds),
|
|
)
|
|
|
|
if response.headers.get('X-Goog-Upload-Status'):
|
|
break
|
|
delay_seconds = INITIAL_RETRY_DELAY * (DELAY_MULTIPLIER**retry_count)
|
|
retry_count += 1
|
|
await asyncio.sleep(delay_seconds)
|
|
|
|
offset += chunk_size
|
|
if (
|
|
response is not None
|
|
and response.headers.get('X-Goog-Upload-Status') != 'active'
|
|
):
|
|
break # upload is complete or it has been interrupted.
|
|
|
|
if upload_size <= offset: # Status is not finalized.
|
|
raise ValueError(
|
|
f'All content has been uploaded, but the upload status is not'
|
|
f' finalized.'
|
|
)
|
|
|
|
await errors.APIError.raise_for_async_response(response)
|
|
if (
|
|
response is not None
|
|
and response.headers.get('X-Goog-Upload-Status') != 'final'
|
|
):
|
|
raise ValueError(
|
|
'Failed to upload file: Upload status is not finalized.'
|
|
)
|
|
return HttpResponse(
|
|
response.headers, response_stream=[await response.text()] # type: ignore[union-attr]
|
|
)
|
|
else:
|
|
# aiohttp is not available. Fall back to httpx.
|
|
while True:
|
|
if isinstance(file, io.IOBase):
|
|
file_chunk = file.read(CHUNK_SIZE)
|
|
else:
|
|
file_chunk = await file.read(CHUNK_SIZE)
|
|
chunk_size = 0
|
|
if file_chunk:
|
|
chunk_size = len(file_chunk)
|
|
upload_command = 'upload'
|
|
# If last chunk, finalize the upload.
|
|
if chunk_size + offset >= upload_size:
|
|
upload_command += ', finalize'
|
|
timeout = (
|
|
http_options.get('timeout')
|
|
if isinstance(http_options, dict)
|
|
else http_options.timeout
|
|
)
|
|
if timeout is None:
|
|
# Per request timeout is not configured. Check the global timeout.
|
|
timeout = (
|
|
self._http_options.timeout
|
|
if isinstance(self._http_options, dict)
|
|
else self._http_options.timeout
|
|
)
|
|
timeout_in_seconds = get_timeout_in_seconds(timeout)
|
|
user_headers = (
|
|
http_options.get('headers', {})
|
|
if isinstance(http_options, dict)
|
|
else (getattr(http_options, 'headers', {}) or {})
|
|
)
|
|
upload_headers = dict(user_headers) if user_headers else {}
|
|
upload_headers.update({
|
|
'X-Goog-Upload-Command': upload_command,
|
|
'X-Goog-Upload-Offset': str(offset),
|
|
'Content-Length': str(chunk_size),
|
|
})
|
|
populate_server_timeout_header(upload_headers, timeout_in_seconds)
|
|
|
|
retry_count = 0
|
|
client_response = None
|
|
while retry_count < MAX_RETRY_COUNT:
|
|
client_response = await self._async_httpx_client.request( # type: ignore[union-attr]
|
|
method='POST',
|
|
url=upload_url,
|
|
content=file_chunk,
|
|
headers=upload_headers,
|
|
timeout=timeout_in_seconds,
|
|
)
|
|
if (
|
|
client_response is not None
|
|
and client_response.headers
|
|
and client_response.headers.get('x-goog-upload-status')
|
|
):
|
|
break
|
|
delay_seconds = INITIAL_RETRY_DELAY * (DELAY_MULTIPLIER**retry_count)
|
|
retry_count += 1
|
|
time.sleep(delay_seconds)
|
|
|
|
offset += chunk_size
|
|
if (
|
|
client_response is not None
|
|
and client_response.headers.get('x-goog-upload-status') != 'active'
|
|
):
|
|
break # upload is complete or it has been interrupted.
|
|
|
|
if upload_size <= offset: # Status is not finalized.
|
|
raise ValueError(
|
|
'All content has been uploaded, but the upload status is not'
|
|
' finalized.'
|
|
)
|
|
|
|
await errors.APIError.raise_for_async_response(client_response)
|
|
if (
|
|
client_response is not None
|
|
and client_response.headers.get('x-goog-upload-status') != 'final'
|
|
):
|
|
raise ValueError(
|
|
'Failed to upload file: Upload status is not finalized.'
|
|
)
|
|
return HttpResponse(
|
|
client_response.headers, response_stream=[client_response.text]
|
|
)
|
|
|
|
async def async_download_file(
|
|
self,
|
|
path: str,
|
|
*,
|
|
http_options: Optional[HttpOptionsOrDict] = None,
|
|
) -> Union[Any, bytes]:
|
|
"""Downloads the file data.
|
|
|
|
Args:
|
|
path: The request path with query params.
|
|
http_options: The http options to use for the request.
|
|
|
|
returns:
|
|
The file bytes
|
|
"""
|
|
http_request = self._build_request(
|
|
'get', path=path, request_dict={}, http_options=http_options
|
|
)
|
|
|
|
data: Optional[bytes] = None
|
|
if http_request.data:
|
|
if not isinstance(http_request.data, bytes):
|
|
data = json.dumps(http_request.data).encode('utf-8')
|
|
else:
|
|
data = http_request.data
|
|
|
|
if self._use_aiohttp():
|
|
self._aiohttp_session = await self._get_aiohttp_session() # type: ignore[assignment]
|
|
response = await self._aiohttp_session.request( # type: ignore[union-attr]
|
|
method=http_request.method,
|
|
url=http_request.url,
|
|
headers=http_request.headers,
|
|
data=data,
|
|
timeout=aiohttp.ClientTimeout(total=http_request.timeout),
|
|
)
|
|
await errors.APIError.raise_for_async_response(response)
|
|
|
|
return HttpResponse(
|
|
response.headers, byte_stream=[await response.read()]
|
|
).byte_stream[0]
|
|
else:
|
|
# aiohttp is not available. Fall back to httpx.
|
|
client_response = await self._async_httpx_client.request( # type: ignore[union-attr]
|
|
method=http_request.method,
|
|
url=http_request.url,
|
|
headers=http_request.headers,
|
|
content=data,
|
|
timeout=http_request.timeout,
|
|
)
|
|
await errors.APIError.raise_for_async_response(client_response)
|
|
|
|
return HttpResponse(
|
|
client_response.headers, byte_stream=[client_response.read()]
|
|
).byte_stream[0]
|
|
|
|
# This method does nothing in the real api client. It is used in the
|
|
# replay_api_client to verify the response from the SDK method matches the
|
|
# recorded response.
|
|
def _verify_response(self, response_model: _common.BaseModel) -> None:
|
|
pass
|
|
|
|
def close(self) -> None:
|
|
"""Closes the API client."""
|
|
# Let users close the custom client explicitly by themselves. Otherwise,
|
|
# close the client when the object is garbage collected.
|
|
if not self._http_options.httpx_client and self._httpx_client:
|
|
self._httpx_client.close()
|
|
if self._authorized_session:
|
|
self._authorized_session.close() # type: ignore[no-untyped-call]
|
|
|
|
async def aclose(self) -> None:
|
|
"""Closes the API async client."""
|
|
# Let users close the custom client explicitly by themselves. Otherwise,
|
|
# close the client when the object is garbage collected.
|
|
if not self._http_options.httpx_async_client:
|
|
await self._async_httpx_client.aclose() # type: ignore[union-attr]
|
|
if self._aiohttp_session and not self._http_options.aiohttp_client:
|
|
await self._aiohttp_session.close()
|
|
|
|
def __del__(self) -> None:
|
|
"""Closes the API client when the object is garbage collected.
|
|
|
|
ADK uses this client so cannot rely on the genai.[Async]Client.__del__
|
|
for cleanup.
|
|
"""
|
|
|
|
try:
|
|
if not self._http_options.httpx_client:
|
|
self.close()
|
|
except Exception: # pylint: disable=broad-except
|
|
pass
|
|
|
|
try:
|
|
asyncio.get_running_loop().create_task(self.aclose())
|
|
except Exception: # pylint: disable=broad-except
|
|
pass
|
|
|
|
|
|
def get_token_from_credentials(
|
|
client: 'BaseApiClient',
|
|
credentials: google.auth.credentials.Credentials
|
|
) -> str:
|
|
"""Refreshes the authentication token for the given credentials."""
|
|
if credentials.expired or not credentials.token:
|
|
# Only refresh when it needs to. Default expiration is 3600 seconds.
|
|
refresh_auth(credentials)
|
|
if not credentials.token:
|
|
raise RuntimeError('Could not resolve API token from the environment')
|
|
return credentials.token # type: ignore[no-any-return]
|
|
|
|
|
|
async def async_get_token_from_credentials(
|
|
client: 'BaseApiClient',
|
|
credentials: google.auth.credentials.Credentials
|
|
) -> str:
|
|
"""Refreshes the authentication token for the given credentials."""
|
|
if credentials.expired or not credentials.token:
|
|
# Only refresh when it needs to. Default expiration is 3600 seconds.
|
|
async_auth_lock = await client._get_async_auth_lock()
|
|
async with async_auth_lock:
|
|
if credentials.expired or not credentials.token:
|
|
# Double check that the credentials expired before refreshing.
|
|
await asyncio.to_thread(refresh_auth, credentials)
|
|
|
|
if not credentials.token:
|
|
raise RuntimeError('Could not resolve API token from the environment')
|
|
|
|
return credentials.token # type: ignore[no-any-return]
|