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

1"""Models used only by the local durable execution runner.""" 

2 

3from __future__ import annotations 

4 

5import base64 

6import json 

7from collections.abc import Mapping 

8from dataclasses import dataclass, field 

9from typing import Any, Protocol, cast 

10 

11from ..._core import ( 

12 AwsApiModel, 

13 CheckpointUpdatedExecutionState, 

14 DurableExecutionInvocationInput, 

15 LambdaContext as LambdaContextProtocol, 

16 Operation, 

17) 

18from ..exceptions import InvalidParameterValueException 

19from ..model import InvokeResponse 

20 

21 

22@dataclass(frozen=True) 

23class LambdaContext(LambdaContextProtocol): 

24 """Lambda context for local testing.""" 

25 

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 

36 

37 def get_remaining_time_in_millis(self) -> int: 

38 return 900000 

39 

40 def log(self, msg) -> None: 

41 pass 

42 

43 

44@dataclass(frozen=True) 

45class StartDurableExecutionInput(AwsApiModel): 

46 """Input for starting a local durable execution.""" 

47 

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 ) 

65 

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 ] 

76 

77 for field in required_fields: 

78 if field not in data: 

79 msg = f"Missing required field: {field}" 

80 raise InvalidParameterValueException(msg) 

81 

82 return super().from_dict(data) 

83 

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) 

91 

92 

93@dataclass(frozen=True) 

94class StartDurableExecutionOutput(AwsApiModel): 

95 """Output from starting a local durable execution.""" 

96 

97 execution_arn: str | None = field(default=None, metadata={"alias": "ExecutionArn"}) 

98 

99 

100@dataclass(frozen=True) 

101class GetDurableExecutionStateResponse(AwsApiModel): 

102 """Local response containing durable execution state operations.""" 

103 

104 operations: list[Operation] = field( 

105 default_factory=list, metadata={"alias": "Operations"} 

106 ) 

107 next_marker: str | None = field(default=None, metadata={"alias": "NextMarker"}) 

108 

109 

110@dataclass(frozen=True) 

111class SendDurableExecutionCallbackSuccessResponse(AwsApiModel): 

112 """Response from sending local callback success.""" 

113 

114 

115@dataclass(frozen=True) 

116class SendDurableExecutionCallbackFailureResponse(AwsApiModel): 

117 """Response from sending local callback failure.""" 

118 

119 

120@dataclass(frozen=True) 

121class SendDurableExecutionCallbackHeartbeatResponse(AwsApiModel): 

122 """Response from sending local callback heartbeat.""" 

123 

124 

125@dataclass(frozen=True) 

126class CheckpointDurableExecutionResponse(AwsApiModel): 

127 """Local response from checkpointing a durable execution.""" 

128 

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 ) 

135 

136 

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: ... 

146 

147 async def invoke( 

148 self, 

149 function_name: str, 

150 input: DurableExecutionInvocationInput, 

151 endpoint_url: str | None = None, 

152 ) -> InvokeResponse: ... 

153 

154 

155@dataclass(frozen=True) 

156class CheckpointToken: 

157 """Model a local checkpoint token.""" 

158 

159 execution_arn: str 

160 token_sequence: int 

161 

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() 

166 

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"]) 

172 

173 

174@dataclass(frozen=True) 

175class CallbackToken: 

176 """Model a local callback token.""" 

177 

178 execution_arn: str 

179 operation_id: str 

180 

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() 

185 

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"]) 

191 

192 

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]