capo-sagemaker-runtime-http2 0.1.0__py3-none-any.whl
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- capo_sagemaker_runtime_http2/__init__.py +74 -0
- capo_sagemaker_runtime_http2/_async.py +25 -0
- capo_sagemaker_runtime_http2/_auth/_identity.py +16 -0
- capo_sagemaker_runtime_http2/_auth/_providers.py +886 -0
- capo_sagemaker_runtime_http2/_auth/_signers.py +88 -0
- capo_sagemaker_runtime_http2/_auth/_sigv4.py +434 -0
- capo_sagemaker_runtime_http2/_auth/_zapros_handler.py +80 -0
- capo_sagemaker_runtime_http2/_iter.py +113 -0
- capo_sagemaker_runtime_http2/_operations/amazon_sage_maker_runtime_http2/invoke_endpoint_with_bidirectional_stream.py +264 -0
- capo_sagemaker_runtime_http2/_pagination.py +21 -0
- capo_sagemaker_runtime_http2/_protocol/__init__.py +1 -0
- capo_sagemaker_runtime_http2/_protocol/errors.py +93 -0
- capo_sagemaker_runtime_http2/_protocol/eventstream.py +238 -0
- capo_sagemaker_runtime_http2/_protocol/serialize.py +47 -0
- capo_sagemaker_runtime_http2/_protocol/xml.py +33 -0
- capo_sagemaker_runtime_http2/_rule_engine/__init__.py +0 -0
- capo_sagemaker_runtime_http2/_rule_engine/_aws_partition.py +160 -0
- capo_sagemaker_runtime_http2/_rule_engine/_endpoint_rule_set.py +560 -0
- capo_sagemaker_runtime_http2/_rule_engine/_endpoint_runtime.py +389 -0
- capo_sagemaker_runtime_http2/_services/_aws_config.py +160 -0
- capo_sagemaker_runtime_http2/_services/_pipeline.py +198 -0
- capo_sagemaker_runtime_http2/_services/async_sage_maker_runtime_http2.py +204 -0
- capo_sagemaker_runtime_http2/_services/sage_maker_runtime_http2.py +203 -0
- capo_sagemaker_runtime_http2/errors/__init__.py +31 -0
- capo_sagemaker_runtime_http2/errors/_base.py +94 -0
- capo_sagemaker_runtime_http2/errors/input_validation_error.py +53 -0
- capo_sagemaker_runtime_http2/errors/internal_server_error.py +51 -0
- capo_sagemaker_runtime_http2/errors/internal_stream_failure.py +61 -0
- capo_sagemaker_runtime_http2/errors/model_error.py +69 -0
- capo_sagemaker_runtime_http2/errors/model_stream_error.py +65 -0
- capo_sagemaker_runtime_http2/errors/service_unavailable_error.py +53 -0
- capo_sagemaker_runtime_http2/py.typed +0 -0
- capo_sagemaker_runtime_http2/types/_prelude/blob.py +12 -0
- capo_sagemaker_runtime_http2/types/_prelude/timestamp.py +17 -0
- capo_sagemaker_runtime_http2/types/invoke_endpoint_with_bidirectional_stream_input.py +21 -0
- capo_sagemaker_runtime_http2/types/invoke_endpoint_with_bidirectional_stream_output.py +15 -0
- capo_sagemaker_runtime_http2/types/request_payload_part.py +62 -0
- capo_sagemaker_runtime_http2/types/request_stream_event.py +52 -0
- capo_sagemaker_runtime_http2/types/response_payload_part.py +62 -0
- capo_sagemaker_runtime_http2/types/response_stream_event.py +100 -0
- capo_sagemaker_runtime_http2/types/sensitive_blob.py +15 -0
- capo_sagemaker_runtime_http2-0.1.0.dist-info/METADATA +108 -0
- capo_sagemaker_runtime_http2-0.1.0.dist-info/RECORD +46 -0
- capo_sagemaker_runtime_http2-0.1.0.dist-info/WHEEL +5 -0
- capo_sagemaker_runtime_http2-0.1.0.dist-info/licenses/LICENSE +21 -0
- capo_sagemaker_runtime_http2-0.1.0.dist-info/top_level.txt +1 -0
|
@@ -0,0 +1,198 @@
|
|
|
1
|
+
from __future__ import annotations
|
|
2
|
+
|
|
3
|
+
import random
|
|
4
|
+
import time
|
|
5
|
+
from collections.abc import Sequence
|
|
6
|
+
from dataclasses import dataclass
|
|
7
|
+
from typing import TYPE_CHECKING, Awaitable, Callable, Generic, TypeVar
|
|
8
|
+
|
|
9
|
+
from zapros import (
|
|
10
|
+
AsyncClient,
|
|
11
|
+
Client,
|
|
12
|
+
ConnectionError,
|
|
13
|
+
Response,
|
|
14
|
+
SSLError,
|
|
15
|
+
TimeoutError,
|
|
16
|
+
)
|
|
17
|
+
|
|
18
|
+
from capo_sagemaker_runtime_http2._async import anysleep
|
|
19
|
+
from capo_sagemaker_runtime_http2.errors import ServiceError
|
|
20
|
+
|
|
21
|
+
if TYPE_CHECKING:
|
|
22
|
+
from capo_sagemaker_runtime_http2._auth._identity import Credentials
|
|
23
|
+
from capo_sagemaker_runtime_http2._auth._providers import IdentityProvider
|
|
24
|
+
|
|
25
|
+
TInput = TypeVar("TInput")
|
|
26
|
+
TOutput = TypeVar("TOutput")
|
|
27
|
+
|
|
28
|
+
|
|
29
|
+
@dataclass
|
|
30
|
+
class OperationOptions:
|
|
31
|
+
client: Client
|
|
32
|
+
use_dual_stack: bool | None = None
|
|
33
|
+
use_fips: bool | None = None
|
|
34
|
+
endpoint: str | None = None
|
|
35
|
+
region: str | None = None
|
|
36
|
+
retry_max_attempts: int | None = None
|
|
37
|
+
credentials_provider: IdentityProvider[Credentials] | None = None
|
|
38
|
+
|
|
39
|
+
|
|
40
|
+
@dataclass
|
|
41
|
+
class AsyncOperationOptions:
|
|
42
|
+
client: AsyncClient
|
|
43
|
+
use_dual_stack: bool | None = None
|
|
44
|
+
use_fips: bool | None = None
|
|
45
|
+
endpoint: str | None = None
|
|
46
|
+
region: str | None = None
|
|
47
|
+
retry_max_attempts: int | None = None
|
|
48
|
+
credentials_provider: IdentityProvider[Credentials] | None = None
|
|
49
|
+
|
|
50
|
+
|
|
51
|
+
@dataclass
|
|
52
|
+
class OperationRequest(Generic[TInput]):
|
|
53
|
+
input: TInput
|
|
54
|
+
options: OperationOptions
|
|
55
|
+
|
|
56
|
+
|
|
57
|
+
@dataclass
|
|
58
|
+
class AsyncOperationRequest(Generic[TInput]):
|
|
59
|
+
input: TInput
|
|
60
|
+
options: AsyncOperationOptions
|
|
61
|
+
|
|
62
|
+
|
|
63
|
+
@dataclass
|
|
64
|
+
class OperationResponse(Generic[TOutput]):
|
|
65
|
+
output: TOutput
|
|
66
|
+
response: Response
|
|
67
|
+
|
|
68
|
+
|
|
69
|
+
@dataclass
|
|
70
|
+
class AsyncOperationResponse(Generic[TOutput]):
|
|
71
|
+
output: TOutput
|
|
72
|
+
response: Response
|
|
73
|
+
|
|
74
|
+
|
|
75
|
+
NextFn = Callable[
|
|
76
|
+
[OperationRequest[TInput]],
|
|
77
|
+
OperationResponse[TOutput],
|
|
78
|
+
]
|
|
79
|
+
|
|
80
|
+
AsyncNextFn = Callable[
|
|
81
|
+
[AsyncOperationRequest[TInput]],
|
|
82
|
+
Awaitable[AsyncOperationResponse[TOutput]],
|
|
83
|
+
]
|
|
84
|
+
|
|
85
|
+
Interceptor = Callable[
|
|
86
|
+
[OperationRequest[TInput], NextFn[TInput, TOutput]],
|
|
87
|
+
OperationResponse[TOutput],
|
|
88
|
+
]
|
|
89
|
+
|
|
90
|
+
AsyncInterceptor = Callable[
|
|
91
|
+
[AsyncOperationRequest[TInput], AsyncNextFn[TInput, TOutput]],
|
|
92
|
+
Awaitable[AsyncOperationResponse[TOutput]],
|
|
93
|
+
]
|
|
94
|
+
|
|
95
|
+
|
|
96
|
+
def execute_pipeline(
|
|
97
|
+
request: OperationRequest[TInput],
|
|
98
|
+
handler: NextFn[TInput, TOutput],
|
|
99
|
+
interceptors: Sequence[Interceptor[TInput, TOutput]],
|
|
100
|
+
) -> OperationResponse[TOutput]:
|
|
101
|
+
def make_chain(index: int) -> NextFn[TInput, TOutput]:
|
|
102
|
+
if index < len(interceptors):
|
|
103
|
+
interceptor = interceptors[index]
|
|
104
|
+
|
|
105
|
+
def next_fn(req: OperationRequest[TInput]) -> OperationResponse[TOutput]:
|
|
106
|
+
return interceptor(req, make_chain(index + 1))
|
|
107
|
+
|
|
108
|
+
return next_fn
|
|
109
|
+
return handler
|
|
110
|
+
|
|
111
|
+
return make_chain(0)(request)
|
|
112
|
+
|
|
113
|
+
|
|
114
|
+
async def aexecute_pipeline(
|
|
115
|
+
request: AsyncOperationRequest[TInput],
|
|
116
|
+
handler: AsyncNextFn[TInput, TOutput],
|
|
117
|
+
interceptors: Sequence[AsyncInterceptor[TInput, TOutput]],
|
|
118
|
+
) -> AsyncOperationResponse[TOutput]:
|
|
119
|
+
def make_chain(index: int) -> AsyncNextFn[TInput, TOutput]:
|
|
120
|
+
if index < len(interceptors):
|
|
121
|
+
interceptor = interceptors[index]
|
|
122
|
+
|
|
123
|
+
async def next_fn(
|
|
124
|
+
req: AsyncOperationRequest[TInput],
|
|
125
|
+
) -> AsyncOperationResponse[TOutput]:
|
|
126
|
+
return await interceptor(req, make_chain(index + 1))
|
|
127
|
+
|
|
128
|
+
return next_fn
|
|
129
|
+
return handler
|
|
130
|
+
|
|
131
|
+
return await make_chain(0)(request)
|
|
132
|
+
|
|
133
|
+
|
|
134
|
+
def _is_retryable(exc: Exception) -> bool:
|
|
135
|
+
if isinstance(exc, ServiceError):
|
|
136
|
+
return exc.is_retryable
|
|
137
|
+
if isinstance(exc, SSLError):
|
|
138
|
+
return False
|
|
139
|
+
if isinstance(exc, (ConnectionError, TimeoutError)):
|
|
140
|
+
return True
|
|
141
|
+
return False
|
|
142
|
+
|
|
143
|
+
|
|
144
|
+
def _retry_delay(attempt: int, is_throttling: bool) -> float:
|
|
145
|
+
base = 1.0 if is_throttling else 0.5
|
|
146
|
+
delay = base * (2.0 ** min(attempt - 1, 10))
|
|
147
|
+
delay = min(20.0, delay)
|
|
148
|
+
return random.uniform(0.0, delay)
|
|
149
|
+
|
|
150
|
+
|
|
151
|
+
def retry() -> Interceptor[TInput, TOutput]:
|
|
152
|
+
def interceptor(
|
|
153
|
+
request: OperationRequest[TInput], next: NextFn[TInput, TOutput]
|
|
154
|
+
) -> OperationResponse[TOutput]:
|
|
155
|
+
max_attempts = request.options.retry_max_attempts or 3
|
|
156
|
+
last_exc: Exception | None = None
|
|
157
|
+
for attempt in range(1, max_attempts + 1):
|
|
158
|
+
try:
|
|
159
|
+
return next(request)
|
|
160
|
+
except Exception as exc:
|
|
161
|
+
if not _is_retryable(exc):
|
|
162
|
+
raise
|
|
163
|
+
last_exc = exc
|
|
164
|
+
if attempt < max_attempts:
|
|
165
|
+
is_throttling = (
|
|
166
|
+
isinstance(exc, ServiceError) and exc.is_throttling_error
|
|
167
|
+
)
|
|
168
|
+
time.sleep(_retry_delay(attempt, is_throttling))
|
|
169
|
+
|
|
170
|
+
assert last_exc
|
|
171
|
+
raise last_exc
|
|
172
|
+
|
|
173
|
+
return interceptor
|
|
174
|
+
|
|
175
|
+
|
|
176
|
+
def aretry() -> AsyncInterceptor[TInput, TOutput]:
|
|
177
|
+
async def interceptor(
|
|
178
|
+
request: AsyncOperationRequest[TInput], next: AsyncNextFn[TInput, TOutput]
|
|
179
|
+
) -> AsyncOperationResponse[TOutput]:
|
|
180
|
+
max_attempts = request.options.retry_max_attempts or 3
|
|
181
|
+
last_exc: Exception | None = None
|
|
182
|
+
for attempt in range(1, max_attempts + 1):
|
|
183
|
+
try:
|
|
184
|
+
return await next(request)
|
|
185
|
+
except Exception as exc:
|
|
186
|
+
if not _is_retryable(exc):
|
|
187
|
+
raise
|
|
188
|
+
last_exc = exc
|
|
189
|
+
if attempt < max_attempts:
|
|
190
|
+
is_throttling = (
|
|
191
|
+
isinstance(exc, ServiceError) and exc.is_throttling_error
|
|
192
|
+
)
|
|
193
|
+
await anysleep(_retry_delay(attempt, is_throttling))
|
|
194
|
+
|
|
195
|
+
assert last_exc
|
|
196
|
+
raise last_exc
|
|
197
|
+
|
|
198
|
+
return interceptor
|
|
@@ -0,0 +1,204 @@
|
|
|
1
|
+
"""Generated from Smithy shape ``com.amazonaws.sagemakerruntimehttp2#AmazonSageMakerRuntimeHttp2``."""
|
|
2
|
+
|
|
3
|
+
import warnings
|
|
4
|
+
from collections.abc import AsyncGenerator, AsyncIterator
|
|
5
|
+
from contextlib import asynccontextmanager
|
|
6
|
+
from typing import TYPE_CHECKING, Any, Iterable, Optional
|
|
7
|
+
|
|
8
|
+
from typing_extensions import Self, TypedDict
|
|
9
|
+
from zapros import AsyncBaseHandler, AsyncClient
|
|
10
|
+
|
|
11
|
+
import capo_sagemaker_runtime_http2._auth._signers
|
|
12
|
+
import capo_sagemaker_runtime_http2._auth._sigv4
|
|
13
|
+
from capo_sagemaker_runtime_http2._auth._identity import Credentials
|
|
14
|
+
from capo_sagemaker_runtime_http2._auth._providers import (
|
|
15
|
+
CredentialsProvider,
|
|
16
|
+
IdentityProvider,
|
|
17
|
+
StaticAwsCredentialsProvider,
|
|
18
|
+
default_aws_credentials_chain,
|
|
19
|
+
)
|
|
20
|
+
from capo_sagemaker_runtime_http2._auth._zapros_handler import AuthMiddleware
|
|
21
|
+
from capo_sagemaker_runtime_http2._iter import (
|
|
22
|
+
ensure_async_iterator,
|
|
23
|
+
)
|
|
24
|
+
from capo_sagemaker_runtime_http2._services._aws_config import aaws_config
|
|
25
|
+
from capo_sagemaker_runtime_http2._services._pipeline import (
|
|
26
|
+
AsyncInterceptor,
|
|
27
|
+
AsyncOperationOptions,
|
|
28
|
+
AsyncOperationRequest,
|
|
29
|
+
AsyncOperationResponse,
|
|
30
|
+
aexecute_pipeline,
|
|
31
|
+
aretry,
|
|
32
|
+
)
|
|
33
|
+
|
|
34
|
+
if TYPE_CHECKING:
|
|
35
|
+
import capo_sagemaker_runtime_http2.types.invoke_endpoint_with_bidirectional_stream_input
|
|
36
|
+
import capo_sagemaker_runtime_http2.types.invoke_endpoint_with_bidirectional_stream_output
|
|
37
|
+
import capo_sagemaker_runtime_http2.types.request_stream_event
|
|
38
|
+
|
|
39
|
+
|
|
40
|
+
class AsyncSageMakerRuntimeHTTP2ClientConfig(TypedDict, total=False, closed=True):
|
|
41
|
+
operation_interceptors: Iterable[AsyncInterceptor[Any, Any]]
|
|
42
|
+
retry_max_attempts: int | None
|
|
43
|
+
use_dual_stack: bool | None
|
|
44
|
+
use_fips: bool | None
|
|
45
|
+
endpoint: str | None
|
|
46
|
+
region: str | None
|
|
47
|
+
credentials_provider: IdentityProvider[Credentials] | None
|
|
48
|
+
|
|
49
|
+
|
|
50
|
+
class AsyncSageMakerRuntimeHTTP2Client:
|
|
51
|
+
"""A client for the ``SageMakerRuntimeHTTP2`` service.
|
|
52
|
+
|
|
53
|
+
Args:
|
|
54
|
+
http_handler: HTTP handler for sending requests. If not provided, creates a default handler.
|
|
55
|
+
operation_interceptors: Interceptors that wrap every operation call. If not provided, defaults to an empty list.
|
|
56
|
+
retry_max_attempts: Maximum number of times to retry a failed operation. Defaults to 3.
|
|
57
|
+
use_dual_stack: The value of the ``AWS::UseDualStack`` endpoint parameter.
|
|
58
|
+
use_fips: The value of the ``AWS::UseFIPS`` endpoint parameter.
|
|
59
|
+
endpoint: The value of the ``SDK::Endpoint`` endpoint parameter.
|
|
60
|
+
region: The value of the ``AWS::Region`` endpoint parameter.
|
|
61
|
+
credentials: AWS credentials for request signing.
|
|
62
|
+
credentials_provider: Provider that resolves AWS credentials. Takes precedence over ``credentials``.
|
|
63
|
+
"""
|
|
64
|
+
|
|
65
|
+
def __init__(
|
|
66
|
+
self,
|
|
67
|
+
http_handler: AsyncBaseHandler | None = None,
|
|
68
|
+
operation_interceptors: Iterable[AsyncInterceptor[Any, Any]] | None = None,
|
|
69
|
+
retry_max_attempts: int | None = None,
|
|
70
|
+
use_dual_stack: bool | None = None,
|
|
71
|
+
use_fips: bool | None = None,
|
|
72
|
+
endpoint: str | None = None,
|
|
73
|
+
region: str | None = None,
|
|
74
|
+
credentials: Credentials | None = None,
|
|
75
|
+
credentials_provider: CredentialsProvider | None = None,
|
|
76
|
+
):
|
|
77
|
+
self._client = AsyncClient(http_handler).wrap_with_middleware(
|
|
78
|
+
lambda next: AuthMiddleware(next)
|
|
79
|
+
)
|
|
80
|
+
if credentials is not None and credentials_provider is not None:
|
|
81
|
+
warnings.warn(
|
|
82
|
+
"Both credentials and credentials_provider given; provider takes precedence"
|
|
83
|
+
)
|
|
84
|
+
resolved_credentials_provider: IdentityProvider[Credentials] | None = (
|
|
85
|
+
credentials_provider
|
|
86
|
+
)
|
|
87
|
+
if resolved_credentials_provider is None and credentials is not None:
|
|
88
|
+
resolved_credentials_provider = StaticAwsCredentialsProvider(credentials)
|
|
89
|
+
if resolved_credentials_provider is None and credentials is None:
|
|
90
|
+
resolved_credentials_provider = default_aws_credentials_chain(
|
|
91
|
+
AsyncClient(http_handler)
|
|
92
|
+
)
|
|
93
|
+
self._config = AsyncSageMakerRuntimeHTTP2ClientConfig(
|
|
94
|
+
{
|
|
95
|
+
"operation_interceptors": operation_interceptors or [],
|
|
96
|
+
"retry_max_attempts": retry_max_attempts,
|
|
97
|
+
"use_dual_stack": use_dual_stack,
|
|
98
|
+
"use_fips": use_fips,
|
|
99
|
+
"endpoint": endpoint,
|
|
100
|
+
"region": region,
|
|
101
|
+
"credentials_provider": resolved_credentials_provider,
|
|
102
|
+
}
|
|
103
|
+
)
|
|
104
|
+
|
|
105
|
+
def operation_options(
|
|
106
|
+
self, config_overrides: Optional[AsyncSageMakerRuntimeHTTP2ClientConfig] = None
|
|
107
|
+
) -> tuple[Iterable[AsyncInterceptor[Any, Any]], AsyncOperationOptions]:
|
|
108
|
+
overrides: AsyncSageMakerRuntimeHTTP2ClientConfig = config_overrides or {}
|
|
109
|
+
interceptors_: list[AsyncInterceptor[Any, Any]] = [
|
|
110
|
+
*overrides.get(
|
|
111
|
+
"operation_interceptors", self._config.get("operation_interceptors", [])
|
|
112
|
+
),
|
|
113
|
+
aaws_config(),
|
|
114
|
+
aretry(),
|
|
115
|
+
]
|
|
116
|
+
options_: AsyncOperationOptions = AsyncOperationOptions(
|
|
117
|
+
client=self._client,
|
|
118
|
+
retry_max_attempts=overrides.get(
|
|
119
|
+
"retry_max_attempts", self._config.get("retry_max_attempts")
|
|
120
|
+
),
|
|
121
|
+
use_dual_stack=overrides.get(
|
|
122
|
+
"use_dual_stack", self._config.get("use_dual_stack")
|
|
123
|
+
),
|
|
124
|
+
use_fips=overrides.get("use_fips", self._config.get("use_fips")),
|
|
125
|
+
endpoint=overrides.get("endpoint", self._config.get("endpoint")),
|
|
126
|
+
region=overrides.get("region", self._config.get("region")),
|
|
127
|
+
credentials_provider=overrides.get(
|
|
128
|
+
"credentials_provider", self._config.get("credentials_provider")
|
|
129
|
+
),
|
|
130
|
+
)
|
|
131
|
+
return interceptors_, options_
|
|
132
|
+
|
|
133
|
+
@asynccontextmanager
|
|
134
|
+
async def invoke_endpoint_with_bidirectional_stream(
|
|
135
|
+
self,
|
|
136
|
+
endpoint_name: str,
|
|
137
|
+
body: "AsyncIterator[capo_sagemaker_runtime_http2.types.request_stream_event._RequestStreamEvent] | capo_sagemaker_runtime_http2.types.request_stream_event._RequestStreamEvent",
|
|
138
|
+
*,
|
|
139
|
+
config_overrides: Optional[AsyncSageMakerRuntimeHTTP2ClientConfig] = None,
|
|
140
|
+
target_variant: Optional[str] = None,
|
|
141
|
+
model_invocation_path: Optional[str] = None,
|
|
142
|
+
model_query_string: Optional[str] = None,
|
|
143
|
+
) -> "AsyncGenerator[capo_sagemaker_runtime_http2.types.invoke_endpoint_with_bidirectional_stream_output.InvokeEndpointWithBidirectionalStreamOutput]":
|
|
144
|
+
r"""<p>Invokes a model endpoint with bidirectional streaming capabilities. This operation establishes a persistent connection that allows you to send multiple requests and receive streaming responses from the model in real-time.</p> <p>Bidirectional streaming is useful for interactive applications such as chatbots, real-time translation, or any scenario where you need to maintain a conversation-like interaction with the model. The connection remains open, allowing you to send additional input and receive responses without establishing a new connection for each request.</p> <p>For an overview of Amazon SageMaker AI, see <a href=\"https://docs.aws.amazon.com/sagemaker/latest/dg/how-it-works.html\">How It Works</a>.</p> <p>Amazon SageMaker AI strips all POST headers except those supported by the API. Amazon SageMaker AI might add additional headers. You should not rely on the behavior of headers outside those enumerated in the request syntax. </p> <p>Calls to <code>InvokeEndpointWithBidirectionalStream</code> are authenticated by using Amazon Web Services Signature Version 4. For information, see <a href=\"https://docs.aws.amazon.com/AmazonS3/latest/API/sig-v4-authenticating-requests.html\">Authenticating Requests (Amazon Web Services Signature Version 4)</a> in the <i>Amazon S3 API Reference</i>.</p> <p>The bidirectional stream maintains the connection until either the client closes it or the model indicates completion. Each request and response in the stream is sent as an event with optional headers for data type and completion state.</p> <note> <p>Endpoints are scoped to an individual account, and are not public. The URL does not contain the account ID, but Amazon SageMaker AI determines the account ID from the authentication token that is supplied by the caller.</p> </note>
|
|
145
|
+
|
|
146
|
+
Args:
|
|
147
|
+
endpoint_name: <p>The name of the endpoint to invoke.</p>
|
|
148
|
+
body: <p>The request payload stream.</p>
|
|
149
|
+
target_variant: <p>Target variant for the request.</p>
|
|
150
|
+
model_invocation_path: <p>Model invocation path.</p>
|
|
151
|
+
model_query_string: <p>Model query string.</p>
|
|
152
|
+
|
|
153
|
+
Raises:
|
|
154
|
+
capo_sagemaker_runtime_http2.errors.input_validation_error.InputValidationError: <p>The input fails to satisfy the constraints specified by an AWS service.</p>
|
|
155
|
+
capo_sagemaker_runtime_http2.errors.internal_server_error.InternalServerError: <p>The request processing has failed because of an unknown error, exception or failure.</p>
|
|
156
|
+
capo_sagemaker_runtime_http2.errors.internal_stream_failure.InternalStreamFailure: <p>Internal stream failure that occurs during streaming.</p>
|
|
157
|
+
capo_sagemaker_runtime_http2.errors.model_error.ModelError: <p>An error occurred while processing the model.</p>
|
|
158
|
+
capo_sagemaker_runtime_http2.errors.model_stream_error.ModelStreamError: <p>Model stream error that occurs during streaming.</p>
|
|
159
|
+
capo_sagemaker_runtime_http2.errors.service_unavailable_error.ServiceUnavailableError: <p>The request has failed due to a temporary failure of the server.</p>
|
|
160
|
+
capo_sagemaker_runtime_http2.errors.UnknownServiceError: The service returned an error code this client does not model.
|
|
161
|
+
"""
|
|
162
|
+
|
|
163
|
+
async def _handler(
|
|
164
|
+
req: "AsyncOperationRequest[capo_sagemaker_runtime_http2.types.invoke_endpoint_with_bidirectional_stream_input.InvokeEndpointWithBidirectionalStreamInput]",
|
|
165
|
+
) -> AsyncOperationResponse[
|
|
166
|
+
"capo_sagemaker_runtime_http2.types.invoke_endpoint_with_bidirectional_stream_output.InvokeEndpointWithBidirectionalStreamOutput"
|
|
167
|
+
]:
|
|
168
|
+
import capo_sagemaker_runtime_http2._operations.amazon_sage_maker_runtime_http2.invoke_endpoint_with_bidirectional_stream
|
|
169
|
+
|
|
170
|
+
(
|
|
171
|
+
output,
|
|
172
|
+
http_response,
|
|
173
|
+
) = await capo_sagemaker_runtime_http2._operations.amazon_sage_maker_runtime_http2.invoke_endpoint_with_bidirectional_stream.async_invoke_endpoint_with_bidirectional_stream(
|
|
174
|
+
req.options, req.input
|
|
175
|
+
)
|
|
176
|
+
return AsyncOperationResponse(output=output, response=http_response)
|
|
177
|
+
|
|
178
|
+
interceptors_, options_ = self.operation_options(config_overrides)
|
|
179
|
+
input_: capo_sagemaker_runtime_http2.types.invoke_endpoint_with_bidirectional_stream_input.InvokeEndpointWithBidirectionalStreamInput = {
|
|
180
|
+
"endpoint_name": endpoint_name,
|
|
181
|
+
"body": ensure_async_iterator(body),
|
|
182
|
+
}
|
|
183
|
+
if target_variant is not None:
|
|
184
|
+
input_["target_variant"] = target_variant
|
|
185
|
+
if model_invocation_path is not None:
|
|
186
|
+
input_["model_invocation_path"] = model_invocation_path
|
|
187
|
+
if model_query_string is not None:
|
|
188
|
+
input_["model_query_string"] = model_query_string
|
|
189
|
+
|
|
190
|
+
response = await aexecute_pipeline(
|
|
191
|
+
AsyncOperationRequest(input=input_, options=options_),
|
|
192
|
+
handler=_handler,
|
|
193
|
+
interceptors=list(interceptors_),
|
|
194
|
+
)
|
|
195
|
+
try:
|
|
196
|
+
yield response.output
|
|
197
|
+
finally:
|
|
198
|
+
await response.response.aclose()
|
|
199
|
+
|
|
200
|
+
async def __aenter__(self) -> Self:
|
|
201
|
+
return self
|
|
202
|
+
|
|
203
|
+
async def __aexit__(self, exc_type: Any, exc: Any, tb: Any):
|
|
204
|
+
await self._client.aclose()
|
|
@@ -0,0 +1,203 @@
|
|
|
1
|
+
"""Generated from Smithy shape ``com.amazonaws.sagemakerruntimehttp2#AmazonSageMakerRuntimeHttp2``."""
|
|
2
|
+
|
|
3
|
+
import warnings
|
|
4
|
+
from collections.abc import Generator, Iterator
|
|
5
|
+
from contextlib import contextmanager
|
|
6
|
+
from typing import TYPE_CHECKING, Any, Iterable, Optional
|
|
7
|
+
|
|
8
|
+
from typing_extensions import Self, TypedDict
|
|
9
|
+
from zapros import BaseHandler, Client
|
|
10
|
+
|
|
11
|
+
import capo_sagemaker_runtime_http2._auth._signers
|
|
12
|
+
import capo_sagemaker_runtime_http2._auth._sigv4
|
|
13
|
+
from capo_sagemaker_runtime_http2._auth._identity import Credentials
|
|
14
|
+
from capo_sagemaker_runtime_http2._auth._providers import (
|
|
15
|
+
CredentialsProvider,
|
|
16
|
+
IdentityProvider,
|
|
17
|
+
StaticAwsCredentialsProvider,
|
|
18
|
+
default_aws_credentials_chain,
|
|
19
|
+
)
|
|
20
|
+
from capo_sagemaker_runtime_http2._auth._zapros_handler import AuthMiddleware
|
|
21
|
+
from capo_sagemaker_runtime_http2._iter import (
|
|
22
|
+
ensure_sync_iterator,
|
|
23
|
+
)
|
|
24
|
+
from capo_sagemaker_runtime_http2._services._aws_config import aws_config
|
|
25
|
+
from capo_sagemaker_runtime_http2._services._pipeline import (
|
|
26
|
+
Interceptor,
|
|
27
|
+
OperationOptions,
|
|
28
|
+
OperationRequest,
|
|
29
|
+
OperationResponse,
|
|
30
|
+
execute_pipeline,
|
|
31
|
+
retry,
|
|
32
|
+
)
|
|
33
|
+
|
|
34
|
+
if TYPE_CHECKING:
|
|
35
|
+
import capo_sagemaker_runtime_http2.types.invoke_endpoint_with_bidirectional_stream_input
|
|
36
|
+
import capo_sagemaker_runtime_http2.types.invoke_endpoint_with_bidirectional_stream_output
|
|
37
|
+
import capo_sagemaker_runtime_http2.types.request_stream_event
|
|
38
|
+
|
|
39
|
+
|
|
40
|
+
class SageMakerRuntimeHTTP2ClientConfig(TypedDict, total=False, closed=True):
|
|
41
|
+
operation_interceptors: Iterable[Interceptor[Any, Any]]
|
|
42
|
+
retry_max_attempts: int | None
|
|
43
|
+
use_dual_stack: bool | None
|
|
44
|
+
use_fips: bool | None
|
|
45
|
+
endpoint: str | None
|
|
46
|
+
region: str | None
|
|
47
|
+
credentials_provider: IdentityProvider[Credentials] | None
|
|
48
|
+
|
|
49
|
+
|
|
50
|
+
class SageMakerRuntimeHTTP2Client:
|
|
51
|
+
"""A client for the ``SageMakerRuntimeHTTP2`` service.
|
|
52
|
+
|
|
53
|
+
Args:
|
|
54
|
+
http_handler: HTTP handler for sending requests. If not provided, creates a default handler.
|
|
55
|
+
operation_interceptors: Interceptors that wrap every operation call. If not provided, defaults to an empty list.
|
|
56
|
+
retry_max_attempts: Maximum number of times to retry a failed operation. Defaults to 3.
|
|
57
|
+
use_dual_stack: The value of the ``AWS::UseDualStack`` endpoint parameter.
|
|
58
|
+
use_fips: The value of the ``AWS::UseFIPS`` endpoint parameter.
|
|
59
|
+
endpoint: The value of the ``SDK::Endpoint`` endpoint parameter.
|
|
60
|
+
region: The value of the ``AWS::Region`` endpoint parameter.
|
|
61
|
+
credentials: AWS credentials for request signing.
|
|
62
|
+
credentials_provider: Provider that resolves AWS credentials. Takes precedence over ``credentials``.
|
|
63
|
+
"""
|
|
64
|
+
|
|
65
|
+
def __init__(
|
|
66
|
+
self,
|
|
67
|
+
http_handler: BaseHandler | None = None,
|
|
68
|
+
operation_interceptors: Iterable[Interceptor[Any, Any]] | None = None,
|
|
69
|
+
retry_max_attempts: int | None = None,
|
|
70
|
+
use_dual_stack: bool | None = None,
|
|
71
|
+
use_fips: bool | None = None,
|
|
72
|
+
endpoint: str | None = None,
|
|
73
|
+
region: str | None = None,
|
|
74
|
+
credentials: Credentials | None = None,
|
|
75
|
+
credentials_provider: CredentialsProvider | None = None,
|
|
76
|
+
):
|
|
77
|
+
self._client = Client(http_handler).wrap_with_middleware(
|
|
78
|
+
lambda next: AuthMiddleware(next)
|
|
79
|
+
)
|
|
80
|
+
if credentials is not None and credentials_provider is not None:
|
|
81
|
+
warnings.warn(
|
|
82
|
+
"Both credentials and credentials_provider given; provider takes precedence"
|
|
83
|
+
)
|
|
84
|
+
resolved_credentials_provider: IdentityProvider[Credentials] | None = (
|
|
85
|
+
credentials_provider
|
|
86
|
+
)
|
|
87
|
+
if resolved_credentials_provider is None and credentials is not None:
|
|
88
|
+
resolved_credentials_provider = StaticAwsCredentialsProvider(credentials)
|
|
89
|
+
if resolved_credentials_provider is None and credentials is None:
|
|
90
|
+
resolved_credentials_provider = default_aws_credentials_chain(
|
|
91
|
+
Client(http_handler)
|
|
92
|
+
)
|
|
93
|
+
self._config = SageMakerRuntimeHTTP2ClientConfig(
|
|
94
|
+
{
|
|
95
|
+
"operation_interceptors": operation_interceptors or [],
|
|
96
|
+
"retry_max_attempts": retry_max_attempts,
|
|
97
|
+
"use_dual_stack": use_dual_stack,
|
|
98
|
+
"use_fips": use_fips,
|
|
99
|
+
"endpoint": endpoint,
|
|
100
|
+
"region": region,
|
|
101
|
+
"credentials_provider": resolved_credentials_provider,
|
|
102
|
+
}
|
|
103
|
+
)
|
|
104
|
+
|
|
105
|
+
def operation_options(
|
|
106
|
+
self, config_overrides: Optional[SageMakerRuntimeHTTP2ClientConfig] = None
|
|
107
|
+
) -> tuple[Iterable[Interceptor[Any, Any]], OperationOptions]:
|
|
108
|
+
overrides: SageMakerRuntimeHTTP2ClientConfig = config_overrides or {}
|
|
109
|
+
interceptors_: list[Interceptor[Any, Any]] = [
|
|
110
|
+
*overrides.get(
|
|
111
|
+
"operation_interceptors", self._config.get("operation_interceptors", [])
|
|
112
|
+
),
|
|
113
|
+
aws_config(),
|
|
114
|
+
retry(),
|
|
115
|
+
]
|
|
116
|
+
options_: OperationOptions = OperationOptions(
|
|
117
|
+
client=self._client,
|
|
118
|
+
retry_max_attempts=overrides.get(
|
|
119
|
+
"retry_max_attempts", self._config.get("retry_max_attempts")
|
|
120
|
+
),
|
|
121
|
+
use_dual_stack=overrides.get(
|
|
122
|
+
"use_dual_stack", self._config.get("use_dual_stack")
|
|
123
|
+
),
|
|
124
|
+
use_fips=overrides.get("use_fips", self._config.get("use_fips")),
|
|
125
|
+
endpoint=overrides.get("endpoint", self._config.get("endpoint")),
|
|
126
|
+
region=overrides.get("region", self._config.get("region")),
|
|
127
|
+
credentials_provider=overrides.get(
|
|
128
|
+
"credentials_provider", self._config.get("credentials_provider")
|
|
129
|
+
),
|
|
130
|
+
)
|
|
131
|
+
return interceptors_, options_
|
|
132
|
+
|
|
133
|
+
@contextmanager
|
|
134
|
+
def invoke_endpoint_with_bidirectional_stream(
|
|
135
|
+
self,
|
|
136
|
+
endpoint_name: str,
|
|
137
|
+
body: "Iterator[capo_sagemaker_runtime_http2.types.request_stream_event._RequestStreamEvent] | capo_sagemaker_runtime_http2.types.request_stream_event._RequestStreamEvent",
|
|
138
|
+
*,
|
|
139
|
+
config_overrides: Optional[SageMakerRuntimeHTTP2ClientConfig] = None,
|
|
140
|
+
target_variant: Optional[str] = None,
|
|
141
|
+
model_invocation_path: Optional[str] = None,
|
|
142
|
+
model_query_string: Optional[str] = None,
|
|
143
|
+
) -> "Generator[capo_sagemaker_runtime_http2.types.invoke_endpoint_with_bidirectional_stream_output.InvokeEndpointWithBidirectionalStreamOutput]":
|
|
144
|
+
r"""<p>Invokes a model endpoint with bidirectional streaming capabilities. This operation establishes a persistent connection that allows you to send multiple requests and receive streaming responses from the model in real-time.</p> <p>Bidirectional streaming is useful for interactive applications such as chatbots, real-time translation, or any scenario where you need to maintain a conversation-like interaction with the model. The connection remains open, allowing you to send additional input and receive responses without establishing a new connection for each request.</p> <p>For an overview of Amazon SageMaker AI, see <a href=\"https://docs.aws.amazon.com/sagemaker/latest/dg/how-it-works.html\">How It Works</a>.</p> <p>Amazon SageMaker AI strips all POST headers except those supported by the API. Amazon SageMaker AI might add additional headers. You should not rely on the behavior of headers outside those enumerated in the request syntax. </p> <p>Calls to <code>InvokeEndpointWithBidirectionalStream</code> are authenticated by using Amazon Web Services Signature Version 4. For information, see <a href=\"https://docs.aws.amazon.com/AmazonS3/latest/API/sig-v4-authenticating-requests.html\">Authenticating Requests (Amazon Web Services Signature Version 4)</a> in the <i>Amazon S3 API Reference</i>.</p> <p>The bidirectional stream maintains the connection until either the client closes it or the model indicates completion. Each request and response in the stream is sent as an event with optional headers for data type and completion state.</p> <note> <p>Endpoints are scoped to an individual account, and are not public. The URL does not contain the account ID, but Amazon SageMaker AI determines the account ID from the authentication token that is supplied by the caller.</p> </note>
|
|
145
|
+
|
|
146
|
+
Args:
|
|
147
|
+
endpoint_name: <p>The name of the endpoint to invoke.</p>
|
|
148
|
+
body: <p>The request payload stream.</p>
|
|
149
|
+
target_variant: <p>Target variant for the request.</p>
|
|
150
|
+
model_invocation_path: <p>Model invocation path.</p>
|
|
151
|
+
model_query_string: <p>Model query string.</p>
|
|
152
|
+
|
|
153
|
+
Raises:
|
|
154
|
+
capo_sagemaker_runtime_http2.errors.input_validation_error.InputValidationError: <p>The input fails to satisfy the constraints specified by an AWS service.</p>
|
|
155
|
+
capo_sagemaker_runtime_http2.errors.internal_server_error.InternalServerError: <p>The request processing has failed because of an unknown error, exception or failure.</p>
|
|
156
|
+
capo_sagemaker_runtime_http2.errors.internal_stream_failure.InternalStreamFailure: <p>Internal stream failure that occurs during streaming.</p>
|
|
157
|
+
capo_sagemaker_runtime_http2.errors.model_error.ModelError: <p>An error occurred while processing the model.</p>
|
|
158
|
+
capo_sagemaker_runtime_http2.errors.model_stream_error.ModelStreamError: <p>Model stream error that occurs during streaming.</p>
|
|
159
|
+
capo_sagemaker_runtime_http2.errors.service_unavailable_error.ServiceUnavailableError: <p>The request has failed due to a temporary failure of the server.</p>
|
|
160
|
+
capo_sagemaker_runtime_http2.errors.UnknownServiceError: The service returned an error code this client does not model.
|
|
161
|
+
"""
|
|
162
|
+
|
|
163
|
+
def _handler(
|
|
164
|
+
req: "OperationRequest[capo_sagemaker_runtime_http2.types.invoke_endpoint_with_bidirectional_stream_input.InvokeEndpointWithBidirectionalStreamInput]",
|
|
165
|
+
) -> OperationResponse[
|
|
166
|
+
"capo_sagemaker_runtime_http2.types.invoke_endpoint_with_bidirectional_stream_output.InvokeEndpointWithBidirectionalStreamOutput"
|
|
167
|
+
]:
|
|
168
|
+
import capo_sagemaker_runtime_http2._operations.amazon_sage_maker_runtime_http2.invoke_endpoint_with_bidirectional_stream
|
|
169
|
+
|
|
170
|
+
output, http_response = (
|
|
171
|
+
capo_sagemaker_runtime_http2._operations.amazon_sage_maker_runtime_http2.invoke_endpoint_with_bidirectional_stream.invoke_endpoint_with_bidirectional_stream(
|
|
172
|
+
req.options, req.input
|
|
173
|
+
)
|
|
174
|
+
)
|
|
175
|
+
return OperationResponse(output=output, response=http_response)
|
|
176
|
+
|
|
177
|
+
interceptors_, options_ = self.operation_options(config_overrides)
|
|
178
|
+
input_: capo_sagemaker_runtime_http2.types.invoke_endpoint_with_bidirectional_stream_input.InvokeEndpointWithBidirectionalStreamInput = {
|
|
179
|
+
"endpoint_name": endpoint_name,
|
|
180
|
+
"body": ensure_sync_iterator(body),
|
|
181
|
+
}
|
|
182
|
+
if target_variant is not None:
|
|
183
|
+
input_["target_variant"] = target_variant
|
|
184
|
+
if model_invocation_path is not None:
|
|
185
|
+
input_["model_invocation_path"] = model_invocation_path
|
|
186
|
+
if model_query_string is not None:
|
|
187
|
+
input_["model_query_string"] = model_query_string
|
|
188
|
+
|
|
189
|
+
response = execute_pipeline(
|
|
190
|
+
OperationRequest(input=input_, options=options_),
|
|
191
|
+
handler=_handler,
|
|
192
|
+
interceptors=list(interceptors_),
|
|
193
|
+
)
|
|
194
|
+
try:
|
|
195
|
+
yield response.output
|
|
196
|
+
finally:
|
|
197
|
+
response.response.close()
|
|
198
|
+
|
|
199
|
+
def __enter__(self) -> Self:
|
|
200
|
+
return self
|
|
201
|
+
|
|
202
|
+
def __exit__(self, exc_type: Any, exc: Any, tb: Any):
|
|
203
|
+
self._client.close()
|
|
@@ -0,0 +1,31 @@
|
|
|
1
|
+
from __future__ import annotations
|
|
2
|
+
|
|
3
|
+
from ._base import (
|
|
4
|
+
DeserializationError as DeserializationError,
|
|
5
|
+
)
|
|
6
|
+
from ._base import (
|
|
7
|
+
SageMakerRuntimeHTTP2Error as SageMakerRuntimeHTTP2Error,
|
|
8
|
+
)
|
|
9
|
+
from ._base import (
|
|
10
|
+
SerializationError as SerializationError,
|
|
11
|
+
)
|
|
12
|
+
from ._base import (
|
|
13
|
+
ServiceError as ServiceError,
|
|
14
|
+
)
|
|
15
|
+
from ._base import (
|
|
16
|
+
UnknownServiceError as UnknownServiceError,
|
|
17
|
+
)
|
|
18
|
+
from ._base import (
|
|
19
|
+
WaiterFailedError as WaiterFailedError,
|
|
20
|
+
)
|
|
21
|
+
from ._base import (
|
|
22
|
+
WaiterTimeoutError as WaiterTimeoutError,
|
|
23
|
+
)
|
|
24
|
+
from .input_validation_error import InputValidationError as InputValidationError
|
|
25
|
+
from .internal_server_error import InternalServerError as InternalServerError
|
|
26
|
+
from .internal_stream_failure import InternalStreamFailure as InternalStreamFailure
|
|
27
|
+
from .model_error import ModelError as ModelError
|
|
28
|
+
from .model_stream_error import ModelStreamError as ModelStreamError
|
|
29
|
+
from .service_unavailable_error import (
|
|
30
|
+
ServiceUnavailableError as ServiceUnavailableError,
|
|
31
|
+
)
|