use std::sync::Arc;
use std::time::Duration;
use alkcall::core::auth::AuthContext;
use alkcall::core::types::{Connection, HandlerError, StreamError};
use alkcall::registry::registration::OperationRegistry;
use async_trait::async_trait;
use axum::routing::get;
use axum::Router;
use hyper_util::rt::{TokioExecutor, TokioIo, TokioTimer};
use hyper_util::server::conn::auto::Builder as HyperBuilder;
use hyper_util::service::TowerToHyperService;
use parking_lot::Mutex;
use tracing::error;
use super::auth::bearer_auth_middleware;
use super::decoy::{decoy_fallback, decoy_method_not_allowed};
use super::healthz::healthz;
use super::state::{DecoyConfig, RouterState};
pub const ALPN_HTTP1: &[u8] = b"http/1.1";
pub const ALPN_H2: &[u8] = b"h2";
pub const WS_UPGRADE_PATH: &str = "/alk/channels";
pub const RESERVED_PATHS: &[&str] = &[
"/search",
"/schema",
"/call",
"/batch",
"/subscribe",
"/publish",
"/healthz",
"/openapi.json",
"/mcp",
WS_UPGRADE_PATH,
];
pub struct HttpAdapter {
identity_provider: Arc<dyn alkcall::core::auth::IdentityProvider>,
registry: Arc<OperationRegistry>,
decoy: DecoyConfig,
extra_routes: Option<Router>,
alpn: &'static [u8],
router: Router,
openapi_doc: CachedOpenAPIDoc,
ws_sessions: Arc<crate::websocket::WsSessions>,
ws_max_sessions: usize,
ws_session_slots: Arc<tokio::sync::Semaphore>,
ws_idle_timeout: Option<Duration>,
ws_openable_alpns: Option<Arc<[crate::websocket::OpenableAlpn]>>,
ws_op_register_acl: alkcall::registry::spec::AccessControl,
}
impl HttpAdapter {
pub fn new(
identity_provider: Arc<dyn alkcall::core::auth::IdentityProvider>,
registry: Arc<OperationRegistry>,
) -> Self {
Self::for_alpn(identity_provider, registry, ALPN_HTTP1)
}
pub fn h2(
identity_provider: Arc<dyn alkcall::core::auth::IdentityProvider>,
registry: Arc<OperationRegistry>,
) -> Self {
Self::for_alpn(identity_provider, registry, ALPN_H2)
}
fn for_alpn(
identity_provider: Arc<dyn alkcall::core::auth::IdentityProvider>,
registry: Arc<OperationRegistry>,
alpn: &'static [u8],
) -> Self {
let decoy = DecoyConfig::default();
let openapi_doc = CachedOpenAPIDoc::new(®istry);
let ws_sessions = Arc::new(crate::websocket::WsSessions::new());
let ws_max_sessions = crate::websocket::DEFAULT_WS_MAX_SESSIONS;
let ws_session_slots = Arc::new(tokio::sync::Semaphore::new(ws_max_sessions));
let ws_idle_timeout = Some(crate::websocket::DEFAULT_WS_IDLE_TIMEOUT);
let state = RouterState {
registry: Arc::clone(®istry),
identity_provider: Arc::clone(&identity_provider),
decoy: decoy.clone(),
openapi_doc: openapi_doc.clone(),
ws_sessions: Arc::clone(&ws_sessions),
ws_session_slots: Arc::clone(&ws_session_slots),
ws_idle_timeout,
ws_openable_alpns: None,
ws_op_register_acl: alkcall::registry::spec::AccessControl::default(),
};
let router = build_router(state, None);
Self {
identity_provider,
registry,
decoy,
extra_routes: None,
alpn,
router,
openapi_doc,
ws_sessions,
ws_max_sessions,
ws_session_slots,
ws_idle_timeout,
ws_openable_alpns: None,
ws_op_register_acl: alkcall::registry::spec::AccessControl::default(),
}
}
pub fn with_decoy(mut self, decoy: DecoyConfig) -> Self {
self.decoy = decoy.clone();
let state = RouterState {
registry: Arc::clone(&self.registry),
identity_provider: Arc::clone(&self.identity_provider),
decoy,
openapi_doc: self.openapi_doc.clone(),
ws_sessions: Arc::clone(&self.ws_sessions),
ws_session_slots: Arc::clone(&self.ws_session_slots),
ws_idle_timeout: self.ws_idle_timeout,
ws_openable_alpns: self.ws_openable_alpns.clone(),
ws_op_register_acl: self.ws_op_register_acl.clone(),
};
self.router = build_router(state, self.extra_routes.clone());
self
}
pub fn with_extra_routes(mut self, routes: Router) -> Self {
let state = RouterState {
registry: Arc::clone(&self.registry),
identity_provider: Arc::clone(&self.identity_provider),
decoy: self.decoy.clone(),
openapi_doc: self.openapi_doc.clone(),
ws_sessions: Arc::clone(&self.ws_sessions),
ws_session_slots: Arc::clone(&self.ws_session_slots),
ws_idle_timeout: self.ws_idle_timeout,
ws_openable_alpns: self.ws_openable_alpns.clone(),
ws_op_register_acl: self.ws_op_register_acl.clone(),
};
self.router = build_router(state, Some(routes.clone()));
self.extra_routes = Some(routes);
self
}
pub fn with_ws_max_sessions(mut self, max_sessions: usize) -> Self {
self.ws_max_sessions = max_sessions;
let state = RouterState {
registry: Arc::clone(&self.registry),
identity_provider: Arc::clone(&self.identity_provider),
decoy: self.decoy.clone(),
openapi_doc: self.openapi_doc.clone(),
ws_sessions: Arc::clone(&self.ws_sessions),
ws_session_slots: Self::rebuild_session_slots(max_sessions),
ws_idle_timeout: self.ws_idle_timeout,
ws_openable_alpns: self.ws_openable_alpns.clone(),
ws_op_register_acl: self.ws_op_register_acl.clone(),
};
self.router = build_router(state, self.extra_routes.clone());
self
}
pub fn with_ws_idle_timeout(mut self, idle_timeout: Option<Duration>) -> Self {
self.ws_idle_timeout = idle_timeout;
let state = RouterState {
registry: Arc::clone(&self.registry),
identity_provider: Arc::clone(&self.identity_provider),
decoy: self.decoy.clone(),
openapi_doc: self.openapi_doc.clone(),
ws_sessions: Arc::clone(&self.ws_sessions),
ws_session_slots: Arc::clone(&self.ws_session_slots),
ws_idle_timeout: self.ws_idle_timeout,
ws_openable_alpns: self.ws_openable_alpns.clone(),
ws_op_register_acl: self.ws_op_register_acl.clone(),
};
self.router = build_router(state, self.extra_routes.clone());
self
}
pub fn ws_sessions(&self) -> Arc<crate::websocket::WsSessions> {
Arc::clone(&self.ws_sessions)
}
pub fn with_ws_openable_alpns(
mut self,
openables: Vec<crate::websocket::OpenableAlpn>,
) -> Self {
self.ws_openable_alpns = Some(openables.into());
let state = RouterState {
registry: Arc::clone(&self.registry),
identity_provider: Arc::clone(&self.identity_provider),
decoy: self.decoy.clone(),
openapi_doc: self.openapi_doc.clone(),
ws_sessions: Arc::clone(&self.ws_sessions),
ws_session_slots: Arc::clone(&self.ws_session_slots),
ws_idle_timeout: self.ws_idle_timeout,
ws_openable_alpns: self.ws_openable_alpns.clone(),
ws_op_register_acl: self.ws_op_register_acl.clone(),
};
self.router = build_router(state, self.extra_routes.clone());
self
}
pub fn with_ws_op_register_acl(mut self, acl: alkcall::registry::spec::AccessControl) -> Self {
self.ws_op_register_acl = acl;
let state = RouterState {
registry: Arc::clone(&self.registry),
identity_provider: Arc::clone(&self.identity_provider),
decoy: self.decoy.clone(),
openapi_doc: self.openapi_doc.clone(),
ws_sessions: Arc::clone(&self.ws_sessions),
ws_session_slots: Arc::clone(&self.ws_session_slots),
ws_idle_timeout: self.ws_idle_timeout,
ws_openable_alpns: self.ws_openable_alpns.clone(),
ws_op_register_acl: self.ws_op_register_acl.clone(),
};
self.router = build_router(state, self.extra_routes.clone());
self
}
fn rebuild_session_slots(max_sessions: usize) -> Arc<tokio::sync::Semaphore> {
Arc::new(tokio::sync::Semaphore::new(max_sessions))
}
pub fn decoy(&self) -> &DecoyConfig {
&self.decoy
}
pub fn alpn(&self) -> &'static [u8] {
self.alpn
}
pub fn router(&self) -> &Router {
&self.router
}
}
fn build_router(state: RouterState, extra_routes: Option<Router>) -> Router {
let auth_state = Arc::clone(&state.identity_provider);
#[cfg(feature = "mcp")]
let mcp_router: Router<RouterState> = {
let dispatch = crate::gateway::GatewayDispatch::new(Arc::clone(&state.registry));
Router::new()
.nest_service(
"/mcp",
crate::adapters::to_mcp_service(std::sync::Arc::new(dispatch)),
)
.layer(axum::middleware::from_fn(mcp_body_limit))
.layer(from_fn_with_state(
auth_state.clone(),
bearer_auth_middleware,
))
};
#[cfg(not(feature = "mcp"))]
let mcp_router: Router<RouterState> = Router::new();
let default: Router<RouterState> = Router::new()
.merge(crate::gateway::routes::gateway_router())
.route("/openapi.json", get(openapi_json_handler))
.route("/healthz", get(healthz))
.fallback(decoy_fallback)
.method_not_allowed_fallback(decoy_method_not_allowed);
let with_extras = match extra_routes {
Some(extra) => {
enforce_reserved_paths(&extra);
let extra: Router<RouterState> = extra.with_state(());
default.merge(extra)
}
None => default,
};
let with_extras = with_extras.method_not_allowed_fallback(decoy_method_not_allowed);
let with_auth = with_extras.route_layer(from_fn_with_state(
Arc::clone(&auth_state),
bearer_auth_middleware,
));
let with_mcp = with_auth.merge(mcp_router);
with_mcp
.route(
WS_UPGRADE_PATH,
get(crate::websocket::ws_upgrade_handler)
.fallback(decoy_method_not_allowed)
.route_layer(from_fn_with_state(
Arc::clone(&auth_state),
crate::websocket::ws_bearer_auth,
)),
)
.with_state(state)
}
fn enforce_reserved_paths(extra: &Router) {
if !extra.has_routes() {
return;
}
let probe = RESERVED_PATHS.iter().fold(Router::new(), |router, path| {
router.route(
path,
get(rejected_reserved_path)
.post(rejected_reserved_path)
.put(rejected_reserved_path)
.patch(rejected_reserved_path)
.delete(rejected_reserved_path)
.head(rejected_reserved_path)
.options(rejected_reserved_path)
.trace(rejected_reserved_path)
.connect(rejected_reserved_path),
)
});
let _ = Router::new().merge(probe).merge(extra.clone());
}
async fn rejected_reserved_path() -> axum::response::Response {
unreachable!("reserved-path probe handler is never called")
}
use axum::middleware::from_fn_with_state;
#[cfg(feature = "mcp")]
const MCP_BODY_LIMIT: usize = 8 * 1024 * 1024;
#[cfg(feature = "mcp")]
const MCP_BODY_LIMIT_EXCEEDED: &str = "mcp body limit exceeded";
#[cfg(feature = "mcp")]
async fn mcp_body_limit(
req: axum::extract::Request,
next: axum::middleware::Next,
) -> axum::response::Response {
use axum::response::IntoResponse;
let (parts, body) = req.into_parts();
if let Some(len) = parts
.headers
.get(http::header::CONTENT_LENGTH)
.and_then(|v| v.to_str().ok())
.and_then(|v| v.parse::<usize>().ok())
{
if len > MCP_BODY_LIMIT {
return (
axum::http::StatusCode::PAYLOAD_TOO_LARGE,
"Payload Too Large: /mcp body exceeds the 8 MiB limit",
)
.into_response();
}
}
let exceeded = Arc::new(std::sync::atomic::AtomicBool::new(false));
let counting = CountingBody {
inner: body.into_data_stream(),
remaining: MCP_BODY_LIMIT,
exceeded: Arc::clone(&exceeded),
};
let mut limited_req =
axum::extract::Request::from_parts(parts, axum::body::Body::from_stream(counting));
limited_req.extensions_mut().insert(exceeded.clone());
let response = next.run(limited_req).await;
if exceeded.load(std::sync::atomic::Ordering::Relaxed) {
return (
axum::http::StatusCode::PAYLOAD_TOO_LARGE,
"Payload Too Large: /mcp body exceeds the 8 MiB limit",
)
.into_response();
}
response
}
#[cfg(feature = "mcp")]
struct CountingBody {
inner: axum::body::BodyDataStream,
remaining: usize,
exceeded: Arc<std::sync::atomic::AtomicBool>,
}
#[cfg(feature = "mcp")]
impl futures::Stream for CountingBody {
type Item = Result<axum::body::Bytes, std::io::Error>;
fn poll_next(
mut self: std::pin::Pin<&mut Self>,
cx: &mut std::task::Context<'_>,
) -> std::task::Poll<Option<Self::Item>> {
let this = &mut *self;
match std::pin::Pin::new(&mut this.inner).poll_next(cx) {
std::task::Poll::Ready(Some(Ok(data))) => {
let len = data.len();
if len > this.remaining {
this.remaining = 0;
this.exceeded
.store(true, std::sync::atomic::Ordering::Relaxed);
return std::task::Poll::Ready(Some(Err(std::io::Error::other(
MCP_BODY_LIMIT_EXCEEDED,
))));
}
this.remaining -= len;
std::task::Poll::Ready(Some(Ok(data)))
}
std::task::Poll::Ready(Some(Err(e))) => {
let _ = e;
std::task::Poll::Ready(Some(Err(std::io::Error::other(MCP_BODY_LIMIT_EXCEEDED))))
}
std::task::Poll::Pending => std::task::Poll::Pending,
std::task::Poll::Ready(None) => std::task::Poll::Ready(None),
}
}
}
#[async_trait]
impl alkcall::core::types::ProtocolHandler for HttpAdapter {
fn alpn(&self) -> &'static [u8] {
self.alpn
}
async fn handle(&self, connection: Connection, auth: &AuthContext) -> Result<(), HandlerError> {
if let Some(identity) = auth.identity.clone() {
let _ = connection.set_identity(identity);
}
let stream = connection
.accept_bi()
.await
.map_err(stream_error_to_handler)?;
self.serve_io(stream).await
}
}
impl HttpAdapter {
async fn serve_io<I>(&self, io: I) -> Result<(), HandlerError>
where
I: tokio::io::AsyncRead + tokio::io::AsyncWrite + Send + Unpin + 'static,
{
let io = TokioIo::new(io);
let service = TowerToHyperService::new(self.router.clone());
const HEADER_READ_TIMEOUT: Duration = Duration::from_secs(10);
#[cfg(feature = "h2")]
const H2_KEEP_ALIVE_INTERVAL: Duration = Duration::from_secs(30);
#[cfg(feature = "h2")]
const H2_KEEP_ALIVE_TIMEOUT: Duration = Duration::from_secs(10);
#[cfg_attr(not(feature = "h2"), allow(unused_mut))]
let mut builder = HyperBuilder::new(TokioExecutor::new());
#[cfg(feature = "http1")]
{
builder
.http1()
.timer(TokioTimer::new())
.keep_alive(true)
.header_read_timeout(HEADER_READ_TIMEOUT);
}
#[cfg(feature = "h2")]
{
builder
.http2()
.timer(TokioTimer::new())
.enable_connect_protocol()
.keep_alive_interval(H2_KEEP_ALIVE_INTERVAL)
.keep_alive_timeout(H2_KEEP_ALIVE_TIMEOUT);
}
let conn = builder.serve_connection_with_upgrades(io, service);
tokio::pin!(conn);
let result = (&mut conn).await;
if let Err(e) = result {
error!("http adapter: connection closed with error: {e}");
}
Ok(())
}
}
fn stream_error_to_handler(e: StreamError) -> HandlerError {
HandlerError::from(e)
}
#[derive(Clone)]
pub(crate) struct CachedOpenAPIDoc {
inner: Arc<Mutex<Option<CachedOpenAPIDocInner>>>,
}
struct CachedOpenAPIDocInner {
bytes: axum::body::Bytes,
}
impl CachedOpenAPIDoc {
pub(crate) fn new(registry: &OperationRegistry) -> Self {
Self {
inner: Arc::new(Mutex::new(None)),
}
.with_registry(registry)
}
fn with_registry(self, registry: &OperationRegistry) -> Self {
let spec = match crate::adapters::to_openapi(registry) {
Ok(spec) => spec,
Err(e) => {
error!("openapi.json projection failed; endpoint will return 500: {e}");
return self;
}
};
match serde_json::to_vec(&spec.raw) {
Ok(bytes) => {
*self.inner.lock() = Some(CachedOpenAPIDocInner {
bytes: axum::body::Bytes::from(bytes),
});
}
Err(e) => {
error!("openapi.json serialization failed; endpoint will return 500: {e}");
}
}
self
}
fn bytes(&self) -> Option<axum::body::Bytes> {
self.inner.lock().as_ref().map(|c| c.bytes.clone())
}
}
async fn openapi_json_handler(
axum::extract::State(doc): axum::extract::State<CachedOpenAPIDoc>,
) -> axum::response::Response {
use axum::response::IntoResponse;
match doc.bytes() {
Some(bytes) => ([(http::header::CONTENT_TYPE, "application/json")], bytes).into_response(),
None => (
axum::http::StatusCode::INTERNAL_SERVER_ERROR,
"internal server error",
)
.into_response(),
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::server::auth::ResolvedIdentity;
use alkcall::core::auth::IdentityProvider;
use alkcall::core::types::ProtocolHandler;
use tokio::io::{AsyncReadExt, AsyncWriteExt};
struct NoopProvider;
impl IdentityProvider for NoopProvider {
fn resolve_from_fingerprint(&self, _: &str) -> Option<alkcall::core::auth::Identity> {
None
}
fn resolve_from_token(
&self,
_: &alkcall::core::auth::AuthToken,
) -> Option<alkcall::core::auth::Identity> {
None
}
}
fn empty_registry() -> Arc<OperationRegistry> {
Arc::new(OperationRegistry::new())
}
fn provider() -> Arc<dyn IdentityProvider> {
Arc::new(NoopProvider)
}
#[test]
fn alpn_returns_http1_for_default_new() {
let adapter = HttpAdapter::new(provider(), empty_registry());
assert_eq!(adapter.alpn(), ALPN_HTTP1);
assert_eq!(adapter.alpn(), b"http/1.1");
}
#[test]
fn alpn_returns_h2_for_h2_constructor() {
let adapter = HttpAdapter::h2(provider(), empty_registry());
assert_eq!(adapter.alpn(), ALPN_H2);
assert_eq!(adapter.alpn(), b"h2");
}
#[test]
fn decoy_config_default_is_not_found() {
assert!(matches!(DecoyConfig::default(), DecoyConfig::NotFound));
}
#[test]
fn with_decoy_updates_decoy() {
let adapter = HttpAdapter::new(provider(), empty_registry());
let adapter = adapter.with_decoy(DecoyConfig::Redirect {
to: "https://example.com".to_string(),
});
assert!(matches!(adapter.decoy(), DecoyConfig::Redirect { .. }));
}
#[tokio::test]
async fn full_http_request_response_cycle_over_duplex() {
let extra = Router::new().route("/v1/ping", get(|| async { "pong" }));
let adapter = HttpAdapter::new(provider(), empty_registry()).with_extra_routes(extra);
let (client, server) = tokio::io::duplex(64 * 1024);
let conn = Connection::from_bidi(server, b"http/1.1".to_vec(), None);
let auth = AuthContext::anonymous(b"http/1.1");
let server_task =
tokio::spawn(async move { ProtocolHandler::handle(&adapter, conn, &auth).await });
let mut client = client;
client
.write_all(b"GET /v1/ping HTTP/1.1\r\nHost: localhost\r\nConnection: close\r\n\r\n")
.await
.unwrap();
let mut response = Vec::new();
let _ = tokio::time::timeout(
std::time::Duration::from_secs(5),
client.read_to_end(&mut response),
)
.await
.expect("read timed out")
.unwrap();
let text = String::from_utf8_lossy(&response);
assert!(text.starts_with("HTTP/1.1 200 OK"), "got: {text}");
assert!(text.contains("pong"), "got: {text}");
let _ = server_task.await;
}
#[tokio::test]
async fn accept_bi_failure_maps_to_handler_error_connection_closed() {
let adapter = HttpAdapter::new(provider(), empty_registry());
let (client, server) = tokio::io::duplex(64 * 1024);
let conn = Connection::from_bidi(server, b"http/1.1".to_vec(), None);
let auth = AuthContext::anonymous(b"http/1.1");
drop(client);
conn.close(0, "gone before the first stream");
let result = ProtocolHandler::handle(&adapter, conn, &auth).await;
match result {
Err(HandlerError::ConnectionClosed) => {}
other => panic!("expected HandlerError::ConnectionClosed, got {other:?}"),
}
}
#[tokio::test]
async fn healthz_served_by_the_adapter_over_duplex() {
let adapter = HttpAdapter::new(provider(), empty_registry());
let (client, server) = tokio::io::duplex(64 * 1024);
let conn = Connection::from_bidi(server, b"http/1.1".to_vec(), None);
let auth = AuthContext::anonymous(b"http/1.1");
let server_task = tokio::spawn(async move {
let _ = ProtocolHandler::handle(&adapter, conn, &auth).await;
});
let mut client = client;
client
.write_all(b"GET /healthz HTTP/1.1\r\nHost: localhost\r\nConnection: close\r\n\r\n")
.await
.unwrap();
let mut response = Vec::new();
let _ = tokio::time::timeout(
std::time::Duration::from_secs(5),
client.read_to_end(&mut response),
)
.await
.expect("read timed out")
.unwrap();
let text = String::from_utf8_lossy(&response);
assert!(text.starts_with("HTTP/1.1 200 OK"), "got: {text}");
assert!(text.contains("ok"), "got: {text}");
let _ = server_task.await;
}
#[tokio::test]
async fn unknown_path_serves_decoy_404_over_duplex() {
let adapter = HttpAdapter::new(provider(), empty_registry());
let (client, server) = tokio::io::duplex(64 * 1024);
let conn = Connection::from_bidi(server, b"http/1.1".to_vec(), None);
let auth = AuthContext::anonymous(b"http/1.1");
let server_task = tokio::spawn(async move {
let _ = ProtocolHandler::handle(&adapter, conn, &auth).await;
});
let mut client = client;
client
.write_all(b"GET /nowhere HTTP/1.1\r\nHost: localhost\r\nConnection: close\r\n\r\n")
.await
.unwrap();
let mut response = Vec::new();
let _ = tokio::time::timeout(
std::time::Duration::from_secs(5),
client.read_to_end(&mut response),
)
.await
.expect("read timed out")
.unwrap();
let text = String::from_utf8_lossy(&response);
assert!(text.starts_with("HTTP/1.1 404 Not Found"), "got: {text}");
assert!(
text.contains("nginx"),
"decoy should look like nginx: {text}"
);
let _ = server_task.await;
}
#[tokio::test]
async fn openapi_json_serves_the_gateway_projection() {
use tokio::io::{AsyncReadExt, AsyncWriteExt};
let registry = OperationRegistry::new();
let spec = alkcall::registry::spec::OperationSpec::new(
"echo/run",
alkcall::registry::spec::OperationType::Query,
alkcall::registry::spec::Visibility::External,
serde_json::json!({}),
serde_json::json!({}),
vec![],
alkcall::registry::spec::AccessControl::default(),
None,
);
registry
.register(alkcall::registry::registration::HandlerRegistration::new(
spec,
alkcall::registry::registration::HandlerKind::Once(
alkcall::registry::registration::make_handler(|input, ctx| async move {
alkcall::protocol::wire::ResponseEnvelope::ok(ctx.request_id, input)
}),
),
alkcall::registry::registration::OperationProvenance::Local,
None,
None,
alkcall::core::types::Capabilities::new(),
))
.unwrap();
let adapter = HttpAdapter::new(provider(), Arc::new(registry));
let (client, server) = tokio::io::duplex(256 * 1024);
let conn = Connection::from_bidi(server, b"http/1.1".to_vec(), None);
let auth = AuthContext::anonymous(b"http/1.1");
let server_task = tokio::spawn(async move {
let _ = ProtocolHandler::handle(&adapter, conn, &auth).await;
});
let mut client = client;
client
.write_all(
b"GET /openapi.json HTTP/1.1\r\nHost: localhost\r\nConnection: close\r\n\r\n",
)
.await
.unwrap();
let mut response = Vec::new();
let _ = tokio::time::timeout(
std::time::Duration::from_secs(5),
client.read_to_end(&mut response),
)
.await
.expect("read timed out")
.unwrap();
let text = String::from_utf8_lossy(&response);
assert!(text.starts_with("HTTP/1.1 200 OK"), "got: {text}");
assert!(text.contains("application/json"), "got: {text}");
assert!(text.contains("\"/publish\""), "publish path in doc");
assert!(text.contains("1.4.0"), "info.version 1.4.0 in doc");
assert!(text.contains("gatewayPublish"), "publish operationId");
let _ = server_task.await;
}
#[cfg(feature = "mcp")]
#[tokio::test]
async fn mcp_endpoint_serves_four_gateway_tools_bearer_gated() {
use tokio::io::{AsyncReadExt, AsyncWriteExt};
let adapter = HttpAdapter::new(provider(), empty_registry());
let (client, server) = tokio::io::duplex(256 * 1024);
let conn = Connection::from_bidi(server, b"http/1.1".to_vec(), None);
let auth = AuthContext::anonymous(b"http/1.1");
let server_task = tokio::spawn(async move {
let _ = ProtocolHandler::handle(&adapter, conn, &auth).await;
});
let mut client = client;
client
.write_all(
b"POST /mcp HTTP/1.1\r\nHost: localhost\r\nContent-Type: application/json\r\nAccept: application/json, text/event-stream\r\nContent-Length: 175\r\nConnection: close\r\n\r\n{\"jsonrpc\": \"2.0\", \"id\": 1, \"method\": \"initialize\", \"params\": {\"protocolVersion\": \"2025-06-18\", \"capabilities\": {}, \"clientInfo\": {\"name\": \"test-client\", \"version\": \"1.0.0\"}}}",
)
.await
.unwrap();
let mut response = Vec::new();
let _ = tokio::time::timeout(
std::time::Duration::from_secs(5),
client.read_to_end(&mut response),
)
.await
.expect("read timed out")
.unwrap();
let text = String::from_utf8_lossy(&response);
assert!(
text.starts_with("HTTP/1.1 200 OK"),
"initialize over /mcp got: {text}"
);
assert!(
text.contains("alkhttp-to-mcp"),
"server info in initialize response: {text}"
);
let _ = server_task.await;
}
#[cfg(feature = "mcp")]
#[tokio::test]
async fn mcp_rejects_oversized_body_declared_content_length_with_413() {
use tokio::io::{AsyncReadExt, AsyncWriteExt};
let oversized = MCP_BODY_LIMIT + 1;
let head = format!(
"POST /mcp HTTP/1.1\r\nHost: localhost\r\nContent-Type: application/json\r\nAccept: application/json, text/event-stream\r\nContent-Length: {oversized}\r\nConnection: close\r\n\r\n"
);
let adapter = HttpAdapter::new(provider(), empty_registry());
let (client, server) = tokio::io::duplex(64 * 1024);
let conn = Connection::from_bidi(server, b"http/1.1".to_vec(), None);
let auth = AuthContext::anonymous(b"http/1.1");
let server_task = tokio::spawn(async move {
let _ = ProtocolHandler::handle(&adapter, conn, &auth).await;
});
let mut client = client;
client.write_all(head.as_bytes()).await.unwrap();
let mut response = Vec::new();
let _ = tokio::time::timeout(
std::time::Duration::from_secs(5),
client.read_to_end(&mut response),
)
.await
.expect("read timed out")
.unwrap();
let text = String::from_utf8_lossy(&response);
assert!(
text.starts_with("HTTP/1.1 413 Payload Too Large"),
"declared oversized body got: {text}"
);
let _ = server_task.await;
}
#[cfg(feature = "mcp")]
#[tokio::test]
async fn mcp_rejects_oversized_chunked_body_with_413() {
use tokio::io::{AsyncReadExt, AsyncWriteExt};
let adapter = HttpAdapter::new(provider(), empty_registry());
let (client, server) = tokio::io::duplex(64 * 1024);
let conn = Connection::from_bidi(server, b"http/1.1".to_vec(), None);
let auth = AuthContext::anonymous(b"http/1.1");
let server_task = tokio::spawn(async move {
let _ = ProtocolHandler::handle(&adapter, conn, &auth).await;
});
let (mut reader_client, mut writer_client) = tokio::io::split(client);
writer_client
.write_all(
b"POST /mcp HTTP/1.1\r\nHost: localhost\r\nContent-Type: application/json\r\nAccept: application/json, text/event-stream\r\nTransfer-Encoding: chunked\r\nConnection: close\r\n\r\n",
)
.await
.unwrap();
let writer = tokio::spawn(async move {
let chunk = vec![b'a'; 64 * 1024];
let chunk_header = format!("{:x}\r\n", chunk.len());
for _ in 0..(MCP_BODY_LIMIT / chunk.len()) + 1 {
if writer_client
.write_all(chunk_header.as_bytes())
.await
.is_err()
{
return;
}
if writer_client.write_all(&chunk).await.is_err() {
return;
}
if writer_client.write_all(b"\r\n").await.is_err() {
return;
}
}
let _ = writer_client.write_all(b"0\r\n\r\n").await;
});
let mut response = Vec::new();
let read = tokio::time::timeout(
std::time::Duration::from_secs(10),
reader_client.read_to_end(&mut response),
)
.await;
writer.abort();
read.expect("read timed out").unwrap();
let text = String::from_utf8_lossy(&response);
assert!(
text.starts_with("HTTP/1.1 413 Payload Too Large"),
"chunked oversized body got: {text}"
);
let _ = server_task.await;
}
struct StaticProvider;
impl IdentityProvider for StaticProvider {
fn resolve_from_fingerprint(&self, _: &str) -> Option<alkcall::core::auth::Identity> {
None
}
fn resolve_from_token(
&self,
_: &alkcall::core::auth::AuthToken,
) -> Option<alkcall::core::auth::Identity> {
Some(alkcall::core::auth::Identity {
id: "worker-a".to_string(),
scopes: vec![],
resources: std::collections::HashMap::new(),
})
}
}
fn static_provider() -> Arc<dyn IdentityProvider> {
Arc::new(StaticProvider)
}
struct CountingProvider {
resolutions: std::sync::Mutex<usize>,
}
impl CountingProvider {
fn new() -> Self {
Self {
resolutions: std::sync::Mutex::new(0),
}
}
fn resolutions(&self) -> usize {
*self.resolutions.lock().unwrap_or_else(|e| e.into_inner())
}
}
impl IdentityProvider for CountingProvider {
fn resolve_from_fingerprint(&self, _: &str) -> Option<alkcall::core::auth::Identity> {
None
}
fn resolve_from_token(
&self,
_: &alkcall::core::auth::AuthToken,
) -> Option<alkcall::core::auth::Identity> {
*self.resolutions.lock().unwrap_or_else(|e| e.into_inner()) += 1;
Some(alkcall::core::auth::Identity {
id: "worker-a".to_string(),
scopes: vec![],
resources: std::collections::HashMap::new(),
})
}
}
async fn ws_upgrade_oneshot(
app: Router,
authorization: &str,
) -> axum::http::Response<axum::body::Body> {
use tower::ServiceExt;
let request = axum::http::Request::builder()
.method(axum::http::Method::GET)
.uri(WS_UPGRADE_PATH)
.header(axum::http::header::AUTHORIZATION, authorization)
.header(axum::http::header::CONNECTION, "upgrade")
.header(axum::http::header::UPGRADE, "websocket")
.header(axum::http::header::SEC_WEBSOCKET_VERSION, "13")
.header(
axum::http::header::SEC_WEBSOCKET_KEY,
"dGhlIHNhbXBsZSBub25jZQ==",
)
.header(
axum::http::header::HOST,
axum::http::HeaderValue::from_static("localhost"),
);
let mut request = request.body(axum::body::Body::empty()).unwrap();
let on_upgrade = hyper::upgrade::on(&mut request);
request.extensions_mut().insert(on_upgrade);
app.oneshot(request).await.unwrap()
}
fn router_state(idp: Arc<dyn IdentityProvider>) -> RouterState {
RouterState {
registry: empty_registry(),
identity_provider: idp,
decoy: DecoyConfig::default(),
openapi_doc: CachedOpenAPIDoc::new(&OperationRegistry::new()),
ws_sessions: Arc::new(crate::websocket::WsSessions::new()),
ws_session_slots: Arc::new(tokio::sync::Semaphore::new(
crate::websocket::DEFAULT_WS_MAX_SESSIONS,
)),
ws_idle_timeout: Some(crate::websocket::DEFAULT_WS_IDLE_TIMEOUT),
ws_openable_alpns: None,
ws_op_register_acl: alkcall::registry::spec::AccessControl::default(),
}
}
async fn get_with_bearer(
app: Router,
path: &str,
authorization: Option<&str>,
) -> axum::http::Response<axum::body::Body> {
use tower::ServiceExt;
let mut builder = axum::http::Request::builder().uri(path);
if let Some(value) = authorization {
builder = builder.header(axum::http::header::AUTHORIZATION, value);
}
app.oneshot(builder.body(axum::body::Body::empty()).unwrap())
.await
.unwrap()
}
async fn get_with_bearer_with_method(
app: Router,
request: axum::http::Request<axum::body::Body>,
) -> axum::http::Response<axum::body::Body> {
use tower::ServiceExt;
app.oneshot(request).await.unwrap()
}
#[tokio::test]
async fn extra_routes_resolve_bearer_identity_through_the_default_auth() {
let extra = Router::new().route(
"/v1/whoami",
get(|ResolvedIdentity(identity): ResolvedIdentity| async move {
match identity {
Some(id) => id.id,
None => "none".to_string(),
}
}),
);
let app = build_router(router_state(static_provider()), Some(extra));
let response = get_with_bearer(app.clone(), "/v1/whoami", Some("Bearer alk_test")).await;
assert_eq!(response.status(), axum::http::StatusCode::OK);
let body = axum::body::to_bytes(response.into_body(), usize::MAX)
.await
.unwrap();
assert_eq!(&body[..], b"worker-a");
let response = get_with_bearer(app, "/v1/whoami", None).await;
assert_eq!(response.status(), axum::http::StatusCode::OK);
let body = axum::body::to_bytes(response.into_body(), usize::MAX)
.await
.unwrap();
assert_eq!(&body[..], b"none");
}
#[tokio::test]
async fn extra_route_with_own_layer_can_opt_out_of_default_auth() {
async fn public_identity(
mut request: axum::extract::Request,
next: axum::middleware::Next,
) -> axum::response::Response {
request
.extensions_mut()
.insert(Some(alkcall::core::auth::Identity {
id: "public-webhook".to_string(),
scopes: vec![],
resources: std::collections::HashMap::new(),
}));
next.run(request).await
}
let extra = Router::new().route(
"/v1/public",
get(|ResolvedIdentity(identity): ResolvedIdentity| async move {
match identity {
Some(id) => id.id,
None => "none".to_string(),
}
})
.route_layer(axum::middleware::from_fn(public_identity)),
);
let app = build_router(router_state(static_provider()), Some(extra));
let response = get_with_bearer(app, "/v1/public", None).await;
assert_eq!(response.status(), axum::http::StatusCode::OK);
let body = axum::body::to_bytes(response.into_body(), usize::MAX)
.await
.unwrap();
assert_eq!(&body[..], b"public-webhook");
}
#[tokio::test]
#[should_panic(expected = "Overlapping method route")]
async fn extra_route_on_reserved_path_panics_at_construction() {
let extra = Router::new().route(
"/search",
axum::routing::post(|| async { "shadowing the gateway" }),
);
let _ = HttpAdapter::new(static_provider(), empty_registry()).with_extra_routes(extra);
}
#[tokio::test]
async fn extra_route_on_non_reserved_path_merges_cleanly() {
let extra = Router::new().route(
"/v1/ping",
get(|| async { "pong" }).post(|| async { "pong-post" }),
);
let adapter =
HttpAdapter::new(static_provider(), empty_registry()).with_extra_routes(extra);
assert!(adapter.router().has_routes());
}
#[tokio::test]
async fn second_with_decoy_keeps_extra_routes() {
let extra = Router::new().route("/v1/ping", get(|| async { "pong" }));
let adapter = HttpAdapter::new(static_provider(), empty_registry())
.with_extra_routes(extra)
.with_decoy(DecoyConfig::Redirect {
to: "https://example.com".to_string(),
});
let request = axum::http::Request::builder()
.uri("/v1/ping")
.body(axum::body::Body::empty())
.unwrap();
let response = get_with_bearer(adapter.router().clone(), "/v1/ping", None).await;
drop(request);
assert_eq!(
response.status(),
axum::http::StatusCode::OK,
"a second builder call must not drop the extra routes (SRV-05)"
);
}
#[tokio::test]
async fn method_mismatch_on_registered_path_serves_decoy_405() {
let adapter = HttpAdapter::new(static_provider(), empty_registry());
let request = axum::http::Request::builder()
.method(axum::http::Method::OPTIONS)
.uri("/search")
.body(axum::body::Body::empty())
.unwrap();
let response = get_with_bearer_with_method(adapter.router().clone(), request).await;
assert_eq!(
response.status(),
axum::http::StatusCode::METHOD_NOT_ALLOWED
);
let server = response
.headers()
.get(axum::http::header::SERVER)
.map(|v| v.to_str().unwrap().to_string());
assert_eq!(
server.as_deref(),
Some("nginx"),
"405 must carry the decoy Server header, not axum's bare 405 (SRV-07)"
);
let body = axum::body::to_bytes(response.into_body(), usize::MAX)
.await
.unwrap();
let body = String::from_utf8_lossy(&body);
assert!(body.contains("405 Not Allowed"), "got: {body}");
assert!(
!body.contains("axum") && !body.contains("alk"),
"got: {body}"
);
}
#[tokio::test]
async fn method_mismatch_on_extra_route_serves_decoy_405() {
let extra = Router::new().route("/v1/ping", get(|| async { "pong" }));
let adapter =
HttpAdapter::new(static_provider(), empty_registry()).with_extra_routes(extra);
let request = axum::http::Request::builder()
.method(axum::http::Method::DELETE)
.uri("/v1/ping")
.body(axum::body::Body::empty())
.unwrap();
let response = get_with_bearer_with_method(adapter.router().clone(), request).await;
assert_eq!(
response.status(),
axum::http::StatusCode::METHOD_NOT_ALLOWED,
"wrong-method probe on an extra route"
);
let server = response
.headers()
.get(axum::http::header::SERVER)
.map(|v| v.to_str().unwrap().to_string());
assert_eq!(
server.as_deref(),
Some("nginx"),
"extra-route 405 must carry the decoy Server header, not axum's bare 405 (SRV-12)"
);
let body = axum::body::to_bytes(response.into_body(), usize::MAX)
.await
.unwrap();
let body = String::from_utf8_lossy(&body);
assert!(body.contains("405 Not Allowed"), "got: {body}");
assert!(
!body.contains("axum") && !body.contains("alk"),
"got: {body}"
);
}
#[tokio::test]
async fn method_mismatch_on_default_surface_still_serves_decoy_405_after_extras_merge() {
let extra = Router::new().route("/v1/ping", get(|| async { "pong" }));
let adapter =
HttpAdapter::new(static_provider(), empty_registry()).with_extra_routes(extra);
let request = axum::http::Request::builder()
.method(axum::http::Method::OPTIONS)
.uri("/search")
.body(axum::body::Body::empty())
.unwrap();
let response = get_with_bearer_with_method(adapter.router().clone(), request).await;
assert_eq!(
response.status(),
axum::http::StatusCode::METHOD_NOT_ALLOWED
);
let server = response
.headers()
.get(axum::http::header::SERVER)
.map(|v| v.to_str().unwrap().to_string());
assert_eq!(
server.as_deref(),
Some("nginx"),
"the re-applied 405 fallback must not regress the default surface (SRV-07)"
);
let body = axum::body::to_bytes(response.into_body(), usize::MAX)
.await
.unwrap();
let body = String::from_utf8_lossy(&body);
assert!(body.contains("405 Not Allowed"), "got: {body}");
}
#[tokio::test]
async fn ws_upgrade_resolves_the_token_exactly_once() {
let provider = Arc::new(CountingProvider::new());
let app = build_router(
router_state(provider.clone() as Arc<dyn IdentityProvider>),
None,
);
let response = ws_upgrade_oneshot(app, "Bearer alk_test").await;
assert_eq!(
response.status(),
axum::http::StatusCode::SWITCHING_PROTOCOLS,
"a valid bearer token upgrades"
);
assert_eq!(
provider.resolutions(),
1,
"the WS upgrade path must resolve the token exactly once (SRV-11)"
);
let provider = Arc::new(CountingProvider::new());
let app = build_router(
router_state(provider.clone() as Arc<dyn IdentityProvider>),
None,
);
let response = ws_upgrade_oneshot(app, "Bearer alk_test").await;
drop(response);
assert_eq!(
provider.resolutions(),
1,
"exactly one resolution per upgrade request after any builder order"
);
}
#[tokio::test]
async fn ws_upgrade_without_token_is_rejected_401() {
let app = build_router(router_state(static_provider()), None);
let response = ws_upgrade_oneshot(app, "").await;
assert_eq!(
response.status(),
axum::http::StatusCode::UNAUTHORIZED,
"the WS route keeps its own enforced layer"
);
}
#[cfg(feature = "mcp")]
#[tokio::test]
async fn mcp_request_resolves_the_token_exactly_once() {
let provider = Arc::new(CountingProvider::new());
let app = build_router(
router_state(provider.clone() as Arc<dyn IdentityProvider>),
None,
);
let request = axum::http::Request::builder()
.method(axum::http::Method::POST)
.uri("/mcp")
.header(axum::http::header::HOST, "localhost")
.header(axum::http::header::AUTHORIZATION, "Bearer alk_test")
.header(axum::http::header::CONTENT_TYPE, "application/json")
.header(
axum::http::header::ACCEPT,
"application/json, text/event-stream",
)
.body(axum::body::Body::from(
serde_json::to_vec(&serde_json::json!({
"jsonrpc": "2.0", "id": 1, "method": "initialize",
"params": {
"protocolVersion": "2025-06-18",
"capabilities": {},
"clientInfo": { "name": "test-client", "version": "1.0.0" }
}
}))
.unwrap(),
))
.unwrap();
let response = tower::ServiceExt::oneshot(app, request).await.unwrap();
assert_eq!(
response.status(),
axum::http::StatusCode::OK,
"the /mcp initialize response comes back"
);
let body = axum::body::to_bytes(response.into_body(), usize::MAX)
.await
.unwrap();
assert!(
String::from_utf8_lossy(&body).contains("alkhttp-to-mcp"),
"initialize response body, got: {}",
String::from_utf8_lossy(&body)
);
assert_eq!(
provider.resolutions(),
1,
"the /mcp path must resolve the token exactly once (SRV-11)"
);
}
#[tokio::test]
async fn method_mismatch_on_ws_upgrade_path_serves_decoy_405() {
let app = build_router(router_state(static_provider()), None);
let request = axum::http::Request::builder()
.method(axum::http::Method::POST)
.uri(WS_UPGRADE_PATH)
.body(axum::body::Body::empty())
.unwrap();
let response = get_with_bearer_with_method(app, request).await;
assert_eq!(
response.status(),
axum::http::StatusCode::METHOD_NOT_ALLOWED
);
let server = response
.headers()
.get(axum::http::header::SERVER)
.map(|v| v.to_str().unwrap().to_string());
assert_eq!(
server.as_deref(),
Some("nginx"),
"wrong-method probe on the WS path keeps the decoy shape after the reorder"
);
}
#[tokio::test]
async fn openapi_json_is_cached_and_generic_on_cache_miss() {
let adapter = HttpAdapter::new(static_provider(), empty_registry());
let first = openapi_json_handler(axum::extract::State(adapter.openapi_doc.clone())).await;
let second = openapi_json_handler(axum::extract::State(adapter.openapi_doc.clone())).await;
let first_bytes = axum::body::to_bytes(first.into_body(), usize::MAX)
.await
.unwrap();
let second_bytes = axum::body::to_bytes(second.into_body(), usize::MAX)
.await
.unwrap();
assert_eq!(
&first_bytes[..],
&second_bytes[..],
"cached doc is byte-stable"
);
let miss = CachedOpenAPIDoc {
inner: Arc::new(Mutex::new(None)),
};
let response = openapi_json_handler(axum::extract::State(miss)).await;
assert_eq!(
response.status(),
axum::http::StatusCode::INTERNAL_SERVER_ERROR
);
let body = axum::body::to_bytes(response.into_body(), usize::MAX)
.await
.unwrap();
assert_eq!(
&body[..],
b"internal server error",
"the 500 body is the fixed generic string, no serde internals"
);
}
}