use std::future::Future;
use std::marker::PhantomData;
use std::pin::Pin;
use std::sync::Arc;
use std::sync::atomic::{AtomicU64, Ordering};
use std::task::{Context, Poll};
use prost::Message;
use saddle_boundary::{
BoundaryTransport, FakeProfuseContractBoundary as TransportFakeBoundary, InvocationTarget,
ProfuseContractEndpoint, TonicBoundary,
};
use serde::de::DeserializeOwned;
pub use saddle_boundary::{FakeAttempt, FakeExecutionCertainty, FakeStep, FakeTechnicalCode};
pub const MAX_PROFUSE_GW_USER_ID_BYTES: usize = 64;
#[derive(Clone)]
pub struct ProfuseGwContext {
user_id: [u8; MAX_PROFUSE_GW_USER_ID_BYTES],
user_id_len: u8,
}
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
#[allow(dead_code)]
pub(crate) enum ProfuseGwContextError {
MissingUserId,
UserIdTooLong,
}
impl ProfuseGwContext {
#[allow(dead_code)]
pub(crate) fn from_framework(user_id: &str) -> Result<Self, ProfuseGwContextError> {
if user_id.is_empty() {
return Err(ProfuseGwContextError::MissingUserId);
}
if user_id.len() > MAX_PROFUSE_GW_USER_ID_BYTES {
return Err(ProfuseGwContextError::UserIdTooLong);
}
let mut stored = [0; MAX_PROFUSE_GW_USER_ID_BYTES];
stored[..user_id.len()].copy_from_slice(user_id.as_bytes());
Ok(Self {
user_id: stored,
user_id_len: user_id.len() as u8,
})
}
pub fn user_id(&self) -> &str {
std::str::from_utf8(&self.user_id[..usize::from(self.user_id_len)])
.expect("ProfuseGwContext is constructed from validated UTF-8")
}
}
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
pub enum ExecutionCertainty {
NotExecuted,
Executed,
MayHaveExecuted,
}
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
pub enum TechnicalFailureCode {
FunctionNotFound,
FunctionRequestInvalid,
CapacityRejected,
DeadlineExceeded,
DependencyUnavailable,
ContractResultInvalid,
InternalFailure,
TransportFailure,
}
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
pub struct TechnicalFailure {
code: TechnicalFailureCode,
certainty: ExecutionCertainty,
}
impl TechnicalFailure {
#[allow(dead_code)]
pub(crate) const fn from_framework(
code: TechnicalFailureCode,
certainty: ExecutionCertainty,
) -> Self {
Self { code, certainty }
}
pub const fn code(&self) -> TechnicalFailureCode {
self.code
}
pub const fn certainty(&self) -> ExecutionCertainty {
self.certainty
}
}
pub enum ExternalFunctionResult<T> {
Completed(T),
TechnicalFailure(TechnicalFailure),
}
#[doc(hidden)]
pub struct ApplicationContractSeal {
adapter: ContractAdapter,
}
impl ApplicationContractSeal {
fn from_fake(boundary: TransportFakeBoundary, user_id: String, deadline_unix_ms: i64) -> Self {
Self {
adapter: ContractAdapter {
boundary: ContractBoundary::Fake(boundary),
request_id: "alpha1-application".into(),
call_id_prefix: "call".into(),
user_id,
deadline_unix_ms,
next_call: Arc::new(AtomicU64::new(1)),
},
}
}
}
#[derive(Clone)]
struct ContractAdapter {
boundary: ContractBoundary,
request_id: String,
call_id_prefix: String,
user_id: String,
deadline_unix_ms: i64,
next_call: Arc<AtomicU64>,
}
#[derive(Clone)]
enum ContractBoundary {
Fake(TransportFakeBoundary),
Tonic(TonicBoundary),
}
impl ContractBoundary {
async fn invoke(
&self,
request: saddle_boundary::InvokeRequest,
) -> Result<saddle_boundary::InvokeResponse, saddle_boundary::BoundaryError> {
match self {
Self::Fake(boundary) => boundary.invoke(request).await,
Self::Tonic(boundary) => boundary.invoke(request).await,
}
}
}
#[doc(hidden)]
#[derive(Clone)]
pub struct ProfuseContractDeployment {
boundary: ContractBoundary,
}
impl ProfuseContractDeployment {
#[doc(hidden)]
pub async fn connect(endpoint: ProfuseContractEndpoint) -> Result<Self, TechnicalFailure> {
TonicBoundary::connect(endpoint)
.await
.map(|boundary| Self {
boundary: ContractBoundary::Tonic(boundary),
})
.map_err(|error| match boundary_error::<()>(error) {
ExternalFunctionResult::TechnicalFailure(failure) => failure,
ExternalFunctionResult::Completed(()) => unreachable!(),
})
}
#[doc(hidden)]
pub fn bind_accepted(
&self,
accepted: &saddle_boundary::ingress::AcceptedIngress,
) -> ApplicationContractSeal {
ApplicationContractSeal {
adapter: ContractAdapter {
boundary: self.boundary.clone(),
request_id: accepted.identity.request_id.clone(),
call_id_prefix: accepted.identity.call_id.clone(),
user_id: accepted.user_id.clone(),
deadline_unix_ms: accepted.identity.deadline_unix_ms,
next_call: Arc::new(AtomicU64::new(1)),
},
}
}
}
#[must_use = "a declared external-function call must be explicitly handled"]
pub struct DeclaredExternalFunctionCall<Application, Function, Request, Response> {
future: Pin<Box<dyn Future<Output = ExternalFunctionResult<Response>> + Send>>,
_type: PhantomData<fn(Application, Function) -> Response>,
_request: PhantomData<fn(Request)>,
}
impl<Application, Function, Request, Response>
DeclaredExternalFunctionCall<Application, Function, Request, Response>
where
Request: Message + Send + 'static,
Response: Message + Default + Send + 'static,
{
#[doc(hidden)]
pub fn from_declared(
request: Request,
seal: &ApplicationContractSeal,
business_unit: &'static str,
function: &'static str,
) -> Self {
let adapter = seal.adapter.clone();
let call_number = adapter.next_call.fetch_add(1, Ordering::Relaxed);
let request_id = adapter.request_id.clone();
let call_id = format!("{}-{call_number}", adapter.call_id_prefix);
Self::from_declared_identity(
request,
adapter,
request_id,
call_id,
business_unit,
function,
)
}
#[doc(hidden)]
pub fn from_declared_with_identity(
request: Request,
seal: &ApplicationContractSeal,
request_id: impl Into<String>,
call_id: impl Into<String>,
business_unit: &'static str,
function: &'static str,
) -> Self {
Self::from_declared_identity(
request,
seal.adapter.clone(),
request_id.into(),
call_id.into(),
business_unit,
function,
)
}
fn from_declared_identity(
request: Request,
adapter: ContractAdapter,
request_id: String,
call_id: String,
business_unit: &'static str,
function: &'static str,
) -> Self {
let future = Box::pin(async move {
let request = match saddle_boundary::InvokeRequest::unary(
request_id,
call_id,
match InvocationTarget::new(business_unit, function) {
Ok(target) => target,
Err(error) => return boundary_error(error),
},
adapter.deadline_unix_ms,
saddle_boundary::CallerContext {
user_id: adapter.user_id,
},
request.encode_to_vec(),
) {
Ok(request) => request,
Err(error) => return boundary_error(error),
};
match adapter.boundary.invoke(request).await {
Err(error) => boundary_error(error),
Ok(response) => match response.outcome {
Some(saddle_boundary::Outcome::Completed(completed)) => {
match Response::decode(completed.result.as_slice()) {
Ok(result) => ExternalFunctionResult::Completed(result),
Err(_) => ExternalFunctionResult::TechnicalFailure(
TechnicalFailure::from_framework(
TechnicalFailureCode::ContractResultInvalid,
ExecutionCertainty::Executed,
),
),
}
}
Some(saddle_boundary::Outcome::TechnicalFailure(failure)) => {
ExternalFunctionResult::TechnicalFailure(TechnicalFailure::from_framework(
map_code(failure.code),
map_certainty(failure.certainty),
))
}
None => {
ExternalFunctionResult::TechnicalFailure(TechnicalFailure::from_framework(
TechnicalFailureCode::ContractResultInvalid,
ExecutionCertainty::MayHaveExecuted,
))
}
},
}
});
Self {
future,
_type: PhantomData,
_request: PhantomData,
}
}
}
impl<Application, Function, Request, Response> Future
for DeclaredExternalFunctionCall<Application, Function, Request, Response>
{
type Output = ExternalFunctionResult<Response>;
fn poll(mut self: Pin<&mut Self>, context: &mut Context<'_>) -> Poll<Self::Output> {
self.future.as_mut().poll(context)
}
}
fn boundary_error<T>(error: saddle_boundary::BoundaryError) -> ExternalFunctionResult<T> {
ExternalFunctionResult::TechnicalFailure(TechnicalFailure::from_framework(
map_boundary_code(error.code),
map_boundary_certainty(error.certainty),
))
}
fn map_code(code: i32) -> TechnicalFailureCode {
saddle_boundary::TechnicalCode::try_from(code)
.map(map_boundary_code)
.unwrap_or(TechnicalFailureCode::ContractResultInvalid)
}
fn map_certainty(certainty: i32) -> ExecutionCertainty {
saddle_boundary::ExecutionCertainty::try_from(certainty)
.map(map_boundary_certainty)
.unwrap_or(ExecutionCertainty::MayHaveExecuted)
}
fn map_boundary_code(code: saddle_boundary::TechnicalCode) -> TechnicalFailureCode {
match code {
saddle_boundary::TechnicalCode::FunctionNotFound => TechnicalFailureCode::FunctionNotFound,
saddle_boundary::TechnicalCode::FunctionRequestInvalid => {
TechnicalFailureCode::FunctionRequestInvalid
}
saddle_boundary::TechnicalCode::CapacityRejected => TechnicalFailureCode::CapacityRejected,
saddle_boundary::TechnicalCode::DeadlineExceeded => TechnicalFailureCode::DeadlineExceeded,
saddle_boundary::TechnicalCode::DependencyUnavailable => {
TechnicalFailureCode::DependencyUnavailable
}
saddle_boundary::TechnicalCode::ContractResultInvalid => {
TechnicalFailureCode::ContractResultInvalid
}
saddle_boundary::TechnicalCode::InternalFailure => TechnicalFailureCode::InternalFailure,
saddle_boundary::TechnicalCode::TransportFailure => TechnicalFailureCode::TransportFailure,
saddle_boundary::TechnicalCode::Unspecified => TechnicalFailureCode::ContractResultInvalid,
}
}
fn map_boundary_certainty(certainty: saddle_boundary::ExecutionCertainty) -> ExecutionCertainty {
match certainty {
saddle_boundary::ExecutionCertainty::NotExecuted => ExecutionCertainty::NotExecuted,
saddle_boundary::ExecutionCertainty::Executed => ExecutionCertainty::Executed,
saddle_boundary::ExecutionCertainty::MayHaveExecuted
| saddle_boundary::ExecutionCertainty::Unspecified => ExecutionCertainty::MayHaveExecuted,
}
}
#[derive(Clone)]
pub struct FakeProfuseContractBoundary {
boundary: TransportFakeBoundary,
}
pub struct FakeApplicationBinding {
context: ProfuseGwContext,
seal: ApplicationContractSeal,
}
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
pub enum ProfuseGwDispatchError {
InterfaceNotFound,
RequestDataInvalid,
ContextInvalid,
IdentityMismatch,
}
#[doc(hidden)]
pub fn ingress_matches_contract(
accepted: &saddle_boundary::ingress::AcceptedIngress,
seal: &ApplicationContractSeal,
) -> bool {
seal.adapter.request_id == accepted.identity.request_id
&& seal.adapter.call_id_prefix == accepted.identity.call_id
&& seal.adapter.deadline_unix_ms == accepted.identity.deadline_unix_ms
&& seal.adapter.user_id == accepted.user_id
}
#[doc(hidden)]
pub fn decode_accepted_profusegw<Request: DeserializeOwned>(
accepted: saddle_boundary::ingress::AcceptedIngress,
) -> Result<(Request, ProfuseGwContext), ProfuseGwDispatchError> {
let context = ProfuseGwContext::from_framework(&accepted.user_id)
.map_err(|_| ProfuseGwDispatchError::ContextInvalid)?;
let request = serde_json::from_value(accepted.request_data)
.map_err(|_| ProfuseGwDispatchError::RequestDataInvalid)?;
Ok((request, context))
}
impl FakeApplicationBinding {
pub fn profusegw_context(&self) -> ProfuseGwContext {
self.context.clone()
}
pub fn into_contract_seal(self) -> ApplicationContractSeal {
self.seal
}
}
impl FakeProfuseContractBoundary {
pub fn completed<Response: Message>(result: Response) -> Self {
Self {
boundary: TransportFakeBoundary::scripted([FakeStep::completed(&result)]),
}
}
pub fn scripted(steps: impl IntoIterator<Item = FakeStep>) -> Self {
Self {
boundary: TransportFakeBoundary::scripted(steps),
}
}
pub fn attempt_count(&self) -> usize {
self.boundary.attempt_count()
}
pub fn attempts(&self) -> Vec<FakeAttempt> {
self.boundary.attempts()
}
pub fn bind(self, user_id: &str, deadline_unix_ms: i64) -> FakeApplicationBinding {
assert!(deadline_unix_ms > 0, "test deadline must be positive");
let context = ProfuseGwContext::from_framework(user_id)
.expect("test user_id must satisfy profusegw bounds");
let seal = ApplicationContractSeal::from_fake(
self.boundary,
context.user_id().to_owned(),
deadline_unix_ms,
);
FakeApplicationBinding { context, seal }
}
#[doc(hidden)]
pub fn bind_accepted(
self,
accepted: &saddle_boundary::ingress::AcceptedIngress,
) -> Result<FakeApplicationBinding, ProfuseGwDispatchError> {
let context = ProfuseGwContext::from_framework(&accepted.user_id)
.map_err(|_| ProfuseGwDispatchError::ContextInvalid)?;
let seal = ApplicationContractSeal {
adapter: ContractAdapter {
boundary: ContractBoundary::Fake(self.boundary),
request_id: accepted.identity.request_id.clone(),
call_id_prefix: accepted.identity.call_id.clone(),
user_id: accepted.user_id.clone(),
deadline_unix_ms: accepted.identity.deadline_unix_ms,
next_call: Arc::new(AtomicU64::new(1)),
},
};
Ok(FakeApplicationBinding { context, seal })
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn profusegw_context_is_fixed_bounded_and_read_only() {
assert!(matches!(
ProfuseGwContext::from_framework(""),
Err(ProfuseGwContextError::MissingUserId)
));
assert!(matches!(
ProfuseGwContext::from_framework(&"x".repeat(MAX_PROFUSE_GW_USER_ID_BYTES + 1)),
Err(ProfuseGwContextError::UserIdTooLong)
));
let context = ProfuseGwContext::from_framework("2088用户").unwrap();
assert_eq!(context.user_id(), "2088用户");
}
#[test]
fn technical_failure_keeps_code_and_execution_certainty_separate() {
let failure = TechnicalFailure::from_framework(
TechnicalFailureCode::DependencyUnavailable,
ExecutionCertainty::MayHaveExecuted,
);
assert_eq!(failure.code(), TechnicalFailureCode::DependencyUnavailable);
assert_eq!(failure.certainty(), ExecutionCertainty::MayHaveExecuted);
}
#[tokio::test]
async fn declared_call_executes_transport_fake_and_decodes_typed_result() {
#[derive(Clone, PartialEq, Message)]
struct Request {
#[prost(uint64, tag = "1")]
value: u64,
}
#[derive(Clone, PartialEq, Message)]
struct Response {
#[prost(bool, tag = "1")]
accepted: bool,
}
struct Application;
struct Function;
let binding = FakeProfuseContractBoundary::completed(Response { accepted: true })
.bind("user-1", 1_800_000_000_000);
let seal = binding.into_contract_seal();
let observed = match &seal.adapter.boundary {
ContractBoundary::Fake(boundary) => boundary.clone(),
ContractBoundary::Tonic(_) => panic!("unit test binds the fake transport"),
};
let call =
DeclaredExternalFunctionCall::<Application, Function, Request, Response>::from_declared(
Request { value: 7 },
&seal,
"puc",
"query",
);
match call.await {
ExternalFunctionResult::Completed(response) => assert!(response.accepted),
ExternalFunctionResult::TechnicalFailure(_) => panic!("fake should complete"),
}
let requests = observed.attempts();
assert_eq!(requests.len(), 1);
assert_eq!(requests[0].function, "query");
assert_eq!(requests[0].request_id, "alpha1-application");
assert_eq!(requests[0].call_id, "call-1");
assert_eq!(requests[0].user_id, "user-1");
}
#[test]
fn transport_codes_map_to_the_closed_eight_code_set() {
let cases = [
(
saddle_boundary::TechnicalCode::FunctionNotFound,
TechnicalFailureCode::FunctionNotFound,
),
(
saddle_boundary::TechnicalCode::FunctionRequestInvalid,
TechnicalFailureCode::FunctionRequestInvalid,
),
(
saddle_boundary::TechnicalCode::CapacityRejected,
TechnicalFailureCode::CapacityRejected,
),
(
saddle_boundary::TechnicalCode::DeadlineExceeded,
TechnicalFailureCode::DeadlineExceeded,
),
(
saddle_boundary::TechnicalCode::DependencyUnavailable,
TechnicalFailureCode::DependencyUnavailable,
),
(
saddle_boundary::TechnicalCode::ContractResultInvalid,
TechnicalFailureCode::ContractResultInvalid,
),
(
saddle_boundary::TechnicalCode::InternalFailure,
TechnicalFailureCode::InternalFailure,
),
(
saddle_boundary::TechnicalCode::TransportFailure,
TechnicalFailureCode::TransportFailure,
),
];
for (boundary, facade) in cases {
assert_eq!(map_boundary_code(boundary), facade);
}
}
}