use std::net::SocketAddr;
use async_trait::async_trait;
use tonic::transport::{Channel, Endpoint};
use tonic::{Request, Status};
use aion_proto::generated::{self, workflow_service_client::WorkflowServiceClient};
pub const FORWARD_HOPS_METADATA: &str = "x-aion-forward-hops";
pub const MAX_FORWARD_HOPS: u32 = 2;
#[derive(Clone, Debug)]
pub enum ForwardRequest {
Start(generated::StartWorkflowRequest),
Signal(generated::SignalRequest),
Query(generated::QueryRequest),
Cancel(generated::CancelRequest),
Reopen(generated::ReopenRequest),
Pause(generated::PauseRequest),
Resume(generated::ResumeRequest),
}
#[derive(Clone, Debug)]
pub enum ForwardReply {
Start(generated::StartWorkflowResponse),
Signal(generated::SignalResponse),
Query(generated::QueryResponse),
Cancel(generated::CancelResponse),
Reopen(generated::ReopenResponse),
Pause(generated::PauseResponse),
Resume(generated::ResumeResponse),
}
#[async_trait]
pub trait RequestForwarder: Send + Sync {
async fn forward(
&self,
target: SocketAddr,
metadata: tonic::metadata::MetadataMap,
request: ForwardRequest,
) -> Result<ForwardReply, Status>;
}
#[derive(Clone, Default)]
pub struct GrpcRequestForwarder;
impl GrpcRequestForwarder {
#[must_use]
pub const fn new() -> Self {
Self
}
}
#[must_use]
pub fn current_hops(metadata: &tonic::metadata::MetadataMap) -> u32 {
metadata
.get(FORWARD_HOPS_METADATA)
.and_then(|value| value.to_str().ok())
.and_then(|value| value.parse().ok())
.unwrap_or(0)
}
fn stamp_next_hop(metadata: &mut tonic::metadata::MetadataMap) -> Result<(), Status> {
let next = current_hops(metadata)
.checked_add(1)
.ok_or_else(|| Status::internal("forward hop counter overflow"))?;
let value = tonic::metadata::MetadataValue::try_from(next.to_string())
.map_err(|_| Status::internal("invalid forward hop metadata value"))?;
metadata.insert(FORWARD_HOPS_METADATA, value);
Ok(())
}
async fn connect(target: SocketAddr) -> Result<WorkflowServiceClient<Channel>, Status> {
let uri = format!("http://{target}");
let endpoint = Endpoint::try_from(uri)
.map_err(|error| Status::unavailable(format!("invalid forward target: {error}")))?;
let channel = endpoint
.connect()
.await
.map_err(|error| Status::unavailable(format!("forward dial failed: {error}")))?;
Ok(WorkflowServiceClient::new(channel))
}
#[async_trait]
impl RequestForwarder for GrpcRequestForwarder {
async fn forward(
&self,
target: SocketAddr,
mut metadata: tonic::metadata::MetadataMap,
request: ForwardRequest,
) -> Result<ForwardReply, Status> {
stamp_next_hop(&mut metadata)?;
let mut client = connect(target).await?;
match request {
ForwardRequest::Start(message) => {
let mut outbound = Request::new(message);
*outbound.metadata_mut() = metadata;
client
.start_workflow(outbound)
.await
.map(|response| ForwardReply::Start(response.into_inner()))
}
ForwardRequest::Signal(message) => {
let mut outbound = Request::new(message);
*outbound.metadata_mut() = metadata;
client
.signal(outbound)
.await
.map(|response| ForwardReply::Signal(response.into_inner()))
}
ForwardRequest::Query(message) => {
let mut outbound = Request::new(message);
*outbound.metadata_mut() = metadata;
client
.query(outbound)
.await
.map(|response| ForwardReply::Query(response.into_inner()))
}
ForwardRequest::Cancel(message) => {
let mut outbound = Request::new(message);
*outbound.metadata_mut() = metadata;
client
.cancel(outbound)
.await
.map(|response| ForwardReply::Cancel(response.into_inner()))
}
ForwardRequest::Reopen(message) => {
let mut outbound = Request::new(message);
*outbound.metadata_mut() = metadata;
client
.reopen(outbound)
.await
.map(|response| ForwardReply::Reopen(response.into_inner()))
}
ForwardRequest::Pause(message) => {
let mut outbound = Request::new(message);
*outbound.metadata_mut() = metadata;
client
.pause(outbound)
.await
.map(|response| ForwardReply::Pause(response.into_inner()))
}
ForwardRequest::Resume(message) => {
let mut outbound = Request::new(message);
*outbound.metadata_mut() = metadata;
client
.resume(outbound)
.await
.map(|response| ForwardReply::Resume(response.into_inner()))
}
}
}
}
#[cfg(test)]
mod tests {
use super::{FORWARD_HOPS_METADATA, MAX_FORWARD_HOPS, current_hops, stamp_next_hop};
#[test]
fn current_hops_defaults_to_zero_when_absent() {
let metadata = tonic::metadata::MetadataMap::new();
assert_eq!(current_hops(&metadata), 0);
}
#[test]
fn stamp_next_hop_increments_from_zero() -> Result<(), tonic::Status> {
let mut metadata = tonic::metadata::MetadataMap::new();
stamp_next_hop(&mut metadata)?;
assert_eq!(current_hops(&metadata), 1);
stamp_next_hop(&mut metadata)?;
assert_eq!(current_hops(&metadata), 2);
Ok(())
}
#[test]
fn malformed_hop_value_reads_as_zero() -> Result<(), tonic::Status> {
let mut metadata = tonic::metadata::MetadataMap::new();
metadata.insert(
FORWARD_HOPS_METADATA,
tonic::metadata::MetadataValue::try_from("not-a-number")
.map_err(|_| tonic::Status::internal("bad fixture"))?,
);
assert_eq!(current_hops(&metadata), 0);
Ok(())
}
#[test]
fn hop_cap_is_two() {
assert_eq!(MAX_FORWARD_HOPS, 2);
}
}