use super::super::coordination::{ConversationGuardHandle, ConversationOrderGate};
use super::super::types::As4ReceivePushOutput;
use crate::core::{AsxError, ErrorCode, ErrorContext, Result, SessionContext};
pub(super) fn ordered_missing_conversation_id_error(session: &SessionContext) -> AsxError {
AsxError::new(
ErrorCode::PolicyViolation,
"ordered AS4 receive requires eb:ConversationId",
ErrorContext::for_session("as4_receive_push_ordered", session),
)
}
pub(super) fn ordered_conversation_id_mismatch_error(
session: &SessionContext,
reserved: &str,
verified: &str,
) -> AsxError {
AsxError::new(
ErrorCode::SecurityVerificationFailed,
format!(
"AS4 ordered receive reserved the turn for conversation '{reserved}' but the \
verified eb:ConversationId is '{verified}'"
),
ErrorContext::for_session("as4_receive_push_ordered", session),
)
}
pub(super) async fn confirm_and_record_turn(
session: &SessionContext,
gate: &dyn ConversationOrderGate,
reserved_conversation_id: &str,
output: &As4ReceivePushOutput,
) -> Result<()> {
let verified = output
.user_message
.conversation_id
.as_deref()
.ok_or_else(|| ordered_missing_conversation_id_error(session))?;
if verified != reserved_conversation_id {
return Err(ordered_conversation_id_mismatch_error(
session,
reserved_conversation_id,
verified,
));
}
gate.record_message_ordering(
verified,
&output.user_message.message_id,
output.user_message.ref_to_message_id.as_deref(),
)
.await
}
pub struct As4Ordered<T> {
value: T,
turn: Option<Box<dyn ConversationGuardHandle>>,
}
impl<T> As4Ordered<T> {
pub(super) fn held(value: T, turn: Box<dyn ConversationGuardHandle>) -> Self {
Self {
value,
turn: Some(turn),
}
}
pub(super) fn untimed(value: T) -> Self {
Self { value, turn: None }
}
pub fn get(&self) -> &T {
&self.value
}
pub fn holds_turn(&self) -> bool {
self.turn.is_some()
}
pub fn into_inner(mut self) -> T {
if let Some(turn) = self.turn.take() {
turn.release();
}
self.value
}
}
impl<T: std::fmt::Debug> std::fmt::Debug for As4Ordered<T> {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("As4Ordered")
.field("value", &self.value)
.field("holds_turn", &self.turn.is_some())
.finish()
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn mismatch_is_a_security_failure_and_names_both_conversations() {
let session = SessionContext::new("s", "p", "strict").expect("session");
let err =
ordered_conversation_id_mismatch_error(&session, "reserved-conv", "verified-conv");
assert_eq!(err.code, ErrorCode::SecurityVerificationFailed);
assert!(err.message.contains("reserved-conv"), "{}", err.message);
assert!(err.message.contains("verified-conv"), "{}", err.message);
}
#[test]
fn a_result_without_a_turn_reports_that_it_holds_none() {
let ordered = As4Ordered::untimed(7u8);
assert!(!ordered.holds_turn());
assert_eq!(*ordered.get(), 7);
assert_eq!(ordered.into_inner(), 7);
}
}