Coverage for async_durable_execution/_runner/local/model.py: 100%
96 statements
« prev ^ index » next coverage.py v7.16.0, created at 2026-08-30 23:43 +0000
« prev ^ index » next coverage.py v7.16.0, created at 2026-08-30 23:43 +0000
1"""Models used only by the local durable execution runner."""
3from __future__ import annotations
5import base64
6import json
7from collections.abc import Mapping
8from dataclasses import dataclass, field
9from typing import Any, Protocol, cast
11from ..._core import (
12 AwsApiModel,
13 CheckpointUpdatedExecutionState,
14 DurableExecutionInvocationInput,
15 LambdaContext as LambdaContextProtocol,
16 Operation,
17)
18from ..exceptions import InvalidParameterValueException
19from ..model import InvokeResponse
22@dataclass(frozen=True)
23class LambdaContext(LambdaContextProtocol):
24 """Lambda context for local testing."""
26 aws_request_id: str
27 log_group_name: str | None = None
28 log_stream_name: str | None = None
29 function_name: str | None = None
30 memory_limit_in_mb: str | None = None
31 function_version: str | None = None
32 invoked_function_arn: str | None = None
33 tenant_id: str | None = None
34 client_context: dict | None = None
35 identity: dict | None = None
37 def get_remaining_time_in_millis(self) -> int:
38 return 900000
40 def log(self, msg) -> None:
41 pass
44@dataclass(frozen=True)
45class StartDurableExecutionInput(AwsApiModel):
46 """Input for starting a local durable execution."""
48 account_id: str = field(metadata={"alias": "AccountId"})
49 function_name: str = field(metadata={"alias": "FunctionName"})
50 function_qualifier: str = field(metadata={"alias": "FunctionQualifier"})
51 execution_name: str = field(metadata={"alias": "ExecutionName"})
52 execution_timeout_seconds: int = field(
53 metadata={"alias": "ExecutionTimeoutSeconds"}
54 )
55 execution_retention_period_days: int = field(
56 metadata={"alias": "ExecutionRetentionPeriodDays"}
57 )
58 invocation_id: str | None = field(default=None, metadata={"alias": "InvocationId"})
59 trace_fields: dict | None = field(default=None, metadata={"alias": "TraceFields"})
60 tenant_id: str | None = field(default=None, metadata={"alias": "TenantId"})
61 input: str | None = field(default=None, metadata={"alias": "Input"})
62 lambda_endpoint: str | None = field(
63 default=None, metadata={"alias": "LambdaEndpoint"}
64 )
66 @classmethod
67 def from_dict(cls, data: Mapping[str, Any]) -> StartDurableExecutionInput:
68 required_fields = [
69 "AccountId",
70 "FunctionName",
71 "FunctionQualifier",
72 "ExecutionName",
73 "ExecutionTimeoutSeconds",
74 "ExecutionRetentionPeriodDays",
75 ]
77 for field in required_fields:
78 if field not in data:
79 msg = f"Missing required field: {field}"
80 raise InvalidParameterValueException(msg)
82 return super().from_dict(data)
84 def get_normalized_input(self) -> str:
85 """Normalize input string to be JSON deserializable."""
86 try:
87 json.loads(cast(str, self.input))
88 return cast(str, self.input)
89 except (json.JSONDecodeError, TypeError):
90 return json.dumps(self.input)
93@dataclass(frozen=True)
94class StartDurableExecutionOutput(AwsApiModel):
95 """Output from starting a local durable execution."""
97 execution_arn: str | None = field(default=None, metadata={"alias": "ExecutionArn"})
100@dataclass(frozen=True)
101class GetDurableExecutionStateResponse(AwsApiModel):
102 """Local response containing durable execution state operations."""
104 operations: list[Operation] = field(
105 default_factory=list, metadata={"alias": "Operations"}
106 )
107 next_marker: str | None = field(default=None, metadata={"alias": "NextMarker"})
110@dataclass(frozen=True)
111class SendDurableExecutionCallbackSuccessResponse(AwsApiModel):
112 """Response from sending local callback success."""
115@dataclass(frozen=True)
116class SendDurableExecutionCallbackFailureResponse(AwsApiModel):
117 """Response from sending local callback failure."""
120@dataclass(frozen=True)
121class SendDurableExecutionCallbackHeartbeatResponse(AwsApiModel):
122 """Response from sending local callback heartbeat."""
125@dataclass(frozen=True)
126class CheckpointDurableExecutionResponse(AwsApiModel):
127 """Local response from checkpointing a durable execution."""
129 checkpoint_token: str | None = field(
130 default=None, metadata={"alias": "CheckpointToken"}
131 )
132 new_execution_state: CheckpointUpdatedExecutionState | None = field(
133 default=None, metadata={"alias": "NewExecutionState"}
134 )
137class Invoker(Protocol):
138 def create_invocation_input(
139 self,
140 *,
141 start_input: StartDurableExecutionInput,
142 durable_execution_arn: str,
143 checkpoint_token: str,
144 operations: list[Operation],
145 ) -> DurableExecutionInvocationInput: ...
147 async def invoke(
148 self,
149 function_name: str,
150 input: DurableExecutionInvocationInput,
151 endpoint_url: str | None = None,
152 ) -> InvokeResponse: ...
155@dataclass(frozen=True)
156class CheckpointToken:
157 """Model a local checkpoint token."""
159 execution_arn: str
160 token_sequence: int
162 def to_str(self) -> str:
163 data = {"arn": self.execution_arn, "seq": self.token_sequence}
164 json_str = json.dumps(data, separators=(",", ":"))
165 return base64.b64encode(json_str.encode()).decode()
167 @classmethod
168 def from_str(cls, token: str) -> CheckpointToken:
169 decoded = base64.b64decode(token).decode()
170 data = json.loads(decoded)
171 return cls(execution_arn=data["arn"], token_sequence=data["seq"])
174@dataclass(frozen=True)
175class CallbackToken:
176 """Model a local callback token."""
178 execution_arn: str
179 operation_id: str
181 def to_str(self) -> str:
182 data = {"arn": self.execution_arn, "op": self.operation_id}
183 json_str = json.dumps(data, separators=(",", ":"))
184 return base64.b64encode(json_str.encode()).decode()
186 @classmethod
187 def from_str(cls, token: str) -> CallbackToken:
188 decoded = base64.b64decode(token).decode()
189 data = json.loads(decoded)
190 return cls(execution_arn=data["arn"], operation_id=data["op"])
193__all__ = [
194 "CallbackToken",
195 "CheckpointDurableExecutionResponse",
196 "CheckpointToken",
197 "GetDurableExecutionStateResponse",
198 "Invoker",
199 "LambdaContext",
200 "SendDurableExecutionCallbackFailureResponse",
201 "SendDurableExecutionCallbackHeartbeatResponse",
202 "SendDurableExecutionCallbackSuccessResponse",
203 "StartDurableExecutionInput",
204 "StartDurableExecutionOutput",
205]