use anyhow::Result;
use std::collections::HashMap;
use std::pin::Pin;
use std::sync::{Arc, Mutex};
use std::task::{Context, Poll};
use tokio::sync::{mpsc, oneshot};
use tracing::Instrument;
use wasmtime::component::{Component, Linker, Source, StreamConsumer, StreamResult};
use wasmtime::{Engine, Store, StoreContextMut};
use crate::consent;
use crate::info::{ComponentError, ComponentInfo};
use crate::store::{HostState, create_store};
use crate::{act, exports};
use crate::{credentials, fs_policy, sessions};
#[derive(Debug, Clone)]
pub struct AuditContext {
pub component_ref: String,
pub digest: String,
pub transport: crate::audit::Transport,
pub has_prompt_channel: bool,
pub record_args: bool,
}
pub(crate) fn meta_str(metadata: &[(String, String)], key: &str) -> Option<String> {
metadata
.iter()
.find(|(k, _)| k == key)
.map(|(_, v)| v.clone())
.filter(|v| !v.is_empty())
}
pub(crate) fn decode_meta_strings(metadata: &[(String, Vec<u8>)]) -> Vec<(String, String)> {
metadata
.iter()
.filter_map(|(k, v)| {
let value = act_types::cbor::cbor_to_json(v).ok()?;
value.as_str().map(|s| (k.clone(), s.to_string()))
})
.collect()
}
fn args_as_json(arguments: &[u8], record_args: bool) -> Option<String> {
if !record_args {
return None;
}
let value = act_types::cbor::cbor_to_json(arguments).ok()?;
serde_json::to_string(&value).ok()
}
fn events_contain_error(events: &[act::tools::types::ToolEvent]) -> bool {
events
.iter()
.any(|e| matches!(e, act::tools::types::ToolEvent::Error(_)))
}
const REQUEST_ID_COUNTER_BITS: u32 = 9; const REQUEST_ID_SALT_BITS: u32 = 24 - REQUEST_ID_COUNTER_BITS; pub(crate) fn pack_visible_request_id(counter: u64, salt: u32) -> u32 {
#[allow(clippy::cast_possible_truncation)]
let counter_field = (counter as u32) & ((1 << REQUEST_ID_COUNTER_BITS) - 1);
let salt_field = salt & ((1 << REQUEST_ID_SALT_BITS) - 1);
(counter_field << REQUEST_ID_SALT_BITS) | salt_field
}
pub(crate) fn new_request_id() -> String {
use std::collections::hash_map::RandomState;
use std::hash::{BuildHasher, Hasher};
use std::sync::OnceLock;
use std::sync::atomic::{AtomicU64, Ordering};
static SALT: OnceLock<u32> = OnceLock::new();
let salt = *SALT.get_or_init(|| {
let mut hasher = RandomState::new().build_hasher();
hasher.write_u32(std::process::id());
hasher.finish() as u32
});
static N: AtomicU64 = AtomicU64::new(0);
let n = N.fetch_add(1, Ordering::Relaxed);
format!("{:06x}-{n:x}", pack_visible_request_id(n, salt))
}
pub use act_types::Metadata;
pub(crate) enum ComponentRequest {
ListTools {
metadata: Metadata,
reply: oneshot::Sender<Result<act::tools::types::ListToolsResponse, ComponentError>>,
},
CallTool {
name: String,
arguments: Vec<u8>,
metadata: Vec<(String, Vec<u8>)>,
reply: oneshot::Sender<Result<CallToolResult, ComponentError>>,
consent: Option<consent::ConsentSink>,
},
GetOpenSessionArgsSchema {
metadata: Vec<(String, Vec<u8>)>,
reply: oneshot::Sender<Result<String, ComponentError>>,
},
OpenSession {
args: Vec<(String, Vec<u8>)>,
metadata: Vec<(String, Vec<u8>)>,
reply: oneshot::Sender<Result<sessions::Session, ComponentError>>,
consent: Option<consent::ConsentSink>,
},
CloseSession {
session_id: String,
reply: oneshot::Sender<Result<(), ComponentError>>,
},
}
pub struct CallToolResult {
pub events: Vec<act::tools::types::ToolEvent>,
}
#[derive(Clone)]
pub struct ComponentHandle {
tx: mpsc::Sender<ComponentRequest>,
schemas: Arc<Mutex<HashMap<Option<String>, Arc<ToolSchemas>>>>,
}
type ToolSchemas = HashMap<String, Option<Arc<crate::validate::Validator>>>;
impl ComponentHandle {
pub(crate) fn new(tx: mpsc::Sender<ComponentRequest>) -> Self {
Self {
tx,
schemas: Arc::new(Mutex::new(HashMap::new())),
}
}
pub fn disconnected() -> Self {
let (tx, _rx) = mpsc::channel(1);
Self::new(tx)
}
async fn round_trip<T>(
&self,
build: impl FnOnce(oneshot::Sender<Result<T, ComponentError>>) -> ComponentRequest,
) -> Result<T, ComponentError> {
let (reply, answer) = oneshot::channel();
self.tx.send(build(reply)).await.map_err(|_| {
ComponentError::Internal(anyhow::anyhow!("component actor unavailable"))
})?;
answer.await.map_err(|_| {
ComponentError::Internal(anyhow::anyhow!("component actor dropped reply"))
})?
}
pub async fn list_tools(
&self,
metadata: &Metadata,
) -> Result<act::tools::types::ListToolsResponse, ComponentError> {
self.round_trip(|reply| ComponentRequest::ListTools {
metadata: metadata.clone(),
reply,
})
.await
}
pub async fn call_tool(
&self,
name: &str,
arguments: Vec<u8>,
metadata: Vec<(String, Vec<u8>)>,
consent: Option<consent::ConsentSink>,
) -> Result<CallToolResult, ComponentError> {
self.check_arguments(name, &arguments, &metadata).await?;
self.round_trip(|reply| ComponentRequest::CallTool {
name: name.to_string(),
arguments,
metadata,
reply,
consent,
})
.await
}
async fn check_arguments(
&self,
name: &str,
arguments: &[u8],
metadata: &[(String, Vec<u8>)],
) -> Result<(), ComponentError> {
let session = meta_str(&decode_meta_strings(metadata), "std:session-id");
let schemas = self.tool_schemas(session, metadata).await?;
let Some(Some(validator)) = schemas.get(name) else {
return Ok(());
};
let value = crate::validate::arguments_as_json(arguments)
.map_err(|e| ComponentError::Tool(crate::validate::invalid_args(e)))?;
validator.check(&value).map_err(|e| {
tracing::debug!(tool = %name, "arguments rejected before reaching the component");
ComponentError::Tool(crate::validate::invalid_args(format!(
"arguments do not match the schema for '{name}': {e}"
)))
})
}
async fn check_session_args(
&self,
args: &[(String, Vec<u8>)],
metadata: &[(String, Vec<u8>)],
) -> Result<(), ComponentError> {
let schema = self.open_session_args_schema(metadata.to_vec()).await?;
let Some(validator) = crate::validate::Validator::compile("open-session", &schema) else {
return Ok(());
};
let value = crate::validate::session_args_as_json(args)
.map_err(|e| ComponentError::Tool(crate::validate::invalid_args(e)))?;
validator.check(&value).map_err(|e| {
ComponentError::Tool(crate::validate::invalid_args(format!(
"session arguments do not match the component's schema: {e}"
)))
})
}
async fn tool_schemas(
&self,
session: Option<String>,
metadata: &[(String, Vec<u8>)],
) -> Result<Arc<ToolSchemas>, ComponentError> {
if let Some(hit) = self
.schemas
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner)
.get(&session)
{
return Ok(hit.clone());
}
let mut listing_meta = Metadata::new();
for (k, v) in metadata {
if let Ok(value) = act_types::cbor::cbor_to_json(v) {
listing_meta.insert(k.clone(), value);
}
}
let listed = self.list_tools(&listing_meta).await?;
let compiled: ToolSchemas = listed
.tools
.iter()
.map(|td| {
let v = crate::validate::Validator::compile(&td.name, &td.parameters_schema)
.map(Arc::new);
(td.name.clone(), v)
})
.collect();
let compiled = Arc::new(compiled);
self.schemas
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner)
.insert(session, compiled.clone());
Ok(compiled)
}
pub async fn open_session(
&self,
args: Vec<(String, Vec<u8>)>,
metadata: Vec<(String, Vec<u8>)>,
consent: Option<consent::ConsentSink>,
) -> Result<sessions::Session, ComponentError> {
self.check_session_args(&args, &metadata).await?;
self.round_trip(|reply| ComponentRequest::OpenSession {
args,
metadata,
reply,
consent,
})
.await
}
pub async fn close_session(&self, session_id: String) -> Result<(), ComponentError> {
self.round_trip(|reply| ComponentRequest::CloseSession { session_id, reply })
.await
}
async fn round_trip_servicing_consent<T, F, Fut>(
&self,
build: impl FnOnce(
oneshot::Sender<Result<T, ComponentError>>,
consent::ConsentSink,
) -> ComponentRequest,
mut answer: F,
) -> Result<T, ComponentError>
where
F: FnMut(String) -> Fut,
Fut: std::future::Future<Output = bool>,
{
let (reply, mut answer_rx) = oneshot::channel();
let (consent_tx, mut consent_rx) = mpsc::channel::<consent::ConsentRequest>(1);
self.tx.send(build(reply, consent_tx)).await.map_err(|_| {
ComponentError::Internal(anyhow::anyhow!("component actor unavailable"))
})?;
let reply = loop {
tokio::select! {
biased;
Some(ask) = consent_rx.recv() => {
let decision = answer(ask.message).await;
let _ = ask.reply.send(decision);
}
reply = &mut answer_rx => break reply,
}
};
reply.map_err(|_| {
ComponentError::Internal(anyhow::anyhow!("component actor dropped reply"))
})?
}
pub async fn call_tool_servicing_consent<F, Fut>(
&self,
name: &str,
arguments: Vec<u8>,
metadata: Vec<(String, Vec<u8>)>,
answer: F,
) -> Result<CallToolResult, ComponentError>
where
F: FnMut(String) -> Fut,
Fut: std::future::Future<Output = bool>,
{
self.round_trip_servicing_consent(
|reply, consent| ComponentRequest::CallTool {
name: name.to_string(),
arguments,
metadata,
reply,
consent: Some(consent),
},
answer,
)
.await
}
pub async fn open_session_servicing_consent<F, Fut>(
&self,
args: Vec<(String, Vec<u8>)>,
metadata: Vec<(String, Vec<u8>)>,
answer: F,
) -> Result<sessions::Session, ComponentError>
where
F: FnMut(String) -> Fut,
Fut: std::future::Future<Output = bool>,
{
self.round_trip_servicing_consent(
|reply, consent| ComponentRequest::OpenSession {
args,
metadata,
reply,
consent: Some(consent),
},
answer,
)
.await
}
pub async fn open_session_args_schema(
&self,
metadata: Vec<(String, Vec<u8>)>,
) -> Result<String, ComponentError> {
self.round_trip(|reply| ComponentRequest::GetOpenSessionArgsSchema { metadata, reply })
.await
}
}
pub use exports::act::tools::tool_provider::Guest as ToolProvider;
#[allow(clippy::too_many_arguments)]
pub async fn instantiate_component(
engine: &Engine,
component: &Component,
linker: &Linker<HostState>,
preopens: &[fs_policy::Preopen],
grant_policy: &act_policy::grant::GrantPolicy,
info: &ComponentInfo,
max_memory: Option<usize>,
prompter: Arc<dyn act_policy::consent::ConsentPrompter>,
cache: Arc<act_policy::consent::DecisionCache>,
credentials: Option<Arc<credentials::CredentialHost>>,
audit: &AuditContext,
) -> Result<(
ToolProvider,
Option<sessions::SessionProvider>,
Store<HostState>,
)> {
use exports::act::sessions::session_provider::GuestIndices as SessionGuestIndices;
use exports::act::tools::tool_provider::GuestIndices as ToolGuestIndices;
let (mut store, ceilings) = create_store(
engine,
preopens,
grant_policy,
info,
max_memory,
prompter,
cache,
credentials,
&audit.component_ref,
)
.await?;
let pre = linker
.instantiate_pre(component)
.map_err(|e| anyhow::anyhow!("failed to pre-instantiate component: {e}"))?;
let tool_indices =
ToolGuestIndices::new(&pre).map_err(|e| anyhow::anyhow!("tool-provider indices: {e}"))?;
let session_indices = SessionGuestIndices::new(&pre).ok();
let instance = pre
.instantiate_async(&mut store)
.await
.map_err(|e| anyhow::anyhow!("failed to instantiate component: {e}"))?;
let tool_provider = tool_indices
.load(&mut store, &instance)
.map_err(|e| anyhow::anyhow!("failed to load tool-provider: {e}"))?;
let session_provider = match session_indices {
Some(idx) => {
let guest = idx
.load(&mut store, &instance)
.map_err(|e| anyhow::anyhow!("failed to load session-provider: {e}"))?;
Some(sessions::SessionProvider::from_guest(&guest))
}
None => None,
};
let inst_span = crate::audit::instantiation_span(&audit.component_ref, &audit.digest);
{
let _g = inst_span.enter();
for (id, c) in &ceilings {
crate::audit::emit_ceiling_class(&crate::audit::CeilingClassRecord {
cap_id: id.clone(),
mode: c.effective_mode().to_string(),
declared: c.declared(),
has_prompt_channel: audit.has_prompt_channel,
});
}
}
drop(inst_span);
Ok((tool_provider, session_provider, store))
}
pub fn spawn_component_actor(
tool_provider: ToolProvider,
session_provider: Option<sessions::SessionProvider>,
mut store: Store<HostState>,
current_consent: Arc<consent::CurrentConsentSink>,
audit: AuditContext,
) -> ComponentHandle {
let (tx, mut rx) = mpsc::channel::<ComponentRequest>(32);
let mut tracked_sessions: Vec<String> = Vec::new();
let credentials = store.data().credentials.clone();
tokio::spawn(async move {
while let Some(request) = rx.recv().await {
match request {
ComponentRequest::ListTools { metadata, reply } => {
let provider = tool_provider.clone();
let result = store
.run_concurrent(async |accessor| {
provider
.call_list_tools(accessor, metadata.clone().into())
.await
})
.await;
let response = match result {
Ok(Ok(Ok(list_response))) => Ok(list_response),
Ok(Ok(Err(tool_error))) => Err(ComponentError::Tool(tool_error)),
Ok(Err(e)) => Err(ComponentError::Internal(anyhow::anyhow!(
"list-tools failed: {e}"
))),
Err(e) => Err(ComponentError::Internal(anyhow::anyhow!(
"run_concurrent failed: {e}"
))),
};
let _ = reply.send(response);
}
ComponentRequest::CallTool {
name,
arguments,
metadata,
reply,
consent,
} => {
current_consent.set(consent);
let provider = tool_provider.clone();
let started = std::time::Instant::now();
let meta_strings = decode_meta_strings(&metadata);
let audit_span = crate::audit::tool_call_span(&crate::audit::ToolCallStart {
component_ref: audit.component_ref.clone(),
digest: audit.digest.clone(),
tool: name.clone(),
args_sha256: crate::audit::sha256_hex(&arguments),
args_json: args_as_json(&arguments, audit.record_args),
session_id: meta_str(&meta_strings, act_types::constants::META_SESSION_ID),
agent_id: meta_str(&meta_strings, act_types::constants::META_AGENT_ID),
request_id: meta_str(&meta_strings, act_types::constants::META_REQUEST_ID)
.unwrap_or_else(new_request_id),
traceparent: meta_str(
&meta_strings,
act_types::constants::META_TRACEPARENT,
),
tracestate: meta_str(&meta_strings, act_types::constants::META_TRACESTATE),
transport: audit.transport,
});
let collected: Arc<std::sync::Mutex<Vec<act::tools::types::ToolEvent>>> =
Arc::new(std::sync::Mutex::new(Vec::new()));
let collected2 = collected.clone();
let (done_tx, done_rx) = oneshot::channel::<()>();
let result = store
.run_concurrent(async |accessor| {
let tool_result = provider
.call_call_tool(
accessor,
name.clone(),
arguments.clone(),
metadata.clone(),
)
.await?;
accessor.with(|access| match tool_result {
exports::act::tools::tool_provider::ToolResult::Streaming(
stream,
) => {
let consumer = CollectingConsumer {
collected,
done_tx: Some(done_tx),
};
let _ = stream.pipe(access, consumer);
}
exports::act::tools::tool_provider::ToolResult::Immediate(
events,
) => {
collected
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner)
.extend(events);
let _ = done_tx.send(());
}
});
let _ = done_rx.await;
Ok::<_, wasmtime::Error>(())
})
.instrument(audit_span.clone())
.await;
let response = match result {
Ok(Ok(())) => {
let events = collected2
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner)
.drain(..)
.collect();
Ok(CallToolResult { events })
}
Ok(Err(e)) => Err(ComponentError::Internal(anyhow::anyhow!(
"call-tool failed: {e}"
))),
Err(e) => Err(ComponentError::Internal(anyhow::anyhow!(
"run_concurrent failed: {e}"
))),
};
current_consent.set(None);
let outcome = match &response {
Ok(r) if events_contain_error(&r.events) => {
crate::audit::Outcome::ToolError
}
Ok(_) => crate::audit::Outcome::Ok,
Err(ComponentError::Tool(_)) => crate::audit::Outcome::ToolError,
Err(_) => crate::audit::Outcome::HostError,
};
crate::audit::finish_tool_call(&audit_span, outcome, started.elapsed());
let _ = reply.send(response);
}
ComponentRequest::GetOpenSessionArgsSchema { metadata, reply } => {
let response = match &session_provider {
Some(sp) => {
let sp = sp.clone();
let result = store
.run_concurrent(async |accessor| {
sp.get_open_session_args_schema
.call_concurrent(&accessor, (metadata,))
.await
})
.await;
session_call_to_response(result, |(r,)| r)
}
None => Err(ComponentError::Internal(anyhow::anyhow!(
"component does not export act:sessions/session-provider"
))),
};
let _ = reply.send(response);
}
ComponentRequest::OpenSession {
args,
metadata,
reply,
consent,
} => {
current_consent.set(consent);
let response = match &session_provider {
Some(sp) => {
let sp = sp.clone();
let result = store
.run_concurrent(async |accessor| {
sp.open_session
.call_concurrent(&accessor, (args, metadata))
.await
})
.await;
let inner = session_call_to_response(result, |(r,)| r);
if let Ok(s) = &inner {
tracked_sessions.push(s.id.clone());
if let Some(c) = &credentials {
c.note_session_opened(&s.id);
}
}
inner
}
None => Err(ComponentError::Internal(anyhow::anyhow!(
"component does not export act:sessions/session-provider"
))),
};
current_consent.set(None);
let _ = reply.send(response);
}
ComponentRequest::CloseSession { session_id, reply } => {
let response: Result<(), ComponentError> = match &session_provider {
Some(sp) => {
let sp = sp.clone();
let id = session_id.clone();
let result = store
.run_concurrent(async |accessor| {
sp.close_session.call_concurrent(&accessor, (id,)).await
})
.await;
tracked_sessions.retain(|sid| sid != &session_id);
if let Some(c) = &credentials {
c.note_session_closed(&session_id);
}
match result {
Ok(Ok(())) => Ok(()),
Ok(Err(e)) => Err(ComponentError::Internal(anyhow::anyhow!(
"close-session failed: {e}"
))),
Err(e) => Err(ComponentError::Internal(anyhow::anyhow!(
"run_concurrent failed: {e}"
))),
}
}
None => Err(ComponentError::Internal(anyhow::anyhow!(
"component does not export act:sessions/session-provider"
))),
};
let _ = reply.send(response);
}
}
}
if let Some(sp) = &session_provider {
for id in std::mem::take(&mut tracked_sessions) {
if let Some(c) = &credentials {
c.note_session_closed(&id);
}
let sp = sp.clone();
let _ = store
.run_concurrent(async |accessor| {
sp.close_session.call_concurrent(&accessor, (id,)).await
})
.await;
}
}
});
ComponentHandle::new(tx)
}
fn session_call_to_response<R, F>(
raw: wasmtime::Result<wasmtime::Result<(Result<R, act::core::types::Error>,)>>,
extract: F,
) -> Result<R, ComponentError>
where
F: FnOnce((Result<R, act::core::types::Error>,)) -> Result<R, act::core::types::Error>,
{
match raw {
Ok(Ok(tuple)) => match extract(tuple) {
Ok(r) => Ok(r),
Err(e) => Err(ComponentError::Tool(e)),
},
Ok(Err(e)) => Err(ComponentError::Internal(anyhow::anyhow!(
"session-provider call failed: {e}"
))),
Err(e) => Err(ComponentError::Internal(anyhow::anyhow!(
"run_concurrent failed: {e}"
))),
}
}
struct CollectingConsumer {
collected: Arc<std::sync::Mutex<Vec<act::tools::types::ToolEvent>>>,
done_tx: Option<oneshot::Sender<()>>,
}
impl StreamConsumer<HostState> for CollectingConsumer {
type Item = act::tools::types::ToolEvent;
fn poll_consume(
mut self: Pin<&mut Self>,
_cx: &mut Context<'_>,
store: StoreContextMut<HostState>,
mut source: Source<'_, Self::Item>,
finish: bool,
) -> Poll<wasmtime::Result<StreamResult>> {
let mut buffer = Vec::with_capacity(64);
source.read(store, &mut buffer)?;
if !buffer.is_empty() {
self.collected
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner)
.extend(buffer);
}
if finish {
if let Some(tx) = self.done_tx.take() {
let _ = tx.send(());
}
Poll::Ready(Ok(StreamResult::Dropped))
} else {
Poll::Ready(Ok(StreamResult::Completed))
}
}
}