Skip to main content

alien_bindings/providers/queue/
aws_sqs.rs

1use crate::error::{ErrorData, Result};
2use crate::traits::{
3    Binding, MessagePayload, Queue, QueueMessage, MAX_BATCH_SIZE, MAX_MESSAGE_BYTES,
4};
5use alien_aws_clients::sqs::{
6    DeleteMessageRequest, Message, ReceiveMessageRequest, SendMessageRequest, SqsApi, SqsClient,
7};
8use alien_error::{AlienError, Context, ContextError, IntoAlienError};
9use async_trait::async_trait;
10use std::fmt::{Debug, Formatter};
11
12pub struct AwsSqsQueue {
13    queue_url: String,
14    client: SqsClient,
15}
16
17impl Debug for AwsSqsQueue {
18    fn fmt(&self, f: &mut Formatter<'_>) -> std::fmt::Result {
19        f.debug_struct("AwsSqsQueue")
20            .field("queue_url", &self.queue_url)
21            .finish()
22    }
23}
24
25impl AwsSqsQueue {
26    pub fn new(queue_url: String, client: SqsClient) -> Self {
27        Self { queue_url, client }
28    }
29}
30
31impl Binding for AwsSqsQueue {}
32
33#[async_trait]
34impl Queue for AwsSqsQueue {
35    async fn send(&self, _queue: &str, message: MessagePayload) -> Result<()> {
36        let (body, _ct) = match message {
37            MessagePayload::Json(v) => (
38                serde_json::to_string(&v).into_alien_error().context(
39                    ErrorData::BindingSetupFailed {
40                        binding_type: "queue.sqs".to_string(),
41                        reason: "Failed to serialize JSON payload".to_string(),
42                    },
43                )?,
44                "application/json".to_string(),
45            ),
46            MessagePayload::Text(s) => (s, "text/plain; charset=utf-8".to_string()),
47        };
48
49        // Client-side validation: check message size
50        if body.len() > MAX_MESSAGE_BYTES {
51            return Err(alien_error::AlienError::new(
52                ErrorData::BindingSetupFailed {
53                    binding_type: "queue.sqs".to_string(),
54                    reason: format!(
55                        "Message size {} bytes exceeds limit of {} bytes",
56                        body.len(),
57                        MAX_MESSAGE_BYTES
58                    ),
59                },
60            ));
61        }
62
63        let req = SendMessageRequest::builder().message_body(body).build();
64        self.client
65            .send_message(&self.queue_url, req)
66            .await
67            .map(|_| ())
68            .map_err(|e| {
69                e.context(ErrorData::BindingSetupFailed {
70                    binding_type: "queue.sqs".to_string(),
71                    reason: "Failed to send message".to_string(),
72                })
73            })
74    }
75
76    async fn receive(&self, _queue: &str, max_messages: usize) -> Result<Vec<QueueMessage>> {
77        // Client-side validation: check batch size
78        if max_messages == 0 || max_messages > MAX_BATCH_SIZE {
79            return Err(alien_error::AlienError::new(
80                ErrorData::BindingSetupFailed {
81                    binding_type: "queue.sqs".to_string(),
82                    reason: format!(
83                        "Batch size {} is invalid. Must be between 1 and {}",
84                        max_messages, MAX_BATCH_SIZE
85                    ),
86                },
87            ));
88        }
89
90        let req = ReceiveMessageRequest::builder()
91            // The SQS Query-protocol client uses AttributeName.N. AWS keeps
92            // this parameter supported for backward compatibility.
93            .attribute_names(vec!["ApproximateReceiveCount".to_string()])
94            .maybe_max_number_of_messages(Some(max_messages as i32))
95            .maybe_wait_time_seconds(Some(20))
96            .build();
97        let resp = self
98            .client
99            .receive_message(&self.queue_url, req)
100            .await
101            .context(ErrorData::BindingSetupFailed {
102                binding_type: "queue.sqs".to_string(),
103                reason: "Failed to receive".to_string(),
104            })?;
105        resp.receive_message_result
106            .messages
107            .into_iter()
108            .map(|m| {
109                let attempt = receive_count(&m)?;
110                let raw = m.body;
111                let payload = serde_json::from_str::<serde_json::Value>(&raw)
112                    .map(MessagePayload::Json)
113                    .unwrap_or(MessagePayload::Text(raw));
114                Ok(QueueMessage {
115                    payload,
116                    receipt_handle: m.receipt_handle,
117                    attempt,
118                })
119            })
120            .collect()
121    }
122
123    async fn ack(&self, _queue: &str, receipt_handle: &str) -> Result<()> {
124        let req = DeleteMessageRequest::builder()
125            .receipt_handle(receipt_handle.to_string())
126            .build();
127        self.client
128            .delete_message(&self.queue_url, req)
129            .await
130            .context(ErrorData::BindingSetupFailed {
131                binding_type: "queue.sqs".to_string(),
132                reason: "Failed to delete message".to_string(),
133            })
134    }
135
136    async fn nack(&self, _queue: &str, _receipt_handle: &str) -> Result<()> {
137        // SQS nack is ChangeMessageVisibility(VisibilityTimeout=0). The
138        // alien-aws-clients SqsApi wrapper does not expose that call, so we
139        // fail explicitly rather than silently waiting out the lease.
140        Err(alien_error::AlienError::new(
141            ErrorData::OperationNotSupported {
142                operation: "queue.nack".to_string(),
143                reason: "AWS SQS nack requires ChangeMessageVisibility, which the SQS client does not expose".to_string(),
144            },
145        ))
146    }
147
148    async fn purge(&self, _queue: &str) -> Result<()> {
149        self.client
150            .purge_queue(&self.queue_url)
151            .await
152            .context(ErrorData::BindingSetupFailed {
153                binding_type: "queue.sqs".to_string(),
154                reason: "Failed to purge queue".to_string(),
155            })
156    }
157}
158
159fn receive_count(message: &Message) -> Result<u32> {
160    let reason = || ErrorData::QueueProviderResponseInvalid {
161        reason: format!(
162            "SQS message '{}' has no positive ApproximateReceiveCount",
163            message.message_id
164        ),
165    };
166    let raw = message
167        .attributes
168        .as_ref()
169        .and_then(|attributes| attributes.get("ApproximateReceiveCount"))
170        .ok_or_else(|| AlienError::new(reason()))?;
171    let count = raw.parse::<u32>().into_alien_error().context(reason())?;
172    if count == 0 {
173        return Err(AlienError::new(reason()));
174    }
175    Ok(count)
176}
177
178#[cfg(test)]
179mod tests {
180    use super::*;
181    use std::collections::HashMap;
182
183    fn message(count: Option<&str>) -> Message {
184        Message {
185            attributes: count.map(|count| {
186                HashMap::from([("ApproximateReceiveCount".to_string(), count.to_string())])
187            }),
188            body: "payload".to_string(),
189            md5_of_body: "unused".to_string(),
190            md5_of_message_attributes: None,
191            message_attributes: None,
192            message_id: "message-1".to_string(),
193            receipt_handle: "receipt-1".to_string(),
194        }
195    }
196
197    #[test]
198    fn reads_positive_sqs_delivery_attempts() {
199        assert_eq!(receive_count(&message(Some("1"))).expect("first"), 1);
200        assert_eq!(receive_count(&message(Some("2"))).expect("redelivery"), 2);
201    }
202
203    #[test]
204    fn rejects_missing_or_invalid_sqs_delivery_attempts() {
205        for count in [
206            None,
207            Some(""),
208            Some("not-a-number"),
209            Some("0"),
210            Some("4294967296"),
211        ] {
212            let error = receive_count(&message(count)).expect_err("invalid attempt must fail");
213            assert!(matches!(
214                error.error,
215                Some(ErrorData::QueueProviderResponseInvalid { .. })
216            ));
217        }
218    }
219}