use std::future::Future;
use std::sync::{
Arc,
atomic::{AtomicBool, Ordering},
};
use buffa::EnumValue;
use connectrpc::{ConnectError, RequestContext, Response, Router, ServiceRequest, ServiceResult};
use polyc_proto::proto::polychrome::state::v1 as pb;
use polyc_state::{
context::CallContext,
error::StateError,
id::{CommandId, NamespaceId},
revision::Revision,
sessions::{RefreshPreflight, SessionDirectory, SessionRead, SessionWrite},
versioned::{VersionedRead, VersionedTransact, authority::SessionId},
};
use crate::{
admission::{
AudienceBinding, PeerIdentity, check_audience_binding, check_call_context_version,
check_not_draining, check_transport_deadline, slot_wait, state_audience,
},
blocking,
error::to_connect_error,
trace::adopt_caller_trace,
wire::{DeclaredCall, Kernel, declared_call},
};
use super::wire::{
command_from_wire, expiration_to_wire, family_to_wire, principal_from_wire, record_to_wire,
};
pub trait SessionAuthority: SessionRead + SessionWrite {
fn namespace(&self) -> &NamespaceId;
}
impl<V> SessionAuthority for SessionDirectory<V>
where
V: VersionedRead + VersionedTransact,
{
fn namespace(&self) -> &NamespaceId {
self.namespace()
}
}
pub struct SessionSvc {
authority: Arc<dyn SessionAuthority>,
draining: Arc<AtomicBool>,
binding: Arc<AudienceBinding>,
}
impl SessionSvc {
#[must_use]
pub const fn new(
authority: Arc<dyn SessionAuthority>,
draining: Arc<AtomicBool>,
binding: Arc<AudienceBinding>,
) -> Self {
Self {
authority,
draining,
binding,
}
}
#[must_use]
pub fn register_on(self, router: Router) -> Router {
use pb::StateSessionServiceExt as _;
Arc::new(self).register(router)
}
fn peer(ctx: &RequestContext) -> PeerIdentity {
PeerIdentity::from_verified_leaf(
ctx.peer_certs().and_then(<[_]>::first).map(|leaf| &**leaf),
)
}
fn admit(
&self,
ctx: &RequestContext,
context: impl Into<Option<pb::CallContext>>,
method: &'static str,
) -> Result<(CallContext, tracing::Span, std::time::Duration), ConnectError> {
check_not_draining(self.draining.load(Ordering::Relaxed))?;
let declared: DeclaredCall =
declared_call(context).map_err(|error| to_connect_error(&error))?;
let family = polyc_state::versioned::family();
check_call_context_version(declared.version).map_err(|error| to_connect_error(&error))?;
check_audience_binding(
&Self::peer(ctx),
&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),
slot_wait(ctx.time_remaining(), &declared),
))
}
fn namespace(&self, value: &str) -> Result<(), ConnectError> {
if value == self.authority.namespace().as_str() {
Ok(())
} else {
Err(to_connect_error(&StateError::Denied {
family: polyc_state::versioned::family(),
}))
}
}
}
#[allow(refining_impl_trait)]
impl pb::StateSessionService for SessionSvc {
fn transact(
&self,
ctx: RequestContext,
request: ServiceRequest<'_, pb::TransactSessionRequest>,
) -> impl Future<Output = ServiceResult<pb::TransactSessionReply>> {
let message = request.to_owned_message();
async move {
let (context, span, wait) = self.admit(&ctx, message.context, "TransactSession")?;
let metadata = message
.metadata
.into_option()
.ok_or_else(|| malformed_connect("metadata", "session command carries metadata"))?;
self.namespace(&metadata.namespace)?;
let operation = message.operation.into_option().ok_or_else(|| {
malformed_connect("operation", "session command carries one operation")
})?;
let command =
command_from_wire(metadata, operation).map_err(|error| to_connect_error(&error))?;
let authority = Arc::clone(&self.authority);
let receipt =
blocking::mutation(span, polyc_state::versioned::family(), wait, move || {
authority.transact(command, &context)
})
.await?;
Response::ok(pb::TransactSessionReply {
receipt: buffa::MessageField::some(pb::Receipt::from(Kernel(&receipt))),
__buffa_unknown_fields: buffa::UnknownFields::default(),
})
}
}
fn get_family(
&self,
ctx: RequestContext,
request: ServiceRequest<'_, pb::GetSessionFamilyRequest>,
) -> impl Future<Output = ServiceResult<pb::GetSessionFamilyReply>> {
let message = request.to_owned_message();
async move {
let (context, span, wait) = self.admit(&ctx, message.context, "GetSessionFamily")?;
self.namespace(&message.namespace)?;
let family_id = SessionId::new(message.family_id);
let authority = Arc::clone(&self.authority);
let fact = blocking::read(span, polyc_state::versioned::family(), wait, move || {
authority.family(&family_id, &context)
})
.await?;
Response::ok(pb::GetSessionFamilyReply {
family: fact
.value()
.as_ref()
.map_or_else(buffa::MessageField::none, |value| {
buffa::MessageField::some(family_to_wire(value))
}),
snapshot_revision: fact.snapshot_revision().get(),
entry_revision: fact
.entry_revision()
.map(polyc_state::revision::Revision::get),
__buffa_unknown_fields: buffa::UnknownFields::default(),
})
}
}
fn inspect_refresh(
&self,
ctx: RequestContext,
request: ServiceRequest<'_, pb::InspectSessionRefreshRequest>,
) -> impl Future<Output = ServiceResult<pb::InspectSessionRefreshReply>> {
let message = request.to_owned_message();
async move {
let (context, span, wait) =
self.admit(&ctx, message.context, "InspectSessionRefresh")?;
self.namespace(&message.namespace)?;
let family_id = SessionId::new(message.family_id);
let presented_generation = message.presented_generation;
let now_ms = message.now_ms;
let authority = Arc::clone(&self.authority);
let result = blocking::read(span, polyc_state::versioned::family(), wait, move || {
authority.inspect_refresh(&family_id, presented_generation, now_ms, &context)
})
.await?;
let (classification, family, snapshot_revision, entry_revision) = match result {
RefreshPreflight::Rotate(fact) => fact_parts(pb::SessionRefreshClass::Rotate, fact),
RefreshPreflight::Replay(fact) => fact_parts(pb::SessionRefreshClass::Replay, fact),
RefreshPreflight::ReuseDetected(fact) => {
fact_parts(pb::SessionRefreshClass::ReuseDetected, fact)
}
RefreshPreflight::Expired(fact) => {
fact_parts(pb::SessionRefreshClass::Expired, fact)
}
RefreshPreflight::UnknownFamily(fact) => (
pb::SessionRefreshClass::UnknownFamily,
None,
fact.snapshot_revision().get(),
fact.entry_revision().map(Revision::get),
),
};
Response::ok(pb::InspectSessionRefreshReply {
classification: EnumValue::Known(classification),
family: family.into(),
snapshot_revision,
entry_revision,
__buffa_unknown_fields: buffa::UnknownFields::default(),
})
}
}
fn get_expiration_shard(
&self,
ctx: RequestContext,
request: ServiceRequest<'_, pb::GetSessionExpirationShardRequest>,
) -> impl Future<Output = ServiceResult<pb::GetSessionExpirationShardReply>> {
let message = request.to_owned_message();
async move {
let (context, span, wait) =
self.admit(&ctx, message.context, "GetSessionExpirationShard")?;
self.namespace(&message.namespace)?;
let shard = u16::try_from(message.shard)
.map_err(|_| malformed_connect("shard", "expiry shard fits u16"))?;
let authority = Arc::clone(&self.authority);
let fact = blocking::read(span, polyc_state::versioned::family(), wait, move || {
authority.expiration_shard(shard, &context)
})
.await?;
Response::ok(pb::GetSessionExpirationShardReply {
expiration: buffa::MessageField::some(expiration_to_wire(fact.value())),
snapshot_revision: fact.snapshot_revision().get(),
entry_revision: fact
.entry_revision()
.map(polyc_state::revision::Revision::get),
__buffa_unknown_fields: buffa::UnknownFields::default(),
})
}
}
fn get_bearer_authorization(
&self,
ctx: RequestContext,
request: ServiceRequest<'_, pb::GetBearerAuthorizationRequest>,
) -> impl Future<Output = ServiceResult<pb::GetBearerAuthorizationReply>> {
let message = request.to_owned_message();
async move {
let (context, span, wait) =
self.admit(&ctx, message.context, "GetBearerAuthorization")?;
self.namespace(&message.namespace)?;
let principal =
principal_from_wire(message.principal.into_option().ok_or_else(|| {
malformed_connect("principal", "authorization request carries a principal")
})?)
.map_err(|error| to_connect_error(&error))?;
let session = SessionId::new(message.session);
let authority = Arc::clone(&self.authority);
let fact = blocking::read(span, polyc_state::versioned::family(), wait, move || {
authority.bearer_authorization(&principal, &session, &context)
})
.await?;
Response::ok(pb::GetBearerAuthorizationReply {
record: buffa::MessageField::some(record_to_wire(fact.value().record())),
current_epoch: fact.value().current_epoch().get(),
snapshot_revision: fact.snapshot_revision().get(),
entry_revision: fact
.entry_revision()
.map(polyc_state::revision::Revision::get),
__buffa_unknown_fields: buffa::UnknownFields::default(),
})
}
}
fn get_bearer_record(
&self,
ctx: RequestContext,
request: ServiceRequest<'_, pb::GetBearerRecordRequest>,
) -> impl Future<Output = ServiceResult<pb::GetBearerRecordReply>> {
let message = request.to_owned_message();
async move {
let (context, span, wait) = self.admit(&ctx, message.context, "GetBearerRecord")?;
self.namespace(&message.namespace)?;
let session = SessionId::new(message.session);
let authority = Arc::clone(&self.authority);
let fact = blocking::read(span, polyc_state::versioned::family(), wait, move || {
authority.bearer_record(&session, &context)
})
.await?;
Response::ok(pb::GetBearerRecordReply {
record: buffa::MessageField::some(record_to_wire(*fact.value())),
snapshot_revision: fact.snapshot_revision().get(),
entry_revision: fact.entry_revision().map(Revision::get),
__buffa_unknown_fields: buffa::UnknownFields::default(),
})
}
}
fn get_authorization_epoch(
&self,
ctx: RequestContext,
request: ServiceRequest<'_, pb::GetSessionAuthorizationEpochRequest>,
) -> impl Future<Output = ServiceResult<pb::GetSessionAuthorizationEpochReply>> {
let message = request.to_owned_message();
async move {
let (context, span, wait) =
self.admit(&ctx, message.context, "GetSessionAuthorizationEpoch")?;
self.namespace(&message.namespace)?;
let principal =
principal_from_wire(message.principal.into_option().ok_or_else(|| {
malformed_connect("principal", "epoch request carries a principal")
})?)
.map_err(|error| to_connect_error(&error))?;
let authority = Arc::clone(&self.authority);
let fact = blocking::read(span, polyc_state::versioned::family(), wait, move || {
authority.authorization_epoch(&principal, &context)
})
.await?;
Response::ok(pb::GetSessionAuthorizationEpochReply {
epoch: fact.value().get(),
snapshot_revision: fact.snapshot_revision().get(),
entry_revision: fact
.entry_revision()
.map(polyc_state::revision::Revision::get),
__buffa_unknown_fields: buffa::UnknownFields::default(),
})
}
}
fn get_receipt(
&self,
ctx: RequestContext,
request: ServiceRequest<'_, pb::GetSessionReceiptRequest>,
) -> impl Future<Output = ServiceResult<pb::GetSessionReceiptReply>> {
let message = request.to_owned_message();
async move {
let (_context, span, wait) = self.admit(&ctx, message.context, "GetSessionReceipt")?;
self.namespace(&message.namespace)?;
let namespace = NamespaceId::new(message.namespace);
let command_id = CommandId::new(message.command_id);
let authority = Arc::clone(&self.authority);
let receipt = blocking::read(span, polyc_state::versioned::family(), wait, move || {
authority.committed_receipt(&namespace, &command_id)
})
.await?;
Response::ok(pb::GetSessionReceiptReply {
receipt: receipt
.as_ref()
.map_or_else(buffa::MessageField::none, |value| {
buffa::MessageField::some(pb::Receipt::from(Kernel(value)))
}),
__buffa_unknown_fields: buffa::UnknownFields::default(),
})
}
}
}
fn fact_parts(
classification: pb::SessionRefreshClass,
fact: polyc_state::sessions::SessionFact<polyc_state::sessions::SessionFamily>,
) -> (
pb::SessionRefreshClass,
Option<pb::StateSessionFamily>,
u64,
Option<u64>,
) {
let snapshot = fact.snapshot_revision().get();
let entry = fact
.entry_revision()
.map(polyc_state::revision::Revision::get);
let family = Some(family_to_wire(&fact.into_value()));
(classification, family, snapshot, entry)
}
fn malformed_connect(field: &str, reason: &str) -> ConnectError {
to_connect_error(&StateError::Malformed {
field: field.into(),
reason: reason.into(),
})
}