Skip to main content

faucet_source_sqs/
stream.rs

1//! The SQS `Source` implementation: long-poll `ReceiveMessage`, buffer to
2//! `batch_size`, delete each page's receipt handles right before yielding it,
3//! and terminate on `idle_timeout_secs` / `max_messages`.
4
5use crate::config::{MAX_RECEIVE_BATCH, SqsSourceConfig};
6use aws_sdk_sqs::Client;
7use aws_sdk_sqs::types::DeleteMessageBatchRequestEntry;
8use faucet_core::{FaucetError, Stream, StreamPage};
9use serde_json::Value;
10use std::pin::Pin;
11use std::time::{Duration, Instant};
12
13/// AWS SQS source. See the crate README for semantics.
14pub struct SqsSource {
15    config: SqsSourceConfig,
16    client: Client,
17}
18
19/// Decode one SQS message body: the parsed JSON value if the body is valid
20/// JSON, otherwise the raw body wrapped as a JSON string. Pure.
21pub(crate) fn decode_body(body: &str) -> Value {
22    serde_json::from_str::<Value>(body).unwrap_or_else(|_| Value::String(body.to_string()))
23}
24
25impl SqsSource {
26    /// Create a new SQS source. Validates the config and builds the AWS client;
27    /// no queue I/O happens until the first `ReceiveMessage` at stream time
28    /// (construction is offline).
29    pub async fn new(config: SqsSourceConfig) -> Result<Self, FaucetError> {
30        config.validate()?;
31        let client = faucet_common_sqs::build_client(
32            config.region.as_deref(),
33            config.endpoint_url.as_deref(),
34            &config.credentials,
35        )
36        .await?;
37        Ok(Self { config, client })
38    }
39
40    /// Delete a page's receipt handles, chunked to the 10-entry API cap. A
41    /// whole-request failure propagates as a typed error; per-entry failures
42    /// are logged and left for SQS to redeliver (at-least-once).
43    async fn delete_handles(&self, handles: &[String]) -> Result<(), FaucetError> {
44        for chunk in handles.chunks(MAX_RECEIVE_BATCH as usize) {
45            let entries: Vec<DeleteMessageBatchRequestEntry> = chunk
46                .iter()
47                .enumerate()
48                .map(|(i, rh)| {
49                    DeleteMessageBatchRequestEntry::builder()
50                        .id(i.to_string())
51                        .receipt_handle(rh)
52                        .build()
53                        .map_err(|e| {
54                            FaucetError::Source(format!("sqs: delete entry build failed: {e}"))
55                        })
56                })
57                .collect::<Result<_, _>>()?;
58            let out = self
59                .client
60                .delete_message_batch()
61                .queue_url(&self.config.queue_url)
62                .set_entries(Some(entries))
63                .send()
64                .await
65                .map_err(|e| {
66                    FaucetError::Source(format!(
67                        "sqs: DeleteMessageBatch on '{}' failed: {}",
68                        self.config.queue_url,
69                        e.into_service_error()
70                    ))
71                })?;
72            for f in out.failed() {
73                tracing::warn!(
74                    queue = %self.config.queue_url,
75                    id = f.id(),
76                    code = f.code(),
77                    "sqs: message delete failed; message will be redelivered"
78                );
79            }
80        }
81        Ok(())
82    }
83}
84
85#[faucet_core::async_trait]
86impl faucet_core::Source for SqsSource {
87    /// Drain the queue to termination (`idle_timeout_secs` / `max_messages` —
88    /// at least one is enforced at construction) and return every message.
89    async fn fetch_with_context(
90        &self,
91        context: &std::collections::HashMap<String, Value>,
92    ) -> Result<Vec<Value>, FaucetError> {
93        use futures::StreamExt;
94        let mut pages = self.stream_pages(context, self.config.batch_size);
95        let mut all = Vec::new();
96        while let Some(page) = pages.next().await {
97            all.extend(page?.records);
98        }
99        Ok(all)
100    }
101
102    fn stream_pages<'a>(
103        &'a self,
104        _context: &'a std::collections::HashMap<String, Value>,
105        _batch_size: usize,
106    ) -> Pin<Box<dyn Stream<Item = Result<StreamPage, FaucetError>> + Send + 'a>> {
107        let batch_size = self.config.batch_size;
108        let chunk = if batch_size == 0 {
109            usize::MAX
110        } else {
111            batch_size
112        };
113
114        Box::pin(async_stream::try_stream! {
115            let idle = self.config.idle_timeout_secs.map(Duration::from_secs);
116            let max = self.config.max_messages;
117            // `buffer` and `handles` stay index-aligned: one handle slot per
118            // buffered record (None if a message somehow lacked a receipt
119            // handle — such a record is emitted but not deleted, so SQS
120            // redelivers it, preserving at-least-once).
121            let mut buffer: Vec<Value> = Vec::new();
122            let mut handles: Vec<Option<String>> = Vec::new();
123            let mut total = 0usize;
124            let mut last_activity = Instant::now();
125
126            loop {
127                // Reached the message cap → stop.
128                let remaining = match max {
129                    Some(m) if total >= m => break,
130                    Some(m) => Some(m - total),
131                    None => None,
132                };
133                let want = remaining
134                    .map_or(MAX_RECEIVE_BATCH, |r| r.min(MAX_RECEIVE_BATCH as usize) as i32);
135
136                let resp = self
137                    .client
138                    .receive_message()
139                    .queue_url(&self.config.queue_url)
140                    .max_number_of_messages(want)
141                    .wait_time_seconds(self.config.wait_time_seconds)
142                    .send()
143                    .await
144                    .map_err(|e| {
145                        FaucetError::Source(format!(
146                            "sqs: ReceiveMessage on '{}' failed: {}",
147                            self.config.queue_url,
148                            e.into_service_error()
149                        ))
150                    })?;
151
152                let messages = resp.messages();
153                if messages.is_empty() {
154                    if let Some(window) = idle
155                        && last_activity.elapsed() >= window
156                    {
157                        tracing::info!(
158                            queue = %self.config.queue_url,
159                            idle_secs = window.as_secs(),
160                            "sqs: idle termination reached"
161                        );
162                        break;
163                    }
164                    continue;
165                }
166                last_activity = Instant::now();
167
168                for msg in messages {
169                    buffer.push(decode_body(msg.body().unwrap_or("")));
170                    handles.push(msg.receipt_handle().map(str::to_string));
171                    total += 1;
172                }
173
174                // Emit every full page, deleting its handles right before yield.
175                while buffer.len() >= chunk {
176                    let page: Vec<Value> = buffer.drain(..chunk).collect();
177                    let to_delete: Vec<String> =
178                        handles.drain(..chunk).flatten().collect();
179                    self.delete_handles(&to_delete).await?;
180                    yield StreamPage { records: page, bookmark: None };
181                }
182
183                if let Some(m) = max
184                    && total >= m
185                {
186                    tracing::info!(
187                        queue = %self.config.queue_url,
188                        max = m,
189                        "sqs: max_messages reached"
190                    );
191                    break;
192                }
193            }
194
195            // Flush whatever is left as a final (short) page.
196            if !buffer.is_empty() {
197                let to_delete: Vec<String> = handles.into_iter().flatten().collect();
198                self.delete_handles(&to_delete).await?;
199                yield StreamPage { records: buffer, bookmark: None };
200            }
201            tracing::info!(
202                queue = %self.config.queue_url,
203                records = total,
204                "sqs source stream complete"
205            );
206        })
207    }
208
209    fn config_schema(&self) -> Value {
210        serde_json::to_value(faucet_core::schema_for!(SqsSourceConfig))
211            .expect("schema serialization")
212    }
213
214    fn connector_name(&self) -> &'static str {
215        "sqs"
216    }
217
218    fn dataset_uri(&self) -> String {
219        let name = self
220            .config
221            .queue_url
222            .rsplit('/')
223            .find(|s| !s.is_empty())
224            .unwrap_or(self.config.queue_url.as_str());
225        format!(
226            "sqs://{}/{}",
227            self.config.region.as_deref().unwrap_or("default"),
228            name
229        )
230    }
231
232    /// Side-effect-free probe: `GetQueueAttributes` (no messages consumed). The
233    /// default first-page probe could block for the full long-poll window on a
234    /// quiet queue.
235    async fn check(
236        &self,
237        ctx: &faucet_core::CheckContext,
238    ) -> Result<faucet_core::CheckReport, FaucetError> {
239        use faucet_core::{CheckReport, Probe};
240        let start = std::time::Instant::now();
241        let fut = self
242            .client
243            .get_queue_attributes()
244            .queue_url(&self.config.queue_url)
245            .send();
246        let probe = match tokio::time::timeout(ctx.timeout, fut).await {
247            Err(_) => Probe::fail("get_queue_attributes", start.elapsed(), "timed out"),
248            Ok(Ok(_)) => Probe::pass("get_queue_attributes", start.elapsed()),
249            Ok(Err(e)) => Probe::fail(
250                "get_queue_attributes",
251                start.elapsed(),
252                e.into_service_error().to_string(),
253            ),
254        };
255        Ok(CheckReport::single(probe))
256    }
257}
258
259#[cfg(test)]
260mod tests {
261    use super::*;
262    use faucet_core::Source as _;
263
264    fn decode(s: &str) -> Value {
265        decode_body(s)
266    }
267
268    #[test]
269    fn decode_body_parses_json_else_string() {
270        assert_eq!(decode(r#"{"a":1}"#), serde_json::json!({"a": 1}));
271        assert_eq!(decode("[1,2,3]"), serde_json::json!([1, 2, 3]));
272        assert_eq!(decode("not json"), Value::String("not json".into()));
273        assert_eq!(decode(""), Value::String(String::new()));
274    }
275
276    async fn offline_source(mut config: SqsSourceConfig) -> SqsSource {
277        config.endpoint_url = Some("http://127.0.0.1:1".into()); // unroutable
278        config.region = Some("us-east-1".into());
279        config.credentials = faucet_common_sqs::SqsCredentials::AccessKey {
280            access_key_id: "test".into(),
281            secret_access_key: "test".into(),
282            session_token: None,
283        };
284        SqsSource::new(config).await.expect("source builds")
285    }
286
287    #[tokio::test]
288    async fn new_validates_config() {
289        let err = match SqsSource::new(SqsSourceConfig::new("https://q")).await {
290            Err(e) => e,
291            Ok(_) => panic!("config without a termination knob must be rejected"),
292        };
293        assert!(err.to_string().contains("idle_timeout_secs"), "{err}");
294    }
295
296    #[tokio::test]
297    async fn identity_overrides() {
298        let mut cfg = SqsSourceConfig::new("https://sqs.us-east-1.amazonaws.com/1/events");
299        cfg.max_messages = Some(10);
300        let source = offline_source(cfg).await;
301        assert_eq!(source.connector_name(), "sqs");
302        assert_eq!(source.dataset_uri(), "sqs://us-east-1/events");
303        assert_eq!(source.state_key(), None);
304        assert!(!source.supports_exactly_once());
305        let schema = source.config_schema();
306        assert!(
307            schema["properties"]["queue_url"].is_object(),
308            "schema exposes config fields"
309        );
310    }
311
312    #[tokio::test]
313    async fn stream_pages_surfaces_receive_errors() {
314        use futures::StreamExt;
315        let mut cfg = SqsSourceConfig::new("https://q");
316        cfg.max_messages = Some(10);
317        cfg.wait_time_seconds = 0;
318        let source = offline_source(cfg).await;
319        let ctx = std::collections::HashMap::new();
320        let mut pages = source.stream_pages(&ctx, 10);
321        let first = pages.next().await.expect("one item");
322        let err = first.unwrap_err();
323        assert!(err.to_string().contains("ReceiveMessage"), "{err}");
324    }
325
326    #[tokio::test]
327    async fn check_probe_fails_cleanly_offline() {
328        let mut cfg = SqsSourceConfig::new("https://q");
329        cfg.max_messages = Some(10);
330        let source = offline_source(cfg).await;
331        let report = source
332            .check(&faucet_core::CheckContext {
333                timeout: Duration::from_millis(500),
334            })
335            .await
336            .unwrap();
337        assert_eq!(
338            report.failed_count(),
339            1,
340            "unreachable endpoint → fail probe"
341        );
342    }
343}