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.
Files changed (46) hide show
  1. capo_sagemaker_runtime_http2/__init__.py +74 -0
  2. capo_sagemaker_runtime_http2/_async.py +25 -0
  3. capo_sagemaker_runtime_http2/_auth/_identity.py +16 -0
  4. capo_sagemaker_runtime_http2/_auth/_providers.py +886 -0
  5. capo_sagemaker_runtime_http2/_auth/_signers.py +88 -0
  6. capo_sagemaker_runtime_http2/_auth/_sigv4.py +434 -0
  7. capo_sagemaker_runtime_http2/_auth/_zapros_handler.py +80 -0
  8. capo_sagemaker_runtime_http2/_iter.py +113 -0
  9. capo_sagemaker_runtime_http2/_operations/amazon_sage_maker_runtime_http2/invoke_endpoint_with_bidirectional_stream.py +264 -0
  10. capo_sagemaker_runtime_http2/_pagination.py +21 -0
  11. capo_sagemaker_runtime_http2/_protocol/__init__.py +1 -0
  12. capo_sagemaker_runtime_http2/_protocol/errors.py +93 -0
  13. capo_sagemaker_runtime_http2/_protocol/eventstream.py +238 -0
  14. capo_sagemaker_runtime_http2/_protocol/serialize.py +47 -0
  15. capo_sagemaker_runtime_http2/_protocol/xml.py +33 -0
  16. capo_sagemaker_runtime_http2/_rule_engine/__init__.py +0 -0
  17. capo_sagemaker_runtime_http2/_rule_engine/_aws_partition.py +160 -0
  18. capo_sagemaker_runtime_http2/_rule_engine/_endpoint_rule_set.py +560 -0
  19. capo_sagemaker_runtime_http2/_rule_engine/_endpoint_runtime.py +389 -0
  20. capo_sagemaker_runtime_http2/_services/_aws_config.py +160 -0
  21. capo_sagemaker_runtime_http2/_services/_pipeline.py +198 -0
  22. capo_sagemaker_runtime_http2/_services/async_sage_maker_runtime_http2.py +204 -0
  23. capo_sagemaker_runtime_http2/_services/sage_maker_runtime_http2.py +203 -0
  24. capo_sagemaker_runtime_http2/errors/__init__.py +31 -0
  25. capo_sagemaker_runtime_http2/errors/_base.py +94 -0
  26. capo_sagemaker_runtime_http2/errors/input_validation_error.py +53 -0
  27. capo_sagemaker_runtime_http2/errors/internal_server_error.py +51 -0
  28. capo_sagemaker_runtime_http2/errors/internal_stream_failure.py +61 -0
  29. capo_sagemaker_runtime_http2/errors/model_error.py +69 -0
  30. capo_sagemaker_runtime_http2/errors/model_stream_error.py +65 -0
  31. capo_sagemaker_runtime_http2/errors/service_unavailable_error.py +53 -0
  32. capo_sagemaker_runtime_http2/py.typed +0 -0
  33. capo_sagemaker_runtime_http2/types/_prelude/blob.py +12 -0
  34. capo_sagemaker_runtime_http2/types/_prelude/timestamp.py +17 -0
  35. capo_sagemaker_runtime_http2/types/invoke_endpoint_with_bidirectional_stream_input.py +21 -0
  36. capo_sagemaker_runtime_http2/types/invoke_endpoint_with_bidirectional_stream_output.py +15 -0
  37. capo_sagemaker_runtime_http2/types/request_payload_part.py +62 -0
  38. capo_sagemaker_runtime_http2/types/request_stream_event.py +52 -0
  39. capo_sagemaker_runtime_http2/types/response_payload_part.py +62 -0
  40. capo_sagemaker_runtime_http2/types/response_stream_event.py +100 -0
  41. capo_sagemaker_runtime_http2/types/sensitive_blob.py +15 -0
  42. capo_sagemaker_runtime_http2-0.1.0.dist-info/METADATA +108 -0
  43. capo_sagemaker_runtime_http2-0.1.0.dist-info/RECORD +46 -0
  44. capo_sagemaker_runtime_http2-0.1.0.dist-info/WHEEL +5 -0
  45. capo_sagemaker_runtime_http2-0.1.0.dist-info/licenses/LICENSE +21 -0
  46. 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
+ )