1use std::convert::Infallible;
4use std::marker::PhantomData;
5use std::sync::Arc;
6
7use bytes::Bytes;
8use email_kit::transport::transport_option_registry;
9use email_transport::{ErrorKind, TransportError, TransportOptionRegistry};
10use restate_sdk::errors::{HandlerError, TerminalError};
11use restate_sdk::prelude::{Context, ContextSideEffects, HandlerResult, Json, RunFuture};
12use restate_sdk::serde::{Deserialize, InputMetadata, OutputMetadata, PayloadMetadata, Serialize};
13use serde::de::DeserializeSeed as _;
14
15use crate::{SendRequest, SendRequestSeed, SendResponse, TransportLookupError, TransportResolver};
16
17pub struct Service<T> {
24 transports: Arc<T>,
25 transport_options: Arc<TransportOptionRegistry>,
26}
27
28impl<T> Clone for Service<T> {
29 fn clone(&self) -> Self {
30 Self {
31 transports: Arc::clone(&self.transports),
32 transport_options: Arc::clone(&self.transport_options),
33 }
34 }
35}
36
37impl<T> Service<T>
38where
39 T: TransportResolver + Send + Sync + 'static,
40{
41 #[must_use]
43 pub fn new(transports: T) -> Self {
44 Self::from_shared(Arc::new(transports))
45 }
46
47 #[must_use]
49 pub fn from_shared(transports: Arc<T>) -> Self {
50 Self {
51 transports,
52 transport_options: Arc::new(transport_option_registry()),
53 }
54 }
55
56 #[must_use]
63 pub fn with_transport_options(mut self, transport_options: TransportOptionRegistry) -> Self {
64 self.transport_options = Arc::new(transport_options);
65 self
66 }
67
68 pub async fn send_request(&self, request: &SendRequest) -> Result<SendResponse, HandlerError> {
77 let transport = self
78 .transports
79 .resolve(&request.transport)
80 .map_err(TerminalError::from)?;
81
82 transport
83 .send(&request.message, &request.options)
84 .await
85 .map(SendResponse::from)
86 .map_err(transport_error_to_handler_error)
87 }
88}
89
90#[restate_sdk::service(name = "Email")]
96impl<T> Service<T>
97where
98 T: TransportResolver + Send + Sync + 'static,
99{
100 #[handler]
108 async fn send(
109 &self,
110 ctx: Context<'_>,
111 request: SeededJson<SendRequest>,
112 ) -> HandlerResult<Json<SendResponse>> {
113 let request = request
114 .deserialize(self.transport_options.as_ref())
115 .map_err(send_request_deserialize_error_to_handler_error)?;
116
117 Ok(ctx
118 .run(|| async move { self.send_request(&request).await.map(Json) })
119 .name("send_email")
120 .await?)
121 }
122}
123
124struct SeededJson<T> {
125 bytes: Bytes,
126 marker: PhantomData<fn() -> T>,
127}
128
129impl SeededJson<SendRequest> {
130 fn deserialize(
131 self,
132 registry: &TransportOptionRegistry,
133 ) -> Result<SendRequest, SendRequestDeserializeError> {
134 let mut deserializer = serde_json::Deserializer::from_slice(&self.bytes);
135 let mut track = serde_path_to_error::Track::new();
136 let path_deserializer =
137 serde_path_to_error::Deserializer::new(&mut deserializer, &mut track);
138 let request = SendRequestSeed::new(registry)
139 .deserialize(path_deserializer)
140 .map_err(|source| SendRequestDeserializeError {
141 path: track.path().to_string(),
142 source,
143 })?;
144 deserializer
145 .end()
146 .map_err(|source| SendRequestDeserializeError {
147 path: String::from("."),
148 source,
149 })?;
150
151 Ok(request)
152 }
153}
154
155impl<T> Deserialize for SeededJson<T> {
156 type Error = Infallible;
157
158 fn deserialize(bytes: &mut Bytes) -> Result<Self, Self::Error> {
159 Ok(Self {
160 bytes: bytes.clone(),
161 marker: PhantomData,
162 })
163 }
164}
165
166impl<T> Serialize for SeededJson<T> {
167 type Error = Infallible;
168
169 fn serialize(&self) -> Result<Bytes, Self::Error> {
170 Ok(self.bytes.clone())
171 }
172}
173
174impl<T> PayloadMetadata for SeededJson<T>
175where
176 Json<T>: PayloadMetadata,
177{
178 fn json_schema() -> Option<serde_json::Value> {
179 <Json<T> as PayloadMetadata>::json_schema()
180 }
181
182 fn input_metadata() -> InputMetadata {
183 <Json<T> as PayloadMetadata>::input_metadata()
184 }
185
186 fn output_metadata() -> OutputMetadata {
187 <Json<T> as PayloadMetadata>::output_metadata()
188 }
189}
190
191#[derive(Debug, thiserror::Error)]
192#[error("{path}: {source}")]
193struct SendRequestDeserializeError {
194 path: String,
195 source: serde_json::Error,
196}
197
198impl From<TransportLookupError> for TerminalError {
199 fn from(error: TransportLookupError) -> Self {
200 Self::new_with_code(404, error.to_string())
201 }
202}
203
204#[allow(clippy::needless_pass_by_value)]
205fn send_request_deserialize_error_to_handler_error(
206 error: SendRequestDeserializeError,
207) -> HandlerError {
208 TerminalError::new_with_code(400, error.to_string()).into()
209}
210
211fn transport_error_to_handler_error(error: TransportError) -> HandlerError {
212 if error.is_retryable() {
213 return error.into();
214 }
215
216 let code = transport_terminal_code(&error);
217
218 TerminalError::new_with_code(code, error.to_string()).into()
219}
220
221const fn transport_terminal_code(error: &TransportError) -> u16 {
222 match error.kind {
223 ErrorKind::Validation | ErrorKind::UnsupportedFeature => 400,
224 ErrorKind::Authentication => 401,
225 ErrorKind::Authorization => 403,
226 ErrorKind::PermanentProvider => 422,
227 _ => 500,
228 }
229}
230
231#[cfg(test)]
232mod tests {
233 use email_message::ContentType;
234 use email_message::{Address, Attachment, Body, Mailbox, Message, OutboundMessage};
235 use email_transport::{SendOptions, SendReport, TransportError};
236 use restate_sdk::discovery::ServiceType as RestateServiceType;
237 use restate_sdk::endpoint::Endpoint;
238 use restate_sdk::service::{Discoverable, IntoServiceDefinition};
239
240 use crate::{TransportKey, TransportLookupError};
241
242 use super::*;
243
244 fn mailbox(input: &str) -> Mailbox {
245 input.parse::<Mailbox>().expect("mailbox should parse")
246 }
247
248 fn request_with_attachment() -> SendRequest {
249 let message = Message::builder(Body::text("hello"))
250 .from_mailbox(mailbox("from@example.com"))
251 .add_to(Address::Mailbox(mailbox("to@example.com")))
252 .add_attachment(
253 Attachment::bytes(
254 ContentType::try_from("application/pdf").expect("content type should parse"),
255 b"attached".to_vec(),
256 )
257 .with_filename("report.pdf"),
258 )
259 .build()
260 .expect("message should validate");
261
262 SendRequest {
263 transport: TransportKey::new_unchecked("transactional"),
264 message: OutboundMessage::new(message).expect("message should be outbound-valid"),
265 options: SendOptions::default(),
266 }
267 }
268
269 struct StubRegistry {
270 error: Option<TransportLookupError>,
271 }
272
273 impl TransportResolver for StubRegistry {
274 fn resolve(
275 &self,
276 _transport: &TransportKey,
277 ) -> Result<&email_transport::DynTransport, TransportLookupError> {
278 Err(self.error.clone().expect("expected lookup error"))
279 }
280 }
281
282 #[test]
283 fn send_email_response_maps_from_send_report() {
284 let report = SendReport::new("resend")
285 .with_provider_message_id("id-1")
286 .with_accepted(vec!["to@example.com".parse().expect("email parses")]);
287
288 let response = SendResponse::from(report);
289
290 assert_eq!(response.report.provider, "resend");
291 assert_eq!(response.report.provider_message_id.as_deref(), Some("id-1"));
292 assert_eq!(response.report.accepted[0].as_str(), "to@example.com");
293 }
294
295 #[test]
296 fn transport_error_disposition_maps_all_current_error_kinds() {
297 let retryable = [
298 ErrorKind::RateLimited,
299 ErrorKind::Timeout,
300 ErrorKind::TransientNetwork,
301 ErrorKind::TransientProvider,
302 ];
303 for kind in retryable {
304 let label = kind.to_string();
305 let error = TransportError::new(kind, "retryable");
306 assert!(error.is_retryable(), "{label} should remain retryable");
307 }
308
309 let terminal = [
310 (ErrorKind::Validation, 400),
311 (ErrorKind::Authentication, 401),
312 (ErrorKind::Authorization, 403),
313 (ErrorKind::PermanentProvider, 422),
314 (ErrorKind::UnsupportedFeature, 400),
315 (ErrorKind::Internal, 500),
316 ];
317 for (kind, expected_code) in terminal {
318 let label = kind.to_string();
319 let error = TransportError::new(kind, "terminal");
320 assert!(!error.is_retryable(), "{label} should remain terminal");
321 assert_eq!(super::transport_terminal_code(&error), expected_code);
322 }
323 }
324
325 #[tokio::test]
326 async fn service_send_maps_lookup_error_to_terminal() {
327 let service = Service::new(StubRegistry {
328 error: Some(TransportLookupError::UnknownKey {
329 key: "transactional".to_owned(),
330 }),
331 });
332
333 let error = service
334 .send_request(&request_with_attachment())
335 .await
336 .expect_err("request should fail");
337
338 let source: &(dyn std::error::Error + Send + Sync + 'static) = error.as_ref();
339 assert!(source.to_string().contains("transactional"));
340 }
341
342 #[test]
343 fn service_discovers_and_binds() {
344 let service = Service::new(StubRegistry {
345 error: Some(TransportLookupError::UnknownKey {
346 key: "transactional".to_owned(),
347 }),
348 });
349
350 let discovery = <Service<StubRegistry> as Discoverable>::discover();
351 assert_eq!(discovery.name.as_str(), "Email");
352 assert_eq!(discovery.ty, RestateServiceType::Service);
353 assert_eq!(discovery.handlers.len(), 1);
354 assert_eq!(discovery.handlers[0].name.as_str(), "send");
355 let input = discovery.handlers[0]
356 .input
357 .as_ref()
358 .expect("send should have input metadata");
359 assert_eq!(
360 input.json_schema,
361 <Json<SendRequest> as PayloadMetadata>::json_schema()
362 );
363 assert_eq!(
364 input.content_type.as_deref(),
365 Some(<Json<SendRequest> as PayloadMetadata>::input_metadata().accept_content_type)
366 );
367
368 let _endpoint = Endpoint::builder()
369 .bind(service.into_service_definition())
370 .build();
371 }
372}