mirror of
https://github.com/wassname/openrouter-python-sdk-retry-errors.git
synced 2026-07-30 12:20:57 +08:00
Initial commit: OpenRouter Python SDK
Auto-generated Python SDK for OpenRouter API, providing comprehensive client library with: - Chat completions and streaming support - Embeddings API - Model and provider information - API key management - Analytics and usage tracking - OAuth authentication flows - Full type hints and error handling
This commit is contained in:
@@ -0,0 +1,200 @@
|
||||
"""Code generated by Speakeasy (https://speakeasy.com). DO NOT EDIT."""
|
||||
|
||||
from typing import TYPE_CHECKING
|
||||
from importlib import import_module
|
||||
import builtins
|
||||
import sys
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from .annotations import get_discriminator
|
||||
from .datetimes import parse_datetime
|
||||
from .enums import OpenEnumMeta
|
||||
from .headers import get_headers, get_response_headers
|
||||
from .metadata import (
|
||||
FieldMetadata,
|
||||
find_metadata,
|
||||
FormMetadata,
|
||||
HeaderMetadata,
|
||||
MultipartFormMetadata,
|
||||
PathParamMetadata,
|
||||
QueryParamMetadata,
|
||||
RequestMetadata,
|
||||
SecurityMetadata,
|
||||
)
|
||||
from .queryparams import get_query_params
|
||||
from .retries import BackoffStrategy, Retries, retry, retry_async, RetryConfig
|
||||
from .requestbodies import serialize_request_body, SerializedRequestBody
|
||||
from .security import get_security, get_security_from_env
|
||||
|
||||
from .serializers import (
|
||||
get_pydantic_model,
|
||||
marshal_json,
|
||||
unmarshal,
|
||||
unmarshal_json,
|
||||
serialize_decimal,
|
||||
serialize_float,
|
||||
serialize_int,
|
||||
stream_to_text,
|
||||
stream_to_text_async,
|
||||
stream_to_bytes,
|
||||
stream_to_bytes_async,
|
||||
validate_const,
|
||||
validate_decimal,
|
||||
validate_float,
|
||||
validate_int,
|
||||
validate_open_enum,
|
||||
)
|
||||
from .url import generate_url, template_url, remove_suffix
|
||||
from .values import (
|
||||
get_global_from_env,
|
||||
match_content_type,
|
||||
match_status_codes,
|
||||
match_response,
|
||||
cast_partial,
|
||||
)
|
||||
from .logger import Logger, get_body_content, get_default_logger
|
||||
|
||||
__all__ = [
|
||||
"BackoffStrategy",
|
||||
"FieldMetadata",
|
||||
"find_metadata",
|
||||
"FormMetadata",
|
||||
"generate_url",
|
||||
"get_body_content",
|
||||
"get_default_logger",
|
||||
"get_discriminator",
|
||||
"parse_datetime",
|
||||
"get_global_from_env",
|
||||
"get_headers",
|
||||
"get_pydantic_model",
|
||||
"get_query_params",
|
||||
"get_response_headers",
|
||||
"get_security",
|
||||
"get_security_from_env",
|
||||
"HeaderMetadata",
|
||||
"Logger",
|
||||
"marshal_json",
|
||||
"match_content_type",
|
||||
"match_status_codes",
|
||||
"match_response",
|
||||
"MultipartFormMetadata",
|
||||
"OpenEnumMeta",
|
||||
"PathParamMetadata",
|
||||
"QueryParamMetadata",
|
||||
"remove_suffix",
|
||||
"Retries",
|
||||
"retry",
|
||||
"retry_async",
|
||||
"RetryConfig",
|
||||
"RequestMetadata",
|
||||
"SecurityMetadata",
|
||||
"serialize_decimal",
|
||||
"serialize_float",
|
||||
"serialize_int",
|
||||
"serialize_request_body",
|
||||
"SerializedRequestBody",
|
||||
"stream_to_text",
|
||||
"stream_to_text_async",
|
||||
"stream_to_bytes",
|
||||
"stream_to_bytes_async",
|
||||
"template_url",
|
||||
"unmarshal",
|
||||
"unmarshal_json",
|
||||
"validate_decimal",
|
||||
"validate_const",
|
||||
"validate_float",
|
||||
"validate_int",
|
||||
"validate_open_enum",
|
||||
"cast_partial",
|
||||
]
|
||||
|
||||
_dynamic_imports: dict[str, str] = {
|
||||
"BackoffStrategy": ".retries",
|
||||
"FieldMetadata": ".metadata",
|
||||
"find_metadata": ".metadata",
|
||||
"FormMetadata": ".metadata",
|
||||
"generate_url": ".url",
|
||||
"get_body_content": ".logger",
|
||||
"get_default_logger": ".logger",
|
||||
"get_discriminator": ".annotations",
|
||||
"parse_datetime": ".datetimes",
|
||||
"get_global_from_env": ".values",
|
||||
"get_headers": ".headers",
|
||||
"get_pydantic_model": ".serializers",
|
||||
"get_query_params": ".queryparams",
|
||||
"get_response_headers": ".headers",
|
||||
"get_security": ".security",
|
||||
"get_security_from_env": ".security",
|
||||
"HeaderMetadata": ".metadata",
|
||||
"Logger": ".logger",
|
||||
"marshal_json": ".serializers",
|
||||
"match_content_type": ".values",
|
||||
"match_status_codes": ".values",
|
||||
"match_response": ".values",
|
||||
"MultipartFormMetadata": ".metadata",
|
||||
"OpenEnumMeta": ".enums",
|
||||
"PathParamMetadata": ".metadata",
|
||||
"QueryParamMetadata": ".metadata",
|
||||
"remove_suffix": ".url",
|
||||
"Retries": ".retries",
|
||||
"retry": ".retries",
|
||||
"retry_async": ".retries",
|
||||
"RetryConfig": ".retries",
|
||||
"RequestMetadata": ".metadata",
|
||||
"SecurityMetadata": ".metadata",
|
||||
"serialize_decimal": ".serializers",
|
||||
"serialize_float": ".serializers",
|
||||
"serialize_int": ".serializers",
|
||||
"serialize_request_body": ".requestbodies",
|
||||
"SerializedRequestBody": ".requestbodies",
|
||||
"stream_to_text": ".serializers",
|
||||
"stream_to_text_async": ".serializers",
|
||||
"stream_to_bytes": ".serializers",
|
||||
"stream_to_bytes_async": ".serializers",
|
||||
"template_url": ".url",
|
||||
"unmarshal": ".serializers",
|
||||
"unmarshal_json": ".serializers",
|
||||
"validate_decimal": ".serializers",
|
||||
"validate_const": ".serializers",
|
||||
"validate_float": ".serializers",
|
||||
"validate_int": ".serializers",
|
||||
"validate_open_enum": ".serializers",
|
||||
"cast_partial": ".values",
|
||||
}
|
||||
|
||||
|
||||
def dynamic_import(modname, retries=3):
|
||||
for attempt in range(retries):
|
||||
try:
|
||||
return import_module(modname, __package__)
|
||||
except KeyError:
|
||||
# Clear any half-initialized module and retry
|
||||
sys.modules.pop(modname, None)
|
||||
if attempt == retries - 1:
|
||||
break
|
||||
raise KeyError(f"Failed to import module '{modname}' after {retries} attempts")
|
||||
|
||||
|
||||
def __getattr__(attr_name: str) -> object:
|
||||
module_name = _dynamic_imports.get(attr_name)
|
||||
if module_name is None:
|
||||
raise AttributeError(
|
||||
f"no {attr_name} found in _dynamic_imports, module name -> {__name__} "
|
||||
)
|
||||
|
||||
try:
|
||||
module = dynamic_import(module_name)
|
||||
return getattr(module, attr_name)
|
||||
except ImportError as e:
|
||||
raise ImportError(
|
||||
f"Failed to import {attr_name} from {module_name}: {e}"
|
||||
) from e
|
||||
except AttributeError as e:
|
||||
raise AttributeError(
|
||||
f"Failed to get {attr_name} from {module_name}: {e}"
|
||||
) from e
|
||||
|
||||
|
||||
def __dir__():
|
||||
lazy_attrs = builtins.list(_dynamic_imports.keys())
|
||||
return builtins.sorted(lazy_attrs)
|
||||
@@ -0,0 +1,79 @@
|
||||
"""Code generated by Speakeasy (https://speakeasy.com). DO NOT EDIT."""
|
||||
|
||||
from enum import Enum
|
||||
from typing import Any, Optional
|
||||
|
||||
|
||||
def get_discriminator(model: Any, fieldname: str, key: str) -> str:
|
||||
"""
|
||||
Recursively search for the discriminator attribute in a model.
|
||||
|
||||
Args:
|
||||
model (Any): The model to search within.
|
||||
fieldname (str): The name of the field to search for.
|
||||
key (str): The key to search for in dictionaries.
|
||||
|
||||
Returns:
|
||||
str: The name of the discriminator attribute.
|
||||
|
||||
Raises:
|
||||
ValueError: If the discriminator attribute is not found.
|
||||
"""
|
||||
upper_fieldname = fieldname.upper()
|
||||
|
||||
def get_field_discriminator(field: Any) -> Optional[str]:
|
||||
"""Search for the discriminator attribute in a given field."""
|
||||
|
||||
if isinstance(field, dict):
|
||||
if key in field:
|
||||
return f"{field[key]}"
|
||||
|
||||
if hasattr(field, fieldname):
|
||||
attr = getattr(field, fieldname)
|
||||
if isinstance(attr, Enum):
|
||||
return f"{attr.value}"
|
||||
return f"{attr}"
|
||||
|
||||
if hasattr(field, upper_fieldname):
|
||||
attr = getattr(field, upper_fieldname)
|
||||
if isinstance(attr, Enum):
|
||||
return f"{attr.value}"
|
||||
return f"{attr}"
|
||||
|
||||
return None
|
||||
|
||||
def search_nested_discriminator(obj: Any) -> Optional[str]:
|
||||
"""Recursively search for discriminator in nested structures."""
|
||||
# First try direct field lookup
|
||||
discriminator = get_field_discriminator(obj)
|
||||
if discriminator is not None:
|
||||
return discriminator
|
||||
|
||||
# If it's a dict, search in nested values
|
||||
if isinstance(obj, dict):
|
||||
for value in obj.values():
|
||||
if isinstance(value, list):
|
||||
# Search in list items
|
||||
for item in value:
|
||||
nested_discriminator = search_nested_discriminator(item)
|
||||
if nested_discriminator is not None:
|
||||
return nested_discriminator
|
||||
elif isinstance(value, dict):
|
||||
# Search in nested dict
|
||||
nested_discriminator = search_nested_discriminator(value)
|
||||
if nested_discriminator is not None:
|
||||
return nested_discriminator
|
||||
|
||||
return None
|
||||
|
||||
if isinstance(model, list):
|
||||
for field in model:
|
||||
discriminator = search_nested_discriminator(field)
|
||||
if discriminator is not None:
|
||||
return discriminator
|
||||
|
||||
discriminator = search_nested_discriminator(model)
|
||||
if discriminator is not None:
|
||||
return discriminator
|
||||
|
||||
raise ValueError(f"Could not find discriminator field {fieldname} in {model}")
|
||||
@@ -0,0 +1,23 @@
|
||||
"""Code generated by Speakeasy (https://speakeasy.com). DO NOT EDIT."""
|
||||
|
||||
from datetime import datetime
|
||||
import sys
|
||||
|
||||
|
||||
def parse_datetime(datetime_string: str) -> datetime:
|
||||
"""
|
||||
Convert a RFC 3339 / ISO 8601 formatted string into a datetime object.
|
||||
Python versions 3.11 and later support parsing RFC 3339 directly with
|
||||
datetime.fromisoformat(), but for earlier versions, this function
|
||||
encapsulates the necessary extra logic.
|
||||
"""
|
||||
# Python 3.11 and later can parse RFC 3339 directly
|
||||
if sys.version_info >= (3, 11):
|
||||
return datetime.fromisoformat(datetime_string)
|
||||
|
||||
# For Python 3.10 and earlier, a common ValueError is trailing 'Z' suffix,
|
||||
# so fix that upfront.
|
||||
if datetime_string.endswith("Z"):
|
||||
datetime_string = datetime_string[:-1] + "+00:00"
|
||||
|
||||
return datetime.fromisoformat(datetime_string)
|
||||
@@ -0,0 +1,74 @@
|
||||
"""Code generated by Speakeasy (https://speakeasy.com). DO NOT EDIT."""
|
||||
|
||||
import enum
|
||||
import sys
|
||||
|
||||
class OpenEnumMeta(enum.EnumMeta):
|
||||
# The __call__ method `boundary` kwarg was added in 3.11 and must be present
|
||||
# for pyright. Refer also: https://github.com/pylint-dev/pylint/issues/9622
|
||||
# pylint: disable=unexpected-keyword-arg
|
||||
# The __call__ method `values` varg must be named for pyright.
|
||||
# pylint: disable=keyword-arg-before-vararg
|
||||
|
||||
if sys.version_info >= (3, 11):
|
||||
def __call__(
|
||||
cls, value, names=None, *values, module=None, qualname=None, type=None, start=1, boundary=None
|
||||
):
|
||||
# The `type` kwarg also happens to be a built-in that pylint flags as
|
||||
# redeclared. Safe to ignore this lint rule with this scope.
|
||||
# pylint: disable=redefined-builtin
|
||||
|
||||
if names is not None:
|
||||
return super().__call__(
|
||||
value,
|
||||
names=names,
|
||||
*values,
|
||||
module=module,
|
||||
qualname=qualname,
|
||||
type=type,
|
||||
start=start,
|
||||
boundary=boundary,
|
||||
)
|
||||
|
||||
try:
|
||||
return super().__call__(
|
||||
value,
|
||||
names=names, # pyright: ignore[reportArgumentType]
|
||||
*values,
|
||||
module=module,
|
||||
qualname=qualname,
|
||||
type=type,
|
||||
start=start,
|
||||
boundary=boundary,
|
||||
)
|
||||
except ValueError:
|
||||
return value
|
||||
else:
|
||||
def __call__(
|
||||
cls, value, names=None, *, module=None, qualname=None, type=None, start=1
|
||||
):
|
||||
# The `type` kwarg also happens to be a built-in that pylint flags as
|
||||
# redeclared. Safe to ignore this lint rule with this scope.
|
||||
# pylint: disable=redefined-builtin
|
||||
|
||||
if names is not None:
|
||||
return super().__call__(
|
||||
value,
|
||||
names=names,
|
||||
module=module,
|
||||
qualname=qualname,
|
||||
type=type,
|
||||
start=start,
|
||||
)
|
||||
|
||||
try:
|
||||
return super().__call__(
|
||||
value,
|
||||
names=names, # pyright: ignore[reportArgumentType]
|
||||
module=module,
|
||||
qualname=qualname,
|
||||
type=type,
|
||||
start=start,
|
||||
)
|
||||
except ValueError:
|
||||
return value
|
||||
@@ -0,0 +1,248 @@
|
||||
"""Code generated by Speakeasy (https://speakeasy.com). DO NOT EDIT."""
|
||||
|
||||
import re
|
||||
import json
|
||||
from typing import (
|
||||
Callable,
|
||||
Generic,
|
||||
TypeVar,
|
||||
Optional,
|
||||
Generator,
|
||||
AsyncGenerator,
|
||||
Tuple,
|
||||
)
|
||||
import httpx
|
||||
|
||||
T = TypeVar("T")
|
||||
|
||||
|
||||
class EventStream(Generic[T]):
|
||||
# Holds a reference to the SDK client to avoid it being garbage collected
|
||||
# and cause termination of the underlying httpx client.
|
||||
client_ref: Optional[object]
|
||||
response: httpx.Response
|
||||
generator: Generator[T, None, None]
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
response: httpx.Response,
|
||||
decoder: Callable[[str], T],
|
||||
sentinel: Optional[str] = None,
|
||||
client_ref: Optional[object] = None,
|
||||
):
|
||||
self.response = response
|
||||
self.generator = stream_events(response, decoder, sentinel)
|
||||
self.client_ref = client_ref
|
||||
|
||||
def __iter__(self):
|
||||
return self
|
||||
|
||||
def __next__(self):
|
||||
return next(self.generator)
|
||||
|
||||
def __enter__(self):
|
||||
return self
|
||||
|
||||
def __exit__(self, exc_type, exc_val, exc_tb):
|
||||
self.response.close()
|
||||
|
||||
|
||||
class EventStreamAsync(Generic[T]):
|
||||
# Holds a reference to the SDK client to avoid it being garbage collected
|
||||
# and cause termination of the underlying httpx client.
|
||||
client_ref: Optional[object]
|
||||
response: httpx.Response
|
||||
generator: AsyncGenerator[T, None]
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
response: httpx.Response,
|
||||
decoder: Callable[[str], T],
|
||||
sentinel: Optional[str] = None,
|
||||
client_ref: Optional[object] = None,
|
||||
):
|
||||
self.response = response
|
||||
self.generator = stream_events_async(response, decoder, sentinel)
|
||||
self.client_ref = client_ref
|
||||
|
||||
def __aiter__(self):
|
||||
return self
|
||||
|
||||
async def __anext__(self):
|
||||
return await self.generator.__anext__()
|
||||
|
||||
async def __aenter__(self):
|
||||
return self
|
||||
|
||||
async def __aexit__(self, exc_type, exc_val, exc_tb):
|
||||
await self.response.aclose()
|
||||
|
||||
|
||||
class ServerEvent:
|
||||
id: Optional[str] = None
|
||||
event: Optional[str] = None
|
||||
data: Optional[str] = None
|
||||
retry: Optional[int] = None
|
||||
|
||||
|
||||
MESSAGE_BOUNDARIES = [
|
||||
b"\r\n\r\n",
|
||||
b"\n\n",
|
||||
b"\r\r",
|
||||
]
|
||||
|
||||
|
||||
async def stream_events_async(
|
||||
response: httpx.Response,
|
||||
decoder: Callable[[str], T],
|
||||
sentinel: Optional[str] = None,
|
||||
) -> AsyncGenerator[T, None]:
|
||||
buffer = bytearray()
|
||||
position = 0
|
||||
discard = False
|
||||
async for chunk in response.aiter_bytes():
|
||||
# We've encountered the sentinel value and should no longer process
|
||||
# incoming data. Instead we throw new data away until the server closes
|
||||
# the connection.
|
||||
if discard:
|
||||
continue
|
||||
|
||||
buffer += chunk
|
||||
for i in range(position, len(buffer)):
|
||||
char = buffer[i : i + 1]
|
||||
seq: Optional[bytes] = None
|
||||
if char in [b"\r", b"\n"]:
|
||||
for boundary in MESSAGE_BOUNDARIES:
|
||||
seq = _peek_sequence(i, buffer, boundary)
|
||||
if seq is not None:
|
||||
break
|
||||
if seq is None:
|
||||
continue
|
||||
|
||||
block = buffer[position:i]
|
||||
position = i + len(seq)
|
||||
event, discard = _parse_event(block, decoder, sentinel)
|
||||
if event is not None:
|
||||
yield event
|
||||
|
||||
if position > 0:
|
||||
buffer = buffer[position:]
|
||||
position = 0
|
||||
|
||||
event, discard = _parse_event(buffer, decoder, sentinel)
|
||||
if event is not None:
|
||||
yield event
|
||||
|
||||
|
||||
def stream_events(
|
||||
response: httpx.Response,
|
||||
decoder: Callable[[str], T],
|
||||
sentinel: Optional[str] = None,
|
||||
) -> Generator[T, None, None]:
|
||||
buffer = bytearray()
|
||||
position = 0
|
||||
discard = False
|
||||
for chunk in response.iter_bytes():
|
||||
# We've encountered the sentinel value and should no longer process
|
||||
# incoming data. Instead we throw new data away until the server closes
|
||||
# the connection.
|
||||
if discard:
|
||||
continue
|
||||
|
||||
buffer += chunk
|
||||
for i in range(position, len(buffer)):
|
||||
char = buffer[i : i + 1]
|
||||
seq: Optional[bytes] = None
|
||||
if char in [b"\r", b"\n"]:
|
||||
for boundary in MESSAGE_BOUNDARIES:
|
||||
seq = _peek_sequence(i, buffer, boundary)
|
||||
if seq is not None:
|
||||
break
|
||||
if seq is None:
|
||||
continue
|
||||
|
||||
block = buffer[position:i]
|
||||
position = i + len(seq)
|
||||
event, discard = _parse_event(block, decoder, sentinel)
|
||||
if event is not None:
|
||||
yield event
|
||||
|
||||
if position > 0:
|
||||
buffer = buffer[position:]
|
||||
position = 0
|
||||
|
||||
event, discard = _parse_event(buffer, decoder, sentinel)
|
||||
if event is not None:
|
||||
yield event
|
||||
|
||||
|
||||
def _parse_event(
|
||||
raw: bytearray, decoder: Callable[[str], T], sentinel: Optional[str] = None
|
||||
) -> Tuple[Optional[T], bool]:
|
||||
block = raw.decode()
|
||||
lines = re.split(r"\r?\n|\r", block)
|
||||
publish = False
|
||||
event = ServerEvent()
|
||||
data = ""
|
||||
for line in lines:
|
||||
if not line:
|
||||
continue
|
||||
|
||||
delim = line.find(":")
|
||||
if delim <= 0:
|
||||
continue
|
||||
|
||||
field = line[0:delim]
|
||||
value = line[delim + 1 :] if delim < len(line) - 1 else ""
|
||||
if len(value) and value[0] == " ":
|
||||
value = value[1:]
|
||||
|
||||
if field == "event":
|
||||
event.event = value
|
||||
publish = True
|
||||
elif field == "data":
|
||||
data += value + "\n"
|
||||
publish = True
|
||||
elif field == "id":
|
||||
event.id = value
|
||||
publish = True
|
||||
elif field == "retry":
|
||||
event.retry = int(value) if value.isdigit() else None
|
||||
publish = True
|
||||
|
||||
if sentinel and data == f"{sentinel}\n":
|
||||
return None, True
|
||||
|
||||
if data:
|
||||
data = data[:-1]
|
||||
event.data = data
|
||||
|
||||
data_is_primitive = (
|
||||
data.isnumeric() or data == "true" or data == "false" or data == "null"
|
||||
)
|
||||
data_is_json = (
|
||||
data.startswith("{") or data.startswith("[") or data.startswith('"')
|
||||
)
|
||||
|
||||
if data_is_primitive or data_is_json:
|
||||
try:
|
||||
event.data = json.loads(data)
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
out = None
|
||||
if publish:
|
||||
out = decoder(json.dumps(event.__dict__))
|
||||
|
||||
return out, False
|
||||
|
||||
|
||||
def _peek_sequence(position: int, buffer: bytearray, sequence: bytes):
|
||||
if len(sequence) > (len(buffer) - position):
|
||||
return None
|
||||
|
||||
for i, seq in enumerate(sequence):
|
||||
if buffer[position + i] != seq:
|
||||
return None
|
||||
|
||||
return sequence
|
||||
@@ -0,0 +1,223 @@
|
||||
"""Code generated by Speakeasy (https://speakeasy.com). DO NOT EDIT."""
|
||||
|
||||
from typing import (
|
||||
Any,
|
||||
Dict,
|
||||
get_type_hints,
|
||||
List,
|
||||
Tuple,
|
||||
)
|
||||
from pydantic import BaseModel
|
||||
from pydantic.fields import FieldInfo
|
||||
|
||||
from .serializers import marshal_json
|
||||
|
||||
from .metadata import (
|
||||
FormMetadata,
|
||||
MultipartFormMetadata,
|
||||
find_field_metadata,
|
||||
)
|
||||
from .values import _is_set, _val_to_string
|
||||
|
||||
|
||||
def _populate_form(
|
||||
field_name: str,
|
||||
explode: bool,
|
||||
obj: Any,
|
||||
delimiter: str,
|
||||
form: Dict[str, List[str]],
|
||||
):
|
||||
if not _is_set(obj):
|
||||
return form
|
||||
|
||||
if isinstance(obj, BaseModel):
|
||||
items = []
|
||||
|
||||
obj_fields: Dict[str, FieldInfo] = obj.__class__.model_fields
|
||||
for name in obj_fields:
|
||||
obj_field = obj_fields[name]
|
||||
obj_field_name = obj_field.alias if obj_field.alias is not None else name
|
||||
if obj_field_name == "":
|
||||
continue
|
||||
|
||||
val = getattr(obj, name)
|
||||
if not _is_set(val):
|
||||
continue
|
||||
|
||||
if explode:
|
||||
form[obj_field_name] = [_val_to_string(val)]
|
||||
else:
|
||||
items.append(f"{obj_field_name}{delimiter}{_val_to_string(val)}")
|
||||
|
||||
if len(items) > 0:
|
||||
form[field_name] = [delimiter.join(items)]
|
||||
elif isinstance(obj, Dict):
|
||||
items = []
|
||||
for key, value in obj.items():
|
||||
if not _is_set(value):
|
||||
continue
|
||||
|
||||
if explode:
|
||||
form[key] = [_val_to_string(value)]
|
||||
else:
|
||||
items.append(f"{key}{delimiter}{_val_to_string(value)}")
|
||||
|
||||
if len(items) > 0:
|
||||
form[field_name] = [delimiter.join(items)]
|
||||
elif isinstance(obj, List):
|
||||
items = []
|
||||
|
||||
for value in obj:
|
||||
if not _is_set(value):
|
||||
continue
|
||||
|
||||
if explode:
|
||||
if not field_name in form:
|
||||
form[field_name] = []
|
||||
form[field_name].append(_val_to_string(value))
|
||||
else:
|
||||
items.append(_val_to_string(value))
|
||||
|
||||
if len(items) > 0:
|
||||
form[field_name] = [delimiter.join([str(item) for item in items])]
|
||||
else:
|
||||
form[field_name] = [_val_to_string(obj)]
|
||||
|
||||
return form
|
||||
|
||||
|
||||
def _extract_file_properties(file_obj: Any) -> Tuple[str, Any, Any]:
|
||||
"""Extract file name, content, and content type from a file object."""
|
||||
file_fields: Dict[str, FieldInfo] = file_obj.__class__.model_fields
|
||||
|
||||
file_name = ""
|
||||
content = None
|
||||
content_type = None
|
||||
|
||||
for file_field_name in file_fields:
|
||||
file_field = file_fields[file_field_name]
|
||||
|
||||
file_metadata = find_field_metadata(file_field, MultipartFormMetadata)
|
||||
if file_metadata is None:
|
||||
continue
|
||||
|
||||
if file_metadata.content:
|
||||
content = getattr(file_obj, file_field_name, None)
|
||||
elif file_field_name == "content_type":
|
||||
content_type = getattr(file_obj, file_field_name, None)
|
||||
else:
|
||||
file_name = getattr(file_obj, file_field_name)
|
||||
|
||||
if file_name == "" or content is None:
|
||||
raise ValueError("invalid multipart/form-data file")
|
||||
|
||||
return file_name, content, content_type
|
||||
|
||||
|
||||
def serialize_multipart_form(
|
||||
media_type: str, request: Any
|
||||
) -> Tuple[str, Dict[str, Any], List[Tuple[str, Any]]]:
|
||||
form: Dict[str, Any] = {}
|
||||
files: List[Tuple[str, Any]] = []
|
||||
|
||||
if not isinstance(request, BaseModel):
|
||||
raise TypeError("invalid request body type")
|
||||
|
||||
request_fields: Dict[str, FieldInfo] = request.__class__.model_fields
|
||||
request_field_types = get_type_hints(request.__class__)
|
||||
|
||||
for name in request_fields:
|
||||
field = request_fields[name]
|
||||
|
||||
val = getattr(request, name)
|
||||
if not _is_set(val):
|
||||
continue
|
||||
|
||||
field_metadata = find_field_metadata(field, MultipartFormMetadata)
|
||||
if not field_metadata:
|
||||
continue
|
||||
|
||||
f_name = field.alias if field.alias else name
|
||||
|
||||
if field_metadata.file:
|
||||
if isinstance(val, List):
|
||||
# Handle array of files
|
||||
for file_obj in val:
|
||||
if not _is_set(file_obj):
|
||||
continue
|
||||
|
||||
file_name, content, content_type = _extract_file_properties(file_obj)
|
||||
|
||||
if content_type is not None:
|
||||
files.append((f_name + "[]", (file_name, content, content_type)))
|
||||
else:
|
||||
files.append((f_name + "[]", (file_name, content)))
|
||||
else:
|
||||
# Handle single file
|
||||
file_name, content, content_type = _extract_file_properties(val)
|
||||
|
||||
if content_type is not None:
|
||||
files.append((f_name, (file_name, content, content_type)))
|
||||
else:
|
||||
files.append((f_name, (file_name, content)))
|
||||
elif field_metadata.json:
|
||||
files.append((f_name, (
|
||||
None,
|
||||
marshal_json(val, request_field_types[name]),
|
||||
"application/json",
|
||||
)))
|
||||
else:
|
||||
if isinstance(val, List):
|
||||
values = []
|
||||
|
||||
for value in val:
|
||||
if not _is_set(value):
|
||||
continue
|
||||
values.append(_val_to_string(value))
|
||||
|
||||
form[f_name + "[]"] = values
|
||||
else:
|
||||
form[f_name] = _val_to_string(val)
|
||||
return media_type, form, files
|
||||
|
||||
|
||||
def serialize_form_data(data: Any) -> Dict[str, Any]:
|
||||
form: Dict[str, List[str]] = {}
|
||||
|
||||
if isinstance(data, BaseModel):
|
||||
data_fields: Dict[str, FieldInfo] = data.__class__.model_fields
|
||||
data_field_types = get_type_hints(data.__class__)
|
||||
for name in data_fields:
|
||||
field = data_fields[name]
|
||||
|
||||
val = getattr(data, name)
|
||||
if not _is_set(val):
|
||||
continue
|
||||
|
||||
metadata = find_field_metadata(field, FormMetadata)
|
||||
if metadata is None:
|
||||
continue
|
||||
|
||||
f_name = field.alias if field.alias is not None else name
|
||||
|
||||
if metadata.json:
|
||||
form[f_name] = [marshal_json(val, data_field_types[name])]
|
||||
else:
|
||||
if metadata.style == "form":
|
||||
_populate_form(
|
||||
f_name,
|
||||
metadata.explode,
|
||||
val,
|
||||
",",
|
||||
form,
|
||||
)
|
||||
else:
|
||||
raise ValueError(f"Invalid form style for field {name}")
|
||||
elif isinstance(data, Dict):
|
||||
for key, value in data.items():
|
||||
if _is_set(value):
|
||||
form[key] = [_val_to_string(value)]
|
||||
else:
|
||||
raise TypeError(f"Invalid request body type {type(data)} for form data")
|
||||
|
||||
return form
|
||||
@@ -0,0 +1,136 @@
|
||||
"""Code generated by Speakeasy (https://speakeasy.com). DO NOT EDIT."""
|
||||
|
||||
from typing import (
|
||||
Any,
|
||||
Dict,
|
||||
List,
|
||||
Optional,
|
||||
)
|
||||
from httpx import Headers
|
||||
from pydantic import BaseModel
|
||||
from pydantic.fields import FieldInfo
|
||||
|
||||
from .metadata import (
|
||||
HeaderMetadata,
|
||||
find_field_metadata,
|
||||
)
|
||||
|
||||
from .values import _is_set, _populate_from_globals, _val_to_string
|
||||
|
||||
|
||||
def get_headers(headers_params: Any, gbls: Optional[Any] = None) -> Dict[str, str]:
|
||||
headers: Dict[str, str] = {}
|
||||
|
||||
globals_already_populated = []
|
||||
if _is_set(headers_params):
|
||||
globals_already_populated = _populate_headers(headers_params, gbls, headers, [])
|
||||
if _is_set(gbls):
|
||||
_populate_headers(gbls, None, headers, globals_already_populated)
|
||||
|
||||
return headers
|
||||
|
||||
|
||||
def _populate_headers(
|
||||
headers_params: Any,
|
||||
gbls: Any,
|
||||
header_values: Dict[str, str],
|
||||
skip_fields: List[str],
|
||||
) -> List[str]:
|
||||
globals_already_populated: List[str] = []
|
||||
|
||||
if not isinstance(headers_params, BaseModel):
|
||||
return globals_already_populated
|
||||
|
||||
param_fields: Dict[str, FieldInfo] = headers_params.__class__.model_fields
|
||||
for name in param_fields:
|
||||
if name in skip_fields:
|
||||
continue
|
||||
|
||||
field = param_fields[name]
|
||||
f_name = field.alias if field.alias is not None else name
|
||||
|
||||
metadata = find_field_metadata(field, HeaderMetadata)
|
||||
if metadata is None:
|
||||
continue
|
||||
|
||||
value, global_found = _populate_from_globals(
|
||||
name, getattr(headers_params, name), HeaderMetadata, gbls
|
||||
)
|
||||
if global_found:
|
||||
globals_already_populated.append(name)
|
||||
value = _serialize_header(metadata.explode, value)
|
||||
|
||||
if value != "":
|
||||
header_values[f_name] = value
|
||||
|
||||
return globals_already_populated
|
||||
|
||||
|
||||
def _serialize_header(explode: bool, obj: Any) -> str:
|
||||
if not _is_set(obj):
|
||||
return ""
|
||||
|
||||
if isinstance(obj, BaseModel):
|
||||
items = []
|
||||
obj_fields: Dict[str, FieldInfo] = obj.__class__.model_fields
|
||||
for name in obj_fields:
|
||||
obj_field = obj_fields[name]
|
||||
obj_param_metadata = find_field_metadata(obj_field, HeaderMetadata)
|
||||
|
||||
if not obj_param_metadata:
|
||||
continue
|
||||
|
||||
f_name = obj_field.alias if obj_field.alias is not None else name
|
||||
|
||||
val = getattr(obj, name)
|
||||
if not _is_set(val):
|
||||
continue
|
||||
|
||||
if explode:
|
||||
items.append(f"{f_name}={_val_to_string(val)}")
|
||||
else:
|
||||
items.append(f_name)
|
||||
items.append(_val_to_string(val))
|
||||
|
||||
if len(items) > 0:
|
||||
return ",".join(items)
|
||||
elif isinstance(obj, Dict):
|
||||
items = []
|
||||
|
||||
for key, value in obj.items():
|
||||
if not _is_set(value):
|
||||
continue
|
||||
|
||||
if explode:
|
||||
items.append(f"{key}={_val_to_string(value)}")
|
||||
else:
|
||||
items.append(key)
|
||||
items.append(_val_to_string(value))
|
||||
|
||||
if len(items) > 0:
|
||||
return ",".join([str(item) for item in items])
|
||||
elif isinstance(obj, List):
|
||||
items = []
|
||||
|
||||
for value in obj:
|
||||
if not _is_set(value):
|
||||
continue
|
||||
|
||||
items.append(_val_to_string(value))
|
||||
|
||||
if len(items) > 0:
|
||||
return ",".join(items)
|
||||
elif _is_set(obj):
|
||||
return f"{_val_to_string(obj)}"
|
||||
|
||||
return ""
|
||||
|
||||
|
||||
def get_response_headers(headers: Headers) -> Dict[str, List[str]]:
|
||||
res: Dict[str, List[str]] = {}
|
||||
for k, v in headers.items():
|
||||
if not k in res:
|
||||
res[k] = []
|
||||
|
||||
res[k].append(v)
|
||||
return res
|
||||
@@ -0,0 +1,27 @@
|
||||
"""Code generated by Speakeasy (https://speakeasy.com). DO NOT EDIT."""
|
||||
|
||||
import httpx
|
||||
import logging
|
||||
import os
|
||||
from typing import Any, Protocol
|
||||
|
||||
|
||||
class Logger(Protocol):
|
||||
def debug(self, msg: str, *args: Any, **kwargs: Any) -> None:
|
||||
pass
|
||||
|
||||
|
||||
class NoOpLogger:
|
||||
def debug(self, msg: str, *args: Any, **kwargs: Any) -> None:
|
||||
pass
|
||||
|
||||
|
||||
def get_body_content(req: httpx.Request) -> str:
|
||||
return "<streaming body>" if not hasattr(req, "_content") else str(req.content)
|
||||
|
||||
|
||||
def get_default_logger() -> Logger:
|
||||
if os.getenv("OPENROUTER_DEBUG"):
|
||||
logging.basicConfig(level=logging.DEBUG)
|
||||
return logging.getLogger("openrouter")
|
||||
return NoOpLogger()
|
||||
@@ -0,0 +1,118 @@
|
||||
"""Code generated by Speakeasy (https://speakeasy.com). DO NOT EDIT."""
|
||||
|
||||
from typing import Optional, Type, TypeVar, Union
|
||||
from dataclasses import dataclass
|
||||
from pydantic.fields import FieldInfo
|
||||
|
||||
|
||||
T = TypeVar("T")
|
||||
|
||||
|
||||
@dataclass
|
||||
class SecurityMetadata:
|
||||
option: bool = False
|
||||
scheme: bool = False
|
||||
scheme_type: Optional[str] = None
|
||||
sub_type: Optional[str] = None
|
||||
field_name: Optional[str] = None
|
||||
|
||||
def get_field_name(self, default: str) -> str:
|
||||
return self.field_name or default
|
||||
|
||||
|
||||
@dataclass
|
||||
class ParamMetadata:
|
||||
serialization: Optional[str] = None
|
||||
style: str = "simple"
|
||||
explode: bool = False
|
||||
|
||||
|
||||
@dataclass
|
||||
class PathParamMetadata(ParamMetadata):
|
||||
pass
|
||||
|
||||
|
||||
@dataclass
|
||||
class QueryParamMetadata(ParamMetadata):
|
||||
style: str = "form"
|
||||
explode: bool = True
|
||||
|
||||
|
||||
@dataclass
|
||||
class HeaderMetadata(ParamMetadata):
|
||||
pass
|
||||
|
||||
|
||||
@dataclass
|
||||
class RequestMetadata:
|
||||
media_type: str = "application/octet-stream"
|
||||
|
||||
|
||||
@dataclass
|
||||
class MultipartFormMetadata:
|
||||
file: bool = False
|
||||
content: bool = False
|
||||
json: bool = False
|
||||
|
||||
|
||||
@dataclass
|
||||
class FormMetadata:
|
||||
json: bool = False
|
||||
style: str = "form"
|
||||
explode: bool = True
|
||||
|
||||
|
||||
class FieldMetadata:
|
||||
security: Optional[SecurityMetadata] = None
|
||||
path: Optional[PathParamMetadata] = None
|
||||
query: Optional[QueryParamMetadata] = None
|
||||
header: Optional[HeaderMetadata] = None
|
||||
request: Optional[RequestMetadata] = None
|
||||
form: Optional[FormMetadata] = None
|
||||
multipart: Optional[MultipartFormMetadata] = None
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
security: Optional[SecurityMetadata] = None,
|
||||
path: Optional[Union[PathParamMetadata, bool]] = None,
|
||||
query: Optional[Union[QueryParamMetadata, bool]] = None,
|
||||
header: Optional[Union[HeaderMetadata, bool]] = None,
|
||||
request: Optional[Union[RequestMetadata, bool]] = None,
|
||||
form: Optional[Union[FormMetadata, bool]] = None,
|
||||
multipart: Optional[Union[MultipartFormMetadata, bool]] = None,
|
||||
):
|
||||
self.security = security
|
||||
self.path = PathParamMetadata() if isinstance(path, bool) else path
|
||||
self.query = QueryParamMetadata() if isinstance(query, bool) else query
|
||||
self.header = HeaderMetadata() if isinstance(header, bool) else header
|
||||
self.request = RequestMetadata() if isinstance(request, bool) else request
|
||||
self.form = FormMetadata() if isinstance(form, bool) else form
|
||||
self.multipart = (
|
||||
MultipartFormMetadata() if isinstance(multipart, bool) else multipart
|
||||
)
|
||||
|
||||
|
||||
def find_field_metadata(field_info: FieldInfo, metadata_type: Type[T]) -> Optional[T]:
|
||||
metadata = find_metadata(field_info, FieldMetadata)
|
||||
if not metadata:
|
||||
return None
|
||||
|
||||
fields = metadata.__dict__
|
||||
|
||||
for field in fields:
|
||||
if isinstance(fields[field], metadata_type):
|
||||
return fields[field]
|
||||
|
||||
return None
|
||||
|
||||
|
||||
def find_metadata(field_info: FieldInfo, metadata_type: Type[T]) -> Optional[T]:
|
||||
metadata = field_info.metadata
|
||||
if not metadata:
|
||||
return None
|
||||
|
||||
for md in metadata:
|
||||
if isinstance(md, metadata_type):
|
||||
return md
|
||||
|
||||
return None
|
||||
@@ -0,0 +1,205 @@
|
||||
"""Code generated by Speakeasy (https://speakeasy.com). DO NOT EDIT."""
|
||||
|
||||
from typing import (
|
||||
Any,
|
||||
Dict,
|
||||
get_type_hints,
|
||||
List,
|
||||
Optional,
|
||||
)
|
||||
|
||||
from pydantic import BaseModel
|
||||
from pydantic.fields import FieldInfo
|
||||
|
||||
from .metadata import (
|
||||
QueryParamMetadata,
|
||||
find_field_metadata,
|
||||
)
|
||||
from .values import (
|
||||
_get_serialized_params,
|
||||
_is_set,
|
||||
_populate_from_globals,
|
||||
_val_to_string,
|
||||
)
|
||||
from .forms import _populate_form
|
||||
|
||||
|
||||
def get_query_params(
|
||||
query_params: Any,
|
||||
gbls: Optional[Any] = None,
|
||||
) -> Dict[str, List[str]]:
|
||||
params: Dict[str, List[str]] = {}
|
||||
|
||||
globals_already_populated = _populate_query_params(query_params, gbls, params, [])
|
||||
if _is_set(gbls):
|
||||
_populate_query_params(gbls, None, params, globals_already_populated)
|
||||
|
||||
return params
|
||||
|
||||
|
||||
def _populate_query_params(
|
||||
query_params: Any,
|
||||
gbls: Any,
|
||||
query_param_values: Dict[str, List[str]],
|
||||
skip_fields: List[str],
|
||||
) -> List[str]:
|
||||
globals_already_populated: List[str] = []
|
||||
|
||||
if not isinstance(query_params, BaseModel):
|
||||
return globals_already_populated
|
||||
|
||||
param_fields: Dict[str, FieldInfo] = query_params.__class__.model_fields
|
||||
param_field_types = get_type_hints(query_params.__class__)
|
||||
for name in param_fields:
|
||||
if name in skip_fields:
|
||||
continue
|
||||
|
||||
field = param_fields[name]
|
||||
|
||||
metadata = find_field_metadata(field, QueryParamMetadata)
|
||||
if not metadata:
|
||||
continue
|
||||
|
||||
value = getattr(query_params, name) if _is_set(query_params) else None
|
||||
|
||||
value, global_found = _populate_from_globals(
|
||||
name, value, QueryParamMetadata, gbls
|
||||
)
|
||||
if global_found:
|
||||
globals_already_populated.append(name)
|
||||
|
||||
f_name = field.alias if field.alias is not None else name
|
||||
serialization = metadata.serialization
|
||||
if serialization is not None:
|
||||
serialized_parms = _get_serialized_params(
|
||||
metadata, f_name, value, param_field_types[name]
|
||||
)
|
||||
for key, value in serialized_parms.items():
|
||||
if key in query_param_values:
|
||||
query_param_values[key].extend(value)
|
||||
else:
|
||||
query_param_values[key] = [value]
|
||||
else:
|
||||
style = metadata.style
|
||||
if style == "deepObject":
|
||||
_populate_deep_object_query_params(f_name, value, query_param_values)
|
||||
elif style == "form":
|
||||
_populate_delimited_query_params(
|
||||
metadata, f_name, value, ",", query_param_values
|
||||
)
|
||||
elif style == "pipeDelimited":
|
||||
_populate_delimited_query_params(
|
||||
metadata, f_name, value, "|", query_param_values
|
||||
)
|
||||
else:
|
||||
raise NotImplementedError(
|
||||
f"query param style {style} not yet supported"
|
||||
)
|
||||
|
||||
return globals_already_populated
|
||||
|
||||
|
||||
def _populate_deep_object_query_params(
|
||||
field_name: str,
|
||||
obj: Any,
|
||||
params: Dict[str, List[str]],
|
||||
):
|
||||
if not _is_set(obj):
|
||||
return
|
||||
|
||||
if isinstance(obj, BaseModel):
|
||||
_populate_deep_object_query_params_basemodel(field_name, obj, params)
|
||||
elif isinstance(obj, Dict):
|
||||
_populate_deep_object_query_params_dict(field_name, obj, params)
|
||||
|
||||
|
||||
def _populate_deep_object_query_params_basemodel(
|
||||
prior_params_key: str,
|
||||
obj: Any,
|
||||
params: Dict[str, List[str]],
|
||||
):
|
||||
if not _is_set(obj) or not isinstance(obj, BaseModel):
|
||||
return
|
||||
|
||||
obj_fields: Dict[str, FieldInfo] = obj.__class__.model_fields
|
||||
for name in obj_fields:
|
||||
obj_field = obj_fields[name]
|
||||
|
||||
f_name = obj_field.alias if obj_field.alias is not None else name
|
||||
|
||||
params_key = f"{prior_params_key}[{f_name}]"
|
||||
|
||||
obj_param_metadata = find_field_metadata(obj_field, QueryParamMetadata)
|
||||
if not _is_set(obj_param_metadata):
|
||||
continue
|
||||
|
||||
obj_val = getattr(obj, name)
|
||||
if not _is_set(obj_val):
|
||||
continue
|
||||
|
||||
if isinstance(obj_val, BaseModel):
|
||||
_populate_deep_object_query_params_basemodel(params_key, obj_val, params)
|
||||
elif isinstance(obj_val, Dict):
|
||||
_populate_deep_object_query_params_dict(params_key, obj_val, params)
|
||||
elif isinstance(obj_val, List):
|
||||
_populate_deep_object_query_params_list(params_key, obj_val, params)
|
||||
else:
|
||||
params[params_key] = [_val_to_string(obj_val)]
|
||||
|
||||
|
||||
def _populate_deep_object_query_params_dict(
|
||||
prior_params_key: str,
|
||||
value: Dict,
|
||||
params: Dict[str, List[str]],
|
||||
):
|
||||
if not _is_set(value):
|
||||
return
|
||||
|
||||
for key, val in value.items():
|
||||
if not _is_set(val):
|
||||
continue
|
||||
|
||||
params_key = f"{prior_params_key}[{key}]"
|
||||
|
||||
if isinstance(val, BaseModel):
|
||||
_populate_deep_object_query_params_basemodel(params_key, val, params)
|
||||
elif isinstance(val, Dict):
|
||||
_populate_deep_object_query_params_dict(params_key, val, params)
|
||||
elif isinstance(val, List):
|
||||
_populate_deep_object_query_params_list(params_key, val, params)
|
||||
else:
|
||||
params[params_key] = [_val_to_string(val)]
|
||||
|
||||
|
||||
def _populate_deep_object_query_params_list(
|
||||
params_key: str,
|
||||
value: List,
|
||||
params: Dict[str, List[str]],
|
||||
):
|
||||
if not _is_set(value):
|
||||
return
|
||||
|
||||
for val in value:
|
||||
if not _is_set(val):
|
||||
continue
|
||||
|
||||
if params.get(params_key) is None:
|
||||
params[params_key] = []
|
||||
|
||||
params[params_key].append(_val_to_string(val))
|
||||
|
||||
|
||||
def _populate_delimited_query_params(
|
||||
metadata: QueryParamMetadata,
|
||||
field_name: str,
|
||||
obj: Any,
|
||||
delimiter: str,
|
||||
query_param_values: Dict[str, List[str]],
|
||||
):
|
||||
_populate_form(
|
||||
field_name,
|
||||
metadata.explode,
|
||||
obj,
|
||||
delimiter,
|
||||
query_param_values,
|
||||
)
|
||||
@@ -0,0 +1,66 @@
|
||||
"""Code generated by Speakeasy (https://speakeasy.com). DO NOT EDIT."""
|
||||
|
||||
import io
|
||||
from dataclasses import dataclass
|
||||
import re
|
||||
from typing import (
|
||||
Any,
|
||||
Optional,
|
||||
)
|
||||
|
||||
from .forms import serialize_form_data, serialize_multipart_form
|
||||
|
||||
from .serializers import marshal_json
|
||||
|
||||
SERIALIZATION_METHOD_TO_CONTENT_TYPE = {
|
||||
"json": "application/json",
|
||||
"form": "application/x-www-form-urlencoded",
|
||||
"multipart": "multipart/form-data",
|
||||
"raw": "application/octet-stream",
|
||||
"string": "text/plain",
|
||||
}
|
||||
|
||||
|
||||
@dataclass
|
||||
class SerializedRequestBody:
|
||||
media_type: Optional[str] = None
|
||||
content: Optional[Any] = None
|
||||
data: Optional[Any] = None
|
||||
files: Optional[Any] = None
|
||||
|
||||
|
||||
def serialize_request_body(
|
||||
request_body: Any,
|
||||
nullable: bool,
|
||||
optional: bool,
|
||||
serialization_method: str,
|
||||
request_body_type,
|
||||
) -> Optional[SerializedRequestBody]:
|
||||
if request_body is None:
|
||||
if not nullable and optional:
|
||||
return None
|
||||
|
||||
media_type = SERIALIZATION_METHOD_TO_CONTENT_TYPE[serialization_method]
|
||||
|
||||
serialized_request_body = SerializedRequestBody(media_type)
|
||||
|
||||
if re.match(r"(application|text)\/.*?\+*json.*", media_type) is not None:
|
||||
serialized_request_body.content = marshal_json(request_body, request_body_type)
|
||||
elif re.match(r"multipart\/.*", media_type) is not None:
|
||||
(
|
||||
serialized_request_body.media_type,
|
||||
serialized_request_body.data,
|
||||
serialized_request_body.files,
|
||||
) = serialize_multipart_form(media_type, request_body)
|
||||
elif re.match(r"application\/x-www-form-urlencoded.*", media_type) is not None:
|
||||
serialized_request_body.data = serialize_form_data(request_body)
|
||||
elif isinstance(request_body, (bytes, bytearray, io.BytesIO, io.BufferedReader)):
|
||||
serialized_request_body.content = request_body
|
||||
elif isinstance(request_body, str):
|
||||
serialized_request_body.content = request_body
|
||||
else:
|
||||
raise TypeError(
|
||||
f"invalid request body type {type(request_body)} for mediaType {media_type}"
|
||||
)
|
||||
|
||||
return serialized_request_body
|
||||
@@ -0,0 +1,217 @@
|
||||
"""Code generated by Speakeasy (https://speakeasy.com). DO NOT EDIT."""
|
||||
|
||||
import asyncio
|
||||
import random
|
||||
import time
|
||||
from typing import List
|
||||
|
||||
import httpx
|
||||
|
||||
|
||||
class BackoffStrategy:
|
||||
initial_interval: int
|
||||
max_interval: int
|
||||
exponent: float
|
||||
max_elapsed_time: int
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
initial_interval: int,
|
||||
max_interval: int,
|
||||
exponent: float,
|
||||
max_elapsed_time: int,
|
||||
):
|
||||
self.initial_interval = initial_interval
|
||||
self.max_interval = max_interval
|
||||
self.exponent = exponent
|
||||
self.max_elapsed_time = max_elapsed_time
|
||||
|
||||
|
||||
class RetryConfig:
|
||||
strategy: str
|
||||
backoff: BackoffStrategy
|
||||
retry_connection_errors: bool
|
||||
|
||||
def __init__(
|
||||
self, strategy: str, backoff: BackoffStrategy, retry_connection_errors: bool
|
||||
):
|
||||
self.strategy = strategy
|
||||
self.backoff = backoff
|
||||
self.retry_connection_errors = retry_connection_errors
|
||||
|
||||
|
||||
class Retries:
|
||||
config: RetryConfig
|
||||
status_codes: List[str]
|
||||
|
||||
def __init__(self, config: RetryConfig, status_codes: List[str]):
|
||||
self.config = config
|
||||
self.status_codes = status_codes
|
||||
|
||||
|
||||
class TemporaryError(Exception):
|
||||
response: httpx.Response
|
||||
|
||||
def __init__(self, response: httpx.Response):
|
||||
self.response = response
|
||||
|
||||
|
||||
class PermanentError(Exception):
|
||||
inner: Exception
|
||||
|
||||
def __init__(self, inner: Exception):
|
||||
self.inner = inner
|
||||
|
||||
|
||||
def retry(func, retries: Retries):
|
||||
if retries.config.strategy == "backoff":
|
||||
|
||||
def do_request() -> httpx.Response:
|
||||
res: httpx.Response
|
||||
try:
|
||||
res = func()
|
||||
|
||||
for code in retries.status_codes:
|
||||
if "X" in code.upper():
|
||||
code_range = int(code[0])
|
||||
|
||||
status_major = res.status_code / 100
|
||||
|
||||
if code_range <= status_major < code_range + 1:
|
||||
raise TemporaryError(res)
|
||||
else:
|
||||
parsed_code = int(code)
|
||||
|
||||
if res.status_code == parsed_code:
|
||||
raise TemporaryError(res)
|
||||
except httpx.ConnectError as exception:
|
||||
if retries.config.retry_connection_errors:
|
||||
raise
|
||||
|
||||
raise PermanentError(exception) from exception
|
||||
except httpx.TimeoutException as exception:
|
||||
if retries.config.retry_connection_errors:
|
||||
raise
|
||||
|
||||
raise PermanentError(exception) from exception
|
||||
except TemporaryError:
|
||||
raise
|
||||
except Exception as exception:
|
||||
raise PermanentError(exception) from exception
|
||||
|
||||
return res
|
||||
|
||||
return retry_with_backoff(
|
||||
do_request,
|
||||
retries.config.backoff.initial_interval,
|
||||
retries.config.backoff.max_interval,
|
||||
retries.config.backoff.exponent,
|
||||
retries.config.backoff.max_elapsed_time,
|
||||
)
|
||||
|
||||
return func()
|
||||
|
||||
|
||||
async def retry_async(func, retries: Retries):
|
||||
if retries.config.strategy == "backoff":
|
||||
|
||||
async def do_request() -> httpx.Response:
|
||||
res: httpx.Response
|
||||
try:
|
||||
res = await func()
|
||||
|
||||
for code in retries.status_codes:
|
||||
if "X" in code.upper():
|
||||
code_range = int(code[0])
|
||||
|
||||
status_major = res.status_code / 100
|
||||
|
||||
if code_range <= status_major < code_range + 1:
|
||||
raise TemporaryError(res)
|
||||
else:
|
||||
parsed_code = int(code)
|
||||
|
||||
if res.status_code == parsed_code:
|
||||
raise TemporaryError(res)
|
||||
except httpx.ConnectError as exception:
|
||||
if retries.config.retry_connection_errors:
|
||||
raise
|
||||
|
||||
raise PermanentError(exception) from exception
|
||||
except httpx.TimeoutException as exception:
|
||||
if retries.config.retry_connection_errors:
|
||||
raise
|
||||
|
||||
raise PermanentError(exception) from exception
|
||||
except TemporaryError:
|
||||
raise
|
||||
except Exception as exception:
|
||||
raise PermanentError(exception) from exception
|
||||
|
||||
return res
|
||||
|
||||
return await retry_with_backoff_async(
|
||||
do_request,
|
||||
retries.config.backoff.initial_interval,
|
||||
retries.config.backoff.max_interval,
|
||||
retries.config.backoff.exponent,
|
||||
retries.config.backoff.max_elapsed_time,
|
||||
)
|
||||
|
||||
return await func()
|
||||
|
||||
|
||||
def retry_with_backoff(
|
||||
func,
|
||||
initial_interval=500,
|
||||
max_interval=60000,
|
||||
exponent=1.5,
|
||||
max_elapsed_time=3600000,
|
||||
):
|
||||
start = round(time.time() * 1000)
|
||||
retries = 0
|
||||
|
||||
while True:
|
||||
try:
|
||||
return func()
|
||||
except PermanentError as exception:
|
||||
raise exception.inner
|
||||
except Exception as exception: # pylint: disable=broad-exception-caught
|
||||
now = round(time.time() * 1000)
|
||||
if now - start > max_elapsed_time:
|
||||
if isinstance(exception, TemporaryError):
|
||||
return exception.response
|
||||
|
||||
raise
|
||||
sleep = (initial_interval / 1000) * exponent**retries + random.uniform(0, 1)
|
||||
sleep = min(sleep, max_interval / 1000)
|
||||
time.sleep(sleep)
|
||||
retries += 1
|
||||
|
||||
|
||||
async def retry_with_backoff_async(
|
||||
func,
|
||||
initial_interval=500,
|
||||
max_interval=60000,
|
||||
exponent=1.5,
|
||||
max_elapsed_time=3600000,
|
||||
):
|
||||
start = round(time.time() * 1000)
|
||||
retries = 0
|
||||
|
||||
while True:
|
||||
try:
|
||||
return await func()
|
||||
except PermanentError as exception:
|
||||
raise exception.inner
|
||||
except Exception as exception: # pylint: disable=broad-exception-caught
|
||||
now = round(time.time() * 1000)
|
||||
if now - start > max_elapsed_time:
|
||||
if isinstance(exception, TemporaryError):
|
||||
return exception.response
|
||||
|
||||
raise
|
||||
sleep = (initial_interval / 1000) * exponent**retries + random.uniform(0, 1)
|
||||
sleep = min(sleep, max_interval / 1000)
|
||||
await asyncio.sleep(sleep)
|
||||
retries += 1
|
||||
@@ -0,0 +1,192 @@
|
||||
"""Code generated by Speakeasy (https://speakeasy.com). DO NOT EDIT."""
|
||||
|
||||
import base64
|
||||
|
||||
from typing import (
|
||||
Any,
|
||||
Dict,
|
||||
List,
|
||||
Optional,
|
||||
Tuple,
|
||||
)
|
||||
from pydantic import BaseModel
|
||||
from pydantic.fields import FieldInfo
|
||||
|
||||
from .metadata import (
|
||||
SecurityMetadata,
|
||||
find_field_metadata,
|
||||
)
|
||||
import os
|
||||
|
||||
|
||||
def get_security(security: Any) -> Tuple[Dict[str, str], Dict[str, List[str]]]:
|
||||
headers: Dict[str, str] = {}
|
||||
query_params: Dict[str, List[str]] = {}
|
||||
|
||||
if security is None:
|
||||
return headers, query_params
|
||||
|
||||
if not isinstance(security, BaseModel):
|
||||
raise TypeError("security must be a pydantic model")
|
||||
|
||||
sec_fields: Dict[str, FieldInfo] = security.__class__.model_fields
|
||||
for name in sec_fields:
|
||||
sec_field = sec_fields[name]
|
||||
|
||||
value = getattr(security, name)
|
||||
if value is None:
|
||||
continue
|
||||
|
||||
metadata = find_field_metadata(sec_field, SecurityMetadata)
|
||||
if metadata is None:
|
||||
continue
|
||||
if metadata.option:
|
||||
_parse_security_option(headers, query_params, value)
|
||||
return headers, query_params
|
||||
if metadata.scheme:
|
||||
# Special case for basic auth or custom auth which could be a flattened model
|
||||
if metadata.sub_type in ["basic", "custom"] and not isinstance(
|
||||
value, BaseModel
|
||||
):
|
||||
_parse_security_scheme(headers, query_params, metadata, name, security)
|
||||
else:
|
||||
_parse_security_scheme(headers, query_params, metadata, name, value)
|
||||
|
||||
return headers, query_params
|
||||
|
||||
|
||||
def get_security_from_env(security: Any, security_class: Any) -> Optional[BaseModel]:
|
||||
if security is not None:
|
||||
return security
|
||||
|
||||
if not issubclass(security_class, BaseModel):
|
||||
raise TypeError("security_class must be a pydantic model class")
|
||||
|
||||
security_dict: Any = {}
|
||||
|
||||
if os.getenv("OPENROUTER_API_KEY"):
|
||||
security_dict["api_key"] = os.getenv("OPENROUTER_API_KEY")
|
||||
|
||||
return security_class(**security_dict) if security_dict else None
|
||||
|
||||
|
||||
def _parse_security_option(
|
||||
headers: Dict[str, str], query_params: Dict[str, List[str]], option: Any
|
||||
):
|
||||
if not isinstance(option, BaseModel):
|
||||
raise TypeError("security option must be a pydantic model")
|
||||
|
||||
opt_fields: Dict[str, FieldInfo] = option.__class__.model_fields
|
||||
for name in opt_fields:
|
||||
opt_field = opt_fields[name]
|
||||
|
||||
metadata = find_field_metadata(opt_field, SecurityMetadata)
|
||||
if metadata is None or not metadata.scheme:
|
||||
continue
|
||||
_parse_security_scheme(
|
||||
headers, query_params, metadata, name, getattr(option, name)
|
||||
)
|
||||
|
||||
|
||||
def _parse_security_scheme(
|
||||
headers: Dict[str, str],
|
||||
query_params: Dict[str, List[str]],
|
||||
scheme_metadata: SecurityMetadata,
|
||||
field_name: str,
|
||||
scheme: Any,
|
||||
):
|
||||
scheme_type = scheme_metadata.scheme_type
|
||||
sub_type = scheme_metadata.sub_type
|
||||
|
||||
if isinstance(scheme, BaseModel):
|
||||
if scheme_type == "http":
|
||||
if sub_type == "basic":
|
||||
_parse_basic_auth_scheme(headers, scheme)
|
||||
return
|
||||
if sub_type == "custom":
|
||||
return
|
||||
|
||||
scheme_fields: Dict[str, FieldInfo] = scheme.__class__.model_fields
|
||||
for name in scheme_fields:
|
||||
scheme_field = scheme_fields[name]
|
||||
|
||||
metadata = find_field_metadata(scheme_field, SecurityMetadata)
|
||||
if metadata is None or metadata.field_name is None:
|
||||
continue
|
||||
|
||||
value = getattr(scheme, name)
|
||||
|
||||
_parse_security_scheme_value(
|
||||
headers, query_params, scheme_metadata, metadata, name, value
|
||||
)
|
||||
else:
|
||||
_parse_security_scheme_value(
|
||||
headers, query_params, scheme_metadata, scheme_metadata, field_name, scheme
|
||||
)
|
||||
|
||||
|
||||
def _parse_security_scheme_value(
|
||||
headers: Dict[str, str],
|
||||
query_params: Dict[str, List[str]],
|
||||
scheme_metadata: SecurityMetadata,
|
||||
security_metadata: SecurityMetadata,
|
||||
field_name: str,
|
||||
value: Any,
|
||||
):
|
||||
scheme_type = scheme_metadata.scheme_type
|
||||
sub_type = scheme_metadata.sub_type
|
||||
|
||||
header_name = security_metadata.get_field_name(field_name)
|
||||
|
||||
if scheme_type == "apiKey":
|
||||
if sub_type == "header":
|
||||
headers[header_name] = value
|
||||
elif sub_type == "query":
|
||||
query_params[header_name] = [value]
|
||||
else:
|
||||
raise ValueError("sub type {sub_type} not supported")
|
||||
elif scheme_type == "openIdConnect":
|
||||
headers[header_name] = _apply_bearer(value)
|
||||
elif scheme_type == "oauth2":
|
||||
if sub_type != "client_credentials":
|
||||
headers[header_name] = _apply_bearer(value)
|
||||
elif scheme_type == "http":
|
||||
if sub_type == "bearer":
|
||||
headers[header_name] = _apply_bearer(value)
|
||||
elif sub_type == "custom":
|
||||
return
|
||||
else:
|
||||
raise ValueError("sub type {sub_type} not supported")
|
||||
else:
|
||||
raise ValueError("scheme type {scheme_type} not supported")
|
||||
|
||||
|
||||
def _apply_bearer(token: str) -> str:
|
||||
return token.lower().startswith("bearer ") and token or f"Bearer {token}"
|
||||
|
||||
|
||||
def _parse_basic_auth_scheme(headers: Dict[str, str], scheme: Any):
|
||||
username = ""
|
||||
password = ""
|
||||
|
||||
if not isinstance(scheme, BaseModel):
|
||||
raise TypeError("basic auth scheme must be a pydantic model")
|
||||
|
||||
scheme_fields: Dict[str, FieldInfo] = scheme.__class__.model_fields
|
||||
for name in scheme_fields:
|
||||
scheme_field = scheme_fields[name]
|
||||
|
||||
metadata = find_field_metadata(scheme_field, SecurityMetadata)
|
||||
if metadata is None or metadata.field_name is None:
|
||||
continue
|
||||
|
||||
field_name = metadata.field_name
|
||||
value = getattr(scheme, name)
|
||||
|
||||
if field_name == "username":
|
||||
username = value
|
||||
if field_name == "password":
|
||||
password = value
|
||||
|
||||
data = f"{username}:{password}".encode()
|
||||
headers["Authorization"] = f"Basic {base64.b64encode(data).decode()}"
|
||||
@@ -0,0 +1,249 @@
|
||||
"""Code generated by Speakeasy (https://speakeasy.com). DO NOT EDIT."""
|
||||
|
||||
from decimal import Decimal
|
||||
import functools
|
||||
import json
|
||||
import typing
|
||||
from typing import Any, Dict, List, Tuple, Union, get_args
|
||||
import typing_extensions
|
||||
from typing_extensions import get_origin
|
||||
|
||||
import httpx
|
||||
from pydantic import ConfigDict, create_model
|
||||
from pydantic_core import from_json
|
||||
|
||||
from ..types.basemodel import BaseModel, Nullable, OptionalNullable, Unset
|
||||
|
||||
|
||||
def serialize_decimal(as_str: bool):
|
||||
def serialize(d):
|
||||
# Optional[T] is a Union[T, None]
|
||||
if is_union(type(d)) and type(None) in get_args(type(d)) and d is None:
|
||||
return None
|
||||
if isinstance(d, Unset):
|
||||
return d
|
||||
|
||||
if not isinstance(d, Decimal):
|
||||
raise ValueError("Expected Decimal object")
|
||||
|
||||
return str(d) if as_str else float(d)
|
||||
|
||||
return serialize
|
||||
|
||||
|
||||
def validate_decimal(d):
|
||||
if d is None:
|
||||
return None
|
||||
|
||||
if isinstance(d, (Decimal, Unset)):
|
||||
return d
|
||||
|
||||
if not isinstance(d, (str, int, float)):
|
||||
raise ValueError("Expected string, int or float")
|
||||
|
||||
return Decimal(str(d))
|
||||
|
||||
|
||||
def serialize_float(as_str: bool):
|
||||
def serialize(f):
|
||||
# Optional[T] is a Union[T, None]
|
||||
if is_union(type(f)) and type(None) in get_args(type(f)) and f is None:
|
||||
return None
|
||||
if isinstance(f, Unset):
|
||||
return f
|
||||
|
||||
if not isinstance(f, float):
|
||||
raise ValueError("Expected float")
|
||||
|
||||
return str(f) if as_str else f
|
||||
|
||||
return serialize
|
||||
|
||||
|
||||
def validate_float(f):
|
||||
if f is None:
|
||||
return None
|
||||
|
||||
if isinstance(f, (float, Unset)):
|
||||
return f
|
||||
|
||||
if not isinstance(f, str):
|
||||
raise ValueError("Expected string")
|
||||
|
||||
return float(f)
|
||||
|
||||
|
||||
def serialize_int(as_str: bool):
|
||||
def serialize(i):
|
||||
# Optional[T] is a Union[T, None]
|
||||
if is_union(type(i)) and type(None) in get_args(type(i)) and i is None:
|
||||
return None
|
||||
if isinstance(i, Unset):
|
||||
return i
|
||||
|
||||
if not isinstance(i, int):
|
||||
raise ValueError("Expected int")
|
||||
|
||||
return str(i) if as_str else i
|
||||
|
||||
return serialize
|
||||
|
||||
|
||||
def validate_int(b):
|
||||
if b is None:
|
||||
return None
|
||||
|
||||
if isinstance(b, (int, Unset)):
|
||||
return b
|
||||
|
||||
if not isinstance(b, str):
|
||||
raise ValueError("Expected string")
|
||||
|
||||
return int(b)
|
||||
|
||||
|
||||
def validate_open_enum(is_int: bool):
|
||||
def validate(e):
|
||||
if e is None:
|
||||
return None
|
||||
|
||||
if isinstance(e, Unset):
|
||||
return e
|
||||
|
||||
if is_int:
|
||||
if not isinstance(e, int):
|
||||
raise ValueError("Expected int")
|
||||
else:
|
||||
if not isinstance(e, str):
|
||||
raise ValueError("Expected string")
|
||||
|
||||
return e
|
||||
|
||||
return validate
|
||||
|
||||
|
||||
def validate_const(v):
|
||||
def validate(c):
|
||||
# Optional[T] is a Union[T, None]
|
||||
if is_union(type(c)) and type(None) in get_args(type(c)) and c is None:
|
||||
return None
|
||||
|
||||
if v != c:
|
||||
raise ValueError(f"Expected {v}")
|
||||
|
||||
return c
|
||||
|
||||
return validate
|
||||
|
||||
|
||||
def unmarshal_json(raw, typ: Any) -> Any:
|
||||
return unmarshal(from_json(raw), typ)
|
||||
|
||||
|
||||
def unmarshal(val, typ: Any) -> Any:
|
||||
unmarshaller = create_model(
|
||||
"Unmarshaller",
|
||||
body=(typ, ...),
|
||||
__config__=ConfigDict(populate_by_name=True, arbitrary_types_allowed=True),
|
||||
)
|
||||
|
||||
m = unmarshaller(body=val)
|
||||
|
||||
# pyright: ignore[reportAttributeAccessIssue]
|
||||
return m.body # type: ignore
|
||||
|
||||
|
||||
def marshal_json(val, typ):
|
||||
if is_nullable(typ) and val is None:
|
||||
return "null"
|
||||
|
||||
marshaller = create_model(
|
||||
"Marshaller",
|
||||
body=(typ, ...),
|
||||
__config__=ConfigDict(populate_by_name=True, arbitrary_types_allowed=True),
|
||||
)
|
||||
|
||||
m = marshaller(body=val)
|
||||
|
||||
d = m.model_dump(by_alias=True, mode="json", exclude_none=True)
|
||||
|
||||
if len(d) == 0:
|
||||
return ""
|
||||
|
||||
return json.dumps(d[next(iter(d))], separators=(",", ":"))
|
||||
|
||||
|
||||
def is_nullable(field):
|
||||
origin = get_origin(field)
|
||||
if origin is Nullable or origin is OptionalNullable:
|
||||
return True
|
||||
|
||||
if not origin is Union or type(None) not in get_args(field):
|
||||
return False
|
||||
|
||||
for arg in get_args(field):
|
||||
if get_origin(arg) is Nullable or get_origin(arg) is OptionalNullable:
|
||||
return True
|
||||
|
||||
return False
|
||||
|
||||
|
||||
def is_union(obj: object) -> bool:
|
||||
"""
|
||||
Returns True if the given object is a typing.Union or typing_extensions.Union.
|
||||
"""
|
||||
return any(
|
||||
obj is typing_obj for typing_obj in _get_typing_objects_by_name_of("Union")
|
||||
)
|
||||
|
||||
|
||||
def stream_to_text(stream: httpx.Response) -> str:
|
||||
return "".join(stream.iter_text())
|
||||
|
||||
|
||||
async def stream_to_text_async(stream: httpx.Response) -> str:
|
||||
return "".join([chunk async for chunk in stream.aiter_text()])
|
||||
|
||||
|
||||
def stream_to_bytes(stream: httpx.Response) -> bytes:
|
||||
return stream.content
|
||||
|
||||
|
||||
async def stream_to_bytes_async(stream: httpx.Response) -> bytes:
|
||||
return await stream.aread()
|
||||
|
||||
|
||||
def get_pydantic_model(data: Any, typ: Any) -> Any:
|
||||
if not _contains_pydantic_model(data):
|
||||
return unmarshal(data, typ)
|
||||
|
||||
return data
|
||||
|
||||
|
||||
def _contains_pydantic_model(data: Any) -> bool:
|
||||
if isinstance(data, BaseModel):
|
||||
return True
|
||||
if isinstance(data, List):
|
||||
return any(_contains_pydantic_model(item) for item in data)
|
||||
if isinstance(data, Dict):
|
||||
return any(_contains_pydantic_model(value) for value in data.values())
|
||||
|
||||
return False
|
||||
|
||||
|
||||
@functools.cache
|
||||
def _get_typing_objects_by_name_of(name: str) -> Tuple[Any, ...]:
|
||||
"""
|
||||
Get typing objects by name from typing and typing_extensions.
|
||||
Reference: https://typing-extensions.readthedocs.io/en/latest/#runtime-use-of-types
|
||||
"""
|
||||
result = tuple(
|
||||
getattr(module, name)
|
||||
for module in (typing, typing_extensions)
|
||||
if hasattr(module, name)
|
||||
)
|
||||
if not result:
|
||||
raise ValueError(
|
||||
f"Neither typing nor typing_extensions has an object called {name!r}"
|
||||
)
|
||||
return result
|
||||
@@ -0,0 +1,24 @@
|
||||
"""Code generated by Speakeasy (https://speakeasy.com). DO NOT EDIT."""
|
||||
|
||||
from typing import Any, Optional
|
||||
|
||||
import httpx
|
||||
|
||||
from .serializers import unmarshal_json
|
||||
from openrouter import errors
|
||||
|
||||
|
||||
def unmarshal_json_response(
|
||||
typ: Any, http_res: httpx.Response, body: Optional[str] = None
|
||||
) -> Any:
|
||||
if body is None:
|
||||
body = http_res.text
|
||||
try:
|
||||
return unmarshal_json(body, typ)
|
||||
except Exception as e:
|
||||
raise errors.ResponseValidationError(
|
||||
"Response validation failed",
|
||||
http_res,
|
||||
e,
|
||||
body,
|
||||
) from e
|
||||
@@ -0,0 +1,155 @@
|
||||
"""Code generated by Speakeasy (https://speakeasy.com). DO NOT EDIT."""
|
||||
|
||||
from decimal import Decimal
|
||||
from typing import (
|
||||
Any,
|
||||
Dict,
|
||||
get_type_hints,
|
||||
List,
|
||||
Optional,
|
||||
Union,
|
||||
get_args,
|
||||
get_origin,
|
||||
)
|
||||
from pydantic import BaseModel
|
||||
from pydantic.fields import FieldInfo
|
||||
|
||||
from .metadata import (
|
||||
PathParamMetadata,
|
||||
find_field_metadata,
|
||||
)
|
||||
from .values import (
|
||||
_get_serialized_params,
|
||||
_is_set,
|
||||
_populate_from_globals,
|
||||
_val_to_string,
|
||||
)
|
||||
|
||||
|
||||
def generate_url(
|
||||
server_url: str,
|
||||
path: str,
|
||||
path_params: Any,
|
||||
gbls: Optional[Any] = None,
|
||||
) -> str:
|
||||
path_param_values: Dict[str, str] = {}
|
||||
|
||||
globals_already_populated = _populate_path_params(
|
||||
path_params, gbls, path_param_values, []
|
||||
)
|
||||
if _is_set(gbls):
|
||||
_populate_path_params(gbls, None, path_param_values, globals_already_populated)
|
||||
|
||||
for key, value in path_param_values.items():
|
||||
path = path.replace("{" + key + "}", value, 1)
|
||||
|
||||
return remove_suffix(server_url, "/") + path
|
||||
|
||||
|
||||
def _populate_path_params(
|
||||
path_params: Any,
|
||||
gbls: Any,
|
||||
path_param_values: Dict[str, str],
|
||||
skip_fields: List[str],
|
||||
) -> List[str]:
|
||||
globals_already_populated: List[str] = []
|
||||
|
||||
if not isinstance(path_params, BaseModel):
|
||||
return globals_already_populated
|
||||
|
||||
path_param_fields: Dict[str, FieldInfo] = path_params.__class__.model_fields
|
||||
path_param_field_types = get_type_hints(path_params.__class__)
|
||||
for name in path_param_fields:
|
||||
if name in skip_fields:
|
||||
continue
|
||||
|
||||
field = path_param_fields[name]
|
||||
|
||||
param_metadata = find_field_metadata(field, PathParamMetadata)
|
||||
if param_metadata is None:
|
||||
continue
|
||||
|
||||
param = getattr(path_params, name) if _is_set(path_params) else None
|
||||
param, global_found = _populate_from_globals(
|
||||
name, param, PathParamMetadata, gbls
|
||||
)
|
||||
if global_found:
|
||||
globals_already_populated.append(name)
|
||||
|
||||
if not _is_set(param):
|
||||
continue
|
||||
|
||||
f_name = field.alias if field.alias is not None else name
|
||||
serialization = param_metadata.serialization
|
||||
if serialization is not None:
|
||||
serialized_params = _get_serialized_params(
|
||||
param_metadata, f_name, param, path_param_field_types[name]
|
||||
)
|
||||
for key, value in serialized_params.items():
|
||||
path_param_values[key] = value
|
||||
else:
|
||||
pp_vals: List[str] = []
|
||||
if param_metadata.style == "simple":
|
||||
if isinstance(param, List):
|
||||
for pp_val in param:
|
||||
if not _is_set(pp_val):
|
||||
continue
|
||||
pp_vals.append(_val_to_string(pp_val))
|
||||
path_param_values[f_name] = ",".join(pp_vals)
|
||||
elif isinstance(param, Dict):
|
||||
for pp_key in param:
|
||||
if not _is_set(param[pp_key]):
|
||||
continue
|
||||
if param_metadata.explode:
|
||||
pp_vals.append(f"{pp_key}={_val_to_string(param[pp_key])}")
|
||||
else:
|
||||
pp_vals.append(f"{pp_key},{_val_to_string(param[pp_key])}")
|
||||
path_param_values[f_name] = ",".join(pp_vals)
|
||||
elif not isinstance(param, (str, int, float, complex, bool, Decimal)):
|
||||
param_fields: Dict[str, FieldInfo] = param.__class__.model_fields
|
||||
for name in param_fields:
|
||||
param_field = param_fields[name]
|
||||
|
||||
param_value_metadata = find_field_metadata(
|
||||
param_field, PathParamMetadata
|
||||
)
|
||||
if param_value_metadata is None:
|
||||
continue
|
||||
|
||||
param_name = (
|
||||
param_field.alias if param_field.alias is not None else name
|
||||
)
|
||||
|
||||
param_field_val = getattr(param, name)
|
||||
if not _is_set(param_field_val):
|
||||
continue
|
||||
if param_metadata.explode:
|
||||
pp_vals.append(
|
||||
f"{param_name}={_val_to_string(param_field_val)}"
|
||||
)
|
||||
else:
|
||||
pp_vals.append(
|
||||
f"{param_name},{_val_to_string(param_field_val)}"
|
||||
)
|
||||
path_param_values[f_name] = ",".join(pp_vals)
|
||||
elif _is_set(param):
|
||||
path_param_values[f_name] = _val_to_string(param)
|
||||
|
||||
return globals_already_populated
|
||||
|
||||
|
||||
def is_optional(field):
|
||||
return get_origin(field) is Union and type(None) in get_args(field)
|
||||
|
||||
|
||||
def template_url(url_with_params: str, params: Dict[str, str]) -> str:
|
||||
for key, value in params.items():
|
||||
url_with_params = url_with_params.replace("{" + key + "}", value)
|
||||
|
||||
return url_with_params
|
||||
|
||||
|
||||
def remove_suffix(input_string, suffix):
|
||||
if suffix and input_string.endswith(suffix):
|
||||
return input_string[: -len(suffix)]
|
||||
return input_string
|
||||
@@ -0,0 +1,137 @@
|
||||
"""Code generated by Speakeasy (https://speakeasy.com). DO NOT EDIT."""
|
||||
|
||||
from datetime import datetime
|
||||
from enum import Enum
|
||||
from email.message import Message
|
||||
from functools import partial
|
||||
import os
|
||||
from typing import Any, Callable, Dict, List, Optional, Tuple, TypeVar, Union, cast
|
||||
|
||||
from httpx import Response
|
||||
from pydantic import BaseModel
|
||||
from pydantic.fields import FieldInfo
|
||||
|
||||
from ..types.basemodel import Unset
|
||||
|
||||
from .serializers import marshal_json
|
||||
|
||||
from .metadata import ParamMetadata, find_field_metadata
|
||||
|
||||
|
||||
def match_content_type(content_type: str, pattern: str) -> bool:
|
||||
if pattern in (content_type, "*", "*/*"):
|
||||
return True
|
||||
|
||||
msg = Message()
|
||||
msg["content-type"] = content_type
|
||||
media_type = msg.get_content_type()
|
||||
|
||||
if media_type == pattern:
|
||||
return True
|
||||
|
||||
parts = media_type.split("/")
|
||||
if len(parts) == 2:
|
||||
if pattern in (f"{parts[0]}/*", f"*/{parts[1]}"):
|
||||
return True
|
||||
|
||||
return False
|
||||
|
||||
|
||||
def match_status_codes(status_codes: List[str], status_code: int) -> bool:
|
||||
if "default" in status_codes:
|
||||
return True
|
||||
|
||||
for code in status_codes:
|
||||
if code == str(status_code):
|
||||
return True
|
||||
|
||||
if code.endswith("XX") and code.startswith(str(status_code)[:1]):
|
||||
return True
|
||||
return False
|
||||
|
||||
|
||||
T = TypeVar("T")
|
||||
|
||||
def cast_partial(typ):
|
||||
return partial(cast, typ)
|
||||
|
||||
def get_global_from_env(
|
||||
value: Optional[T], env_key: str, type_cast: Callable[[str], T]
|
||||
) -> Optional[T]:
|
||||
if value is not None:
|
||||
return value
|
||||
env_value = os.getenv(env_key)
|
||||
if env_value is not None:
|
||||
try:
|
||||
return type_cast(env_value)
|
||||
except ValueError:
|
||||
pass
|
||||
return None
|
||||
|
||||
|
||||
def match_response(
|
||||
response: Response, code: Union[str, List[str]], content_type: str
|
||||
) -> bool:
|
||||
codes = code if isinstance(code, list) else [code]
|
||||
return match_status_codes(codes, response.status_code) and match_content_type(
|
||||
response.headers.get("content-type", "application/octet-stream"), content_type
|
||||
)
|
||||
|
||||
|
||||
def _populate_from_globals(
|
||||
param_name: str, value: Any, param_metadata_type: type, gbls: Any
|
||||
) -> Tuple[Any, bool]:
|
||||
if gbls is None:
|
||||
return value, False
|
||||
|
||||
if not isinstance(gbls, BaseModel):
|
||||
raise TypeError("globals must be a pydantic model")
|
||||
|
||||
global_fields: Dict[str, FieldInfo] = gbls.__class__.model_fields
|
||||
found = False
|
||||
for name in global_fields:
|
||||
field = global_fields[name]
|
||||
if name is not param_name:
|
||||
continue
|
||||
|
||||
found = True
|
||||
|
||||
if value is not None:
|
||||
return value, True
|
||||
|
||||
global_value = getattr(gbls, name)
|
||||
|
||||
param_metadata = find_field_metadata(field, param_metadata_type)
|
||||
if param_metadata is None:
|
||||
return value, True
|
||||
|
||||
return global_value, True
|
||||
|
||||
return value, found
|
||||
|
||||
|
||||
def _val_to_string(val) -> str:
|
||||
if isinstance(val, bool):
|
||||
return str(val).lower()
|
||||
if isinstance(val, datetime):
|
||||
return str(val.isoformat().replace("+00:00", "Z"))
|
||||
if isinstance(val, Enum):
|
||||
return str(val.value)
|
||||
|
||||
return str(val)
|
||||
|
||||
|
||||
def _get_serialized_params(
|
||||
metadata: ParamMetadata, field_name: str, obj: Any, typ: type
|
||||
) -> Dict[str, str]:
|
||||
params: Dict[str, str] = {}
|
||||
|
||||
serialization = metadata.serialization
|
||||
if serialization == "json":
|
||||
params[field_name] = marshal_json(obj, typ)
|
||||
|
||||
return params
|
||||
|
||||
|
||||
def _is_set(value: Any) -> bool:
|
||||
return value is not None and not isinstance(value, Unset)
|
||||
Reference in New Issue
Block a user