aion_server/routing/
forwarder.rs1use std::net::SocketAddr;
13
14use async_trait::async_trait;
15use tonic::transport::{Channel, Endpoint};
16use tonic::{Request, Status};
17
18use aion_proto::generated::{self, workflow_service_client::WorkflowServiceClient};
19
20pub const FORWARD_HOPS_METADATA: &str = "x-aion-forward-hops";
25
26pub const MAX_FORWARD_HOPS: u32 = 2;
30
31#[derive(Clone, Debug)]
36pub enum ForwardRequest {
37 Start(generated::StartWorkflowRequest),
39 Signal(generated::SignalRequest),
41 Query(generated::QueryRequest),
43 Cancel(generated::CancelRequest),
45 Reopen(generated::ReopenRequest),
47 Pause(generated::PauseRequest),
49 Resume(generated::ResumeRequest),
51 MintNamespace(generated::MintNamespaceRequest),
54}
55
56#[derive(Clone, Debug)]
58pub enum ForwardReply {
59 Start(generated::StartWorkflowResponse),
61 Signal(generated::SignalResponse),
63 Query(generated::QueryResponse),
65 Cancel(generated::CancelResponse),
67 Reopen(generated::ReopenResponse),
69 Pause(generated::PauseResponse),
71 Resume(generated::ResumeResponse),
73 MintNamespace(generated::MintNamespaceResponse),
75}
76
77#[async_trait]
84pub trait RequestForwarder: Send + Sync {
85 async fn forward(
88 &self,
89 target: SocketAddr,
90 metadata: tonic::metadata::MetadataMap,
91 request: ForwardRequest,
92 ) -> Result<ForwardReply, Status>;
93}
94
95#[derive(Clone, Default)]
98pub struct GrpcRequestForwarder;
99
100impl GrpcRequestForwarder {
101 #[must_use]
103 pub const fn new() -> Self {
104 Self
105 }
106}
107
108#[must_use]
110pub fn current_hops(metadata: &tonic::metadata::MetadataMap) -> u32 {
111 metadata
112 .get(FORWARD_HOPS_METADATA)
113 .and_then(|value| value.to_str().ok())
114 .and_then(|value| value.parse().ok())
115 .unwrap_or(0)
116}
117
118fn stamp_next_hop(metadata: &mut tonic::metadata::MetadataMap) -> Result<(), Status> {
120 let next = current_hops(metadata)
121 .checked_add(1)
122 .ok_or_else(|| Status::internal("forward hop counter overflow"))?;
123 let value = tonic::metadata::MetadataValue::try_from(next.to_string())
124 .map_err(|_| Status::internal("invalid forward hop metadata value"))?;
125 metadata.insert(FORWARD_HOPS_METADATA, value);
126 Ok(())
127}
128
129async fn connect(target: SocketAddr) -> Result<WorkflowServiceClient<Channel>, Status> {
130 let uri = format!("http://{target}");
131 let endpoint = Endpoint::try_from(uri)
132 .map_err(|error| Status::unavailable(format!("invalid forward target: {error}")))?;
133 let channel = endpoint
134 .connect()
135 .await
136 .map_err(|error| Status::unavailable(format!("forward dial failed: {error}")))?;
137 Ok(WorkflowServiceClient::new(channel))
138}
139
140#[async_trait]
141impl RequestForwarder for GrpcRequestForwarder {
142 async fn forward(
143 &self,
144 target: SocketAddr,
145 mut metadata: tonic::metadata::MetadataMap,
146 request: ForwardRequest,
147 ) -> Result<ForwardReply, Status> {
148 stamp_next_hop(&mut metadata)?;
149 let mut client = connect(target).await?;
150 match request {
151 ForwardRequest::Start(message) => {
152 let mut outbound = Request::new(message);
153 *outbound.metadata_mut() = metadata;
154 client
155 .start_workflow(outbound)
156 .await
157 .map(|response| ForwardReply::Start(response.into_inner()))
158 }
159 ForwardRequest::Signal(message) => {
160 let mut outbound = Request::new(message);
161 *outbound.metadata_mut() = metadata;
162 client
163 .signal(outbound)
164 .await
165 .map(|response| ForwardReply::Signal(response.into_inner()))
166 }
167 ForwardRequest::Query(message) => {
168 let mut outbound = Request::new(message);
169 *outbound.metadata_mut() = metadata;
170 client
171 .query(outbound)
172 .await
173 .map(|response| ForwardReply::Query(response.into_inner()))
174 }
175 ForwardRequest::Cancel(message) => {
176 let mut outbound = Request::new(message);
177 *outbound.metadata_mut() = metadata;
178 client
179 .cancel(outbound)
180 .await
181 .map(|response| ForwardReply::Cancel(response.into_inner()))
182 }
183 ForwardRequest::Reopen(message) => {
184 let mut outbound = Request::new(message);
185 *outbound.metadata_mut() = metadata;
186 client
187 .reopen(outbound)
188 .await
189 .map(|response| ForwardReply::Reopen(response.into_inner()))
190 }
191 ForwardRequest::Pause(message) => {
192 let mut outbound = Request::new(message);
193 *outbound.metadata_mut() = metadata;
194 client
195 .pause(outbound)
196 .await
197 .map(|response| ForwardReply::Pause(response.into_inner()))
198 }
199 ForwardRequest::Resume(message) => {
200 let mut outbound = Request::new(message);
201 *outbound.metadata_mut() = metadata;
202 client
203 .resume(outbound)
204 .await
205 .map(|response| ForwardReply::Resume(response.into_inner()))
206 }
207 ForwardRequest::MintNamespace(message) => {
208 let mut outbound = Request::new(message);
209 *outbound.metadata_mut() = metadata;
210 client
211 .mint_namespace(outbound)
212 .await
213 .map(|response| ForwardReply::MintNamespace(response.into_inner()))
214 }
215 }
216 }
217}
218
219#[cfg(test)]
220mod tests {
221 use super::{FORWARD_HOPS_METADATA, MAX_FORWARD_HOPS, current_hops, stamp_next_hop};
222
223 #[test]
224 fn current_hops_defaults_to_zero_when_absent() {
225 let metadata = tonic::metadata::MetadataMap::new();
226 assert_eq!(current_hops(&metadata), 0);
227 }
228
229 #[test]
230 fn stamp_next_hop_increments_from_zero() -> Result<(), tonic::Status> {
231 let mut metadata = tonic::metadata::MetadataMap::new();
232 stamp_next_hop(&mut metadata)?;
233 assert_eq!(current_hops(&metadata), 1);
234 stamp_next_hop(&mut metadata)?;
235 assert_eq!(current_hops(&metadata), 2);
236 Ok(())
237 }
238
239 #[test]
240 fn malformed_hop_value_reads_as_zero() -> Result<(), tonic::Status> {
241 let mut metadata = tonic::metadata::MetadataMap::new();
242 metadata.insert(
243 FORWARD_HOPS_METADATA,
244 tonic::metadata::MetadataValue::try_from("not-a-number")
245 .map_err(|_| tonic::Status::internal("bad fixture"))?,
246 );
247 assert_eq!(current_hops(&metadata), 0);
248 Ok(())
249 }
250
251 #[test]
252 fn hop_cap_is_two() {
253 assert_eq!(MAX_FORWARD_HOPS, 2);
254 }
255}