use crate::aws::aws_request_utils::{AwsRequestBuilderExt, AwsSignConfig};
use crate::aws::AwsClientConfig;
use crate::aws::AwsClientConfigExt;
use alien_client_core::{ErrorData, Result};
use alien_error::{Context, ContextError, IntoAlienError};
use async_trait::async_trait;
use bon::Builder;
use reqwest::{Client, Method, StatusCode};
use serde::de::DeserializeOwned;
use serde::{Deserialize, Serialize};
use std::collections::HashMap;
#[cfg(feature = "test-utils")]
use mockall::automock;
#[cfg_attr(feature = "test-utils", automock)]
#[cfg_attr(target_arch = "wasm32", async_trait(?Send))]
#[cfg_attr(not(target_arch = "wasm32"), async_trait)]
pub trait SsmApi: Send + Sync + std::fmt::Debug {
async fn put_parameter(&self, request: PutParameterRequest) -> Result<PutParameterResponse>;
async fn get_parameter(&self, request: GetParameterRequest) -> Result<GetParameterResponse>;
async fn delete_parameter(&self, name: &str) -> Result<()>;
async fn get_parameters(&self, request: GetParametersRequest) -> Result<GetParametersResponse>;
async fn describe_parameters(
&self,
request: DescribeParametersRequest,
) -> Result<DescribeParametersResponse>;
async fn send_command(&self, request: SendCommandRequest) -> Result<SendCommandResponse>;
async fn get_command_invocation(
&self,
request: GetCommandInvocationRequest,
) -> Result<GetCommandInvocationResponse>;
async fn list_command_invocations(
&self,
request: ListCommandInvocationsRequest,
) -> Result<ListCommandInvocationsResponse>;
}
#[derive(Debug, Clone)]
pub struct SsmClient {
client: Client,
config: AwsClientConfig,
}
impl SsmClient {
pub fn new(client: Client, config: AwsClientConfig) -> Self {
Self { client, config }
}
fn sign_config(&self) -> AwsSignConfig {
AwsSignConfig {
service_name: "ssm".into(),
region: self.config.region.clone(),
credentials: self.config.get_credentials(),
signing_region: None,
}
}
fn get_base_url(&self) -> String {
if let Some(override_url) = self.config.get_service_endpoint_option("ssm") {
override_url.to_string()
} else {
format!("https://ssm.{}.amazonaws.com", self.config.region)
}
}
fn get_host(&self) -> String {
format!("ssm.{}.amazonaws.com", self.config.region)
}
async fn send_json<T: DeserializeOwned + Send + 'static>(
&self,
target: &str,
body: String,
operation: &str,
resource: &str,
) -> Result<T> {
let url = self.get_base_url();
let builder = self
.client
.request(Method::POST, &url)
.host(&self.get_host())
.header("X-Amz-Target", format!("AmazonSSM.{}", target))
.header("Content-Type", "application/x-amz-json-1.1")
.content_sha256(&body)
.body(body.clone());
let result =
crate::aws::aws_request_utils::sign_send_json(builder, &self.sign_config()).await;
Self::map_result(result, operation, resource, Some(&body))
}
async fn send_no_response(
&self,
target: &str,
body: String,
operation: &str,
resource: &str,
) -> Result<()> {
let url = self.get_base_url();
let builder = self
.client
.request(Method::POST, &url)
.host(&self.get_host())
.header("X-Amz-Target", format!("AmazonSSM.{}", target))
.header("Content-Type", "application/x-amz-json-1.1")
.content_sha256(&body)
.body(body.clone());
let result =
crate::aws::aws_request_utils::sign_send_no_response(builder, &self.sign_config())
.await;
Self::map_result(result, operation, resource, Some(&body))
}
fn map_result<T>(
result: Result<T>,
operation: &str,
resource: &str,
request_body: Option<&str>,
) -> Result<T> {
match result {
Ok(v) => Ok(v),
Err(e) => {
if let Some(ErrorData::HttpResponseError {
http_status,
http_response_text: Some(ref text),
..
}) = &e.error
{
let status = StatusCode::from_u16(*http_status)
.unwrap_or(StatusCode::INTERNAL_SERVER_ERROR);
if let Some(mapped) =
Self::map_ssm_error(status, text, operation, resource, request_body)
{
Err(e.context(mapped))
} else {
Err(e)
}
} else {
Err(e)
}
}
}
}
fn map_ssm_error(
status: StatusCode,
body: &str,
_operation: &str,
resource: &str,
request_body: Option<&str>,
) -> Option<ErrorData> {
let parsed: std::result::Result<SsmErrorResponse, _> = serde_json::from_str(body);
let (code, message) = match parsed {
Ok(e) => {
let c = e
.type_field
.or(e.code)
.unwrap_or_else(|| "UnknownErrorCode".into());
let m = e.message.unwrap_or_else(|| "Unknown error".into());
(c, m)
}
Err(_) => {
return None;
}
};
let error_code = code.split('#').last().unwrap_or(&code);
Some(match error_code {
"AccessDeniedException" | "UnauthorizedOperation" => ErrorData::RemoteAccessDenied {
resource_type: "SSM Resource".into(),
resource_name: resource.into(),
},
"ThrottlingException" | "TooManyUpdates" => ErrorData::RateLimitExceeded { message },
"InternalServerError" | "ServiceUnavailable" => {
ErrorData::RemoteServiceUnavailable { message }
}
"ParameterNotFound" | "ParameterVersionNotFound" => ErrorData::RemoteResourceNotFound {
resource_type: "SSM Parameter".into(),
resource_name: resource.into(),
},
"InvocationDoesNotExist" => ErrorData::RemoteResourceNotFound {
resource_type: "CommandInvocation".into(),
resource_name: resource.into(),
},
"ParameterAlreadyExists" => ErrorData::RemoteResourceConflict {
message,
resource_type: "SSM Parameter".into(),
resource_name: resource.into(),
},
"InvalidParameters" | "ParameterLimitExceeded" | "ParameterMaxVersionLimitExceeded" => {
ErrorData::InvalidInput {
message,
field_name: None,
}
}
"InvalidDocument" | "InvalidDocumentContent" | "InvalidDocumentVersion" => {
ErrorData::InvalidInput {
message,
field_name: Some("document_name".into()),
}
}
"InvalidInstanceId" | "InvalidTarget" => ErrorData::InvalidInput {
message,
field_name: Some("instance_id".into()),
},
"DuplicateInstanceId" => ErrorData::RemoteResourceConflict {
message,
resource_type: "Instance".into(),
resource_name: resource.into(),
},
"HierarchyLevelLimitExceededException" | "PoliciesLimitExceededException" => {
ErrorData::QuotaExceeded { message }
}
_ => match status {
StatusCode::NOT_FOUND => ErrorData::RemoteResourceNotFound {
resource_type: "SSM Resource".into(),
resource_name: resource.into(),
},
StatusCode::CONFLICT => ErrorData::RemoteResourceConflict {
message,
resource_type: "SSM Resource".into(),
resource_name: resource.into(),
},
StatusCode::FORBIDDEN | StatusCode::UNAUTHORIZED => ErrorData::RemoteAccessDenied {
resource_type: "SSM Resource".into(),
resource_name: resource.into(),
},
StatusCode::TOO_MANY_REQUESTS => ErrorData::RateLimitExceeded { message },
StatusCode::SERVICE_UNAVAILABLE
| StatusCode::BAD_GATEWAY
| StatusCode::GATEWAY_TIMEOUT => ErrorData::RemoteServiceUnavailable { message },
_ => ErrorData::HttpResponseError {
message: format!("SSM operation failed: {}", message),
url: "ssm.amazonaws.com".into(),
http_status: status.as_u16(),
http_response_text: Some(body.into()),
http_request_text: request_body.map(|s| s.to_string()),
},
},
})
}
}
#[cfg_attr(target_arch = "wasm32", async_trait(?Send))]
#[cfg_attr(not(target_arch = "wasm32"), async_trait)]
impl SsmApi for SsmClient {
async fn put_parameter(&self, request: PutParameterRequest) -> Result<PutParameterResponse> {
let body = serde_json::to_string(&request).into_alien_error().context(
ErrorData::SerializationError {
message: format!(
"Failed to serialize PutParameterRequest for '{}'",
request.name
),
},
)?;
self.send_json("PutParameter", body, "PutParameter", &request.name)
.await
}
async fn get_parameter(&self, request: GetParameterRequest) -> Result<GetParameterResponse> {
let body = serde_json::to_string(&request).into_alien_error().context(
ErrorData::SerializationError {
message: format!(
"Failed to serialize GetParameterRequest for '{}'",
request.name
),
},
)?;
self.send_json("GetParameter", body, "GetParameter", &request.name)
.await
}
async fn delete_parameter(&self, name: &str) -> Result<()> {
let request = DeleteParameterRequest {
name: name.to_string(),
};
let body = serde_json::to_string(&request).into_alien_error().context(
ErrorData::SerializationError {
message: format!("Failed to serialize DeleteParameterRequest for '{}'", name),
},
)?;
let _: serde_json::Value = self
.send_json("DeleteParameter", body, "DeleteParameter", name)
.await?;
Ok(())
}
async fn get_parameters(&self, request: GetParametersRequest) -> Result<GetParametersResponse> {
let names_str = request.names.join(", ");
let body = serde_json::to_string(&request).into_alien_error().context(
ErrorData::SerializationError {
message: format!(
"Failed to serialize GetParametersRequest for [{}]",
names_str
),
},
)?;
self.send_json("GetParameters", body, "GetParameters", &names_str)
.await
}
async fn describe_parameters(
&self,
request: DescribeParametersRequest,
) -> Result<DescribeParametersResponse> {
let body = serde_json::to_string(&request).into_alien_error().context(
ErrorData::SerializationError {
message: "Failed to serialize DescribeParametersRequest".to_string(),
},
)?;
self.send_json(
"DescribeParameters",
body,
"DescribeParameters",
"parameters",
)
.await
}
async fn send_command(&self, request: SendCommandRequest) -> Result<SendCommandResponse> {
let body = serde_json::to_string(&request).into_alien_error().context(
ErrorData::SerializationError {
message: format!(
"Failed to serialize SendCommandRequest for document '{}'",
request.document_name
),
},
)?;
self.send_json("SendCommand", body, "SendCommand", &request.document_name)
.await
}
async fn get_command_invocation(
&self,
request: GetCommandInvocationRequest,
) -> Result<GetCommandInvocationResponse> {
let resource = format!("{}:{}", request.command_id, request.instance_id);
let body = serde_json::to_string(&request).into_alien_error().context(
ErrorData::SerializationError {
message: format!(
"Failed to serialize GetCommandInvocationRequest for '{}'",
resource
),
},
)?;
self.send_json(
"GetCommandInvocation",
body,
"GetCommandInvocation",
&resource,
)
.await
}
async fn list_command_invocations(
&self,
request: ListCommandInvocationsRequest,
) -> Result<ListCommandInvocationsResponse> {
let resource = request.command_id.as_deref().unwrap_or("all");
let body = serde_json::to_string(&request).into_alien_error().context(
ErrorData::SerializationError {
message: "Failed to serialize ListCommandInvocationsRequest".to_string(),
},
)?;
self.send_json(
"ListCommandInvocations",
body,
"ListCommandInvocations",
resource,
)
.await
}
}
#[derive(Debug, Deserialize)]
struct SsmErrorResponse {
#[serde(rename = "__type")]
type_field: Option<String>,
#[serde(rename = "Code")]
code: Option<String>,
#[serde(rename = "Message", alias = "message")]
message: Option<String>,
}
#[derive(Debug, Clone, Serialize, Builder)]
#[serde(rename_all = "PascalCase")]
pub struct PutParameterRequest {
pub name: String,
pub value: String,
#[serde(rename = "Type")]
pub parameter_type: String,
#[serde(skip_serializing_if = "Option::is_none")]
pub description: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub overwrite: Option<bool>,
#[serde(skip_serializing_if = "Option::is_none")]
pub key_id: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub tags: Option<Vec<SsmTag>>,
#[serde(skip_serializing_if = "Option::is_none")]
pub tier: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub data_type: Option<String>,
}
#[derive(Debug, Deserialize)]
#[serde(rename_all = "PascalCase")]
pub struct PutParameterResponse {
pub version: Option<i64>,
pub tier: Option<String>,
}
#[derive(Debug, Clone, Serialize, Builder)]
#[serde(rename_all = "PascalCase")]
pub struct GetParameterRequest {
pub name: String,
#[serde(skip_serializing_if = "Option::is_none")]
pub with_decryption: Option<bool>,
}
#[derive(Debug, Deserialize)]
#[serde(rename_all = "PascalCase")]
pub struct GetParameterResponse {
pub parameter: Option<Parameter>,
}
#[derive(Debug, Clone, Serialize)]
#[serde(rename_all = "PascalCase")]
struct DeleteParameterRequest {
pub name: String,
}
#[derive(Debug, Clone, Serialize, Builder)]
#[serde(rename_all = "PascalCase")]
pub struct GetParametersRequest {
pub names: Vec<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub with_decryption: Option<bool>,
}
#[derive(Debug, Deserialize)]
#[serde(rename_all = "PascalCase")]
pub struct GetParametersResponse {
pub parameters: Option<Vec<Parameter>>,
pub invalid_parameters: Option<Vec<String>>,
}
#[derive(Debug, Clone, Serialize, Builder, Default)]
#[serde(rename_all = "PascalCase")]
pub struct DescribeParametersRequest {
#[serde(skip_serializing_if = "Option::is_none")]
pub parameter_filters: Option<Vec<ParameterStringFilter>>,
#[serde(skip_serializing_if = "Option::is_none")]
pub max_results: Option<i32>,
#[serde(skip_serializing_if = "Option::is_none")]
pub next_token: Option<String>,
}
#[derive(Debug, Deserialize)]
#[serde(rename_all = "PascalCase")]
pub struct DescribeParametersResponse {
pub parameters: Option<Vec<ParameterMetadata>>,
pub next_token: Option<String>,
}
#[derive(Debug, Clone, Deserialize)]
#[serde(rename_all = "PascalCase")]
pub struct Parameter {
pub name: Option<String>,
pub value: Option<String>,
#[serde(rename = "Type")]
pub parameter_type: Option<String>,
pub version: Option<i64>,
#[serde(rename = "ARN")]
pub arn: Option<String>,
pub last_modified_date: Option<f64>,
pub data_type: Option<String>,
}
#[derive(Debug, Clone, Deserialize)]
#[serde(rename_all = "PascalCase")]
pub struct ParameterMetadata {
pub name: Option<String>,
#[serde(rename = "Type")]
pub parameter_type: Option<String>,
pub key_id: Option<String>,
pub last_modified_date: Option<f64>,
pub last_modified_user: Option<String>,
pub description: Option<String>,
pub version: Option<i64>,
pub tier: Option<String>,
pub data_type: Option<String>,
}
#[derive(Debug, Clone, Serialize, Builder)]
#[serde(rename_all = "PascalCase")]
pub struct ParameterStringFilter {
pub key: String,
#[serde(skip_serializing_if = "Option::is_none")]
pub option: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub values: Option<Vec<String>>,
}
#[derive(Debug, Clone, Serialize, Deserialize, Builder)]
#[serde(rename_all = "PascalCase")]
pub struct SsmTag {
pub key: String,
pub value: String,
}
#[derive(Debug, Clone, Serialize, Builder)]
#[serde(rename_all = "PascalCase")]
pub struct SendCommandRequest {
pub document_name: String,
#[serde(skip_serializing_if = "Option::is_none")]
pub instance_ids: Option<Vec<String>>,
#[serde(skip_serializing_if = "Option::is_none")]
pub targets: Option<Vec<Target>>,
#[serde(skip_serializing_if = "Option::is_none")]
pub parameters: Option<HashMap<String, Vec<String>>>,
#[serde(skip_serializing_if = "Option::is_none")]
pub comment: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub timeout_seconds: Option<i32>,
#[serde(skip_serializing_if = "Option::is_none")]
pub output_s3_bucket_name: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub output_s3_key_prefix: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub cloud_watch_output_config: Option<CloudWatchOutputConfig>,
}
#[derive(Debug, Deserialize)]
#[serde(rename_all = "PascalCase")]
pub struct SendCommandResponse {
pub command: Option<Command>,
}
#[derive(Debug, Clone, Serialize, Builder)]
#[serde(rename_all = "PascalCase")]
pub struct GetCommandInvocationRequest {
pub command_id: String,
pub instance_id: String,
#[serde(skip_serializing_if = "Option::is_none")]
pub plugin_name: Option<String>,
}
#[derive(Debug, Deserialize)]
#[serde(rename_all = "PascalCase")]
pub struct GetCommandInvocationResponse {
pub command_id: Option<String>,
pub instance_id: Option<String>,
pub comment: Option<String>,
pub document_name: Option<String>,
pub document_version: Option<String>,
pub plugin_name: Option<String>,
pub response_code: Option<i32>,
pub execution_start_date_time: Option<String>,
pub execution_elapsed_time: Option<String>,
pub execution_end_date_time: Option<String>,
pub status: Option<String>,
pub status_details: Option<String>,
pub standard_output_content: Option<String>,
pub standard_output_url: Option<String>,
pub standard_error_content: Option<String>,
pub standard_error_url: Option<String>,
pub cloud_watch_output_config: Option<CloudWatchOutputConfig>,
}
#[derive(Debug, Clone, Serialize, Builder, Default)]
#[serde(rename_all = "PascalCase")]
pub struct ListCommandInvocationsRequest {
#[serde(skip_serializing_if = "Option::is_none")]
pub command_id: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub instance_id: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub max_results: Option<i32>,
#[serde(skip_serializing_if = "Option::is_none")]
pub next_token: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub filters: Option<Vec<CommandFilter>>,
#[serde(skip_serializing_if = "Option::is_none")]
pub details: Option<bool>,
}
#[derive(Debug, Deserialize)]
#[serde(rename_all = "PascalCase")]
pub struct ListCommandInvocationsResponse {
pub command_invocations: Option<Vec<CommandInvocation>>,
pub next_token: Option<String>,
}
#[derive(Debug, Clone, Deserialize)]
#[serde(rename_all = "PascalCase")]
pub struct Command {
pub command_id: Option<String>,
pub document_name: Option<String>,
pub document_version: Option<String>,
pub comment: Option<String>,
pub expires_after: Option<f64>,
pub parameters: Option<HashMap<String, Vec<String>>>,
pub instance_ids: Option<Vec<String>>,
pub targets: Option<Vec<Target>>,
pub requested_date_time: Option<f64>,
pub status: Option<String>,
pub status_details: Option<String>,
pub output_s3_bucket_name: Option<String>,
pub output_s3_key_prefix: Option<String>,
pub max_concurrency: Option<String>,
pub max_errors: Option<String>,
pub target_count: Option<i32>,
pub completed_count: Option<i32>,
pub error_count: Option<i32>,
pub delivery_timed_out_count: Option<i32>,
}
#[derive(Debug, Clone, Deserialize)]
#[serde(rename_all = "PascalCase")]
pub struct CommandInvocation {
pub command_id: Option<String>,
pub instance_id: Option<String>,
pub instance_name: Option<String>,
pub comment: Option<String>,
pub document_name: Option<String>,
pub document_version: Option<String>,
pub requested_date_time: Option<f64>,
pub status: Option<String>,
pub status_details: Option<String>,
pub trace_output: Option<String>,
pub standard_output_url: Option<String>,
pub standard_error_url: Option<String>,
pub command_plugins: Option<Vec<CommandPlugin>>,
}
#[derive(Debug, Clone, Deserialize)]
#[serde(rename_all = "PascalCase")]
pub struct CommandPlugin {
pub name: Option<String>,
pub status: Option<String>,
pub status_details: Option<String>,
pub response_code: Option<i32>,
pub response_start_date_time: Option<f64>,
pub response_finish_date_time: Option<f64>,
pub output: Option<String>,
pub standard_output_url: Option<String>,
pub standard_error_url: Option<String>,
pub output_s3_bucket_name: Option<String>,
pub output_s3_key_prefix: Option<String>,
}
#[derive(Debug, Clone, Serialize, Deserialize, Builder)]
#[serde(rename_all = "PascalCase")]
pub struct Target {
pub key: Option<String>,
pub values: Option<Vec<String>>,
}
#[derive(Debug, Clone, Serialize, Builder)]
pub struct CommandFilter {
pub key: String,
pub value: String,
}
#[derive(Debug, Clone, Serialize, Deserialize, Builder)]
#[serde(rename_all = "PascalCase")]
pub struct CloudWatchOutputConfig {
pub cloud_watch_log_group_name: Option<String>,
pub cloud_watch_output_enabled: Option<bool>,
}