Skip to main content

alien_aws_clients/aws/
cloudwatch.rs

1use crate::aws::aws_request_utils::{AwsRequestBuilderExt, AwsSignConfig};
2use crate::aws::credential_provider::AwsCredentialProvider;
3use alien_client_core::{ErrorData, Result};
4use alien_error::ContextError;
5use bon::Builder;
6use reqwest::{Client, StatusCode};
7use serde::de::DeserializeOwned;
8use serde::{Deserialize, Serialize};
9
10#[cfg(feature = "test-utils")]
11use mockall::automock;
12
13pub const GET_METRIC_DATA_TARGET: &str = "GraniteServiceVersion20100801.GetMetricData";
14pub const LIST_METRICS_TARGET: &str = "GraniteServiceVersion20100801.ListMetrics";
15
16#[cfg_attr(feature = "test-utils", automock)]
17#[cfg_attr(target_arch = "wasm32", async_trait::async_trait(?Send))]
18#[cfg_attr(not(target_arch = "wasm32"), async_trait::async_trait)]
19pub trait CloudWatchApi: Send + Sync + std::fmt::Debug {
20    async fn get_metric_data(&self, request: GetMetricDataRequest)
21        -> Result<GetMetricDataResponse>;
22
23    async fn list_metrics(&self, request: ListMetricsRequest) -> Result<ListMetricsResponse>;
24}
25
26#[derive(Debug, Clone)]
27pub struct CloudWatchClient {
28    client: Client,
29    credentials: AwsCredentialProvider,
30}
31
32impl CloudWatchClient {
33    pub fn new(client: Client, credentials: AwsCredentialProvider) -> Self {
34        Self {
35            client,
36            credentials,
37        }
38    }
39
40    fn sign_config(&self) -> AwsSignConfig {
41        AwsSignConfig {
42            service_name: "monitoring".into(),
43            region: self.credentials.region().to_string(),
44            credentials: self.credentials.get_credentials(),
45            signing_region: None,
46        }
47    }
48
49    fn host(&self) -> String {
50        format!("monitoring.{}.amazonaws.com", self.credentials.region())
51    }
52
53    fn get_base_url(&self) -> String {
54        if let Some(override_url) = self
55            .credentials
56            .get_service_endpoint_option("monitoring")
57            .or_else(|| self.credentials.get_service_endpoint_option("cloudwatch"))
58        {
59            override_url.to_string()
60        } else {
61            format!("https://{}", self.host())
62        }
63    }
64
65    async fn send_json<T: DeserializeOwned + Send + 'static>(
66        &self,
67        target: &str,
68        body: String,
69        operation: &str,
70        resource: &str,
71    ) -> Result<T> {
72        self.credentials.ensure_fresh().await?;
73        let url = format!("{}/", self.get_base_url().trim_end_matches('/'));
74
75        let builder = self
76            .client
77            .post(&url)
78            .host(&self.host())
79            .header("X-Amz-Target", target)
80            .content_type_amz_json()
81            .content_sha256(&body)
82            .body(body.clone());
83
84        let result =
85            crate::aws::aws_request_utils::sign_send_json(builder, &self.sign_config()).await;
86
87        Self::map_result(result, operation, resource, Some(&body))
88    }
89
90    fn map_result<T>(
91        result: Result<T>,
92        operation: &str,
93        resource: &str,
94        request_body: Option<&str>,
95    ) -> Result<T> {
96        match result {
97            Ok(value) => Ok(value),
98            Err(error) => {
99                if let Some(ErrorData::HttpResponseError {
100                    http_status,
101                    http_response_text: Some(ref text),
102                    ..
103                }) = &error.error
104                {
105                    let status = StatusCode::from_u16(*http_status)
106                        .unwrap_or(StatusCode::INTERNAL_SERVER_ERROR);
107                    if let Some(mapped) =
108                        Self::map_error(status, text, operation, resource, request_body)
109                    {
110                        return Err(error.context(mapped));
111                    }
112                }
113                Err(error)
114            }
115        }
116    }
117
118    fn map_error(
119        status: StatusCode,
120        body: &str,
121        operation: &str,
122        resource: &str,
123        request_body: Option<&str>,
124    ) -> Option<ErrorData> {
125        let parsed: Option<CloudWatchErrorResponse> = serde_json::from_str(body).ok();
126        let code = parsed
127            .as_ref()
128            .map(|error| error.code.trim_start_matches('#'))
129            .unwrap_or_default();
130        let message = parsed
131            .as_ref()
132            .and_then(|error| error.message.clone())
133            .unwrap_or_else(|| body.to_string());
134
135        match status {
136            StatusCode::UNAUTHORIZED | StatusCode::FORBIDDEN => {
137                Some(ErrorData::AuthenticationError { message })
138            }
139            StatusCode::TOO_MANY_REQUESTS => Some(ErrorData::RateLimitExceeded { message }),
140            StatusCode::BAD_REQUEST
141                if matches!(code, "InvalidParameterValue" | "InvalidNextToken") =>
142            {
143                Some(ErrorData::InvalidClientConfig {
144                    message,
145                    errors: None,
146                })
147            }
148            StatusCode::INTERNAL_SERVER_ERROR if code == "InternalServiceError" => {
149                Some(ErrorData::RemoteServiceUnavailable { message })
150            }
151            _ if !body.trim().is_empty() => Some(ErrorData::HttpResponseError {
152                message: format!("{} failed for '{}': {}", operation, resource, message),
153                url: String::new(),
154                http_status: status.as_u16(),
155                http_request_text: request_body.map(ToOwned::to_owned),
156                http_response_text: Some(body.to_string()),
157            }),
158            _ => None,
159        }
160    }
161}
162
163#[cfg_attr(target_arch = "wasm32", async_trait::async_trait(?Send))]
164#[cfg_attr(not(target_arch = "wasm32"), async_trait::async_trait)]
165impl CloudWatchApi for CloudWatchClient {
166    async fn get_metric_data(
167        &self,
168        request: GetMetricDataRequest,
169    ) -> Result<GetMetricDataResponse> {
170        let body = serde_json::to_string(&request).map_err(|error| {
171            alien_error::AlienError::new(ErrorData::InvalidClientConfig {
172                message: format!("Failed to serialize GetMetricData request: {}", error),
173                errors: None,
174            })
175        })?;
176
177        self.send_json(
178            GET_METRIC_DATA_TARGET,
179            body,
180            "GetMetricData",
181            "cloudwatch metrics",
182        )
183        .await
184    }
185
186    async fn list_metrics(&self, request: ListMetricsRequest) -> Result<ListMetricsResponse> {
187        let body = serde_json::to_string(&request).map_err(|error| {
188            alien_error::AlienError::new(ErrorData::InvalidClientConfig {
189                message: format!("Failed to serialize ListMetrics request: {}", error),
190                errors: None,
191            })
192        })?;
193
194        self.send_json(
195            LIST_METRICS_TARGET,
196            body,
197            "ListMetrics",
198            request.namespace.as_deref().unwrap_or("all namespaces"),
199        )
200        .await
201    }
202}
203
204#[derive(Debug, Clone, Builder, Serialize)]
205#[serde(rename_all = "PascalCase")]
206pub struct GetMetricDataRequest {
207    pub start_time: i64,
208    pub end_time: i64,
209    pub metric_data_queries: Vec<MetricDataQuery>,
210    #[serde(skip_serializing_if = "Option::is_none")]
211    pub max_datapoints: Option<i32>,
212    #[serde(skip_serializing_if = "Option::is_none")]
213    pub next_token: Option<String>,
214    #[serde(skip_serializing_if = "Option::is_none")]
215    pub scan_by: Option<String>,
216}
217
218#[derive(Debug, Clone, Builder, Serialize)]
219#[serde(rename_all = "PascalCase")]
220pub struct MetricDataQuery {
221    pub id: String,
222    #[serde(skip_serializing_if = "Option::is_none")]
223    pub expression: Option<String>,
224    #[serde(skip_serializing_if = "Option::is_none")]
225    pub label: Option<String>,
226    #[serde(skip_serializing_if = "Option::is_none")]
227    pub metric_stat: Option<MetricStat>,
228    #[serde(skip_serializing_if = "Option::is_none")]
229    pub period: Option<i32>,
230    #[serde(skip_serializing_if = "Option::is_none")]
231    pub return_data: Option<bool>,
232}
233
234#[derive(Debug, Clone, Builder, Serialize)]
235#[serde(rename_all = "PascalCase")]
236pub struct MetricStat {
237    pub metric: Metric,
238    pub period: i32,
239    pub stat: String,
240    #[serde(skip_serializing_if = "Option::is_none")]
241    pub unit: Option<String>,
242}
243
244#[derive(Debug, Clone, Builder, Serialize, Deserialize)]
245#[serde(rename_all = "PascalCase")]
246pub struct Metric {
247    #[serde(skip_serializing_if = "Option::is_none")]
248    pub namespace: Option<String>,
249    #[serde(skip_serializing_if = "Option::is_none")]
250    pub metric_name: Option<String>,
251    #[serde(default, skip_serializing_if = "Vec::is_empty")]
252    #[builder(default)]
253    pub dimensions: Vec<Dimension>,
254}
255
256#[derive(Debug, Clone, Builder, Serialize, Deserialize)]
257#[serde(rename_all = "PascalCase")]
258pub struct Dimension {
259    pub name: String,
260    pub value: String,
261}
262
263#[derive(Debug, Clone, Builder, Serialize)]
264#[serde(rename_all = "PascalCase")]
265pub struct ListMetricsRequest {
266    #[serde(skip_serializing_if = "Option::is_none")]
267    pub namespace: Option<String>,
268    #[serde(skip_serializing_if = "Option::is_none")]
269    pub metric_name: Option<String>,
270    #[serde(default, skip_serializing_if = "Vec::is_empty")]
271    #[builder(default)]
272    pub dimensions: Vec<DimensionFilter>,
273    #[serde(skip_serializing_if = "Option::is_none")]
274    pub next_token: Option<String>,
275    #[serde(skip_serializing_if = "Option::is_none")]
276    pub recently_active: Option<String>,
277    #[serde(skip_serializing_if = "Option::is_none")]
278    pub include_linked_accounts: Option<bool>,
279    #[serde(skip_serializing_if = "Option::is_none")]
280    pub owning_account: Option<String>,
281}
282
283#[derive(Debug, Clone, Builder, Serialize)]
284#[serde(rename_all = "PascalCase")]
285pub struct DimensionFilter {
286    pub name: String,
287    #[serde(skip_serializing_if = "Option::is_none")]
288    pub value: Option<String>,
289}
290
291#[derive(Debug, Clone, Deserialize)]
292#[serde(rename_all = "PascalCase")]
293pub struct GetMetricDataResponse {
294    #[serde(default)]
295    pub metric_data_results: Vec<MetricDataResult>,
296    #[serde(default)]
297    pub messages: Vec<MessageData>,
298    #[serde(default)]
299    pub next_token: Option<String>,
300}
301
302#[derive(Debug, Clone, Deserialize)]
303#[serde(rename_all = "PascalCase")]
304pub struct MetricDataResult {
305    pub id: String,
306    #[serde(default)]
307    pub label: Option<String>,
308    #[serde(default)]
309    pub status_code: Option<String>,
310    #[serde(default)]
311    pub timestamps: Vec<i64>,
312    #[serde(default)]
313    pub values: Vec<f64>,
314    #[serde(default)]
315    pub messages: Vec<MessageData>,
316}
317
318#[derive(Debug, Clone, Deserialize)]
319#[serde(rename_all = "PascalCase")]
320pub struct MessageData {
321    #[serde(default)]
322    pub code: Option<String>,
323    #[serde(default)]
324    pub value: Option<String>,
325}
326
327#[derive(Debug, Clone, Deserialize)]
328#[serde(rename_all = "PascalCase")]
329pub struct ListMetricsResponse {
330    #[serde(default)]
331    pub metrics: Vec<Metric>,
332    #[serde(default)]
333    pub next_token: Option<String>,
334    #[serde(default)]
335    pub owning_accounts: Vec<String>,
336}
337
338#[derive(Debug, Deserialize)]
339#[serde(rename_all = "PascalCase")]
340struct CloudWatchErrorResponse {
341    #[serde(rename = "__type", alias = "Code")]
342    code: String,
343    #[serde(rename = "Message", alias = "message")]
344    message: Option<String>,
345}
346
347#[cfg(test)]
348mod tests {
349    use super::*;
350    use serde_json::json;
351
352    #[test]
353    fn get_metric_data_request_matches_aws_json_shape() {
354        let request = GetMetricDataRequest::builder()
355            .start_time(1_637_061_900)
356            .end_time(1_637_074_500)
357            .metric_data_queries(vec![MetricDataQuery::builder()
358                .id("m1".to_string())
359                .label("CPU".to_string())
360                .metric_stat(
361                    MetricStat::builder()
362                        .metric(
363                            Metric::builder()
364                                .namespace("AWS/EC2".to_string())
365                                .metric_name("CPUUtilization".to_string())
366                                .dimensions(vec![Dimension::builder()
367                                    .name("InstanceId".to_string())
368                                    .value("i-123".to_string())
369                                    .build()])
370                                .build(),
371                        )
372                        .period(300)
373                        .stat("Average".to_string())
374                        .build(),
375                )
376                .return_data(true)
377                .build()])
378            .build();
379
380        let encoded = serde_json::to_value(request).unwrap();
381
382        assert_eq!(
383            encoded,
384            json!({
385                "StartTime": 1637061900,
386                "EndTime": 1637074500,
387                "MetricDataQueries": [{
388                    "Id": "m1",
389                    "Label": "CPU",
390                    "MetricStat": {
391                        "Metric": {
392                            "Namespace": "AWS/EC2",
393                            "MetricName": "CPUUtilization",
394                            "Dimensions": [{ "Name": "InstanceId", "Value": "i-123" }]
395                        },
396                        "Period": 300,
397                        "Stat": "Average"
398                    },
399                    "ReturnData": true
400                }]
401            })
402        );
403    }
404
405    #[test]
406    fn list_metrics_request_matches_aws_json_shape() {
407        let request = ListMetricsRequest::builder()
408            .namespace("AWS/EC2".to_string())
409            .dimensions(vec![DimensionFilter::builder()
410                .name("InstanceId".to_string())
411                .build()])
412            .build();
413
414        let encoded = serde_json::to_value(request).unwrap();
415
416        assert_eq!(
417            encoded,
418            json!({
419                "Namespace": "AWS/EC2",
420                "Dimensions": [{ "Name": "InstanceId" }]
421            })
422        );
423    }
424
425    #[test]
426    fn get_metric_data_response_parses_values() {
427        let response: GetMetricDataResponse = serde_json::from_value(json!({
428            "NextToken": "next",
429            "MetricDataResults": [{
430                "Id": "m1",
431                "Label": "CPU",
432                "StatusCode": "Complete",
433                "Timestamps": [1637074200],
434                "Values": [0.5]
435            }]
436        }))
437        .unwrap();
438
439        assert_eq!(response.next_token.as_deref(), Some("next"));
440        assert_eq!(response.metric_data_results[0].id, "m1");
441        assert_eq!(
442            response.metric_data_results[0].timestamps,
443            vec![1_637_074_200]
444        );
445        assert_eq!(response.metric_data_results[0].values, vec![0.5]);
446    }
447
448    #[test]
449    fn list_metrics_response_parses_metrics() {
450        let response: ListMetricsResponse = serde_json::from_value(json!({
451            "Metrics": [{
452                "Namespace": "AWS/EC2",
453                "MetricName": "CPUUtilization",
454                "Dimensions": [{ "Name": "InstanceId", "Value": "i-123" }]
455            }],
456            "OwningAccounts": ["111111111111"]
457        }))
458        .unwrap();
459
460        assert_eq!(response.metrics[0].namespace.as_deref(), Some("AWS/EC2"));
461        assert_eq!(
462            response.metrics[0].metric_name.as_deref(),
463            Some("CPUUtilization")
464        );
465        assert_eq!(response.owning_accounts, vec!["111111111111"]);
466    }
467}