mod auth;
mod serde_errors;
mod serialization_checks;
use auth::validate_auth;
use serde_errors::{deserialize, serialize};
use serialization_checks::pre_check_serialization_types;
use crate::Database;
use crate::net::bind_exclusive;
use axum::{
Router,
body::Body,
extract::{DefaultBodyLimit, State},
http::{HeaderMap, HeaderName, HeaderValue, Method, StatusCode, Uri, header::SERVER},
response::Response,
routing::any,
};
use std::net::{IpAddr, Ipv4Addr, Ipv6Addr, SocketAddr, TcpStream};
use std::time::Duration;
use tower_http::set_header::SetResponseHeaderLayer;
const AWS_REGION: &str = "us-east-1";
const CONTENT_TYPE: &str = "application/x-amz-json-1.0";
const TARGET_PREFIX: &str = "DynamoDB_20120810.";
const STREAMS_TARGET_PREFIX: &str = "DynamoDBStreams_20120810.";
fn check_port_available(addr: SocketAddr) -> Result<(), String> {
let timeout = Duration::from_millis(100);
let port = addr.port();
let cross = SocketAddr::new(
match addr.ip() {
IpAddr::V4(ip) if ip.is_loopback() => IpAddr::V4(Ipv4Addr::UNSPECIFIED),
IpAddr::V4(_) => IpAddr::V4(Ipv4Addr::LOCALHOST),
IpAddr::V6(ip) if ip.is_loopback() => IpAddr::V6(Ipv6Addr::UNSPECIFIED),
IpAddr::V6(_) => IpAddr::V6(Ipv6Addr::LOCALHOST),
},
port,
);
for probe in [addr, cross] {
if TcpStream::connect_timeout(&probe, timeout).is_ok() {
return Err(format!(
"port {port} is already in use (detected listener on {probe})"
));
}
}
Ok(())
}
pub async fn start(host: &str, port: u16, db: Database) -> Result<(), String> {
let addr: SocketAddr = format!("{host}:{port}")
.parse()
.map_err(|e| format!("invalid address {host}:{port}: {e}"))?;
check_port_available(addr)?;
let std_listener = bind_exclusive(addr)?;
let listener = tokio::net::TcpListener::from_std(std_listener)
.map_err(|e| format!("failed to create async listener: {e}"))?;
let app = build_router(db);
eprintln!("Dynoxide listening on http://{addr}");
axum::serve(listener, app)
.with_graceful_shutdown(shutdown_signal())
.await
.map_err(|e| format!("server failed: {e}"))
}
pub async fn serve_on(listener: tokio::net::TcpListener, db: Database) {
let app = build_router(db);
axum::serve(listener, app).await.unwrap();
}
const MAX_BODY_SIZE: usize = 16 * 1024 * 1024;
fn build_router(db: Database) -> Router {
Router::new()
.route("/", any(handle_root))
.fallback(handle_fallback)
.layer(DefaultBodyLimit::max(MAX_BODY_SIZE))
.layer(SetResponseHeaderLayer::overriding(
SERVER,
HeaderValue::from_static(concat!("Dynoxide/", env!("CARGO_PKG_VERSION"))),
))
.layer(SetResponseHeaderLayer::overriding(
HeaderName::from_static("x-dynoxide-version"),
HeaderValue::from_static(env!("CARGO_PKG_VERSION")),
))
.with_state(db)
}
const NOT_FOUND_BODY: &str = "<UnknownOperationException/>\n";
async fn handle_root(
method: Method,
uri: Uri,
State(db): State<Database>,
headers: HeaderMap,
body: String,
) -> Response {
let has_origin = headers.get("origin").is_some();
let mut resp = match method {
Method::GET => {
let body_str = format!("healthy: dynamodb.{AWS_REGION}.amazonaws.com ");
dynamo_response_raw(StatusCode::OK, &body_str)
}
Method::OPTIONS if has_origin => {
let mut r = Response::builder()
.status(StatusCode::OK)
.body(Body::from(""))
.unwrap();
add_dynamo_headers(&mut r, b"");
r.headers_mut().insert(
HeaderName::from_static("content-length"),
HeaderValue::from_static("0"),
);
r.headers_mut().insert(
HeaderName::from_static("access-control-allow-origin"),
HeaderValue::from_static("*"),
);
r.headers_mut().insert(
HeaderName::from_static("access-control-max-age"),
HeaderValue::from_static("172800"),
);
if let Some(req_headers) = headers.get("access-control-request-headers") {
r.headers_mut().insert(
HeaderName::from_static("access-control-allow-headers"),
req_headers.clone(),
);
}
if let Some(req_method) = headers.get("access-control-request-method") {
r.headers_mut().insert(
HeaderName::from_static("access-control-allow-methods"),
req_method.clone(),
);
}
return r;
}
Method::POST => handle_request(uri, State(db), headers.clone(), body).await,
_ => {
dynamo_response_raw(StatusCode::NOT_FOUND, NOT_FOUND_BODY)
}
};
if has_origin {
resp.headers_mut().insert(
HeaderName::from_static("access-control-allow-origin"),
HeaderValue::from_static("*"),
);
}
resp
}
async fn handle_fallback() -> Response {
dynamo_response_raw(StatusCode::NOT_FOUND, NOT_FOUND_BODY)
}
async fn shutdown_signal() {
#[cfg(unix)]
{
use tokio::signal::unix::{SignalKind, signal};
let mut sigterm =
signal(SignalKind::terminate()).expect("failed to install SIGTERM handler");
tokio::select! {
_ = tokio::signal::ctrl_c() => {},
_ = sigterm.recv() => {},
}
}
#[cfg(not(unix))]
{
tokio::signal::ctrl_c()
.await
.expect("failed to install CTRL+C handler");
}
eprintln!("\nShutting down...");
}
async fn handle_request(
uri: Uri,
State(db): State<Database>,
headers: HeaderMap,
body: String,
) -> Response {
let raw_ct = headers
.get("content-type")
.and_then(|v| v.to_str().ok())
.unwrap_or("");
let base_ct = raw_ct.split(';').next().unwrap_or("").trim();
let is_amz_json = base_ct.eq_ignore_ascii_case(CONTENT_TYPE);
let is_plain_json = base_ct.eq_ignore_ascii_case("application/json");
if !is_amz_json && !is_plain_json && (!body.is_empty() || !raw_ct.is_empty()) {
return dynamo_response_raw(StatusCode::NOT_FOUND, NOT_FOUND_BODY);
}
let response_ct = if is_amz_json {
CONTENT_TYPE
} else {
"application/json"
};
if !body.is_empty() && serde_json::from_str::<serde_json::Value>(&body).is_err() {
return serialization_exception_bare(response_ct);
}
let target = match headers.get("x-amz-target").and_then(|v| v.to_str().ok()) {
Some(t) => t,
None => {
return unknown_operation_response(response_ct);
}
};
let operation = target
.strip_prefix(TARGET_PREFIX)
.or_else(|| target.strip_prefix(STREAMS_TARGET_PREFIX));
let operation = match operation {
Some(op) if crate::dynamo_ops::is_known_operation(op) => op,
_ => {
return unknown_operation_response(response_ct);
}
};
if let Some(auth_error) = validate_auth(&headers, &uri, response_ct) {
return auth_error;
}
if body.is_empty() {
return serialization_exception_bare(response_ct);
}
tracing::debug!(operation, body_len = body.len(), "request");
tracing::trace!(operation, body = %body, "request body");
match dispatch(&db, operation, &body) {
Ok(json) => {
tracing::debug!(operation, body_len = json.len(), "response");
tracing::trace!(operation, body = %json, "response body");
dynamo_response(StatusCode::OK, response_ct, json)
}
Err(e) => {
let status = StatusCode::from_u16(e.status_code()).unwrap_or(StatusCode::BAD_REQUEST);
let json = e.to_json();
tracing::warn!(operation, status = %status, "error response");
tracing::trace!(operation, body = %json, "error response body");
dynamo_response(status, response_ct, json)
}
}
}
fn serialization_exception_bare(content_type: &str) -> Response {
let body = r#"{"__type":"com.amazon.coral.service#SerializationException"}"#.to_string();
dynamo_response(StatusCode::BAD_REQUEST, content_type, body)
}
fn unknown_operation_response(content_type: &str) -> Response {
let body = r#"{"__type":"com.amazon.coral.service#UnknownOperationException"}"#.to_string();
dynamo_response(StatusCode::BAD_REQUEST, content_type, body)
}
fn dispatch(db: &Database, operation: &str, body: &str) -> crate::Result<String> {
pre_check_serialization_types(operation, body)?;
match operation {
"CreateTable" => {
let req = deserialize(body)?;
let resp = db.create_table(req)?;
serialize(&resp)
}
"DeleteTable" => {
let req = deserialize(body)?;
let resp = db.delete_table(req)?;
serialize(&resp)
}
"DescribeTable" => {
let req = deserialize(body)?;
let resp = db.describe_table(req)?;
serialize(&resp)
}
"ListTables" => {
let req = deserialize(body)?;
let resp = db.list_tables(req)?;
serialize(&resp)
}
"UpdateTable" => {
let req = deserialize(body)?;
let resp = db.update_table(req)?;
serialize(&resp)
}
"PutItem" => {
let req = deserialize(body)?;
let resp = db.put_item(req)?;
serialize(&resp)
}
"GetItem" => {
let req = deserialize(body)?;
let resp = db.get_item(req)?;
serialize(&resp)
}
"DeleteItem" => {
let req = deserialize(body)?;
let resp = db.delete_item(req)?;
serialize(&resp)
}
"UpdateItem" => {
let req = deserialize(body)?;
let resp = db.update_item(req)?;
serialize(&resp)
}
"Query" => {
let req = deserialize(body)?;
let resp = db.query(req)?;
serialize(&resp)
}
"Scan" => {
let req = deserialize(body)?;
let resp = db.scan(req)?;
serialize(&resp)
}
"BatchGetItem" => {
let req = deserialize(body)?;
let resp = db.batch_get_item(req)?;
serialize(&resp)
}
"BatchWriteItem" => {
let req = deserialize(body)?;
let resp = db.batch_write_item(req)?;
serialize(&resp)
}
"TransactWriteItems" => {
let req = deserialize(body)?;
let resp = db.transact_write_items(req)?;
serialize(&resp)
}
"TransactGetItems" => {
let req = deserialize(body)?;
let resp = db.transact_get_items(req)?;
serialize(&resp)
}
"ListStreams" => {
let req = deserialize(body)?;
let resp = db.list_streams(req)?;
serialize(&resp)
}
"DescribeStream" => {
let req = deserialize(body)?;
let resp = db.describe_stream(req)?;
serialize(&resp)
}
"GetShardIterator" => {
let req = deserialize(body)?;
let resp = db.get_shard_iterator(req)?;
serialize(&resp)
}
"GetRecords" => {
let req = deserialize(body)?;
let resp = db.get_records(req)?;
serialize(&resp)
}
"UpdateTimeToLive" => {
let req = deserialize(body)?;
let resp = db.update_time_to_live(req)?;
serialize(&resp)
}
"DescribeTimeToLive" => {
let req = deserialize(body)?;
let resp = db.describe_time_to_live(req)?;
serialize(&resp)
}
"ExecuteStatement" => {
let req = deserialize(body)?;
let resp = db.execute_statement(req)?;
serialize(&resp)
}
"ExecuteTransaction" => {
let req = deserialize(body)?;
let resp = db.execute_transaction(req)?;
serialize(&resp)
}
"BatchExecuteStatement" => {
let req = deserialize(body)?;
let resp = db.batch_execute_statement(req)?;
serialize(&resp)
}
"TagResource" => {
let req = deserialize(body)?;
let resp = db.tag_resource(req)?;
serialize(&resp)
}
"UntagResource" => {
let req = deserialize(body)?;
let resp = db.untag_resource(req)?;
serialize(&resp)
}
"ListTagsOfResource" => {
let req = deserialize(body)?;
let resp = db.list_tags_of_resource(req)?;
serialize(&resp)
}
_ => {
Err(crate::DynoxideError::SerializationException(
"UnknownOperationException".to_string(),
))
}
}
}
fn generate_request_id() -> String {
use uuid::Uuid;
let u1 = Uuid::now_v7();
let u2 = Uuid::now_v7();
let hex = format!(
"{}{}",
u1.as_simple().to_string().to_ascii_uppercase(),
u2.as_simple().to_string().to_ascii_uppercase()
);
hex[..52].to_string()
}
fn compute_crc32(body: &[u8]) -> String {
crc32fast::hash(body).to_string()
}
fn add_dynamo_headers(response: &mut Response, body_bytes: &[u8]) {
let headers = response.headers_mut();
headers.insert(
HeaderName::from_static("x-amzn-requestid"),
HeaderValue::from_str(&generate_request_id()).unwrap(),
);
headers.insert(
HeaderName::from_static("x-amz-crc32"),
HeaderValue::from_str(&compute_crc32(body_bytes)).unwrap(),
);
headers.insert(
HeaderName::from_static("content-length"),
HeaderValue::from_str(&body_bytes.len().to_string()).unwrap(),
);
}
fn dynamo_response(status: StatusCode, content_type: &str, body_str: String) -> Response {
let body_bytes = body_str.as_bytes();
let mut resp = Response::builder()
.status(status)
.header("content-type", content_type)
.body(Body::from(body_str.clone()))
.unwrap();
add_dynamo_headers(&mut resp, body_bytes);
resp
}
fn dynamo_response_raw(status: StatusCode, body_str: &str) -> Response {
let body_bytes = body_str.as_bytes();
let mut resp = Response::builder()
.status(status)
.body(Body::from(body_str.to_string()))
.unwrap();
add_dynamo_headers(&mut resp, body_bytes);
resp
}