use std::net::SocketAddr;
use axum::http::HeaderMap;
use tonic::metadata::MetadataMap;
use aion_proto::{ProtoRenameRequest, ProtoRenameResponse, WireError};
use crate::ServerState;
use crate::api::grpc::convert::{decode_rename_response, encode_rename_request};
use crate::routing::{
ForwardReply, ForwardRequest, RequestForwarder, RouteDecision, ShardDirectory, not_owner_wire,
owner_refusal, route_mutation,
};
use super::auth::caller_credentials_metadata;
pub(crate) async fn forward_rename(
state: &ServerState,
headers: &HeaderMap,
request: &ProtoRenameRequest,
) -> Result<Option<ProtoRenameResponse>, WireError> {
let Some(cluster_store) = state.cluster_store() else {
return Ok(None);
};
let Some(proto_id) = request.workflow_id.clone() else {
return Ok(None);
};
let Ok(workflow_id) = aion_core::WorkflowId::try_from(proto_id) else {
return Ok(None);
};
let directory = state
.shard_directory()
.map(|directory| directory.as_ref() as &dyn ShardDirectory);
match route_mutation(Some(cluster_store.as_ref()), directory, &workflow_id) {
RouteDecision::Local => Ok(None),
RouteDecision::NotOwner { shard } => Err(not_owner_wire(shard)),
RouteDecision::Forward { owner, shard } => {
let Some(target) = owner.grpc_addr else {
return Err(not_owner_wire(shard));
};
let Some(forwarder) = state.request_forwarder() else {
return Err(not_owner_wire(shard));
};
let metadata = caller_credentials_metadata(headers)?;
relay_rename(forwarder.as_ref(), target, metadata, request, shard)
.await
.map(Some)
}
}
}
async fn relay_rename(
forwarder: &dyn RequestForwarder,
target: SocketAddr,
metadata: MetadataMap,
request: &ProtoRenameRequest,
shard: usize,
) -> Result<ProtoRenameResponse, WireError> {
match forwarder
.forward(
target,
metadata,
ForwardRequest::Rename(encode_rename_request(request.clone())),
)
.await
{
Ok(ForwardReply::Rename(reply)) => Ok(decode_rename_response(reply)),
Ok(_) => Err(WireError::backend(
"the forwarder returned a reply for a different request",
)),
Err(status) => Err(owner_refusal(&status).unwrap_or_else(|| {
tracing::warn!(
shard,
%target,
grpc_code = %status.code(),
grpc_message = status.message(),
"forwarded rename never reached the shard owner; answering NotOwner so the \
caller re-resolves"
);
not_owner_wire(shard)
})),
}
}
#[cfg(test)]
mod tests {
use std::net::SocketAddr;
use std::sync::Arc;
use aion::EngineBuilder;
use aion_proto::{ProtoRenameRequest, ProtoWorkflowId, WireError};
use aion_store::{EventStore, InMemoryStore};
use axum::http::HeaderMap;
use tonic::metadata::MetadataMap;
use super::super::test_support::{NAMESPACE, runtime_config, server_state};
use super::{ForwardReply, ForwardRequest, RequestForwarder, forward_rename, relay_rename};
use crate::NamespaceResolver;
use crate::config::NamespaceMode;
use crate::namespace::{StaticScheduleNamespaces, StaticWorkflowNamespaces};
use crate::test_support::{EngineUnderTest, StateUnderTest};
async fn non_clustered_state() -> Result<StateUnderTest, Box<dyn std::error::Error>> {
let store: Arc<dyn EventStore> = Arc::new(InMemoryStore::default());
let engine = EngineUnderTest::new(Arc::new(
EngineBuilder::new()
.stop_drain_timeout(std::time::Duration::from_secs(5))
.store_arc(store)
.in_memory_visibility()
.scheduler_threads(1)
.build()
.await?,
));
let resolver = NamespaceResolver::from_parts(
NamespaceMode::SharedEngine,
Some(engine.handle()),
Arc::new(StaticWorkflowNamespaces::default()),
Arc::new(StaticScheduleNamespaces::default()),
);
server_state(engine, resolver, runtime_config()).await
}
fn rename_request(workflow_id: Option<ProtoWorkflowId>) -> ProtoRenameRequest {
ProtoRenameRequest {
namespace: NAMESPACE.to_owned(),
workflow_id,
run_id: None,
display_name: String::from("Nightly settlement"),
}
}
#[tokio::test]
async fn a_non_clustered_boot_always_serves_the_rename_locally()
-> Result<(), Box<dyn std::error::Error>> {
let state = non_clustered_state().await?;
let workflow_id = ProtoWorkflowId::from(aion_core::WorkflowId::new_v4());
let routed = forward_rename(
&state,
&HeaderMap::new(),
&rename_request(Some(workflow_id)),
)
.await?;
assert!(
routed.is_none(),
"with no cluster store there is no owner but this node, so nothing may be forwarded"
);
state.shutdown()?;
Ok(())
}
#[tokio::test]
async fn a_target_routing_cannot_read_is_left_to_the_handler()
-> Result<(), Box<dyn std::error::Error>> {
let state = non_clustered_state().await?;
assert!(
forward_rename(&state, &HeaderMap::new(), &rename_request(None))
.await?
.is_none()
);
assert!(
forward_rename(
&state,
&HeaderMap::new(),
&rename_request(Some(ProtoWorkflowId {
uuid: String::from("not-a-uuid"),
})),
)
.await?
.is_none()
);
state.shutdown()?;
Ok(())
}
struct ScriptedForwarder(Result<ForwardReply, tonic::Status>);
#[async_trait::async_trait]
impl RequestForwarder for ScriptedForwarder {
async fn forward(
&self,
_target: SocketAddr,
_metadata: MetadataMap,
_request: ForwardRequest,
) -> Result<ForwardReply, tonic::Status> {
match &self.0 {
Ok(reply) => Ok(reply.clone()),
Err(status) => Err(tonic::Status::with_details(
status.code(),
status.message(),
status.details().to_vec().into(),
)),
}
}
}
fn target() -> SocketAddr {
SocketAddr::from(([127, 0, 0, 1], 7233))
}
const SHARD: usize = 7;
async fn relayed(answer: Result<ForwardReply, tonic::Status>) -> Result<(), WireError> {
relay_rename(
&ScriptedForwarder(answer),
target(),
MetadataMap::new(),
&rename_request(Some(ProtoWorkflowId::from(aion_core::WorkflowId::new_v4()))),
SHARD,
)
.await
.map(|_response| ())
}
#[tokio::test]
async fn a_forwarded_refusal_reaches_the_caller_as_the_owners_own_error()
-> Result<(), Box<dyn std::error::Error>> {
let refusals = [
WireError::invalid_input("display_name must not be blank"),
WireError::not_found("workflow 7 has no recorded history"),
WireError::namespace_denied("tenant-b is not visible to this caller"),
WireError::invalid_state(
"run is Running but is not resident on this node; retry once it is resident",
),
];
for refusal in refusals {
let status = crate::api::grpc::status_from_wire_error(refusal.clone());
let relayed = relayed(Err(status))
.await
.err()
.ok_or("a refused forward must not be reported as success")?;
assert_eq!(
relayed.code, refusal.code,
"the owner's code must reach the caller, not NotOwner: {relayed:?}"
);
assert_eq!(relayed.message, refusal.message);
assert_ne!(
relayed.code,
aion_proto::WireErrorCode::NotOwner,
"flattening the owner's answer is the defect this pins"
);
}
Ok(())
}
#[tokio::test]
async fn a_forward_that_never_reached_the_owner_is_not_owner()
-> Result<(), Box<dyn std::error::Error>> {
let relayed = relayed(Err(tonic::Status::unavailable("forward dial failed")))
.await
.err()
.ok_or("a failed dial must not be reported as success")?;
assert_eq!(relayed.code, aion_proto::WireErrorCode::NotOwner);
assert_eq!(relayed.error_type.as_deref(), Some("NotOwner"));
assert!(
relayed.message.contains(&format!("shard {SHARD}")),
"the refusal names the shard to re-resolve: {}",
relayed.message
);
Ok(())
}
#[tokio::test]
async fn the_owners_reply_is_relayed_verbatim() -> Result<(), Box<dyn std::error::Error>> {
let run_id = aion_core::RunId::new_v4();
let reply = ForwardReply::Rename(aion_proto::generated::RenameResponse {
run_id: Some(aion_proto::generated::RunId {
uuid: run_id.to_string(),
}),
display_name: String::from("Nightly settlement"),
});
let response = relay_rename(
&ScriptedForwarder(Ok(reply)),
target(),
MetadataMap::new(),
&rename_request(Some(ProtoWorkflowId::from(aion_core::WorkflowId::new_v4()))),
SHARD,
)
.await?;
assert_eq!(response.display_name, "Nightly settlement");
assert_eq!(
response.run_id.map(|run| run.uuid),
Some(run_id.to_string())
);
Ok(())
}
#[tokio::test]
async fn a_mismatched_reply_is_a_backend_error() -> Result<(), Box<dyn std::error::Error>> {
let reply = ForwardReply::Cancel(aion_proto::generated::CancelResponse {});
let relayed = relayed(Ok(reply))
.await
.err()
.ok_or("a mismatched reply must not be reported as success")?;
assert_eq!(relayed.code, aion_proto::WireErrorCode::Backend);
Ok(())
}
}