Skip to main content

restate_email/
service.rs

1//! Restate service adapter for `restate-email`.
2
3use 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
17/// Concrete Restate service implementation over a transport resolver.
18///
19/// Most applications construct this once at worker startup, register one or
20/// more transports in a [`StaticTransportRegistry`](crate::StaticTransportRegistry),
21/// and bind it to a Restate endpoint through
22/// [`IntoServiceDefinition`](restate_sdk::service::IntoServiceDefinition).
23pub 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    /// Build a service around an owned transport resolver.
42    #[must_use]
43    pub fn new(transports: T) -> Self {
44        Self::from_shared(Arc::new(transports))
45    }
46
47    /// Build a service around a shared transport resolver.
48    #[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    /// Override the provider-option registry used to hydrate queued
57    /// `transport_options`.
58    ///
59    /// The default registry is `email_kit::transport::transport_option_registry()`;
60    /// use this method when a worker has additional provider-specific option
61    /// types outside `email-kit`.
62    #[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    /// Send one email request through the configured worker dependencies.
69    ///
70    /// # Errors
71    ///
72    /// Returns [`HandlerError`] when the requested transport key is unknown or
73    /// sending fails. Unknown transport keys and non-retryable transport
74    /// failures become Restate terminal errors; retryable transport failures
75    /// remain retryable handler errors.
76    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 service for queued email delivery.
91///
92/// The service is exposed as `Email.send` through Restate ingress. Callers that
93/// are not running behind Restate should use [`Service::send_request`] to
94/// exercise the same dispatch path without the service protocol.
95#[restate_sdk::service(name = "Email")]
96impl<T> Service<T>
97where
98    T: TransportResolver + Send + Sync + 'static,
99{
100    /// Dispatch one queued email request through its selected transport.
101    ///
102    /// # Errors
103    ///
104    /// Returns [`HandlerError`] when the request cannot be decoded, the
105    /// transport key cannot be resolved, or the selected transport fails.
106    /// Unregistered provider option keys are ignored during request decoding.
107    #[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}