use std::collections::HashMap;
use std::sync::{Arc, Mutex, RwLock};
use meerkat_core::service::{CreateSessionRequest, SessionBuildOptions};
use meerkat_core::types::{AssistantBlock, Message, ToolDef};
use meerkat_core::{
AgentError, AgentLlmClient, AgentLlmFallbackSwitch, CompiledSchema, LlmStreamResult,
OutputSchema, ProviderParamsOverride, ProviderRequestPressure, SchemaError, SessionLlmIdentity,
};
use crate::member_comms_id;
use crate::memory::taint::SessionTaintTracker;
#[derive(Clone, Default)]
pub struct DispatchTaintSlot {
inner: Arc<RwLock<Option<SessionTaintTracker>>>,
}
impl DispatchTaintSlot {
pub fn fill(&self, tracker: SessionTaintTracker) {
*self
.inner
.write()
.unwrap_or_else(std::sync::PoisonError::into_inner) = Some(tracker);
}
fn tracker(&self) -> Option<SessionTaintTracker> {
self.inner
.read()
.unwrap_or_else(std::sync::PoisonError::into_inner)
.clone()
}
}
impl std::fmt::Debug for DispatchTaintSlot {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("DispatchTaintSlot")
.field("filled", &self.tracker().is_some())
.finish()
}
}
fn member_taint_identity(req: &CreateSessionRequest) -> Option<String> {
let binding = req.build.as_ref()?.mob_member_binding.as_ref()?;
Some(member_comms_id::logical_memory_identity(&binding.member))
}
pub(crate) fn attach_member_taint_decorator(
req: &mut CreateSessionRequest,
slot: &DispatchTaintSlot,
) {
let Some(identity) = member_taint_identity(req) else {
return;
};
let build = req.build.get_or_insert_with(SessionBuildOptions::default);
let prior = build.agent_llm_client_decorator.take();
let slot = slot.clone();
build.agent_llm_client_decorator = Some(Arc::new(move |client| {
let client = match prior.as_ref() {
Some(prior) => prior(client),
None => client,
};
Arc::new(TaintObservingLlmClient::new(
client,
identity.clone(),
slot.clone(),
))
}));
}
pub struct TaintObservingLlmClient {
inner: Arc<dyn AgentLlmClient>,
identity: String,
slot: DispatchTaintSlot,
scanned: Mutex<usize>,
}
impl TaintObservingLlmClient {
pub fn new(inner: Arc<dyn AgentLlmClient>, identity: String, slot: DispatchTaintSlot) -> Self {
Self {
inner,
identity,
slot,
scanned: Mutex::new(0),
}
}
fn mark_request_ingestions(
&self,
tracker: &SessionTaintTracker,
messages: &[Message],
tools: &[Arc<ToolDef>],
) {
let mut scanned = self
.scanned
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner);
let start = if *scanned > messages.len() {
0
} else {
*scanned
};
let mut names: HashMap<&str, &str> = HashMap::new();
for message in &messages[start..] {
match message {
Message::BlockAssistant(assistant) => {
for block in &assistant.blocks {
match block {
AssistantBlock::ToolUse { id, name, .. } => {
names.insert(id.as_str(), name.as_str());
}
AssistantBlock::ServerToolContent { kind, .. } => {
tracker.observe_dispatched_server_tool(&self.identity, kind);
}
_ => {}
}
}
}
Message::ToolResults { results, .. } => {
for result in results {
let Some(name) = names.get(result.tool_use_id.as_str()) else {
continue;
};
let provenance = tools
.iter()
.find(|tool| tool.name.as_ref() == *name)
.and_then(|tool| tool.provenance.as_ref());
tracker.observe_dispatched_tool_result(&self.identity, name, provenance);
}
}
_ => {}
}
}
*scanned = messages.len();
}
}
#[async_trait::async_trait]
impl AgentLlmClient for TaintObservingLlmClient {
async fn stream_response(
&self,
messages: &[Message],
tools: &[Arc<ToolDef>],
max_tokens: u32,
temperature: Option<f32>,
provider_params: Option<&ProviderParamsOverride>,
) -> Result<LlmStreamResult, AgentError> {
if let Some(tracker) = self.slot.tracker() {
self.mark_request_ingestions(&tracker, messages, tools);
}
let result = self
.inner
.stream_response(messages, tools, max_tokens, temperature, provider_params)
.await?;
if let Some(tracker) = self.slot.tracker() {
for block in result.blocks() {
if let AssistantBlock::ServerToolContent { kind, .. } = block {
tracker.observe_dispatched_server_tool(&self.identity, kind);
}
}
}
Ok(result)
}
fn request_pressure(
&self,
messages: &[Message],
tools: &[Arc<ToolDef>],
max_tokens: u32,
temperature: Option<f32>,
provider_params: Option<&ProviderParamsOverride>,
) -> Result<Option<ProviderRequestPressure>, AgentError> {
self.inner
.request_pressure(messages, tools, max_tokens, temperature, provider_params)
}
fn provider(&self) -> meerkat_core::Provider {
self.inner.provider()
}
fn model(&self) -> &str {
self.inner.model()
}
fn prepare_model_fallback(&self, failure: &AgentError) -> Option<AgentLlmFallbackSwitch> {
self.inner.prepare_model_fallback(failure)
}
fn commit_model_fallback(
&self,
previous_identity: &SessionLlmIdentity,
target_identity: &SessionLlmIdentity,
) -> Result<(), AgentError> {
self.inner
.commit_model_fallback(previous_identity, target_identity)
}
fn active_model_fallback_identity(&self) -> Option<SessionLlmIdentity> {
self.inner.active_model_fallback_identity()
}
fn compile_model_fallback_schema(
&self,
target_identity: &SessionLlmIdentity,
output_schema: &OutputSchema,
) -> Result<CompiledSchema, AgentError> {
self.inner
.compile_model_fallback_schema(target_identity, output_schema)
}
fn begin_stream_output_observation(&self) {
self.inner.begin_stream_output_observation();
}
fn stream_output_observed(&self) -> bool {
self.inner.stream_output_observed()
}
fn stream_activity_count(&self) -> Option<u64> {
self.inner.stream_activity_count()
}
fn compile_schema(&self, output_schema: &OutputSchema) -> Result<CompiledSchema, SchemaError> {
self.inner.compile_schema(output_schema)
}
}
#[cfg(test)]
#[allow(clippy::expect_used, clippy::unwrap_used, clippy::panic)]
mod tests {
use super::*;
use crate::memory::taint::ContentTrustConfig;
use meerkat_core::types::{
BlockAssistantMessage, ContentBlock, ServerToolKind, StopReason, ToolProvenance,
ToolResult, ToolSourceKind, Usage,
};
use serde_json::value::RawValue;
struct ScriptedInner {
blocks: Vec<AssistantBlock>,
}
#[async_trait::async_trait]
impl AgentLlmClient for ScriptedInner {
async fn stream_response(
&self,
_messages: &[Message],
_tools: &[Arc<ToolDef>],
_max_tokens: u32,
_temperature: Option<f32>,
_provider_params: Option<&ProviderParamsOverride>,
) -> Result<LlmStreamResult, AgentError> {
Ok(LlmStreamResult::new(
self.blocks.clone(),
StopReason::EndTurn,
Usage::default(),
))
}
fn provider(&self) -> meerkat_core::Provider {
meerkat_core::Provider::OpenAI
}
fn model(&self) -> &'static str {
"gpt-5.5"
}
}
fn tool_use(id: &str, name: &str) -> AssistantBlock {
AssistantBlock::ToolUse {
id: id.to_string(),
name: name.to_string(),
args: RawValue::from_string("{}".to_string()).expect("raw args"),
meta: None,
}
}
fn assistant(blocks: Vec<AssistantBlock>) -> Message {
Message::BlockAssistant(BlockAssistantMessage::new(blocks, StopReason::ToolUse))
}
fn tool_results(id: &str, text: &str) -> Message {
Message::tool_results(vec![ToolResult {
tool_use_id: id.to_string(),
content: vec![ContentBlock::Text {
text: text.to_string(),
}],
is_error: false,
}])
}
fn mcp_tool(name: &str, server: &str) -> Arc<ToolDef> {
Arc::new(ToolDef {
name: name.into(),
description: String::new(),
input_schema: serde_json::json!({"type": "object"}),
provenance: Some(ToolProvenance {
kind: ToolSourceKind::Mcp,
source_id: server.into(),
}),
})
}
async fn drive(client: &TaintObservingLlmClient, messages: &[Message], tools: &[Arc<ToolDef>]) {
client
.stream_response(messages, tools, 128, None, None)
.await
.expect("scripted call succeeds");
}
#[tokio::test]
async fn marks_unqualified_mcp_tool_via_request_catalog_provenance() {
let tracker = SessionTaintTracker::new(ContentTrustConfig::default());
let slot = DispatchTaintSlot::default();
slot.fill(tracker.clone());
let client = TaintObservingLlmClient::new(
Arc::new(ScriptedInner { blocks: vec![] }),
"identity:a".to_string(),
slot,
);
let messages = vec![
assistant(vec![tool_use("call-1", "scrape_page")]),
tool_results("call-1", "attacker text"),
];
let tools = vec![mcp_tool("scrape_page", "scraper")];
drive(&client, &messages, &tools).await;
let taint = tracker
.identity_taint("identity:a")
.expect("MCP result must mark before the call proceeds");
assert!(
taint.source.contains("MCP server 'scraper'"),
"{}",
taint.source
);
}
#[tokio::test]
async fn falls_back_to_name_classification_without_provenance() {
let tracker = SessionTaintTracker::new(ContentTrustConfig::default());
let slot = DispatchTaintSlot::default();
slot.fill(tracker.clone());
let client = TaintObservingLlmClient::new(
Arc::new(ScriptedInner { blocks: vec![] }),
"identity:a".to_string(),
slot,
);
let plain = Arc::new(ToolDef {
name: "lookup".into(),
description: String::new(),
input_schema: serde_json::json!({"type": "object"}),
provenance: None,
});
let messages = vec![
assistant(vec![tool_use("call-1", "lookup")]),
tool_results("call-1", "fine"),
];
drive(&client, &messages, std::slice::from_ref(&plain)).await;
assert!(tracker.identity_taint("identity:a").is_none());
let messages = vec![
assistant(vec![tool_use("call-1", "lookup")]),
tool_results("call-1", "fine"),
assistant(vec![tool_use("call-2", "web_fetch")]),
tool_results("call-2", "attacker text"),
];
drive(&client, &messages, std::slice::from_ref(&plain)).await;
assert!(
tracker.identity_taint("identity:a").is_some(),
"web builtins classify untrusted with no catalog entry at all"
);
}
#[tokio::test]
async fn marks_server_tool_content_from_the_response() {
let tracker = SessionTaintTracker::new(ContentTrustConfig::default());
let slot = DispatchTaintSlot::default();
slot.fill(tracker.clone());
let client = TaintObservingLlmClient::new(
Arc::new(ScriptedInner {
blocks: vec![AssistantBlock::ServerToolContent {
id: None,
kind: ServerToolKind::WebSearch,
content: serde_json::json!({"results": []}),
meta: None,
}],
}),
"identity:a".to_string(),
slot,
);
drive(&client, &[], &[]).await;
let taint = tracker.identity_taint("identity:a").expect("marks");
assert!(taint.source.contains("web_search"), "{}", taint.source);
}
#[tokio::test]
async fn unfilled_slot_is_inert_and_late_fill_activates() {
let slot = DispatchTaintSlot::default();
let client = TaintObservingLlmClient::new(
Arc::new(ScriptedInner { blocks: vec![] }),
"identity:a".to_string(),
slot.clone(),
);
let messages = vec![
assistant(vec![tool_use("call-1", "web_fetch")]),
tool_results("call-1", "attacker text"),
];
drive(&client, &messages, &[]).await;
let tracker = SessionTaintTracker::new(ContentTrustConfig::default());
slot.fill(tracker.clone());
assert!(tracker.identity_taint("identity:a").is_none());
drive(&client, &messages, &[]).await;
assert!(tracker.identity_taint("identity:a").is_some());
}
#[test]
fn member_identity_resolves_binding_to_the_write_gate_spelling() {
let mut req = CreateSessionRequest {
model: "gpt-5.5".to_string(),
prompt: meerkat_core::ContentInput::Text("hi".to_string()),
injected_context: Vec::new(),
system_prompt: meerkat_core::config::SystemPromptOverride::Inherit,
max_tokens: None,
event_tx: None,
initial_turn: meerkat_core::service::InitialTurnPolicy::Defer,
deferred_prompt_policy: meerkat_core::service::DeferredPromptPolicy::default(),
build: Some(SessionBuildOptions {
mob_member_binding: Some(meerkat_core::MobMemberBinding {
mob_id: "mob-1".to_string(),
role: "worker".to_string(),
member: member_comms_id::mob_member_id_str("rt:review:singleton:0")
.into_owned(),
}),
..SessionBuildOptions::default()
}),
labels: None,
};
assert_eq!(
member_taint_identity(&req).as_deref(),
Some("review:singleton"),
"rt:{{identity}}:{{generation}} normalizes to the durable identity"
);
let binding = req
.build
.as_mut()
.and_then(|build| build.mob_member_binding.as_mut())
.expect("binding");
binding.member = "helper".to_string();
assert_eq!(member_taint_identity(&req).as_deref(), Some("helper"));
let binding = req
.build
.as_mut()
.and_then(|build| build.mob_member_binding.as_mut())
.expect("binding");
binding.member = member_comms_id::mob_member_id_str("review:singleton").into_owned();
assert_eq!(
member_taint_identity(&req).as_deref(),
Some("review:singleton")
);
let binding = req
.build
.as_mut()
.and_then(|build| build.mob_member_binding.as_mut())
.expect("binding");
binding.member = member_comms_id::mob_member_id_str("rt:oddly:named").into_owned();
assert_eq!(
member_taint_identity(&req).as_deref(),
Some("rt:oddly:named")
);
req.build = None;
assert_eq!(member_taint_identity(&req), None);
let mut req = req;
attach_member_taint_decorator(&mut req, &DispatchTaintSlot::default());
assert!(req.build.is_none(), "non-member requests stay untouched");
}
}