use std::sync::{
Arc,
atomic::{AtomicBool, Ordering},
};
use connectrpc::{ConnectError, RequestContext, Response, Router, ServiceRequest, ServiceResult};
use polyc_proto::proto::polychrome::state::v1::{
BeginQueryAuditReply, BeginQueryAuditRequest, CompleteQueryAuditReply,
CompleteQueryAuditRequest, GetQueryAuditReceiptReply, GetQueryAuditReceiptRequest,
GetQueryAuditReply, GetQueryAuditRequest, ListUnmatchedQueryAuditsReply,
ListUnmatchedQueryAuditsRequest, StateQueryAuditService, StateQueryAuditServiceExt,
};
use polyc_state::{
command::CommandEnvelope,
context::CallContext,
error::StateError,
id::{Audience, NamespaceId, Purpose},
query_audit::{
self, AuditPhase, BeginQueryAudit, CompleteQueryAudit, ExecutionPermit,
ListUnmatchedIntents, QueryAuditRead, QueryAuditWrite, QueryId, ReadQueryAudit,
RequesterId,
},
};
use crate::{
admission::{
AudienceBinding, PeerIdentity, check_audience_binding, check_call_context_version,
check_not_draining, check_transport_deadline, state_audience,
},
error::to_connect_error as state_to_connect_error,
query_audit::{error::to_connect_error, wire::digest},
trace::adopt_caller_trace,
wire::{DeclaredCall, Kernel, declared_call},
};
pub trait QueryAuditAuthority: QueryAuditRead + QueryAuditWrite {}
impl<T> QueryAuditAuthority for T where T: QueryAuditRead + QueryAuditWrite {}
pub struct QueryAuditSvc {
audit: Arc<dyn QueryAuditAuthority>,
draining: Arc<AtomicBool>,
binding: Arc<AudienceBinding>,
}
impl QueryAuditSvc {
#[must_use]
pub const fn new(
audit: Arc<dyn QueryAuditAuthority>,
draining: Arc<AtomicBool>,
binding: Arc<AudienceBinding>,
) -> Self {
Self {
audit,
draining,
binding,
}
}
#[must_use]
pub fn register_on(self, router: Router) -> Router {
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<polyc_proto::proto::polychrome::state::v1::CallContext>>,
method: &'static str,
) -> Result<(CallContext, tracing::Span), ConnectError> {
check_not_draining(self.draining.load(Ordering::Relaxed))?;
let declared: DeclaredCall =
declared_call(context).map_err(|error| state_to_connect_error(&error))?;
let family = query_audit::family();
check_call_context_version(declared.version)
.map_err(|error| state_to_connect_error(&error))?;
check_audience_binding(
&Self::peer(ctx),
&declared.audience,
&state_audience(),
&self.binding,
&family,
)
.map_err(|error| state_to_connect_error(&error))?;
check_transport_deadline(ctx.time_remaining(), &family)
.map_err(|error| state_to_connect_error(&error))?;
Ok((
declared.origin_relative_context(),
adopt_caller_trace(ctx.headers(), method),
))
}
}
fn envelope(purpose: String, audience: String) -> CommandEnvelope {
CommandEnvelope::new(
Purpose::new(purpose),
Audience::new(audience),
query_audit::command_bounds(),
)
}
fn required_field(field: &str, reason: &str) -> ConnectError {
to_connect_error(
&StateError::Malformed {
field: field.to_owned(),
reason: reason.to_owned(),
}
.into(),
)
}
#[allow(refining_impl_trait)]
impl StateQueryAuditService for QueryAuditSvc {
async fn begin(
&self,
ctx: RequestContext,
request: ServiceRequest<'_, BeginQueryAuditRequest>,
) -> ServiceResult<BeginQueryAuditReply> {
let message = request.to_owned_message();
let (context, span) = self.admit(&ctx, message.context, "BeginQueryAudit")?;
let _entered = span.enter();
let command = BeginQueryAudit::new(
QueryId::new(message.query),
NamespaceId::new(message.namespace),
RequesterId::new(message.requester),
digest("shape", &message.shape)
.map_err(|error| required_field("shape", &error.to_string()))?,
digest("digest", &message.digest)
.map_err(|error| required_field("digest", &error.to_string()))?,
envelope(message.purpose, message.command_audience),
);
let permit = self
.audit
.begin(command, &context)
.map_err(|error| to_connect_error(&error))?;
Response::ok(BeginQueryAuditReply {
receipt: buffa::MessageField::some(Kernel(permit.receipt()).into()),
__buffa_unknown_fields: buffa::UnknownFields::default(),
})
}
async fn complete(
&self,
ctx: RequestContext,
request: ServiceRequest<'_, CompleteQueryAuditRequest>,
) -> ServiceResult<CompleteQueryAuditReply> {
let message = request.to_owned_message();
let (context, span) = self.admit(&ctx, message.context, "CompleteQueryAudit")?;
let _entered = span.enter();
let query = QueryId::new(message.query);
let namespace = NamespaceId::new(message.namespace);
let intent = self
.audit
.recorded_receipt(&query, &namespace, AuditPhase::Intent)
.map_err(|error| to_connect_error(&error))?
.ok_or_else(|| {
to_connect_error(
&polyc_state::query_audit::QueryAuditError::NoRecordedIntent {
query: query.clone(),
},
)
})?;
let completion = Kernel::try_from(message.completion.into_option().ok_or_else(|| {
required_field("completion", "a completion call carries what the query did")
})?)
.map_err(|error: StateError| to_connect_error(&error.into()))?
.into_inner();
let command = CompleteQueryAudit::new(
ExecutionPermit::granted(query, namespace, intent),
completion,
digest("digest", &message.digest)
.map_err(|error| required_field("digest", &error.to_string()))?,
envelope(message.purpose, message.command_audience),
);
let receipt = self
.audit
.complete(command, &context)
.map_err(|error| to_connect_error(&error))?;
Response::ok(CompleteQueryAuditReply {
receipt: buffa::MessageField::some(Kernel(&receipt).into()),
__buffa_unknown_fields: buffa::UnknownFields::default(),
})
}
async fn get_audit(
&self,
ctx: RequestContext,
request: ServiceRequest<'_, GetQueryAuditRequest>,
) -> ServiceResult<GetQueryAuditReply> {
let message = request.to_owned_message();
let (context, span) = self.admit(&ctx, message.context, "GetQueryAudit")?;
let _entered = span.enter();
let audit = self
.audit
.audit(
ReadQueryAudit::new(
QueryId::new(message.query),
NamespaceId::new(message.namespace),
),
&context,
)
.map_err(|error| to_connect_error(&error))?;
Response::ok(GetQueryAuditReply {
audit: audit
.as_ref()
.map_or_else(buffa::MessageField::default, |audit| {
buffa::MessageField::some(Kernel(audit).into())
}),
__buffa_unknown_fields: buffa::UnknownFields::default(),
})
}
async fn list_unmatched(
&self,
ctx: RequestContext,
request: ServiceRequest<'_, ListUnmatchedQueryAuditsRequest>,
) -> ServiceResult<ListUnmatchedQueryAuditsReply> {
let message = request.to_owned_message();
let (context, span) = self.admit(&ctx, message.context, "ListUnmatchedQueryAudits")?;
let _entered = span.enter();
let page =
Kernel::try_from(message.page.into_option().ok_or_else(|| {
required_field("page", "a bounded listing carries its page request")
})?)
.map_err(|error: StateError| to_connect_error(&error.into()))?
.into_inner();
let result = self
.audit
.unmatched_intents(
ListUnmatchedIntents::new(NamespaceId::new(message.namespace), page),
&context,
)
.map_err(|error| to_connect_error(&error))?;
Response::ok(ListUnmatchedQueryAuditsReply {
page: buffa::MessageField::some(Kernel(&result).into()),
__buffa_unknown_fields: buffa::UnknownFields::default(),
})
}
async fn get_receipt(
&self,
ctx: RequestContext,
request: ServiceRequest<'_, GetQueryAuditReceiptRequest>,
) -> ServiceResult<GetQueryAuditReceiptReply> {
let message = request.to_owned_message();
let (_context, span) = self.admit(&ctx, message.context, "GetQueryAuditReceipt")?;
let _entered = span.enter();
let phase = if message.completion {
AuditPhase::Completion
} else {
AuditPhase::Intent
};
let receipt = self
.audit
.recorded_receipt(
&QueryId::new(message.query),
&NamespaceId::new(message.namespace),
phase,
)
.map_err(|error| to_connect_error(&error))?;
Response::ok(GetQueryAuditReceiptReply {
receipt: receipt
.as_ref()
.map_or_else(buffa::MessageField::default, |receipt| {
buffa::MessageField::some(Kernel(receipt).into())
}),
__buffa_unknown_fields: buffa::UnknownFields::default(),
})
}
}