use std::collections::HashMap;
use std::net::{IpAddr, Ipv4Addr, SocketAddr};
use std::sync::atomic::{AtomicU64, Ordering};
use std::sync::{Arc, PoisonError, RwLock, Weak};
use async_trait::async_trait;
use dravr_tronc::mcp::auth::{AuthError, AuthHook};
use dravr_tronc::mcp::host::ToolDispatcher;
use dravr_tronc::mcp::protocol::JsonRpcRequest;
use dravr_tronc::mcp::schema::{Tool, ToolResponse};
use dravr_tronc::mcp::server::McpServer;
use dravr_tronc::mcp::tool::{ToolContext, ToolRegistry};
use dravr_tronc::mcp::transport::http::mcp_router;
use embacle::types::{McpHeader, McpServerConfig, McpTransport, RunnerError};
use embacle::{McpToolDefinition, McpToolExecutor};
use rand::RngCore;
use serde_json::{json, Value};
use subtle::ConstantTimeEq;
use tokio::net::TcpListener;
use tokio::sync::oneshot;
use tracing::{debug, warn};
const AUTHORIZATION: &str = "authorization";
#[derive(Debug, Clone)]
pub struct ToolHostConfig {
pub server_name: String,
pub bind_addr: IpAddr,
pub port: u16,
pub instructions: Option<String>,
}
impl Default for ToolHostConfig {
fn default() -> Self {
Self {
server_name: "tools".to_owned(),
bind_addr: IpAddr::V4(Ipv4Addr::LOCALHOST),
port: 0,
instructions: None,
}
}
}
#[derive(Debug, Clone)]
pub struct ToolOutcome {
pub text: String,
pub structured: Option<Value>,
pub is_error: bool,
}
impl ToolOutcome {
#[must_use]
pub fn json(value: Value) -> Self {
Self {
text: value.to_string(),
structured: Some(value),
is_error: false,
}
}
#[must_use]
pub fn text(text: impl Into<String>) -> Self {
Self {
text: text.into(),
structured: None,
is_error: false,
}
}
#[must_use]
pub fn refused(reason: impl Into<String>) -> Self {
Self {
text: reason.into(),
structured: None,
is_error: true,
}
}
#[must_use]
pub fn with_structured(mut self, value: Value) -> Self {
self.structured = Some(value);
self
}
}
#[async_trait]
pub trait ToolSurface: Send + Sync {
async fn list_tools(&self) -> Vec<McpToolDefinition>;
async fn call(&self, tool_name: &str, arguments: &Value) -> ToolOutcome;
}
pub struct StaticSurface {
tools: Vec<McpToolDefinition>,
executor: Arc<dyn McpToolExecutor>,
}
impl StaticSurface {
#[must_use]
pub const fn new(tools: Vec<McpToolDefinition>, executor: Arc<dyn McpToolExecutor>) -> Self {
Self { tools, executor }
}
}
#[async_trait]
impl ToolSurface for StaticSurface {
async fn list_tools(&self) -> Vec<McpToolDefinition> {
self.tools.clone()
}
async fn call(&self, tool_name: &str, arguments: &Value) -> ToolOutcome {
match self.executor.execute(tool_name, arguments).await {
Ok(value) => ToolOutcome::json(value),
Err(e) => ToolOutcome::refused(e.message.clone())
.with_structured(json!({ "error_kind": format!("{:?}", e.kind) })),
}
}
}
struct SessionState {
id: String,
bearer: String,
surface: Arc<dyn ToolSurface>,
calls_served: AtomicU64,
}
struct Inner {
server_name: String,
addr: SocketAddr,
sessions: RwLock<HashMap<String, Arc<SessionState>>>,
shutdown: RwLock<Option<oneshot::Sender<()>>>,
}
impl Inner {
fn session_for(&self, bearer: &str) -> Option<Arc<SessionState>> {
let sessions = self.sessions.read().unwrap_or_else(PoisonError::into_inner);
sessions
.values()
.find(|s| s.bearer.as_bytes().ct_eq(bearer.as_bytes()).into())
.map(Arc::clone)
}
}
#[derive(Clone)]
pub struct ToolHost {
inner: Arc<Inner>,
}
impl ToolHost {
pub async fn bind(config: ToolHostConfig) -> Result<Self, RunnerError> {
let listener = TcpListener::bind(SocketAddr::new(config.bind_addr, config.port))
.await
.map_err(|e| RunnerError::config(format!("tool host could not bind: {e}")))?;
let addr = listener
.local_addr()
.map_err(|e| RunnerError::config(format!("tool host bound but has no address: {e}")))?;
let instructions = config.instructions;
let (tx, rx) = oneshot::channel();
let inner = Arc::new(Inner {
server_name: config.server_name,
addr,
sessions: RwLock::new(HashMap::new()),
shutdown: RwLock::new(Some(tx)),
});
let mut server = McpServer::new(
"embacle-tool-host",
env!("CARGO_PKG_VERSION"),
ToolRegistry::new(),
Arc::clone(&inner),
)
.with_tool_dispatcher(Arc::new(Forwarding))
.with_auth_hook(Arc::new(BearerSessions));
if let Some(text) = instructions {
server = server.with_instructions(text);
}
let server = Arc::new(server);
let router = mcp_router(server);
tokio::spawn(async move {
let outcome = axum::serve(listener, router)
.with_graceful_shutdown(async {
let _ = rx.await;
})
.await;
if let Err(e) = outcome {
warn!(error = %e, "tool host listener stopped");
}
});
debug!(%addr, "tool host listening");
Ok(Self { inner })
}
#[must_use]
pub fn local_addr(&self) -> SocketAddr {
self.inner.addr
}
#[must_use]
pub fn open_session(&self, surface: Arc<dyn ToolSurface>) -> ToolSession {
let session_id = uuid::Uuid::new_v4().to_string();
let bearer = mint_bearer();
let state = Arc::new(SessionState {
id: session_id.clone(),
bearer: bearer.clone(),
surface,
calls_served: AtomicU64::new(0),
});
self.inner
.sessions
.write()
.unwrap_or_else(PoisonError::into_inner)
.insert(session_id.clone(), Arc::clone(&state));
ToolSession {
session_id,
bearer,
server_name: self.inner.server_name.clone(),
addr: self.inner.addr,
state,
host: Arc::downgrade(&self.inner),
}
}
#[must_use]
pub fn open_sessions(&self) -> usize {
self.inner
.sessions
.read()
.unwrap_or_else(PoisonError::into_inner)
.len()
}
pub fn shutdown(&self) {
let signal = self
.inner
.shutdown
.write()
.unwrap_or_else(PoisonError::into_inner)
.take();
if let Some(tx) = signal {
let _ = tx.send(());
}
}
}
fn mint_bearer() -> String {
let mut raw = [0_u8; 32];
rand::thread_rng().fill_bytes(&mut raw);
raw.iter().fold(String::with_capacity(64), |mut acc, b| {
use std::fmt::Write;
let _ = write!(acc, "{b:02x}");
acc
})
}
pub struct ToolSession {
session_id: String,
bearer: String,
server_name: String,
addr: SocketAddr,
state: Arc<SessionState>,
host: Weak<Inner>,
}
impl ToolSession {
#[must_use]
pub fn mcp_servers(&self) -> Vec<McpServerConfig> {
vec![McpServerConfig {
name: self.server_name.clone(),
transport: McpTransport::Http {
url: format!("http://{}/mcp", self.addr),
headers: vec![McpHeader {
name: "Authorization".to_owned(),
value: format!("Bearer {}", self.bearer),
}],
},
}]
}
#[must_use]
pub fn session_id(&self) -> &str {
&self.session_id
}
#[must_use]
pub fn calls_served(&self) -> u64 {
self.state.calls_served.load(Ordering::SeqCst)
}
}
impl Drop for ToolSession {
fn drop(&mut self) {
if let Some(inner) = self.host.upgrade() {
inner
.sessions
.write()
.unwrap_or_else(PoisonError::into_inner)
.remove(&self.session_id);
}
}
}
struct BearerSessions;
#[async_trait]
impl AuthHook<Inner> for BearerSessions {
async fn authenticate(
&self,
request: &JsonRpcRequest,
state: &Arc<Inner>,
) -> Result<ToolContext, AuthError> {
let bearer = request
.auth_token
.as_deref()
.or_else(|| {
request
.headers
.as_ref()
.and_then(|h| h.get(AUTHORIZATION))
.and_then(Value::as_str)
})
.map(|raw| raw.trim_start_matches("Bearer ").trim())
.unwrap_or_default();
state.session_for(bearer).map_or_else(
|| {
Err(AuthError::Unauthorized {
www_authenticate: "Bearer".to_owned(),
})
},
|session| Ok(ToolContext::new().with_request_id(Value::from(session.id.clone()))),
)
}
}
struct Forwarding;
#[async_trait]
impl ToolDispatcher<Inner> for Forwarding {
async fn list_tools(&self, state: &Arc<Inner>, ctx: &ToolContext) -> Vec<Tool> {
let Some(session) = resolve(state, ctx) else {
return Vec::new();
};
session
.surface
.list_tools()
.await
.into_iter()
.map(|t| Tool {
name: t.name,
description: t.description,
input_schema: t.input_schema,
annotations: None,
})
.collect()
}
async fn call_tool(
&self,
name: &str,
state: &Arc<Inner>,
ctx: &ToolContext,
arguments: Value,
) -> ToolResponse {
let Some(session) = resolve(state, ctx) else {
return ToolResponse::error("session is no longer open".to_owned());
};
if !session
.surface
.list_tools()
.await
.iter()
.any(|t| t.name == name)
{
return ToolResponse::error(format!("unknown tool: {name}"));
}
session.calls_served.fetch_add(1, Ordering::SeqCst);
let outcome = session.surface.call(name, &arguments).await;
let mut response = if outcome.is_error {
ToolResponse::error(outcome.text)
} else {
ToolResponse::text(outcome.text)
};
response.structured_content = outcome.structured;
response
}
}
fn resolve(state: &Arc<Inner>, ctx: &ToolContext) -> Option<Arc<SessionState>> {
let id = ctx.request_id.as_ref()?.as_str()?;
let sessions = state
.sessions
.read()
.unwrap_or_else(PoisonError::into_inner);
sessions.get(id).map(Arc::clone)
}