use std::any::{Any, TypeId};
use std::collections::BTreeMap;
use std::future::Future;
use std::ops::Deref;
use std::pin::Pin;
use std::sync::{Arc, Weak};
use futures_util::{Stream, StreamExt};
use serde::de::DeserializeOwned;
use serde::Serialize;
use serde_json::Value;
use unb_core::{Envelope, ErrorCode};
use crate::handler::HandlerError;
use crate::layer::{ErasedCall, Origin, ServiceBody};
use crate::node::{Node, NodeSnapshot};
#[derive(Clone, Copy, Debug, PartialEq, Eq, PartialOrd, Ord)]
pub enum Operation {
Unary,
Streaming,
}
impl Operation {
pub fn of<T>(request: &http::Request<T>) -> Option<Operation> {
let kind = request
.headers()
.get(unb_core::UNB_KIND)
.map(|value| value.to_str().ok());
match kind {
None | Some(Some("request")) | Some(Some("discover")) => Some(Operation::Unary),
Some(Some("subscribe")) | Some(Some("channel")) => Some(Operation::Streaming),
_ => None,
}
}
}
pub struct Request<T> {
payload: T,
parts: http::request::Parts,
origin: Origin,
node: Weak<Node>,
_snapshot: Arc<NodeSnapshot>,
}
impl<T> Request<T> {
pub fn payload(&self) -> &T {
&self.payload
}
pub fn into_payload(self) -> T {
self.payload
}
pub fn subject(&self) -> String {
Envelope::subject_of(&self.parts.uri)
}
pub fn method(&self) -> &http::Method {
&self.parts.method
}
pub fn headers(&self) -> &http::HeaderMap {
&self.parts.headers
}
pub fn extensions(&self) -> &http::Extensions {
&self.parts.extensions
}
pub fn take_body_stream(&self) -> Option<unb_runtime::BodyStream> {
self.parts
.extensions
.get::<StreamingBody>()?
.0
.lock()
.expect("streaming body slot")
.take()
}
pub fn origin(&self) -> &Origin {
&self.origin
}
pub async fn call(&self, subject: &str, payload: Value) -> Result<Value, HandlerError> {
let Some(node) = self.node.upgrade() else {
return Err(HandlerError::new(
ErrorCode::Internal,
"the node behind this request has shut down",
));
};
let mut headers = serde_json::Map::new();
for (name, value) in &self.parts.headers {
if name.as_str().starts_with("unb-") {
continue;
}
let Ok(value) = std::str::from_utf8(value.as_bytes()) else {
continue;
};
headers.insert(name.as_str().to_string(), Value::String(value.to_string()));
}
node.call_nested(subject, payload, headers).await
}
}
impl<T: DeserializeOwned> Request<T> {
pub(crate) fn decode(request: http::Request<bytes::Bytes>) -> Result<Request<T>, HandlerError> {
let (parts, body) = request.into_parts();
let decoded = if body.is_empty() {
serde_json::from_value(Value::Null)
} else {
serde_json::from_slice(&body)
};
let payload: T = decoded.map_err(|error| {
let input = std::any::type_name::<T>()
.rsplit("::")
.next()
.unwrap_or("input");
let subject = Envelope::subject_of(&parts.uri);
HandlerError::new(
ErrorCode::InvalidInput,
format!(
"payload does not match {input}, the declared input for {subject:?}: {error}"
),
)
})?;
let origin = parts
.extensions
.get::<Origin>()
.cloned()
.unwrap_or(Origin::Local);
let node = parts
.extensions
.get::<Weak<Node>>()
.cloned()
.unwrap_or_default();
let snapshot = parts
.extensions
.get::<Arc<NodeSnapshot>>()
.cloned()
.ok_or_else(|| HandlerError::new(ErrorCode::Internal, "missing invocation snapshot"))?;
Ok(Request {
payload,
parts,
origin,
node,
_snapshot: snapshot,
})
}
}
#[derive(Clone)]
pub(crate) struct StreamingBody(pub(crate) Arc<std::sync::Mutex<Option<unb_runtime::BodyStream>>>);
pub struct Reply<T>(T);
impl<T> Reply<T> {
pub fn new(value: T) -> Reply<T> {
Reply(value)
}
}
enum StreamBody<T, E> {
Typed(Pin<Box<dyn Stream<Item = Result<T, E>> + Send>>),
Raw(crate::handler::EventStream),
}
pub struct Streaming<T, E> {
body: StreamBody<T, E>,
}
impl<T, E> Streaming<T, E> {
pub fn new(stream: impl Stream<Item = Result<T, E>> + Send + 'static) -> Streaming<T, E> {
Streaming {
body: StreamBody::Typed(Box::pin(stream)),
}
}
pub fn raw(
stream: impl Stream<Item = Result<bytes::Bytes, HandlerError>> + Send + 'static,
) -> Streaming<T, E> {
Streaming {
body: StreamBody::Raw(Box::pin(stream)),
}
}
}
mod sealed {
pub trait Sealed {}
}
pub trait HandlerOutput: sealed::Sealed {
const OPERATION: Operation;
fn into_response(self) -> Result<http::Response<ServiceBody>, HandlerError>;
}
fn respond(body: ServiceBody) -> Result<http::Response<ServiceBody>, HandlerError> {
http::Response::builder().body(body).map_err(|error| {
HandlerError::new(
ErrorCode::Internal,
format!("response construction failed: {error}"),
)
})
}
impl<T: Serialize> sealed::Sealed for Reply<T> {}
impl<T: Serialize> HandlerOutput for Reply<T> {
const OPERATION: Operation = Operation::Unary;
fn into_response(self) -> Result<http::Response<ServiceBody>, HandlerError> {
let value = serde_json::to_value(self.0).map_err(|error| {
HandlerError::new(
ErrorCode::Internal,
format!("response serialization failed: {error}"),
)
})?;
respond(ServiceBody::Unary(Envelope::encode_payload(&value)))
}
}
impl<T, E> sealed::Sealed for Streaming<T, E>
where
T: Serialize + Send + 'static,
E: Into<HandlerError> + Send + 'static,
{
}
impl<T, E> HandlerOutput for Streaming<T, E>
where
T: Serialize + Send + 'static,
E: Into<HandlerError> + Send + 'static,
{
const OPERATION: Operation = Operation::Streaming;
fn into_response(self) -> Result<http::Response<ServiceBody>, HandlerError> {
let body = match self.body {
StreamBody::Typed(stream) => ServiceBody::Stream(Box::pin(stream.map(|item| {
match item {
Ok(event) => serde_json::to_value(event)
.map(|value| Envelope::encode_payload(&value))
.map_err(|error| {
HandlerError::new(
ErrorCode::Internal,
format!("event serialization failed: {error}"),
)
}),
Err(error) => Err(error.into()),
}
}))),
StreamBody::Raw(stream) => ServiceBody::Stream(stream),
};
respond(body)
}
}
pub struct State<T>(Arc<T>);
impl<T> State<T> {
pub fn new(value: T) -> State<T> {
State(Arc::new(value))
}
}
impl<T> Clone for State<T> {
fn clone(&self) -> State<T> {
State(self.0.clone())
}
}
impl<T> Deref for State<T> {
type Target = T;
fn deref(&self) -> &T {
&self.0
}
}
#[derive(Default, Clone)]
pub(crate) struct StateMap {
values: BTreeMap<TypeId, Arc<dyn Any + Send + Sync>>,
}
impl StateMap {
pub(crate) fn insert<T: Send + Sync + 'static>(&mut self, value: T) {
self.values.insert(TypeId::of::<T>(), Arc::new(value));
}
pub(crate) fn get<T: Send + Sync + 'static>(&self) -> Option<State<T>> {
self.values
.get(&TypeId::of::<T>())
.cloned()
.and_then(|any| any.downcast::<T>().ok())
.map(State)
}
pub(crate) fn merged_over(&self, outer: &StateMap) -> StateMap {
let mut merged = outer.clone();
for (key, value) in &self.values {
merged.values.insert(*key, value.clone());
}
merged
}
}
pub struct States<'a>(pub(crate) &'a StateMap);
impl States<'_> {
pub fn state<T: Send + Sync + 'static>(&self) -> Result<State<T>, String> {
self.0.get::<T>().ok_or_else(|| {
format!(
"no registered state provides {}; register it with state(...) on the node or scope",
std::any::type_name::<T>()
)
})
}
}
pub enum ContractSchema {
Static(fn() -> Value),
Owned(Value),
}
impl From<fn() -> Value> for ContractSchema {
fn from(factory: fn() -> Value) -> ContractSchema {
ContractSchema::Static(factory)
}
}
impl From<Value> for ContractSchema {
fn from(value: Value) -> ContractSchema {
ContractSchema::Owned(value)
}
}
impl ContractSchema {
fn value(&self) -> Value {
match self {
ContractSchema::Static(factory) => factory(),
ContractSchema::Owned(value) => value.clone(),
}
}
}
pub struct OperationContract {
pub input: Option<ContractSchema>,
pub output: Option<ContractSchema>,
pub event: Option<ContractSchema>,
pub error: Option<ContractSchema>,
}
impl OperationContract {
pub fn unknown() -> OperationContract {
OperationContract {
input: None,
output: None,
event: None,
error: None,
}
}
pub(crate) fn to_json(&self, operation: Operation) -> Value {
let render = |schema: &Option<ContractSchema>| {
schema
.as_ref()
.map(ContractSchema::value)
.unwrap_or_else(|| serde_json::json!({ "unknown": true }))
};
match operation {
Operation::Unary => serde_json::json!({
"input_schema": render(&self.input),
"output_schema": render(&self.output),
}),
Operation::Streaming => serde_json::json!({
"input_schema": render(&self.input),
"event_schema": render(&self.event),
"error_schema": render(&self.error),
}),
}
}
}
type BuildFn = Box<dyn FnOnce(&States<'_>) -> Result<ErasedCall, String> + Send>;
pub struct HandlerService {
pub(crate) local_name: String,
pub(crate) subject_override: Option<String>,
pub(crate) one_line: Option<String>,
pub(crate) operation: Operation,
pub(crate) metadata: Option<Value>,
pub(crate) contract: OperationContract,
pub(crate) build: BuildFn,
}
impl HandlerService {
pub fn declare(
local_name: &str,
subject_override: Option<&str>,
one_line: Option<&str>,
operation: Operation,
contract: OperationContract,
build: impl FnOnce(&States<'_>) -> Result<ErasedCall, String> + Send + 'static,
) -> HandlerService {
HandlerService {
local_name: local_name.into(),
subject_override: subject_override.map(Into::into),
one_line: one_line.map(Into::into),
operation,
metadata: None,
contract,
build: Box::new(build),
}
}
pub fn at_subject(mut self, subject: impl Into<String>) -> HandlerService {
self.subject_override = Some(subject.into());
self
}
pub fn describe(mut self, metadata: Value) -> HandlerService {
self.metadata = Some(metadata);
self
}
pub(crate) fn effective_subject(&self, scopes: &[String]) -> Result<String, String> {
let local = self.subject_override.as_deref().unwrap_or(&self.local_name);
let subject = if scopes.is_empty() {
local.to_string()
} else {
format!("{}.{local}", scopes.join("."))
};
let valid = !subject.is_empty()
&& subject.len() <= unb_core::MAX_SUBJECT_LEN
&& subject.split('.').all(|segment| !segment.is_empty());
if valid {
Ok(subject)
} else {
Err(format!(
"subject {subject:?} needs non-empty dot-separated segments within {} bytes",
unb_core::MAX_SUBJECT_LEN
))
}
}
}
pub trait Handler: Sized {
fn into_service(self) -> HandlerService;
fn at_subject(self, subject: impl Into<String>) -> HandlerService {
self.into_service().at_subject(subject)
}
fn describe(self, metadata: Value) -> HandlerService {
self.into_service().describe(metadata)
}
}
impl Handler for HandlerService {
fn into_service(self) -> HandlerService {
self
}
}
pub fn erase_unary<F, Fut, In, Out, E>(f: F) -> ErasedCall
where
F: Fn(Request<In>) -> Fut + Send + Sync + 'static,
Fut: Future<Output = Result<Reply<Out>, E>> + Send + 'static,
In: DeserializeOwned + Send + 'static,
Out: Serialize + Send + 'static,
E: Into<HandlerError> + Send + 'static,
{
let f = Arc::new(f);
Arc::new(move |request: http::Request<bytes::Bytes>| {
let f = f.clone();
Box::pin(async move {
let request = Request::<In>::decode(request)?;
match f(request).await {
Ok(reply) => reply.into_response(),
Err(error) => Err(error.into()),
}
})
})
}
pub fn erase_streaming<F, Fut, In, Event, StreamError, E>(f: F) -> ErasedCall
where
F: Fn(Request<In>) -> Fut + Send + Sync + 'static,
Fut: Future<Output = Result<Streaming<Event, StreamError>, E>> + Send + 'static,
In: DeserializeOwned + Send + 'static,
Event: Serialize + Send + 'static,
StreamError: Into<HandlerError> + Send + 'static,
E: Into<HandlerError> + Send + 'static,
{
let f = Arc::new(f);
Arc::new(move |request: http::Request<bytes::Bytes>| {
let f = f.clone();
Box::pin(async move {
let request = Request::<In>::decode(request)?;
match f(request).await {
Ok(streaming) => streaming.into_response(),
Err(error) => Err(error.into()),
}
})
})
}
#[cfg(test)]
mod tests {
use serde::Deserialize;
use serde_json::json;
use super::*;
fn service_request(payload: Value) -> http::Request<bytes::Bytes> {
let node = Node::builder("service-test")
.insecure_accept_declared_peer_identities()
.build()
.expect("test node builds");
let mut request = http::Request::builder()
.method("POST")
.uri("/probe")
.body(Envelope::encode_payload(&payload))
.expect("test request is well formed");
request.extensions_mut().insert(Origin::Local);
request.extensions_mut().insert(node.snapshot.load_full());
request
}
#[derive(Deserialize)]
struct Input {
n: u32,
}
#[derive(Serialize)]
struct Output {
doubled: u32,
}
#[tokio::test]
async fn a_unary_reply_serializes_through_the_erased_adapter() {
let call = erase_unary(|request: Request<Input>| async move {
Ok::<_, HandlerError>(Reply::new(Output {
doubled: request.payload().n * 2,
}))
});
let response = call(service_request(json!({ "n": 21 })))
.await
.unwrap_or_else(|error| panic!("call failed: {error}"));
let ServiceBody::Unary(payload) = response.into_body() else {
panic!("expected a unary response");
};
let value: Value = serde_json::from_slice(&payload).expect("unary payload is json");
assert_eq!(value["doubled"], 42);
}
#[tokio::test]
async fn an_undeclared_payload_fails_decode_before_the_handler_runs() {
let call = erase_unary(|_request: Request<Input>| async move {
panic!("the handler must not run on an invalid payload");
#[allow(unreachable_code)]
Ok::<_, HandlerError>(Reply::new(Value::Null))
});
let error = match call(service_request(json!({ "n": "not-a-number" }))).await {
Err(error) => error,
Ok(_) => panic!("decode must fail"),
};
assert_eq!(error.code, ErrorCode::InvalidInput);
assert!(error.message.contains("probe"));
}
#[tokio::test]
async fn a_streaming_output_maps_events_and_errors_into_the_event_stream() {
let call = erase_streaming(|_request: Request<Input>| async move {
let events = futures_util::stream::iter(vec![
Ok(Output { doubled: 2 }),
Err(HandlerError::new(ErrorCode::Internal, "stream broke")),
]);
Ok::<_, HandlerError>(Streaming::new(events))
});
let response = call(service_request(json!({ "n": 1 })))
.await
.unwrap_or_else(|error| panic!("call failed: {error}"));
let ServiceBody::Stream(mut stream) = response.into_body() else {
panic!("expected a stream response");
};
let first = stream.next().await.unwrap().unwrap();
let first: Value = serde_json::from_slice(&first).unwrap();
assert_eq!(first["doubled"], 2);
let second = stream.next().await.unwrap().unwrap_err();
assert_eq!(second.code, ErrorCode::Internal);
assert!(stream.next().await.is_none());
}
#[test]
fn nearest_scope_state_wins_over_outer_state() {
let mut node_states = StateMap::default();
node_states.insert(7u32);
node_states.insert("node".to_string());
let mut scope_states = StateMap::default();
scope_states.insert("scope".to_string());
let merged = scope_states.merged_over(&node_states);
assert_eq!(*merged.get::<String>().unwrap(), "scope");
assert_eq!(*merged.get::<u32>().unwrap(), 7);
let missing = match States(&merged).state::<bool>() {
Err(missing) => missing,
Ok(_) => panic!("bool state must be absent"),
};
assert!(missing.contains("bool"));
}
}