use std::sync::{
Arc,
atomic::{AtomicBool, Ordering},
};
use connectrpc::{ConnectError, RequestContext, Response, Router, ServiceRequest, ServiceResult};
use polyc_proto::proto::polychrome::state::v1 as pb;
use polyc_state::{
error::StateError,
model_attempt::{DispatchDisposition, ModelAttemptAuthority},
};
use tracing::Instrument as _;
use super::wire::{lifecycle_from_wire, receipt_to_wire, request_from_wire};
use crate::{
admission::{
AudienceBinding, PeerIdentity, check_audience_binding, check_call_context_version,
check_not_draining, check_transport_deadline, state_audience,
},
error::to_connect_error,
trace::adopt_caller_trace,
wire::{DeclaredCall, declared_call},
};
pub struct ModelAttemptSvc {
authority: Arc<dyn ModelAttemptAuthority>,
draining: Arc<AtomicBool>,
binding: Arc<AudienceBinding>,
}
impl ModelAttemptSvc {
#[must_use]
pub const fn new(
authority: Arc<dyn ModelAttemptAuthority>,
draining: Arc<AtomicBool>,
binding: Arc<AudienceBinding>,
) -> Self {
Self {
authority,
draining,
binding,
}
}
#[must_use]
pub fn register_on(self, router: Router) -> Router {
use pb::StateModelAttemptServiceExt as _;
Arc::new(self).register(router)
}
fn admit(
&self,
ctx: &RequestContext,
value: impl Into<Option<pb::CallContext>>,
method: &'static str,
) -> Result<(polyc_state::context::CallContext, tracing::Span), ConnectError> {
check_not_draining(self.draining.load(Ordering::Relaxed))?;
let declared: DeclaredCall =
declared_call(value).map_err(|error| to_connect_error(&error))?;
let family = polyc_state::model_attempt::family();
check_call_context_version(declared.version).map_err(|error| to_connect_error(&error))?;
let peer = PeerIdentity::from_verified_leaf(
ctx.peer_certs().and_then(<[_]>::first).map(|leaf| &**leaf),
);
check_audience_binding(
&peer,
&declared.audience,
&state_audience(),
&self.binding,
&family,
)
.map_err(|error| to_connect_error(&error))?;
check_transport_deadline(ctx.time_remaining(), &family)
.map_err(|error| to_connect_error(&error))?;
Ok((
declared.origin_relative_context(),
adopt_caller_trace(ctx.headers(), method),
))
}
async fn blocking<R: Send + 'static>(
operation: impl FnOnce() -> Result<R, StateError> + Send + 'static,
) -> Result<R, ConnectError> {
tokio::task::spawn_blocking(operation)
.await
.map_err(|_| {
to_connect_error(&StateError::Unavailable {
family: polyc_state::model_attempt::family(),
reach: polyc_state::error::OutageReach::PossiblyApplied,
})
})?
.map_err(|error| to_connect_error(&error))
}
}
#[allow(refining_impl_trait)]
impl pb::StateModelAttemptService for ModelAttemptSvc {
async fn reserve(
&self,
ctx: RequestContext,
request: ServiceRequest<'_, pb::ReserveModelAttemptRequest>,
) -> ServiceResult<pb::ReserveModelAttemptReply> {
let message = request.to_owned_message();
let (context, span) = self.admit(&ctx, message.context, "ReserveModelAttempt")?;
let request = request_from_wire(message.request.into_option().ok_or_else(|| {
to_connect_error(&StateError::Malformed {
field: "request".to_owned(),
reason: "reserve carries request".to_owned(),
})
})?)
.map_err(|error| to_connect_error(&error))?;
let authority = Arc::clone(&self.authority);
let receipt = Self::blocking(move || authority.reserve(request, &context))
.instrument(span)
.await?;
Response::ok(pb::ReserveModelAttemptReply {
receipt: buffa::MessageField::some(receipt_to_wire(&receipt)),
__buffa_unknown_fields: buffa::UnknownFields::default(),
})
}
async fn dispatch(
&self,
ctx: RequestContext,
request: ServiceRequest<'_, pb::DispatchModelAttemptRequest>,
) -> ServiceResult<pb::DispatchModelAttemptReply> {
let message = request.to_owned_message();
let (context, span) = self.admit(&ctx, message.context, "DispatchModelAttempt")?;
let request = request_from_wire(message.request.into_option().ok_or_else(|| {
to_connect_error(&StateError::Malformed {
field: "request".to_owned(),
reason: "dispatch carries request".to_owned(),
})
})?)
.map_err(|error| to_connect_error(&error))?;
let authority = Arc::clone(&self.authority);
let (begin, receipt) = Self::blocking(move || {
let disposition = authority.dispatch(&request, &context)?;
let begin = matches!(&disposition, DispatchDisposition::Begin);
let receipt = match disposition {
DispatchDisposition::Begin => authority.reserve(request, &context)?,
DispatchDisposition::Existing(receipt) => receipt,
};
Ok((begin, receipt))
})
.instrument(span)
.await?;
Response::ok(pb::DispatchModelAttemptReply {
begin,
receipt: buffa::MessageField::some(receipt_to_wire(&receipt)),
__buffa_unknown_fields: buffa::UnknownFields::default(),
})
}
async fn transition(
&self,
ctx: RequestContext,
request: ServiceRequest<'_, pb::TransitionModelAttemptRequest>,
) -> ServiceResult<pb::TransitionModelAttemptReply> {
let message = request.to_owned_message();
let (context, span) = self.admit(&ctx, message.context, "TransitionModelAttempt")?;
let request = request_from_wire(message.request.into_option().ok_or_else(|| {
to_connect_error(&StateError::Malformed {
field: "request".to_owned(),
reason: "transition carries request".to_owned(),
})
})?)
.map_err(|error| to_connect_error(&error))?;
let next = lifecycle_from_wire(message.next).map_err(|error| to_connect_error(&error))?;
let authority = Arc::clone(&self.authority);
let receipt = Self::blocking(move || authority.transition(&request, next, &context))
.instrument(span)
.await?;
Response::ok(pb::TransitionModelAttemptReply {
receipt: buffa::MessageField::some(receipt_to_wire(&receipt)),
__buffa_unknown_fields: buffa::UnknownFields::default(),
})
}
}