use saddle_boundary::standard_wire::{Encode, MethodSpec};
#[doc(hidden)]
pub trait Method: 'static {
type Fields<'a>: Encode;
type View<'a>;
const SPEC: MethodSpec;
fn view<'a>(bytes: &'a [u8]) -> Self::View<'a>;
}
pub struct Request<M: Method> {
pub(crate) wire: saddle_admission::ReadOnlyInput<Vec<u8>>,
_method: std::marker::PhantomData<fn() -> M>,
}
pub struct Response<M: Method> {
pub(crate) wire: saddle_admission::ReadOnlyInput<Vec<u8>>,
_method: std::marker::PhantomData<fn() -> M>,
}
impl<M: Method> Response<M> {
pub fn message(&self) -> M::View<'_> {
M::view(self.wire.get())
}
pub fn storage_bytes(&self) -> usize {
self.wire.storage_bytes()
}
pub(crate) fn from_wire(wire: saddle_admission::ReadOnlyInput<Vec<u8>>) -> Self {
Self {
wire,
_method: std::marker::PhantomData,
}
}
}
impl<M: Method> Request<M> {
pub fn storage_bytes(&self) -> usize {
self.wire.storage_bytes()
}
}
#[derive(Debug)]
#[doc(hidden)]
pub enum BuildError {
Unregistered,
TooLarge,
Wire(saddle_boundary::standard_wire::WireError),
Resource(saddle_admission::AdmissionError),
}
impl std::fmt::Display for BuildError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
Self::Unregistered => f.write_str("gRPC method is not registered in the application"),
Self::TooLarge => f.write_str("gRPC request exceeds the one-megabyte wire limit"),
Self::Wire(error) => error.fmt(f),
Self::Resource(error) => error.fmt(f),
}
}
}
impl std::error::Error for BuildError {
fn source(&self) -> Option<&(dyn std::error::Error + 'static)> {
match self {
Self::Wire(error) => Some(error),
Self::Resource(error) => Some(error),
_ => None,
}
}
}
pub(crate) fn registered<M: Method>(registry: &[&MethodSpec]) -> bool {
registry.iter().any(|entry| {
(entry.marker)() == std::any::TypeId::of::<M>()
&& entry.dependency == M::SPEC.dependency
&& entry.contract == M::SPEC.contract
&& entry.service == M::SPEC.service
&& entry.method == M::SPEC.method
&& entry.path == M::SPEC.path
&& std::ptr::eq(entry.request, M::SPEC.request)
&& std::ptr::eq(entry.response, M::SPEC.response)
})
}
#[doc(hidden)]
pub fn prepare<M: Method>(
memory: &saddle_admission::RequestMemory,
registry: &[&MethodSpec],
fields: &M::Fields<'_>,
) -> Result<Request<M>, BuildError> {
if !registered::<M>(registry) {
return Err(BuildError::Unregistered);
}
let len = fields.encoded_len().map_err(BuildError::Wire)?;
if len > saddle_boundary::MAX_UNARY_PAYLOAD_BYTES {
return Err(BuildError::TooLarge);
}
let mut original = None;
let wire = memory
.framework_output(|builder| {
builder.write_bytes(len, |bytes| {
let mut writer = saddle_boundary::standard_wire::Writer::new(bytes);
fields
.encode(&mut writer)
.and_then(|_| writer.finish())
.map_err(|error| {
original = Some(error);
saddle_admission::AdmissionError::ResponseCapacityExceeded
})
})
})
.map_err(|error| match original {
Some(error) => BuildError::Wire(error),
None => BuildError::Resource(error),
})?;
M::SPEC
.request
.validate(wire.get())
.map_err(BuildError::Wire)?;
Ok(Request {
wire,
_method: std::marker::PhantomData,
})
}
#[derive(serde::Deserialize)]
#[serde(deny_unknown_fields)]
#[doc(hidden)]
pub struct DependencyConfig {
pub authority: String,
#[serde(default)]
pub token: Option<String>,
#[serde(default)]
pub metadata: Vec<saddle_boundary::standard_rpc::MetadataConfig>,
}
#[derive(Clone)]
#[doc(hidden)]
pub struct FrozenRegistry {
methods: &'static [&'static MethodSpec],
endpoints: std::sync::Arc<[FrozenEndpoint]>,
identities: std::sync::Arc<[FrozenMethodIdentity]>,
}
struct FrozenMethodIdentity {
method: &'static MethodSpec,
module: saddle_core::ModuleId,
service: saddle_core::ServiceId,
operation: saddle_core::OperationId,
}
struct FrozenEndpoint {
alias: &'static str,
authority: saddle_boundary::ProfuseContractAuthorityTemplate,
metadata: Box<[saddle_boundary::standard_rpc::FrozenMetadata]>,
}
#[derive(Debug)]
pub(crate) enum RegistrationError {
Empty,
DuplicateMethod,
MissingDependency(&'static str),
UnexpectedDependency(String),
InvalidAuthority(saddle_boundary::BoundaryError),
InvalidMetadata(&'static str),
}
impl std::fmt::Display for RegistrationError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
Self::Empty => f.write_str("standard gRPC registry must select a method"),
Self::DuplicateMethod => f.write_str("duplicate selected gRPC method identity"),
Self::MissingDependency(alias) => write!(f, "missing gRPC dependency `{alias}`"),
Self::UnexpectedDependency(alias) => {
write!(f, "unregistered gRPC dependency `{alias}`")
}
Self::InvalidAuthority(error) => error.fmt(f),
Self::InvalidMetadata(reason) => f.write_str(reason),
}
}
}
impl std::error::Error for RegistrationError {
fn source(&self) -> Option<&(dyn std::error::Error + 'static)> {
match self {
Self::InvalidAuthority(error) => Some(error),
_ => None,
}
}
}
impl FrozenRegistry {
pub(crate) fn freeze(
methods: &'static [&'static MethodSpec],
mut config: std::collections::BTreeMap<String, DependencyConfig>,
) -> Result<Self, RegistrationError> {
if methods.is_empty() {
return Err(RegistrationError::Empty);
}
for (index, method) in methods.iter().enumerate() {
if methods[..index].iter().any(|prior| {
(prior.marker)() == (method.marker)()
|| (prior.dependency == method.dependency
&& prior.service == method.service
&& prior.method == method.method)
}) {
return Err(RegistrationError::DuplicateMethod);
}
}
let mut endpoints = Vec::new();
for method in methods {
if endpoints
.iter()
.any(|entry: &FrozenEndpoint| entry.alias == method.dependency)
{
continue;
}
let input = config
.remove(method.dependency)
.ok_or(RegistrationError::MissingDependency(method.dependency))?;
let mut headers = std::collections::BTreeSet::new();
let metadata = input
.metadata
.into_iter()
.map(|mapping| {
if !headers.insert(mapping.header.clone()) {
return Err(RegistrationError::InvalidMetadata(
"duplicate trace metadata header",
));
}
saddle_boundary::standard_rpc::FrozenMetadata::freeze(mapping)
.map_err(RegistrationError::InvalidMetadata)
})
.collect::<Result<Vec<_>, _>>()?
.into_boxed_slice();
let authority = saddle_boundary::ProfuseContractAuthorityTemplate::new(input.authority)
.and_then(|value| value.with_token(input.token))
.map_err(RegistrationError::InvalidAuthority)?;
endpoints.push(FrozenEndpoint {
alias: method.dependency,
authority,
metadata,
});
}
if let Some((alias, _)) = config.into_iter().next() {
return Err(RegistrationError::UnexpectedDependency(alias));
}
Ok(Self {
methods,
endpoints: endpoints.into(),
identities: methods
.iter()
.map(|method| FrozenMethodIdentity {
method,
module: saddle_core::ModuleId::new(method.dependency),
service: saddle_core::ServiceId::new(method.service),
operation: saddle_core::OperationId::new(method.method),
})
.collect::<Vec<_>>()
.into(),
})
}
pub(crate) fn endpoint<M: Method>(
&self,
) -> Option<&saddle_boundary::ProfuseContractAuthorityTemplate> {
if !registered::<M>(self.methods) {
return None;
}
self.endpoints
.iter()
.find(|entry| entry.alias == M::SPEC.dependency)
.map(|entry| &entry.authority)
}
pub(crate) fn metadata<M: Method>(&self) -> &[saddle_boundary::standard_rpc::FrozenMetadata] {
if !registered::<M>(self.methods) {
return &[];
}
self.endpoints
.iter()
.find(|entry| entry.alias == M::SPEC.dependency)
.map(|entry| entry.metadata.as_ref())
.unwrap_or(&[])
}
pub(crate) fn child_context<M: Method>(
&self,
observer: &saddle_observability::Observer,
parent: &saddle_core::CallContext,
) -> Option<saddle_core::CallContext> {
if !registered::<M>(self.methods) {
return None;
}
let identity = self
.identities
.iter()
.find(|identity| (identity.method.marker)() == std::any::TypeId::of::<M>())?;
Some(observer.managed_child_context(
parent,
identity.module.clone(),
identity.service.clone(),
identity.operation.clone(),
))
}
pub(crate) fn methods(&self) -> &'static [&'static MethodSpec] {
self.methods
}
}
#[cfg(test)]
mod tests {
use super::*;
crate::grpc_bindings! {
mod protocol {
proto_root "../../tests/fixtures/standard-grpc";
dependency mirror { service "alpha.Mirror"; methods { Echo => "Echo"; } }
dependency alternate { service "beta.Alternate"; methods { Echo => "Echo"; } }
}
}
fn config() -> std::collections::BTreeMap<String, DependencyConfig> {
[
("mirror", "http://mirror-{zone}:50051"),
("alternate", "http://127.0.0.1:50052"),
]
.into_iter()
.map(|(alias, authority)| {
(
alias.to_owned(),
DependencyConfig {
authority: authority.to_owned(),
token: None,
metadata: Vec::new(),
},
)
})
.collect()
}
struct Forged;
impl Method for Forged {
type Fields<'a> = <protocol::methods::mirror::Echo as Method>::Fields<'a>;
type View<'a> = <protocol::methods::mirror::Echo as Method>::View<'a>;
const SPEC: MethodSpec = <protocol::methods::mirror::Echo as Method>::SPEC;
fn view<'a>(bytes: &'a [u8]) -> Self::View<'a> {
<protocol::methods::mirror::Echo as Method>::view(bytes)
}
}
#[test]
fn frozen_aliases_and_actual_marker_identity_guard_endpoints() {
let registry = FrozenRegistry::freeze(protocol::METHODS, config()).unwrap();
assert_eq!(
registry
.endpoint::<protocol::methods::mirror::Echo>()
.unwrap()
.resolve("zone-a")
.unwrap()
.as_str(),
"http://mirror-zone-a:50051"
);
assert_eq!(
registry
.endpoint::<protocol::methods::alternate::Echo>()
.unwrap()
.as_str(),
"http://127.0.0.1:50052"
);
assert!(registry.endpoint::<Forged>().is_none());
assert_eq!(registry.methods().len(), 2);
}
#[test]
fn startup_rejects_missing_extra_and_invalid_endpoint_credentials() {
let mut input = config();
input.remove("alternate");
assert!(matches!(
FrozenRegistry::freeze(protocol::METHODS, input),
Err(RegistrationError::MissingDependency("alternate"))
));
let mut input = config();
input.insert(
"unselected".into(),
DependencyConfig {
authority: "http://unused:1".into(),
token: None,
metadata: Vec::new(),
},
);
assert!(matches!(
FrozenRegistry::freeze(protocol::METHODS, input),
Err(RegistrationError::UnexpectedDependency(_))
));
let mut input = config();
input.get_mut("mirror").unwrap().token = Some("invalid\ncredential".into());
assert!(matches!(
FrozenRegistry::freeze(protocol::METHODS, input),
Err(RegistrationError::InvalidAuthority(_))
));
assert!(matches!(
FrozenRegistry::freeze(&[], config()),
Err(RegistrationError::Empty)
));
static DUPLICATE: &[&MethodSpec] = &[
&<protocol::methods::mirror::Echo as Method>::SPEC,
&<protocol::methods::mirror::Echo as Method>::SPEC,
];
assert!(matches!(
FrozenRegistry::freeze(DUPLICATE, config()),
Err(RegistrationError::DuplicateMethod)
));
}
#[test]
fn trace_mapping_rejects_reserved_headers_and_duplicate_destinations() {
use saddle_boundary::standard_rpc::{MetadataConfig, MetadataSource};
for header in [
"authorization",
"grpc-timeout",
"content-type",
"host",
"X-UPPER",
] {
let mut input = config();
input
.get_mut("mirror")
.unwrap()
.metadata
.push(MetadataConfig {
header: header.into(),
source: MetadataSource::TraceId,
});
assert!(matches!(
FrozenRegistry::freeze(protocol::METHODS, input),
Err(RegistrationError::InvalidMetadata(_))
));
}
let mut input = config();
input.get_mut("mirror").unwrap().metadata = vec![
MetadataConfig {
header: "x-trace".into(),
source: MetadataSource::TraceId,
},
MetadataConfig {
header: "x-trace".into(),
source: MetadataSource::RpcId,
},
];
assert!(matches!(
FrozenRegistry::freeze(protocol::METHODS, input),
Err(RegistrationError::InvalidMetadata(_))
));
}
}