use std::{
collections::HashMap,
future::Future,
io::{self, Write},
pin::Pin,
sync::Arc,
};
use saddle_core::{ApplicationId, ErrorKind, Result, SaddleError, TraceId};
use saddle_observability::Observer;
use saddle_runtime::RequestLifecycle;
use serde::{Serialize, de::DeserializeOwned};
use crate::{Service, ServiceClient, ServiceRegistry};
pub const MAX_EXTERNAL_REQUEST_BYTES: usize = 1_048_576;
pub const MAX_EXTERNAL_RESPONSE_BYTES: usize = 1_048_576;
pub const MAX_EXTERNAL_ROUTE_BYTES: usize = 256;
type DispatchFuture<'a> = Pin<Box<dyn Future<Output = Result<Vec<u8>>> + Send + 'a>>;
trait Endpoint: Send + Sync {
fn dispatch<'a>(
&'a self,
context: &'a saddle_core::CallContext,
body: &'a [u8],
) -> DispatchFuture<'a>;
}
struct JsonEndpoint<S: Service> {
client: ServiceClient<S>,
}
impl<S> Endpoint for JsonEndpoint<S>
where
S: Service,
S::Request: DeserializeOwned,
S::Response: Serialize,
{
fn dispatch<'a>(
&'a self,
context: &'a saddle_core::CallContext,
body: &'a [u8],
) -> DispatchFuture<'a> {
Box::pin(async move {
let request = serde_json::from_slice(body).map_err(|_| {
SaddleError::new(
ErrorKind::InvalidArgument,
"service.invalid_request",
"the external request body is not valid for this Service",
)
})?;
let response = self.client.call(context, request).await?;
encode_response(&response)
})
}
}
fn encode_response(response: &impl Serialize) -> Result<Vec<u8>> {
let mut output = BoundedBuffer::new(MAX_EXTERNAL_RESPONSE_BYTES);
if serde_json::to_writer(&mut output, response).is_err() {
return Err(if output.exceeded {
SaddleError::new(
ErrorKind::Internal,
"service.response_too_large",
"the Service response exceeds the external response limit",
)
} else {
SaddleError::new(
ErrorKind::Internal,
"service.response_encoding_failed",
"the Service response could not be encoded",
)
});
}
Ok(output.into_bytes())
}
struct BoundedBuffer {
storage: Vec<u8>,
limit: usize,
exceeded: bool,
}
impl BoundedBuffer {
fn new(limit: usize) -> Self {
Self {
storage: Vec::new(),
limit,
exceeded: false,
}
}
fn into_bytes(self) -> Vec<u8> {
self.storage
}
}
impl Write for BoundedBuffer {
fn write(&mut self, buffer: &[u8]) -> io::Result<usize> {
let Some(end) = self.storage.len().checked_add(buffer.len()) else {
self.exceeded = true;
return Err(io::Error::other("encoded response exceeds limit"));
};
if end > self.limit {
self.exceeded = true;
return Err(io::Error::other("encoded response exceeds limit"));
}
self.storage.extend_from_slice(buffer);
Ok(buffer.len())
}
fn flush(&mut self) -> io::Result<()> {
Ok(())
}
}
pub struct ExternalRequest<'a> {
pub route: &'a str,
pub inbound_trace_id: Option<&'a str>,
pub body: &'a [u8],
}
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
pub enum ExternalStatus {
Ok,
InvalidArgument,
NotFound,
Conflict,
Business,
Unavailable,
Internal,
}
#[derive(Debug)]
pub struct ExternalResponse {
status: ExternalStatus,
body: Option<Vec<u8>>,
error_code: Option<&'static str>,
trace_id: Option<TraceId>,
}
impl ExternalResponse {
pub fn status(&self) -> ExternalStatus {
self.status
}
pub fn body(&self) -> Option<&[u8]> {
self.body.as_deref()
}
pub fn error_code(&self) -> Option<&'static str> {
self.error_code
}
pub fn trace_id(&self) -> Option<TraceId> {
self.trace_id
}
fn success(body: Vec<u8>, trace_id: TraceId) -> Self {
Self {
status: ExternalStatus::Ok,
body: Some(body),
error_code: None,
trace_id: Some(trace_id),
}
}
fn failure(error: &SaddleError, trace_id: Option<TraceId>) -> Self {
Self {
status: ExternalStatus::from_error(error),
body: None,
error_code: Some(error.code()),
trace_id,
}
}
}
impl ExternalStatus {
pub fn from_error(error: &SaddleError) -> Self {
match error.kind() {
ErrorKind::InvalidArgument => Self::InvalidArgument,
ErrorKind::NotFound => Self::NotFound,
ErrorKind::Conflict => Self::Conflict,
ErrorKind::Business => Self::Business,
ErrorKind::Unavailable => Self::Unavailable,
ErrorKind::Infrastructure | ErrorKind::Internal => Self::Internal,
_ => Self::Internal,
}
}
}
trait Admission: Send + Sync {
fn try_accept(&self) -> Result<Box<dyn Send>>;
}
struct RuntimeAdmission(RequestLifecycle);
impl Admission for RuntimeAdmission {
fn try_accept(&self) -> Result<Box<dyn Send>> {
self.0
.try_accept()
.map(|guard| Box::new(guard) as Box<dyn Send>)
}
}
pub struct ExternalDispatcherBuilder {
application: ApplicationId,
registry: ServiceRegistry,
admission: Arc<dyn Admission>,
observer: Observer,
routes: HashMap<String, Arc<dyn Endpoint>>,
}
impl ExternalDispatcherBuilder {
pub fn new(
application: impl Into<ApplicationId>,
registry: ServiceRegistry,
requests: RequestLifecycle,
) -> Self {
let observer = registry.observer();
Self {
application: application.into(),
registry,
admission: Arc::new(RuntimeAdmission(requests)),
observer,
routes: HashMap::new(),
}
}
#[cfg(test)]
fn with_admission(
application: impl Into<ApplicationId>,
registry: ServiceRegistry,
admission: Arc<dyn Admission>,
) -> Self {
let observer = registry.observer();
Self {
application: application.into(),
registry,
admission,
observer,
routes: HashMap::new(),
}
}
pub fn expose_json<S>(&mut self, route: impl Into<String>) -> Result<()>
where
S: Service,
S::Request: DeserializeOwned,
S::Response: Serialize,
{
let route = route.into();
if route.is_empty() {
return Err(SaddleError::new(
ErrorKind::InvalidArgument,
"service.empty_route",
"an external Service route must not be empty",
));
}
if route.len() > MAX_EXTERNAL_ROUTE_BYTES {
return Err(SaddleError::new(
ErrorKind::InvalidArgument,
"service.route_too_large",
"an external Service route exceeds the route limit",
));
}
if self.routes.contains_key(&route) {
return Err(SaddleError::new(
ErrorKind::Conflict,
"service.duplicate_route",
"the external Service route is already registered",
));
}
self.routes.insert(
route,
Arc::new(JsonEndpoint::<S> {
client: self.registry.client::<S>()?,
}),
);
Ok(())
}
pub fn build(self) -> ExternalDispatcher {
ExternalDispatcher {
application: self.application,
admission: self.admission,
observer: self.observer,
routes: self.routes,
}
}
}
pub struct ExternalDispatcher {
application: ApplicationId,
admission: Arc<dyn Admission>,
observer: Observer,
routes: HashMap<String, Arc<dyn Endpoint>>,
}
impl ExternalDispatcher {
pub fn builder(
application: impl Into<ApplicationId>,
registry: ServiceRegistry,
requests: RequestLifecycle,
) -> ExternalDispatcherBuilder {
ExternalDispatcherBuilder::new(application, registry, requests)
}
pub async fn dispatch(&self, request: ExternalRequest<'_>) -> ExternalResponse {
let _request_guard = match self.admission.try_accept() {
Ok(guard) => guard,
Err(error) => return ExternalResponse::failure(&error, None),
};
let (call, _) = self.observer.start_external_call(
self.application.clone(),
"saddle",
"external",
"dispatch",
request.inbound_trace_id,
);
let trace_id = call.context().trace_id();
let result = self.dispatch_admitted(call.context(), request).await;
match &result {
Ok(_) => call.succeed(),
Err(error) => call.fail(error),
}
match result {
Ok(body) => ExternalResponse::success(body, trace_id),
Err(error) => ExternalResponse::failure(&error, Some(trace_id)),
}
}
async fn dispatch_admitted(
&self,
context: &saddle_core::CallContext,
request: ExternalRequest<'_>,
) -> Result<Vec<u8>> {
if request.body.len() > MAX_EXTERNAL_REQUEST_BYTES {
return Err(SaddleError::new(
ErrorKind::InvalidArgument,
"service.request_too_large",
"the external request exceeds the request limit",
));
}
if request.route.len() > MAX_EXTERNAL_ROUTE_BYTES {
return Err(SaddleError::new(
ErrorKind::InvalidArgument,
"service.route_too_large",
"the external Service route exceeds the route limit",
));
}
let endpoint = self.routes.get(request.route).ok_or_else(|| {
SaddleError::new(
ErrorKind::NotFound,
"service.route_not_found",
"no external Service is registered for this route",
)
})?;
endpoint.dispatch(context, request.body).await
}
}
#[cfg(test)]
mod tests {
use std::{
io::{self, Write},
sync::{
Arc, Mutex,
atomic::{AtomicUsize, Ordering},
},
};
use saddle_core::CallContext;
use saddle_observability::ObserverConfig;
use serde::{Deserialize, Serialize};
use tokio::sync::oneshot;
use super::*;
use crate::{ServiceDescriptor, ServiceFuture, ServiceHandler, ServiceRegistryBuilder};
#[derive(Clone, Default)]
struct Capture(Arc<Mutex<Vec<u8>>>);
impl Write for Capture {
fn write(&mut self, buffer: &[u8]) -> io::Result<usize> {
self.0.lock().unwrap().extend_from_slice(buffer);
Ok(buffer.len())
}
fn flush(&mut self) -> io::Result<()> {
Ok(())
}
}
#[derive(Default)]
struct TestAdmission {
active: Arc<AtomicUsize>,
}
struct TestGuard(Arc<AtomicUsize>);
impl Drop for TestGuard {
fn drop(&mut self) {
self.0.fetch_sub(1, Ordering::AcqRel);
}
}
impl Admission for TestAdmission {
fn try_accept(&self) -> Result<Box<dyn Send>> {
self.active.fetch_add(1, Ordering::AcqRel);
Ok(Box::new(TestGuard(Arc::clone(&self.active))))
}
}
struct Echo;
#[derive(Deserialize)]
struct EchoRequest {
value: String,
}
#[derive(Serialize)]
struct EchoResponse {
value: String,
}
impl Service for Echo {
type Request = EchoRequest;
type Response = EchoResponse;
}
struct EchoImpl;
impl ServiceHandler<Echo> for EchoImpl {
fn call<'a>(
&'a self,
_context: &'a CallContext,
request: EchoRequest,
) -> ServiceFuture<'a, EchoResponse> {
Box::pin(async move {
if request.value == "fail" {
Err(SaddleError::new(
ErrorKind::Business,
"echo.rejected",
"sensitive internal reason",
))
} else {
Ok(EchoResponse {
value: request.value,
})
}
})
}
}
struct Large;
impl Service for Large {
type Request = ();
type Response = String;
}
struct LargeImpl;
impl ServiceHandler<Large> for LargeImpl {
fn call<'a>(
&'a self,
_context: &'a CallContext,
_request: (),
) -> ServiceFuture<'a, String> {
Box::pin(async { Ok("x".repeat(MAX_EXTERNAL_RESPONSE_BYTES * 10)) })
}
}
struct Pending;
impl Service for Pending {
type Request = ();
type Response = ();
}
struct PendingImpl {
started: Mutex<Option<oneshot::Sender<()>>>,
release: Mutex<Option<oneshot::Receiver<()>>>,
}
impl ServiceHandler<Pending> for PendingImpl {
fn call<'a>(&'a self, _context: &'a CallContext, _request: ()) -> ServiceFuture<'a, ()> {
Box::pin(async move {
let started = self.started.lock().unwrap().take().unwrap();
let release = self.release.lock().unwrap().take().unwrap();
started.send(()).unwrap();
release.await.unwrap();
Ok(())
})
}
}
fn observer() -> (Observer, Capture) {
let capture = Capture::default();
let observer = Observer::with_writer(ObserverConfig::default(), capture.clone()).unwrap();
(observer, capture)
}
fn builder(observer: Observer, admission: Arc<TestAdmission>) -> ExternalDispatcherBuilder {
let mut registry = ServiceRegistryBuilder::new();
registry
.register::<Echo, _>(ServiceDescriptor::new("test", "echo", "call"), EchoImpl)
.unwrap();
registry
.register::<Large, _>(ServiceDescriptor::new("test", "large", "call"), LargeImpl)
.unwrap();
let registry = registry.build(observer.clone()).unwrap();
ExternalDispatcherBuilder::with_admission("app", registry, admission)
}
#[tokio::test]
async fn dispatcher_wires_success_and_safe_business_failure() {
let (observer, _) = observer();
let admission = Arc::new(TestAdmission::default());
let mut builder = builder(observer, admission);
builder.expose_json::<Echo>("/echo").unwrap();
let dispatcher = builder.build();
let success = dispatcher
.dispatch(ExternalRequest {
route: "/echo",
inbound_trace_id: None,
body: br#"{"value":"hello"}"#,
})
.await;
assert_eq!(success.status(), ExternalStatus::Ok);
assert_eq!(success.body(), Some(br#"{"value":"hello"}"#.as_slice()));
let failure = dispatcher
.dispatch(ExternalRequest {
route: "/echo",
inbound_trace_id: None,
body: br#"{"value":"fail"}"#,
})
.await;
assert_eq!(failure.status(), ExternalStatus::Business);
assert_eq!(failure.error_code(), Some("echo.rejected"));
assert_eq!(failure.body(), None);
assert!(failure.trace_id().is_some());
let framework_failure = dispatcher
.dispatch(ExternalRequest {
route: "/echo",
inbound_trace_id: None,
body: b"not json",
})
.await;
assert_eq!(framework_failure.status(), ExternalStatus::InvalidArgument);
assert_eq!(
framework_failure.error_code(),
Some("service.invalid_request")
);
assert_eq!(framework_failure.body(), None);
}
#[tokio::test]
async fn unknown_route_still_has_balanced_external_trace() {
let (observer, capture) = observer();
let admission = Arc::new(TestAdmission::default());
let dispatcher = builder(observer.clone(), admission).build();
let response = dispatcher
.dispatch(ExternalRequest {
route: "/missing",
inbound_trace_id: None,
body: b"{}",
})
.await;
assert_eq!(response.error_code(), Some("service.route_not_found"));
let response_trace = response.trace_id().unwrap().to_string();
observer.flush().await.unwrap();
let bytes = capture.0.lock().unwrap().clone();
let records: Vec<serde_json::Value> = bytes
.split(|byte| *byte == b'\n')
.filter(|line| !line.is_empty())
.map(|line| serde_json::from_slice(line).unwrap())
.collect();
let external: Vec<_> = records
.iter()
.filter(|record| record["call_kind"] == "external_request")
.collect();
assert_eq!(external.len(), 2);
assert_eq!(external[0]["trace_id"], external[1]["trace_id"]);
assert_eq!(external[0]["trace_id"], response_trace);
assert_eq!(external[1]["outcome"], "failure");
}
#[tokio::test]
async fn rejects_oversized_request_and_response() {
let (observer, _) = observer();
let admission = Arc::new(TestAdmission::default());
let mut builder = builder(observer, admission);
builder.expose_json::<Echo>("/echo").unwrap();
builder.expose_json::<Large>("/large").unwrap();
let dispatcher = builder.build();
let request = vec![b'x'; MAX_EXTERNAL_REQUEST_BYTES + 1];
assert_eq!(
dispatcher
.dispatch(ExternalRequest {
route: "/echo",
inbound_trace_id: None,
body: &request,
})
.await
.error_code(),
Some("service.request_too_large")
);
assert_eq!(
dispatcher
.dispatch(ExternalRequest {
route: "/large",
inbound_trace_id: None,
body: b"null",
})
.await
.error_code(),
Some("service.response_too_large")
);
}
#[test]
fn response_sink_stops_at_fixed_limit_and_preserves_other_encoding_errors() {
let mut output = BoundedBuffer::new(64);
let large = "x".repeat(1_024);
assert!(serde_json::to_writer(&mut output, &large).is_err());
assert!(output.exceeded);
assert!(output.storage.len() <= 64);
let mut small = BoundedBuffer::new(MAX_EXTERNAL_RESPONSE_BYTES);
serde_json::to_writer(&mut small, "ok").unwrap();
assert_eq!(small.storage, br#""ok""#);
assert!(small.storage.capacity() < 1_024);
struct Broken;
impl Serialize for Broken {
fn serialize<S>(&self, _serializer: S) -> std::result::Result<S::Ok, S::Error>
where
S: serde::Serializer,
{
Err(serde::ser::Error::custom("broken"))
}
}
assert_eq!(
encode_response(&Broken).unwrap_err().code(),
"service.response_encoding_failed"
);
}
#[tokio::test]
async fn route_limit_is_identical_at_registration_and_dispatch() {
let (observer, _) = observer();
let admission = Arc::new(TestAdmission::default());
let mut builder = builder(observer, admission);
let oversized = "x".repeat(MAX_EXTERNAL_ROUTE_BYTES + 1);
assert_eq!(
builder
.expose_json::<Echo>(oversized.clone())
.unwrap_err()
.code(),
"service.route_too_large"
);
let dispatcher = builder.build();
let response = dispatcher
.dispatch(ExternalRequest {
route: &oversized,
inbound_trace_id: None,
body: b"{}",
})
.await;
assert_eq!(response.error_code(), Some("service.route_too_large"));
assert!(response.trace_id().is_some());
}
#[tokio::test]
async fn pending_and_cancelled_dispatch_hold_then_release_admission_guard() {
let (observer, _) = observer();
let admission = Arc::new(TestAdmission::default());
let active = Arc::clone(&admission.active);
let (started_tx, started_rx) = oneshot::channel();
let (release_tx, release_rx) = oneshot::channel();
let mut registry = ServiceRegistryBuilder::new();
registry
.register::<Pending, _>(
ServiceDescriptor::new("test", "pending", "call"),
PendingImpl {
started: Mutex::new(Some(started_tx)),
release: Mutex::new(Some(release_rx)),
},
)
.unwrap();
let registry = registry.build(observer.clone()).unwrap();
let mut builder = ExternalDispatcherBuilder::with_admission("app", registry, admission);
builder.expose_json::<Pending>("/pending").unwrap();
let dispatcher = Arc::new(builder.build());
let task = {
let dispatcher = Arc::clone(&dispatcher);
tokio::spawn(async move {
dispatcher
.dispatch(ExternalRequest {
route: "/pending",
inbound_trace_id: None,
body: b"null",
})
.await
})
};
started_rx.await.unwrap();
assert_eq!(active.load(Ordering::Acquire), 1);
task.abort();
assert!(task.await.unwrap_err().is_cancelled());
assert_eq!(active.load(Ordering::Acquire), 0);
drop(release_tx);
}
}