use crate::{DrivenError, Result};
use std::collections::HashMap;
pub use dcp::binary::{
ArgType, BinaryMessageEnvelope, ChunkFlags, Flags, MessageType, SignedInvocation,
SignedToolDef, StreamChunk, ToolInvocation,
};
pub use dcp::capability::CapabilityManifest;
pub use dcp::compat::adapter::AdapterError;
pub use dcp::compat::json_rpc::RequestId;
pub use dcp::compat::{JsonRpcError, JsonRpcParser, JsonRpcRequest, JsonRpcResponse, McpAdapter};
pub use dcp::security::{Signer, Verifier};
#[derive(Debug, Clone)]
pub struct ZctiBuilder {
tool_id: u32,
arg_layout: u64,
args_data: Vec<u8>,
arg_index: usize,
}
impl ZctiBuilder {
pub fn new(tool_id: u32) -> Self {
Self {
tool_id,
arg_layout: 0,
args_data: Vec::new(),
arg_index: 0,
}
}
pub fn add_null(mut self) -> Self {
self.set_arg_type(ArgType::Null);
self
}
pub fn add_bool(mut self, value: bool) -> Self {
self.set_arg_type(ArgType::Bool);
self.args_data.push(if value { 1 } else { 0 });
self
}
pub fn add_i32(mut self, value: i32) -> Self {
self.set_arg_type(ArgType::I32);
self.args_data.extend_from_slice(&value.to_le_bytes());
self
}
pub fn add_i64(mut self, value: i64) -> Self {
self.set_arg_type(ArgType::I64);
self.args_data.extend_from_slice(&value.to_le_bytes());
self
}
pub fn add_f64(mut self, value: f64) -> Self {
self.set_arg_type(ArgType::F64);
self.args_data.extend_from_slice(&value.to_le_bytes());
self
}
pub fn add_string(mut self, value: &str) -> Self {
self.set_arg_type(ArgType::String);
let len = value.len() as u32;
self.args_data.extend_from_slice(&len.to_le_bytes());
self.args_data.extend_from_slice(value.as_bytes());
self
}
pub fn add_bytes(mut self, value: &[u8]) -> Self {
self.set_arg_type(ArgType::Bytes);
let len = value.len() as u32;
self.args_data.extend_from_slice(&len.to_le_bytes());
self.args_data.extend_from_slice(value);
self
}
pub fn add_raw(mut self, data: &[u8]) -> Self {
self.args_data.extend_from_slice(data);
self
}
fn set_arg_type(&mut self, arg_type: ArgType) {
if self.arg_index < ToolInvocation::MAX_ARGS {
let shift = self.arg_index * 4;
self.arg_layout |= (arg_type as u64) << shift;
self.arg_index += 1;
}
}
pub fn build(self) -> ToolInvocation {
ToolInvocation::new(
self.tool_id,
self.arg_layout,
0, self.args_data.len() as u32,
)
}
pub fn build_with_data(self) -> (ToolInvocation, Vec<u8>) {
let invocation = ToolInvocation::new(
self.tool_id,
self.arg_layout,
0,
self.args_data.len() as u32,
);
(invocation, self.args_data)
}
pub fn arg_count(&self) -> usize {
self.arg_index
}
pub fn data_size(&self) -> usize {
self.args_data.len()
}
}
#[derive(Debug)]
pub struct ZctiReader<'a> {
invocation: ToolInvocation,
data: &'a [u8],
offset: usize,
arg_index: usize,
}
impl<'a> ZctiReader<'a> {
pub fn new(invocation: &ToolInvocation, data: &'a [u8]) -> Self {
Self {
invocation: *invocation,
data,
offset: 0,
arg_index: 0,
}
}
pub fn tool_id(&self) -> u32 {
self.invocation.tool_id
}
pub fn arg_count(&self) -> usize {
self.invocation.arg_count()
}
pub fn peek_type(&self) -> Option<ArgType> {
self.invocation.get_arg_type(self.arg_index)
}
pub fn read_bool(&mut self) -> Result<bool> {
self.check_type(ArgType::Bool)?;
if self.offset >= self.data.len() {
return Err(DrivenError::InvalidBinary(
"Insufficient data for bool".to_string(),
));
}
let value = self.data[self.offset] != 0;
self.offset += 1;
self.arg_index += 1;
Ok(value)
}
pub fn read_i32(&mut self) -> Result<i32> {
self.check_type(ArgType::I32)?;
if self.offset + 4 > self.data.len() {
return Err(DrivenError::InvalidBinary(
"Insufficient data for i32".to_string(),
));
}
let bytes: [u8; 4] = self.data[self.offset..self.offset + 4].try_into().unwrap();
let value = i32::from_le_bytes(bytes);
self.offset += 4;
self.arg_index += 1;
Ok(value)
}
pub fn read_i64(&mut self) -> Result<i64> {
self.check_type(ArgType::I64)?;
if self.offset + 8 > self.data.len() {
return Err(DrivenError::InvalidBinary(
"Insufficient data for i64".to_string(),
));
}
let bytes: [u8; 8] = self.data[self.offset..self.offset + 8].try_into().unwrap();
let value = i64::from_le_bytes(bytes);
self.offset += 8;
self.arg_index += 1;
Ok(value)
}
pub fn read_f64(&mut self) -> Result<f64> {
self.check_type(ArgType::F64)?;
if self.offset + 8 > self.data.len() {
return Err(DrivenError::InvalidBinary(
"Insufficient data for f64".to_string(),
));
}
let bytes: [u8; 8] = self.data[self.offset..self.offset + 8].try_into().unwrap();
let value = f64::from_le_bytes(bytes);
self.offset += 8;
self.arg_index += 1;
Ok(value)
}
pub fn read_string(&mut self) -> Result<&'a str> {
self.check_type(ArgType::String)?;
if self.offset + 4 > self.data.len() {
return Err(DrivenError::InvalidBinary(
"Insufficient data for string length".to_string(),
));
}
let len_bytes: [u8; 4] = self.data[self.offset..self.offset + 4].try_into().unwrap();
let len = u32::from_le_bytes(len_bytes) as usize;
self.offset += 4;
if self.offset + len > self.data.len() {
return Err(DrivenError::InvalidBinary(
"Insufficient data for string content".to_string(),
));
}
let value = std::str::from_utf8(&self.data[self.offset..self.offset + len])
.map_err(|e| DrivenError::InvalidBinary(format!("Invalid UTF-8: {}", e)))?;
self.offset += len;
self.arg_index += 1;
Ok(value)
}
pub fn read_bytes(&mut self) -> Result<&'a [u8]> {
self.check_type(ArgType::Bytes)?;
if self.offset + 4 > self.data.len() {
return Err(DrivenError::InvalidBinary(
"Insufficient data for bytes length".to_string(),
));
}
let len_bytes: [u8; 4] = self.data[self.offset..self.offset + 4].try_into().unwrap();
let len = u32::from_le_bytes(len_bytes) as usize;
self.offset += 4;
if self.offset + len > self.data.len() {
return Err(DrivenError::InvalidBinary(
"Insufficient data for bytes content".to_string(),
));
}
let value = &self.data[self.offset..self.offset + len];
self.offset += len;
self.arg_index += 1;
Ok(value)
}
pub fn skip(&mut self) -> Result<()> {
let arg_type = self
.peek_type()
.ok_or_else(|| DrivenError::InvalidBinary("No more arguments".to_string()))?;
match arg_type {
ArgType::Null => {
self.arg_index += 1;
}
ArgType::Bool => {
self.offset += 1;
self.arg_index += 1;
}
ArgType::I32 => {
self.offset += 4;
self.arg_index += 1;
}
ArgType::I64 | ArgType::F64 => {
self.offset += 8;
self.arg_index += 1;
}
ArgType::String | ArgType::Bytes => {
if self.offset + 4 > self.data.len() {
return Err(DrivenError::InvalidBinary(
"Insufficient data for length".to_string(),
));
}
let len_bytes: [u8; 4] =
self.data[self.offset..self.offset + 4].try_into().unwrap();
let len = u32::from_le_bytes(len_bytes) as usize;
self.offset += 4 + len;
self.arg_index += 1;
}
ArgType::Array | ArgType::Object => {
return Err(DrivenError::InvalidBinary(
"Complex types not yet supported".to_string(),
));
}
}
Ok(())
}
pub fn has_more(&self) -> bool {
self.arg_index < self.arg_count()
}
pub fn remaining_data(&self) -> &'a [u8] {
&self.data[self.offset..]
}
fn check_type(&self, expected: ArgType) -> Result<()> {
let actual = self
.peek_type()
.ok_or_else(|| DrivenError::InvalidBinary("No more arguments".to_string()))?;
if actual != expected {
return Err(DrivenError::InvalidBinary(format!(
"Type mismatch: expected {:?}, got {:?}",
expected, actual
)));
}
Ok(())
}
}
#[derive(Debug)]
pub struct SharedArgBuffer {
data: Vec<u8>,
write_offset: usize,
}
impl SharedArgBuffer {
pub fn new(capacity: usize) -> Self {
Self {
data: vec![0u8; capacity],
write_offset: 0,
}
}
pub fn with_default_capacity() -> Self {
Self::new(65536)
}
pub fn write(&mut self, data: &[u8]) -> Result<u32> {
if self.write_offset + data.len() > self.data.len() {
return Err(DrivenError::InvalidBinary("Buffer overflow".to_string()));
}
let offset = self.write_offset as u32;
self.data[self.write_offset..self.write_offset + data.len()].copy_from_slice(data);
self.write_offset += data.len();
Ok(offset)
}
pub fn read(&self, offset: u32, len: u32) -> Result<&[u8]> {
let start = offset as usize;
let end = start + len as usize;
if end > self.data.len() {
return Err(DrivenError::InvalidBinary("Read out of bounds".to_string()));
}
Ok(&self.data[start..end])
}
pub fn offset(&self) -> u32 {
self.write_offset as u32
}
pub fn reset(&mut self) {
self.write_offset = 0;
}
pub fn data(&self) -> &[u8] {
&self.data[..self.write_offset]
}
pub fn capacity(&self) -> usize {
self.data.len()
}
pub fn remaining(&self) -> usize {
self.data.len() - self.write_offset
}
}
#[derive(Debug, Clone)]
pub struct DcpConfig {
pub enabled: bool,
pub prefer_dcp: bool,
pub endpoint: Option<String>,
pub timeout_ms: u64,
pub signing_enabled: bool,
pub signing_seed: Option<[u8; 32]>,
}
impl Default for DcpConfig {
fn default() -> Self {
Self {
enabled: true,
prefer_dcp: true,
endpoint: None,
timeout_ms: 5000,
signing_enabled: false,
signing_seed: None,
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum ConnectionState {
Disconnected,
Connecting,
ConnectedDcp,
ConnectedMcp,
Failed,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum Protocol {
Dcp,
Mcp,
}
#[derive(Debug)]
pub enum InvocationResult {
Dcp(Vec<u8>),
Mcp(String),
}
impl InvocationResult {
pub fn is_dcp(&self) -> bool {
matches!(self, InvocationResult::Dcp(_))
}
pub fn is_mcp(&self) -> bool {
matches!(self, InvocationResult::Mcp(_))
}
pub fn as_dcp(&self) -> Option<&[u8]> {
match self {
InvocationResult::Dcp(bytes) => Some(bytes),
_ => None,
}
}
pub fn as_mcp(&self) -> Option<&str> {
match self {
InvocationResult::Mcp(json) => Some(json),
_ => None,
}
}
}
pub struct DcpClient {
config: DcpConfig,
state: ConnectionState,
mcp_adapter: Option<McpAdapter>,
capabilities: CapabilityManifest,
signer: Option<Signer>,
tools: HashMap<u32, ToolDefinition>,
next_tool_id: u32,
}
#[derive(Debug, Clone)]
pub struct ToolDefinition {
pub id: u32,
pub name: String,
pub description: String,
pub schema_hash: [u8; 32],
pub capabilities: u64,
}
#[derive(Debug, Clone, Copy)]
pub struct CapabilityStats {
pub tool_count: u32,
pub resource_count: u32,
pub prompt_count: u32,
pub extension_count: u32,
pub version: u16,
}
#[derive(Debug, Clone)]
pub struct DcpMessage {
pub message_type: MessageType,
pub flags: u8,
pub payload: Vec<u8>,
}
impl DcpMessage {
pub fn new(message_type: MessageType, flags: u8, payload: Vec<u8>) -> Self {
Self {
message_type,
flags,
payload,
}
}
pub fn tool(payload: Vec<u8>) -> Self {
Self::new(MessageType::Tool, 0, payload)
}
pub fn resource(payload: Vec<u8>) -> Self {
Self::new(MessageType::Resource, 0, payload)
}
pub fn prompt(payload: Vec<u8>) -> Self {
Self::new(MessageType::Prompt, 0, payload)
}
pub fn response(payload: Vec<u8>) -> Self {
Self::new(MessageType::Response, 0, payload)
}
pub fn error(payload: Vec<u8>) -> Self {
Self::new(MessageType::Error, 0, payload)
}
pub fn stream(payload: Vec<u8>) -> Self {
Self::new(MessageType::Stream, Flags::STREAMING, payload)
}
pub fn is_streaming(&self) -> bool {
self.flags & Flags::STREAMING != 0
}
pub fn is_compressed(&self) -> bool {
self.flags & Flags::COMPRESSED != 0
}
pub fn is_signed(&self) -> bool {
self.flags & Flags::SIGNED != 0
}
pub fn with_streaming(mut self) -> Self {
self.flags |= Flags::STREAMING;
self
}
pub fn with_compressed(mut self) -> Self {
self.flags |= Flags::COMPRESSED;
self
}
pub fn with_signed(mut self) -> Self {
self.flags |= Flags::SIGNED;
self
}
pub fn encode(&self) -> Vec<u8> {
let envelope =
BinaryMessageEnvelope::new(self.message_type, self.flags, self.payload.len() as u32);
let mut result = Vec::with_capacity(BinaryMessageEnvelope::SIZE + self.payload.len());
result.extend_from_slice(envelope.as_bytes());
result.extend_from_slice(&self.payload);
result
}
pub fn decode(bytes: &[u8]) -> Result<Self> {
if bytes.len() < BinaryMessageEnvelope::SIZE {
return Err(DrivenError::InvalidBinary(
"Insufficient data for BME header".to_string(),
));
}
let envelope = BinaryMessageEnvelope::from_bytes(bytes)
.map_err(|e| DrivenError::InvalidBinary(format!("Invalid BME: {:?}", e)))?;
let message_type = envelope.get_message_type().ok_or_else(|| {
DrivenError::InvalidBinary(format!("Unknown message type: {}", envelope.message_type))
})?;
let payload_len = envelope.payload_len as usize;
let total_len = BinaryMessageEnvelope::SIZE + payload_len;
if bytes.len() < total_len {
return Err(DrivenError::InvalidBinary(format!(
"Insufficient data: expected {} bytes, got {}",
total_len,
bytes.len()
)));
}
let payload = bytes[BinaryMessageEnvelope::SIZE..total_len].to_vec();
Ok(Self {
message_type,
flags: envelope.flags,
payload,
})
}
}
#[derive(Debug)]
pub struct StreamAssembler {
chunks: Vec<StreamChunk>,
data: Vec<u8>,
next_sequence: u32,
complete: bool,
error: bool,
}
impl StreamAssembler {
pub fn new() -> Self {
Self {
chunks: Vec::new(),
data: Vec::new(),
next_sequence: 0,
complete: false,
error: false,
}
}
pub fn add_chunk(&mut self, chunk_bytes: &[u8], payload: &[u8]) -> Result<()> {
if self.complete {
return Err(DrivenError::InvalidBinary(
"Stream already complete".to_string(),
));
}
let chunk = StreamChunk::from_bytes(chunk_bytes)
.map_err(|e| DrivenError::InvalidBinary(format!("Invalid chunk: {:?}", e)))?;
let sequence = chunk.sequence;
let flags = chunk.flags;
let len = chunk.len;
if sequence != self.next_sequence {
return Err(DrivenError::InvalidBinary(format!(
"Out of order chunk: expected {}, got {}",
self.next_sequence, sequence
)));
}
if flags & ChunkFlags::ERROR != 0 {
self.error = true;
self.complete = true;
return Err(DrivenError::InvalidBinary("Stream error".to_string()));
}
if payload.len() != len as usize {
return Err(DrivenError::InvalidBinary(format!(
"Payload length mismatch: header says {}, got {}",
len,
payload.len()
)));
}
self.data.extend_from_slice(payload);
self.chunks.push(*chunk);
self.next_sequence += 1;
if flags & ChunkFlags::LAST != 0 {
self.complete = true;
}
Ok(())
}
pub fn is_complete(&self) -> bool {
self.complete
}
pub fn has_error(&self) -> bool {
self.error
}
pub fn data(&self) -> Option<&[u8]> {
if self.complete && !self.error {
Some(&self.data)
} else {
None
}
}
pub fn take_data(self) -> Option<Vec<u8>> {
if self.complete && !self.error {
Some(self.data)
} else {
None
}
}
pub fn chunk_count(&self) -> usize {
self.chunks.len()
}
}
impl Default for StreamAssembler {
fn default() -> Self {
Self::new()
}
}
#[derive(Debug)]
pub struct StreamBuilder {
chunk_size: usize,
sequence: u32,
}
impl StreamBuilder {
pub fn new(chunk_size: usize) -> Self {
Self {
chunk_size,
sequence: 0,
}
}
pub fn with_default_chunk_size() -> Self {
Self::new(65536)
}
pub fn build_chunks(&mut self, data: &[u8]) -> Vec<(StreamChunk, Vec<u8>)> {
let mut chunks = Vec::new();
let mut offset = 0;
while offset < data.len() {
let remaining = data.len() - offset;
let chunk_len = remaining.min(self.chunk_size);
let is_first = offset == 0;
let is_last = offset + chunk_len >= data.len();
let flags = if is_first && is_last {
ChunkFlags::FIRST | ChunkFlags::LAST
} else if is_first {
ChunkFlags::FIRST
} else if is_last {
ChunkFlags::LAST
} else {
ChunkFlags::CONTINUE
};
let chunk = StreamChunk::new(self.sequence, flags, chunk_len as u16);
let payload = data[offset..offset + chunk_len].to_vec();
chunks.push((chunk, payload));
self.sequence += 1;
offset += chunk_len;
}
chunks
}
pub fn reset(&mut self) {
self.sequence = 0;
}
}
impl DcpClient {
pub fn new(config: DcpConfig) -> Self {
let signer = if config.signing_enabled {
config.signing_seed.map(|seed| Signer::from_seed(&seed))
} else {
None
};
Self {
config,
state: ConnectionState::Disconnected,
mcp_adapter: None,
capabilities: CapabilityManifest::new(1),
signer,
tools: HashMap::new(),
next_tool_id: 1,
}
}
pub fn with_defaults() -> Self {
Self::new(DcpConfig::default())
}
pub fn state(&self) -> ConnectionState {
self.state
}
pub fn is_connected(&self) -> bool {
matches!(
self.state,
ConnectionState::ConnectedDcp | ConnectionState::ConnectedMcp
)
}
pub fn is_dcp_connected(&self) -> bool {
self.state == ConnectionState::ConnectedDcp
}
pub fn is_mcp_connected(&self) -> bool {
self.state == ConnectionState::ConnectedMcp
}
pub fn prefer_dcp(&self) -> bool {
self.config.prefer_dcp && self.config.enabled
}
pub fn set_prefer_dcp(&mut self, prefer: bool) {
self.config.prefer_dcp = prefer;
}
pub fn enable_dcp(&mut self) {
self.config.enabled = true;
}
pub fn disable_dcp(&mut self) {
self.config.enabled = false;
}
pub fn is_dcp_enabled(&self) -> bool {
self.config.enabled
}
pub fn current_protocol(&self) -> Option<Protocol> {
match self.state {
ConnectionState::ConnectedDcp => Some(Protocol::Dcp),
ConnectionState::ConnectedMcp => Some(Protocol::Mcp),
_ => None,
}
}
pub fn select_protocol(&self) -> Option<Protocol> {
if self.is_dcp_connected() && self.prefer_dcp() {
Some(Protocol::Dcp)
} else if self.is_mcp_connected() || self.has_mcp_fallback() {
Some(Protocol::Mcp)
} else if self.is_dcp_connected() {
Some(Protocol::Dcp)
} else {
None
}
}
pub fn invoke_tool_auto(&self, tool_id: u32, args: &[u8]) -> Result<InvocationResult> {
match self.select_protocol() {
Some(Protocol::Dcp) => {
let encoded = self.invoke_tool(tool_id, args)?;
Ok(InvocationResult::Dcp(encoded))
}
Some(Protocol::Mcp) => {
let tool = self
.tools
.get(&tool_id)
.ok_or_else(|| DrivenError::Config(format!("Tool {} not found", tool_id)))?;
let request = serde_json::json!({
"jsonrpc": "2.0",
"method": "tools/call",
"params": {
"name": tool.name,
"arguments": serde_json::from_slice::<serde_json::Value>(args)
.unwrap_or(serde_json::Value::Null)
},
"id": 1
});
let response = self.handle_mcp_request(&request.to_string())?;
Ok(InvocationResult::Mcp(response))
}
None => Err(DrivenError::Config("No protocol available".to_string())),
}
}
pub fn capabilities(&self) -> &CapabilityManifest {
&self.capabilities
}
pub fn capabilities_mut(&mut self) -> &mut CapabilityManifest {
&mut self.capabilities
}
pub fn connect(&mut self, endpoint: &str) -> Result<()> {
self.state = ConnectionState::Connecting;
if self.config.enabled {
match self.try_dcp_connect(endpoint) {
Ok(()) => {
self.state = ConnectionState::ConnectedDcp;
return Ok(());
}
Err(e) => {
tracing::warn!("DCP connection failed, falling back to MCP: {}", e);
}
}
}
match self.try_mcp_connect(endpoint) {
Ok(()) => {
self.state = ConnectionState::ConnectedMcp;
Ok(())
}
Err(e) => {
self.state = ConnectionState::Failed;
Err(e)
}
}
}
fn try_dcp_connect(&mut self, _endpoint: &str) -> Result<()> {
Err(DrivenError::Config(
"DCP connection not yet implemented".to_string(),
))
}
fn try_mcp_connect(&mut self, _endpoint: &str) -> Result<()> {
self.mcp_adapter = Some(McpAdapter::new());
Ok(())
}
pub fn disconnect(&mut self) {
self.state = ConnectionState::Disconnected;
self.mcp_adapter = None;
}
pub fn register_tool(
&mut self,
name: &str,
description: &str,
schema_hash: [u8; 32],
capabilities: u64,
) -> u32 {
let id = self.next_tool_id;
self.next_tool_id += 1;
let tool = ToolDefinition {
id,
name: name.to_string(),
description: description.to_string(),
schema_hash,
capabilities,
};
self.tools.insert(id, tool);
if id < CapabilityManifest::MAX_TOOLS as u32 {
self.capabilities.set_tool(id as u16);
}
id
}
pub fn get_tool(&self, id: u32) -> Option<&ToolDefinition> {
self.tools.get(&id)
}
pub fn list_tools(&self) -> Vec<&ToolDefinition> {
self.tools.values().collect()
}
pub fn invoke_tool(&self, tool_id: u32, args: &[u8]) -> Result<Vec<u8>> {
if !self.is_connected() {
return Err(DrivenError::Config("Not connected".to_string()));
}
if !self.tools.contains_key(&tool_id) {
return Err(DrivenError::Config(format!("Tool {} not found", tool_id)));
}
let invocation = ToolInvocation::new(tool_id, 0, 0, args.len() as u32);
let mut payload = Vec::with_capacity(ToolInvocation::SIZE + args.len());
payload.extend_from_slice(&invocation.as_bytes());
payload.extend_from_slice(args);
let flags = if self.signer.is_some() {
Flags::SIGNED
} else {
0
};
let message = DcpMessage::new(MessageType::Tool, flags, payload);
let encoded = message.encode();
Ok(encoded)
}
pub fn encode_message(&self, message_type: MessageType, payload: &[u8]) -> Vec<u8> {
let flags = if self.signer.is_some() {
Flags::SIGNED
} else {
0
};
let message = DcpMessage::new(message_type, flags, payload.to_vec());
message.encode()
}
pub fn decode_message(&self, bytes: &[u8]) -> Result<DcpMessage> {
DcpMessage::decode(bytes)
}
pub fn create_tool_message(&self, tool_id: u32, args: &[u8]) -> Result<DcpMessage> {
if !self.tools.contains_key(&tool_id) {
return Err(DrivenError::Config(format!("Tool {} not found", tool_id)));
}
let invocation = ToolInvocation::new(tool_id, 0, 0, args.len() as u32);
let mut payload = Vec::with_capacity(ToolInvocation::SIZE + args.len());
payload.extend_from_slice(&invocation.as_bytes());
payload.extend_from_slice(args);
let flags = if self.signer.is_some() {
Flags::SIGNED
} else {
0
};
Ok(DcpMessage::new(MessageType::Tool, flags, payload))
}
pub fn create_resource_message(&self, resource_uri: &str) -> DcpMessage {
let payload = resource_uri.as_bytes().to_vec();
let flags = if self.signer.is_some() {
Flags::SIGNED
} else {
0
};
DcpMessage::new(MessageType::Resource, flags, payload)
}
pub fn create_prompt_message(&self, prompt: &str) -> DcpMessage {
let payload = prompt.as_bytes().to_vec();
let flags = if self.signer.is_some() {
Flags::SIGNED
} else {
0
};
DcpMessage::new(MessageType::Prompt, flags, payload)
}
pub fn create_response_message(&self, response: &[u8]) -> DcpMessage {
let flags = if self.signer.is_some() {
Flags::SIGNED
} else {
0
};
DcpMessage::new(MessageType::Response, flags, response.to_vec())
}
pub fn create_error_message(&self, error_code: u32, error_msg: &str) -> DcpMessage {
let mut payload = Vec::with_capacity(4 + error_msg.len());
payload.extend_from_slice(&error_code.to_le_bytes());
payload.extend_from_slice(error_msg.as_bytes());
DcpMessage::new(MessageType::Error, 0, payload)
}
pub fn create_stream_builder(&self, chunk_size: usize) -> StreamBuilder {
StreamBuilder::new(chunk_size)
}
pub fn create_stream_assembler(&self) -> StreamAssembler {
StreamAssembler::new()
}
pub fn create_zcti_builder(&self, tool_id: u32) -> Result<ZctiBuilder> {
if !self.tools.contains_key(&tool_id) {
return Err(DrivenError::Config(format!("Tool {} not found", tool_id)));
}
Ok(ZctiBuilder::new(tool_id))
}
pub fn invoke_tool_zcti(&self, builder: ZctiBuilder) -> Result<Vec<u8>> {
if !self.is_connected() {
return Err(DrivenError::Config("Not connected".to_string()));
}
let tool_id = builder.tool_id;
if !self.tools.contains_key(&tool_id) {
return Err(DrivenError::Config(format!("Tool {} not found", tool_id)));
}
let (invocation, args_data) = builder.build_with_data();
let mut payload = Vec::with_capacity(ToolInvocation::SIZE + args_data.len());
payload.extend_from_slice(&invocation.as_bytes());
payload.extend_from_slice(&args_data);
let flags = if self.signer.is_some() {
Flags::SIGNED
} else {
0
};
let message = DcpMessage::new(MessageType::Tool, flags, payload);
Ok(message.encode())
}
pub fn parse_tool_invocation<'a>(&self, message: &'a DcpMessage) -> Result<ZctiReader<'a>> {
if message.message_type != MessageType::Tool {
return Err(DrivenError::InvalidBinary(
"Message is not a tool invocation".to_string(),
));
}
if message.payload.len() < ToolInvocation::SIZE {
return Err(DrivenError::InvalidBinary(
"Payload too small for tool invocation".to_string(),
));
}
let invocation = ToolInvocation::from_bytes(&message.payload)
.map_err(|e| DrivenError::InvalidBinary(format!("Invalid invocation: {:?}", e)))?;
let args_data = &message.payload[ToolInvocation::SIZE..];
Ok(ZctiReader::new(&invocation, args_data))
}
pub fn create_shared_buffer(&self, capacity: usize) -> SharedArgBuffer {
SharedArgBuffer::new(capacity)
}
pub fn mcp_adapter(&self) -> Option<&McpAdapter> {
self.mcp_adapter.as_ref()
}
pub fn mcp_adapter_mut(&mut self) -> Option<&mut McpAdapter> {
self.mcp_adapter.as_mut()
}
pub fn register_tool_mcp(&mut self, name: &str, tool_id: u32) {
if let Some(ref mut adapter) = self.mcp_adapter {
adapter.register_tool(name, tool_id as u16);
}
}
pub fn handle_mcp_request(&self, json: &str) -> Result<String> {
let adapter = self
.mcp_adapter
.as_ref()
.ok_or_else(|| DrivenError::Config("MCP adapter not available".to_string()))?;
let request = adapter
.parse_request(json)
.map_err(|e| DrivenError::Parse(format!("Invalid JSON-RPC request: {:?}", e)))?;
match request.method.as_str() {
"initialize" => adapter
.handle_initialize(&request)
.map_err(|e| DrivenError::Format(format!("Failed to handle initialize: {:?}", e))),
"tools/list" => adapter
.handle_tools_list(&request)
.map_err(|e| DrivenError::Format(format!("Failed to handle tools/list: {:?}", e))),
"tools/call" => {
self.handle_mcp_tools_call(&request)
}
"resources/list" => {
let result = serde_json::json!({ "resources": [] });
adapter
.format_success_response(request.id, result)
.map_err(|e| DrivenError::Format(format!("Failed to format response: {:?}", e)))
}
"prompts/list" => {
let result = serde_json::json!({ "prompts": [] });
adapter
.format_success_response(request.id, result)
.map_err(|e| DrivenError::Format(format!("Failed to format response: {:?}", e)))
}
_ => {
adapter
.format_error_response(request.id, JsonRpcError::method_not_found())
.map_err(|e| DrivenError::Format(format!("Failed to format error: {:?}", e)))
}
}
}
fn handle_mcp_tools_call(&self, request: &JsonRpcRequest) -> Result<String> {
let adapter = self
.mcp_adapter
.as_ref()
.ok_or_else(|| DrivenError::Config("MCP adapter not available".to_string()))?;
let params = request
.params
.as_ref()
.ok_or_else(|| DrivenError::Parse("Missing params".to_string()))?;
let tool_name = params
.get("name")
.and_then(|v| v.as_str())
.ok_or_else(|| DrivenError::Parse("Missing tool name".to_string()))?;
let arguments = params.get("arguments").cloned();
let tool_id = adapter
.resolve_tool_name(tool_name)
.ok_or_else(|| DrivenError::Config(format!("Unknown tool: {}", tool_name)))?;
if !self.tools.contains_key(&(tool_id as u32)) {
return adapter
.format_error_response(
request.id.clone(),
JsonRpcError::new(-32602, format!("Tool not found: {}", tool_name)),
)
.map_err(|e| DrivenError::Format(format!("Failed to format error: {:?}", e)));
}
let args_bytes = adapter.translate_params(&arguments);
let response_result = serde_json::json!({
"content": [{
"type": "text",
"text": format!("Tool {} called with {} bytes of arguments", tool_name, args_bytes.len())
}]
});
adapter
.format_success_response(request.id.clone(), response_result)
.map_err(|e| DrivenError::Format(format!("Failed to format response: {:?}", e)))
}
pub fn has_mcp_fallback(&self) -> bool {
self.mcp_adapter.is_some()
}
pub fn dcp_to_mcp(&self, message: &DcpMessage) -> Result<String> {
match message.message_type {
MessageType::Response => {
let result: serde_json::Value = serde_json::from_slice(&message.payload)
.unwrap_or_else(|_| {
serde_json::Value::String(
String::from_utf8_lossy(&message.payload).to_string(),
)
});
let response = serde_json::json!({
"jsonrpc": "2.0",
"result": result,
"id": null
});
serde_json::to_string(&response)
.map_err(|e| DrivenError::Format(format!("Failed to serialize: {}", e)))
}
MessageType::Error => {
let (code, msg) = if message.payload.len() >= 4 {
let code_bytes: [u8; 4] = message.payload[0..4].try_into().unwrap();
let code = i32::from_le_bytes(code_bytes);
let msg = String::from_utf8_lossy(&message.payload[4..]).to_string();
(code, msg)
} else {
(-32000, "Unknown error".to_string())
};
let response = serde_json::json!({
"jsonrpc": "2.0",
"error": {
"code": code,
"message": msg
},
"id": null
});
serde_json::to_string(&response)
.map_err(|e| DrivenError::Format(format!("Failed to serialize: {}", e)))
}
_ => Err(DrivenError::Format(format!(
"Cannot convert {:?} message to MCP format",
message.message_type
))),
}
}
pub fn mcp_to_dcp(&self, json: &str) -> Result<DcpMessage> {
let adapter = self
.mcp_adapter
.as_ref()
.ok_or_else(|| DrivenError::Config("MCP adapter not available".to_string()))?;
let request = adapter
.parse_request(json)
.map_err(|e| DrivenError::Parse(format!("Invalid JSON-RPC: {:?}", e)))?;
let message_type = match request.method.as_str() {
"tools/call" => MessageType::Tool,
"resources/read" => MessageType::Resource,
"prompts/get" => MessageType::Prompt,
_ => MessageType::Tool, };
let payload = serde_json::to_vec(&request)
.map_err(|e| DrivenError::Format(format!("Failed to serialize: {}", e)))?;
Ok(DcpMessage::new(message_type, 0, payload))
}
pub fn invoke_tool_mcp(&self, request: &str) -> Result<String> {
if !self.is_connected() {
return Err(DrivenError::Config("Not connected".to_string()));
}
let _parsed = JsonRpcParser::parse_request(request)
.map_err(|e| DrivenError::Parse(format!("Invalid JSON-RPC request: {:?}", e)))?;
if let Some(ref adapter) = self.mcp_adapter {
let response = adapter
.format_error_response(RequestId::Null, JsonRpcError::method_not_found())
.map_err(|e| DrivenError::Format(format!("Failed to format response: {:?}", e)))?;
Ok(response)
} else {
Err(DrivenError::Config("MCP adapter not available".to_string()))
}
}
pub fn sign_tool_def(&self, tool_id: u32) -> Result<dcp::binary::SignedToolDef> {
let signer = self
.signer
.as_ref()
.ok_or_else(|| DrivenError::Security("Signing not enabled".to_string()))?;
let tool = self
.tools
.get(&tool_id)
.ok_or_else(|| DrivenError::Config(format!("Tool {} not found", tool_id)))?;
Ok(signer.sign_tool_def(tool_id, tool.schema_hash, tool.capabilities))
}
pub fn verify_tool_def(&self, def: &dcp::binary::SignedToolDef) -> Result<()> {
Verifier::verify_tool_def(def)
.map_err(|e| DrivenError::Security(format!("Signature verification failed: {:?}", e)))
}
pub fn sign_invocation(
&self,
tool_id: u32,
nonce: u64,
timestamp: u64,
args: &[u8],
) -> Result<SignedInvocation> {
let signer = self
.signer
.as_ref()
.ok_or_else(|| DrivenError::Security("Signing not enabled".to_string()))?;
if !self.tools.contains_key(&tool_id) {
return Err(DrivenError::Config(format!("Tool {} not found", tool_id)));
}
Ok(signer.sign_invocation(tool_id, nonce, timestamp, args))
}
pub fn verify_invocation(&self, inv: &SignedInvocation, public_key: &[u8; 32]) -> Result<()> {
Verifier::verify_invocation(inv, public_key)
.map_err(|e| DrivenError::Security(format!("Invocation verification failed: {:?}", e)))
}
pub fn verify_args_hash(&self, inv: &SignedInvocation, args: &[u8]) -> bool {
Verifier::verify_args_hash(inv, args)
}
pub fn generate_signer() -> Signer {
Signer::generate()
}
pub fn parse_capability_manifest(bytes: &[u8]) -> Result<&CapabilityManifest> {
CapabilityManifest::from_bytes(bytes).map_err(|e| {
DrivenError::InvalidBinary(format!("Failed to parse capability manifest: {:?}", e))
})
}
pub fn serialize_capabilities(&self) -> Vec<u8> {
self.capabilities.as_bytes().to_vec()
}
pub fn register_resource(&mut self, resource_id: u16) {
self.capabilities.set_resource(resource_id);
}
pub fn unregister_resource(&mut self, resource_id: u16) {
self.capabilities.clear_resource(resource_id);
}
pub fn has_resource(&self, resource_id: u16) -> bool {
self.capabilities.has_resource(resource_id)
}
pub fn register_prompt(&mut self, prompt_id: u16) {
self.capabilities.set_prompt(prompt_id);
}
pub fn unregister_prompt(&mut self, prompt_id: u16) {
self.capabilities.clear_prompt(prompt_id);
}
pub fn has_prompt(&self, prompt_id: u16) -> bool {
self.capabilities.has_prompt(prompt_id)
}
pub fn set_extension(&mut self, bit: u8) {
self.capabilities.set_extension(bit);
}
pub fn clear_extension(&mut self, bit: u8) {
self.capabilities.clear_extension(bit);
}
pub fn has_extension(&self, bit: u8) -> bool {
self.capabilities.has_extension(bit)
}
pub fn enforce_tool(&self, tool_id: u32) -> Result<()> {
if tool_id >= CapabilityManifest::MAX_TOOLS as u32 {
return Err(DrivenError::Config(format!(
"Tool ID {} exceeds maximum ({})",
tool_id,
CapabilityManifest::MAX_TOOLS
)));
}
if !self.capabilities.has_tool(tool_id as u16) {
return Err(DrivenError::Config(format!(
"Tool {} is not available in capability manifest",
tool_id
)));
}
Ok(())
}
pub fn enforce_resource(&self, resource_id: u16) -> Result<()> {
if !self.capabilities.has_resource(resource_id) {
return Err(DrivenError::Config(format!(
"Resource {} is not available in capability manifest",
resource_id
)));
}
Ok(())
}
pub fn enforce_prompt(&self, prompt_id: u16) -> Result<()> {
if !self.capabilities.has_prompt(prompt_id) {
return Err(DrivenError::Config(format!(
"Prompt {} is not available in capability manifest",
prompt_id
)));
}
Ok(())
}
pub fn negotiate_capabilities(
&mut self,
server_manifest: &CapabilityManifest,
) -> CapabilityManifest {
let negotiated = self.capabilities.intersect(server_manifest);
self.capabilities = negotiated.clone();
negotiated
}
pub fn capability_stats(&self) -> CapabilityStats {
CapabilityStats {
tool_count: self.capabilities.tool_count(),
resource_count: self.capabilities.resource_count(),
prompt_count: self.capabilities.prompt_count(),
extension_count: self.capabilities.extension_count(),
version: self.capabilities.version,
}
}
pub fn tool_ids(&self) -> impl Iterator<Item = u16> + '_ {
self.capabilities.tool_ids()
}
pub fn resource_ids(&self) -> impl Iterator<Item = u16> + '_ {
self.capabilities.resource_ids()
}
pub fn prompt_ids(&self) -> impl Iterator<Item = u16> + '_ {
self.capabilities.prompt_ids()
}
pub fn intersect_capabilities(&self, other: &CapabilityManifest) -> CapabilityManifest {
self.capabilities.intersect(other)
}
pub fn enable_signing(&mut self, seed: [u8; 32]) {
self.signer = Some(Signer::from_seed(&seed));
self.config.signing_enabled = true;
self.config.signing_seed = Some(seed);
}
pub fn disable_signing(&mut self) {
self.signer = None;
self.config.signing_enabled = false;
self.config.signing_seed = None;
}
pub fn public_key(&self) -> Option<[u8; 32]> {
self.signer.as_ref().map(|s| s.public_key_bytes())
}
}
pub fn create_envelope(
message_type: MessageType,
flags: u8,
payload_len: u32,
) -> BinaryMessageEnvelope {
BinaryMessageEnvelope::new(message_type, flags, payload_len)
}
pub fn parse_envelope(bytes: &[u8]) -> Result<&BinaryMessageEnvelope> {
BinaryMessageEnvelope::from_bytes(bytes)
.map_err(|e| DrivenError::InvalidBinary(format!("Failed to parse BME: {:?}", e)))
}
#[cfg(test)]
mod tests {
use super::*;
use proptest::prelude::*;
#[test]
fn test_dcp_client_creation() {
let client = DcpClient::with_defaults();
assert_eq!(client.state(), ConnectionState::Disconnected);
assert!(!client.is_connected());
assert!(client.prefer_dcp());
}
#[test]
fn test_dcp_client_config() {
let config = DcpConfig {
enabled: false,
prefer_dcp: false,
endpoint: Some("tcp://localhost:9000".to_string()),
timeout_ms: 10000,
signing_enabled: true,
signing_seed: Some([42u8; 32]),
};
let client = DcpClient::new(config);
assert!(!client.prefer_dcp());
assert!(client.signer.is_some());
}
#[test]
fn test_tool_registration() {
let mut client = DcpClient::with_defaults();
let tool_id = client.register_tool("test_tool", "A test tool", [0xAB; 32], 0x1234);
assert_eq!(tool_id, 1);
assert!(client.get_tool(tool_id).is_some());
assert!(client.capabilities().has_tool(tool_id as u16));
let tool = client.get_tool(tool_id).unwrap();
assert_eq!(tool.name, "test_tool");
assert_eq!(tool.description, "A test tool");
}
#[test]
fn test_signing() {
let mut client = DcpClient::with_defaults();
client.enable_signing([42u8; 32]);
let tool_id = client.register_tool("signed_tool", "A signed tool", [0xCD; 32], 0x5678);
let signed_def = client.sign_tool_def(tool_id).unwrap();
assert!(client.verify_tool_def(&signed_def).is_ok());
}
#[test]
fn test_sign_invocation() {
let mut client = DcpClient::with_defaults();
client.enable_signing([42u8; 32]);
let tool_id = client.register_tool("test_tool", "Test", [0; 32], 0);
let args = b"test arguments";
let signed_inv = client
.sign_invocation(tool_id, 12345, 1234567890, args)
.unwrap();
let public_key = client.public_key().unwrap();
assert!(client.verify_invocation(&signed_inv, &public_key).is_ok());
assert!(client.verify_args_hash(&signed_inv, args));
assert!(!client.verify_args_hash(&signed_inv, b"wrong args"));
}
#[test]
fn test_sign_invocation_without_signing_enabled() {
let mut client = DcpClient::with_defaults();
let tool_id = client.register_tool("test_tool", "Test", [0; 32], 0);
let result = client.sign_invocation(tool_id, 12345, 1234567890, b"args");
assert!(result.is_err());
}
#[test]
fn test_sign_invocation_unknown_tool() {
let mut client = DcpClient::with_defaults();
client.enable_signing([42u8; 32]);
let result = client.sign_invocation(999, 12345, 1234567890, b"args");
assert!(result.is_err());
}
#[test]
fn test_generate_signer() {
let signer1 = DcpClient::generate_signer();
let signer2 = DcpClient::generate_signer();
assert_ne!(signer1.public_key_bytes(), signer2.public_key_bytes());
}
#[test]
fn test_capability_intersection() {
let mut client = DcpClient::with_defaults();
client.register_tool("tool1", "Tool 1", [0; 32], 0);
client.register_tool("tool2", "Tool 2", [0; 32], 0);
client.register_tool("tool3", "Tool 3", [0; 32], 0);
let mut other = CapabilityManifest::new(1);
other.set_tool(2);
other.set_tool(3);
other.set_tool(4);
let intersection = client.intersect_capabilities(&other);
assert!(!intersection.has_tool(1));
assert!(intersection.has_tool(2));
assert!(intersection.has_tool(3));
assert!(!intersection.has_tool(4));
}
#[test]
fn test_envelope_creation() {
let envelope = create_envelope(MessageType::Tool, Flags::STREAMING, 1024);
assert!(envelope.is_streaming());
assert!(!envelope.is_compressed());
assert!(!envelope.is_signed());
let bytes = envelope.as_bytes();
let parsed = parse_envelope(bytes).unwrap();
assert_eq!(parsed.get_message_type(), Some(MessageType::Tool));
}
#[test]
fn test_mcp_fallback() {
let mut client = DcpClient::with_defaults();
let result = client.connect("tcp://localhost:9000");
assert!(result.is_ok());
assert_eq!(client.state(), ConnectionState::ConnectedMcp);
}
#[test]
fn test_protocol_preference() {
let mut client = DcpClient::with_defaults();
assert!(client.prefer_dcp());
assert!(client.is_dcp_enabled());
client.set_prefer_dcp(false);
assert!(!client.prefer_dcp());
client.set_prefer_dcp(true);
assert!(client.prefer_dcp());
}
#[test]
fn test_enable_disable_dcp() {
let mut client = DcpClient::with_defaults();
assert!(client.is_dcp_enabled());
client.disable_dcp();
assert!(!client.is_dcp_enabled());
assert!(!client.prefer_dcp());
client.enable_dcp();
assert!(client.is_dcp_enabled());
}
#[test]
fn test_current_protocol() {
let mut client = DcpClient::with_defaults();
assert!(client.current_protocol().is_none());
client.connect("tcp://localhost:9000").unwrap();
assert_eq!(client.current_protocol(), Some(Protocol::Mcp));
}
#[test]
fn test_select_protocol() {
let mut client = DcpClient::with_defaults();
assert!(client.select_protocol().is_none());
client.connect("tcp://localhost:9000").unwrap();
assert_eq!(client.select_protocol(), Some(Protocol::Mcp));
}
#[test]
fn test_is_mcp_connected() {
let mut client = DcpClient::with_defaults();
assert!(!client.is_mcp_connected());
client.connect("tcp://localhost:9000").unwrap();
assert!(client.is_mcp_connected());
assert!(!client.is_dcp_connected());
}
#[test]
fn test_invoke_tool_auto() {
let mut client = DcpClient::with_defaults();
client.connect("tcp://localhost:9000").unwrap();
let tool_id = client.register_tool("test_tool", "Test", [0; 32], 0);
client.register_tool_mcp("test_tool", tool_id);
let args = b"{}";
let result = client.invoke_tool_auto(tool_id, args);
assert!(result.is_ok());
let result = result.unwrap();
assert!(result.is_mcp());
}
#[test]
fn test_invocation_result() {
let dcp_result = InvocationResult::Dcp(vec![1, 2, 3]);
assert!(dcp_result.is_dcp());
assert!(!dcp_result.is_mcp());
assert_eq!(dcp_result.as_dcp(), Some(&[1, 2, 3][..]));
assert!(dcp_result.as_mcp().is_none());
let mcp_result = InvocationResult::Mcp("{}".to_string());
assert!(!mcp_result.is_dcp());
assert!(mcp_result.is_mcp());
assert!(mcp_result.as_dcp().is_none());
assert_eq!(mcp_result.as_mcp(), Some("{}"));
}
#[test]
fn test_dcp_message_tool() {
let payload = vec![1, 2, 3, 4, 5];
let message = DcpMessage::tool(payload.clone());
assert_eq!(message.message_type, MessageType::Tool);
assert_eq!(message.flags, 0);
assert_eq!(message.payload, payload);
assert!(!message.is_streaming());
assert!(!message.is_compressed());
assert!(!message.is_signed());
}
#[test]
fn test_dcp_message_resource() {
let payload = b"resource://test".to_vec();
let message = DcpMessage::resource(payload.clone());
assert_eq!(message.message_type, MessageType::Resource);
assert_eq!(message.payload, payload);
}
#[test]
fn test_dcp_message_prompt() {
let payload = b"Hello, AI!".to_vec();
let message = DcpMessage::prompt(payload.clone());
assert_eq!(message.message_type, MessageType::Prompt);
assert_eq!(message.payload, payload);
}
#[test]
fn test_dcp_message_response() {
let payload = b"Response data".to_vec();
let message = DcpMessage::response(payload.clone());
assert_eq!(message.message_type, MessageType::Response);
assert_eq!(message.payload, payload);
}
#[test]
fn test_dcp_message_error() {
let payload = b"Error occurred".to_vec();
let message = DcpMessage::error(payload.clone());
assert_eq!(message.message_type, MessageType::Error);
assert_eq!(message.payload, payload);
}
#[test]
fn test_dcp_message_stream() {
let payload = b"Streaming data".to_vec();
let message = DcpMessage::stream(payload.clone());
assert_eq!(message.message_type, MessageType::Stream);
assert!(message.is_streaming());
assert_eq!(message.payload, payload);
}
#[test]
fn test_dcp_message_flags() {
let message = DcpMessage::tool(vec![])
.with_streaming()
.with_compressed()
.with_signed();
assert!(message.is_streaming());
assert!(message.is_compressed());
assert!(message.is_signed());
}
#[test]
fn test_dcp_message_encode_decode_roundtrip() {
let original = DcpMessage::new(
MessageType::Tool,
Flags::STREAMING | Flags::SIGNED,
vec![0xDE, 0xAD, 0xBE, 0xEF],
);
let encoded = original.encode();
let decoded = DcpMessage::decode(&encoded).unwrap();
assert_eq!(decoded.message_type, original.message_type);
assert_eq!(decoded.flags, original.flags);
assert_eq!(decoded.payload, original.payload);
}
#[test]
fn test_dcp_message_decode_insufficient_data() {
let bytes = vec![0u8; 4]; let result = DcpMessage::decode(&bytes);
assert!(result.is_err());
}
#[test]
fn test_dcp_message_decode_invalid_magic() {
let mut bytes = vec![0u8; 16];
bytes[0] = 0xFF;
bytes[1] = 0xFF;
let result = DcpMessage::decode(&bytes);
assert!(result.is_err());
}
#[test]
fn test_dcp_message_all_types_roundtrip() {
let test_cases = vec![
DcpMessage::tool(vec![1, 2, 3]),
DcpMessage::resource(b"uri://test".to_vec()),
DcpMessage::prompt(b"test prompt".to_vec()),
DcpMessage::response(b"test response".to_vec()),
DcpMessage::error(b"test error".to_vec()),
DcpMessage::stream(b"streaming".to_vec()),
];
for original in test_cases {
let encoded = original.encode();
let decoded = DcpMessage::decode(&encoded).unwrap();
assert_eq!(decoded.message_type, original.message_type);
assert_eq!(decoded.payload, original.payload);
}
}
#[test]
fn test_stream_assembler_single_chunk() {
let mut assembler = StreamAssembler::new();
let chunk = StreamChunk::new(0, ChunkFlags::FIRST | ChunkFlags::LAST, 5);
let payload = vec![1, 2, 3, 4, 5];
assembler.add_chunk(chunk.as_bytes(), &payload).unwrap();
assert!(assembler.is_complete());
assert!(!assembler.has_error());
assert_eq!(assembler.data(), Some(payload.as_slice()));
assert_eq!(assembler.chunk_count(), 1);
}
#[test]
fn test_stream_assembler_multiple_chunks() {
let mut assembler = StreamAssembler::new();
let chunk1 = StreamChunk::first(0, 3);
assembler.add_chunk(chunk1.as_bytes(), &[1, 2, 3]).unwrap();
assert!(!assembler.is_complete());
let chunk2 = StreamChunk::continuation(1, 3);
assembler.add_chunk(chunk2.as_bytes(), &[4, 5, 6]).unwrap();
assert!(!assembler.is_complete());
let chunk3 = StreamChunk::last(2, 2);
assembler.add_chunk(chunk3.as_bytes(), &[7, 8]).unwrap();
assert!(assembler.is_complete());
assert_eq!(assembler.data(), Some(&[1, 2, 3, 4, 5, 6, 7, 8][..]));
assert_eq!(assembler.chunk_count(), 3);
}
#[test]
fn test_stream_assembler_out_of_order() {
let mut assembler = StreamAssembler::new();
let chunk1 = StreamChunk::first(0, 3);
assembler.add_chunk(chunk1.as_bytes(), &[1, 2, 3]).unwrap();
let chunk_wrong = StreamChunk::continuation(5, 3);
let result = assembler.add_chunk(chunk_wrong.as_bytes(), &[4, 5, 6]);
assert!(result.is_err());
}
#[test]
fn test_stream_assembler_error_chunk() {
let mut assembler = StreamAssembler::new();
let chunk = StreamChunk::error(0, 0);
let result = assembler.add_chunk(chunk.as_bytes(), &[]);
assert!(result.is_err());
assert!(assembler.has_error());
assert!(assembler.is_complete());
assert!(assembler.data().is_none());
}
#[test]
fn test_stream_assembler_take_data() {
let mut assembler = StreamAssembler::new();
let chunk = StreamChunk::new(0, ChunkFlags::FIRST | ChunkFlags::LAST, 3);
assembler.add_chunk(chunk.as_bytes(), &[1, 2, 3]).unwrap();
let data = assembler.take_data();
assert_eq!(data, Some(vec![1, 2, 3]));
}
#[test]
fn test_stream_builder_single_chunk() {
let mut builder = StreamBuilder::new(100);
let data = vec![1, 2, 3, 4, 5];
let chunks = builder.build_chunks(&data);
assert_eq!(chunks.len(), 1);
let (chunk, payload) = &chunks[0];
assert!(chunk.is_first());
assert!(chunk.is_last());
assert_eq!(payload, &data);
}
#[test]
fn test_stream_builder_multiple_chunks() {
let mut builder = StreamBuilder::new(3);
let data = vec![1, 2, 3, 4, 5, 6, 7, 8];
let chunks = builder.build_chunks(&data);
assert_eq!(chunks.len(), 3);
assert!(chunks[0].0.is_first());
assert!(!chunks[0].0.is_last());
assert_eq!(chunks[0].1, vec![1, 2, 3]);
assert!(chunks[1].0.is_continuation());
assert!(!chunks[1].0.is_last());
assert_eq!(chunks[1].1, vec![4, 5, 6]);
assert!(!chunks[2].0.is_first());
assert!(chunks[2].0.is_last());
assert_eq!(chunks[2].1, vec![7, 8]);
}
#[test]
fn test_stream_builder_reset() {
let mut builder = StreamBuilder::new(10);
let _ = builder.build_chunks(&[1, 2, 3]);
builder.reset();
let chunks = builder.build_chunks(&[4, 5, 6]);
let sequence = chunks[0].0.sequence;
assert_eq!(sequence, 0);
}
#[test]
fn test_stream_builder_assembler_roundtrip() {
let mut builder = StreamBuilder::new(5);
let original_data = vec![1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12];
let chunks = builder.build_chunks(&original_data);
let mut assembler = StreamAssembler::new();
for (chunk, payload) in chunks {
assembler.add_chunk(chunk.as_bytes(), &payload).unwrap();
}
assert!(assembler.is_complete());
assert_eq!(assembler.take_data(), Some(original_data));
}
#[test]
fn test_client_create_tool_message() {
let mut client = DcpClient::with_defaults();
client.connect("tcp://localhost:9000").unwrap();
let tool_id = client.register_tool("test", "Test tool", [0; 32], 0);
let args = vec![1, 2, 3, 4];
let message = client.create_tool_message(tool_id, &args).unwrap();
assert_eq!(message.message_type, MessageType::Tool);
assert!(message.payload.len() >= ToolInvocation::SIZE + args.len());
}
#[test]
fn test_client_create_resource_message() {
let client = DcpClient::with_defaults();
let message = client.create_resource_message("file:///test.txt");
assert_eq!(message.message_type, MessageType::Resource);
assert_eq!(message.payload, b"file:///test.txt");
}
#[test]
fn test_client_create_prompt_message() {
let client = DcpClient::with_defaults();
let message = client.create_prompt_message("Hello, AI!");
assert_eq!(message.message_type, MessageType::Prompt);
assert_eq!(message.payload, b"Hello, AI!");
}
#[test]
fn test_client_create_error_message() {
let client = DcpClient::with_defaults();
let message = client.create_error_message(404, "Not found");
assert_eq!(message.message_type, MessageType::Error);
assert_eq!(&message.payload[0..4], &404u32.to_le_bytes());
assert_eq!(&message.payload[4..], b"Not found");
}
#[test]
fn test_client_encode_decode_message() {
let client = DcpClient::with_defaults();
let payload = b"test payload".to_vec();
let encoded = client.encode_message(MessageType::Response, &payload);
let decoded = client.decode_message(&encoded).unwrap();
assert_eq!(decoded.message_type, MessageType::Response);
assert_eq!(decoded.payload, payload);
}
#[test]
fn test_client_stream_builder_assembler() {
let client = DcpClient::with_defaults();
let mut builder = client.create_stream_builder(10);
let assembler = client.create_stream_assembler();
let chunks = builder.build_chunks(&[1, 2, 3]);
assert_eq!(chunks.len(), 1);
assert!(!assembler.is_complete());
}
#[test]
fn test_zcti_builder_basic() {
let builder = ZctiBuilder::new(42)
.add_bool(true)
.add_i32(123)
.add_string("hello");
assert_eq!(builder.arg_count(), 3);
let (invocation, data) = builder.build_with_data();
assert_eq!(invocation.tool_id, 42);
assert_eq!(invocation.arg_count(), 3);
assert!(!data.is_empty());
}
#[test]
fn test_zcti_builder_all_types() {
let builder = ZctiBuilder::new(1)
.add_null()
.add_bool(false)
.add_i32(-42)
.add_i64(i64::MAX)
.add_f64(3.14159)
.add_string("test")
.add_bytes(&[0xDE, 0xAD, 0xBE, 0xEF]);
assert_eq!(builder.arg_count(), 7);
}
#[test]
fn test_zcti_reader_basic() {
let builder = ZctiBuilder::new(42)
.add_bool(true)
.add_i32(123)
.add_string("hello");
let (invocation, data) = builder.build_with_data();
let mut reader = ZctiReader::new(&invocation, &data);
assert_eq!(reader.tool_id(), 42);
assert_eq!(reader.arg_count(), 3);
assert!(reader.has_more());
assert_eq!(reader.read_bool().unwrap(), true);
assert_eq!(reader.read_i32().unwrap(), 123);
assert_eq!(reader.read_string().unwrap(), "hello");
assert!(!reader.has_more());
}
#[test]
fn test_zcti_reader_all_types() {
let builder = ZctiBuilder::new(1)
.add_bool(false)
.add_i32(-42)
.add_i64(i64::MAX)
.add_f64(3.14159)
.add_string("test")
.add_bytes(&[0xDE, 0xAD, 0xBE, 0xEF]);
let (invocation, data) = builder.build_with_data();
let mut reader = ZctiReader::new(&invocation, &data);
assert_eq!(reader.read_bool().unwrap(), false);
assert_eq!(reader.read_i32().unwrap(), -42);
assert_eq!(reader.read_i64().unwrap(), i64::MAX);
assert!((reader.read_f64().unwrap() - 3.14159).abs() < 0.00001);
assert_eq!(reader.read_string().unwrap(), "test");
assert_eq!(reader.read_bytes().unwrap(), &[0xDE, 0xAD, 0xBE, 0xEF]);
}
#[test]
fn test_zcti_reader_skip() {
let builder = ZctiBuilder::new(1)
.add_bool(true)
.add_i32(123)
.add_string("skip me")
.add_i64(456);
let (invocation, data) = builder.build_with_data();
let mut reader = ZctiReader::new(&invocation, &data);
reader.skip().unwrap(); reader.skip().unwrap(); reader.skip().unwrap(); assert_eq!(reader.read_i64().unwrap(), 456);
}
#[test]
fn test_zcti_type_mismatch() {
let builder = ZctiBuilder::new(1).add_bool(true);
let (invocation, data) = builder.build_with_data();
let mut reader = ZctiReader::new(&invocation, &data);
let result = reader.read_i32();
assert!(result.is_err());
}
#[test]
fn test_shared_arg_buffer() {
let mut buffer = SharedArgBuffer::new(1024);
let data1 = b"hello";
let offset1 = buffer.write(data1).unwrap();
assert_eq!(offset1, 0);
let data2 = b"world";
let offset2 = buffer.write(data2).unwrap();
assert_eq!(offset2, 5);
assert_eq!(buffer.read(offset1, 5).unwrap(), b"hello");
assert_eq!(buffer.read(offset2, 5).unwrap(), b"world");
}
#[test]
fn test_shared_arg_buffer_overflow() {
let mut buffer = SharedArgBuffer::new(10);
let result = buffer.write(&[0u8; 20]);
assert!(result.is_err());
}
#[test]
fn test_shared_arg_buffer_reset() {
let mut buffer = SharedArgBuffer::new(100);
buffer.write(b"test").unwrap();
assert_eq!(buffer.offset(), 4);
buffer.reset();
assert_eq!(buffer.offset(), 0);
}
#[test]
fn test_client_zcti_integration() {
let mut client = DcpClient::with_defaults();
client.connect("tcp://localhost:9000").unwrap();
let tool_id = client.register_tool("test_tool", "Test", [0; 32], 0);
let builder = client
.create_zcti_builder(tool_id)
.unwrap()
.add_string("arg1")
.add_i32(42);
let encoded = client.invoke_tool_zcti(builder).unwrap();
let message = client.decode_message(&encoded).unwrap();
assert_eq!(message.message_type, MessageType::Tool);
let mut reader = client.parse_tool_invocation(&message).unwrap();
assert_eq!(reader.tool_id(), tool_id);
assert_eq!(reader.read_string().unwrap(), "arg1");
assert_eq!(reader.read_i32().unwrap(), 42);
}
#[test]
fn test_mcp_adapter_available() {
let mut client = DcpClient::with_defaults();
client.connect("tcp://localhost:9000").unwrap();
assert!(client.has_mcp_fallback());
assert!(client.mcp_adapter().is_some());
}
#[test]
fn test_mcp_register_tool() {
let mut client = DcpClient::with_defaults();
client.connect("tcp://localhost:9000").unwrap();
let tool_id = client.register_tool("read_file", "Read a file", [0; 32], 0);
client.register_tool_mcp("read_file", tool_id);
let adapter = client.mcp_adapter().unwrap();
assert_eq!(adapter.resolve_tool_name("read_file"), Some(tool_id as u16));
}
#[test]
fn test_mcp_handle_initialize() {
let mut client = DcpClient::with_defaults();
client.connect("tcp://localhost:9000").unwrap();
let request = r#"{"jsonrpc":"2.0","method":"initialize","id":1}"#;
let response = client.handle_mcp_request(request).unwrap();
assert!(response.contains("protocolVersion"));
assert!(response.contains("capabilities"));
}
#[test]
fn test_mcp_handle_tools_list() {
let mut client = DcpClient::with_defaults();
client.connect("tcp://localhost:9000").unwrap();
let tool_id = client.register_tool("test_tool", "Test tool", [0; 32], 0);
client.register_tool_mcp("test_tool", tool_id);
let request = r#"{"jsonrpc":"2.0","method":"tools/list","id":2}"#;
let response = client.handle_mcp_request(request).unwrap();
assert!(response.contains("tools"));
assert!(response.contains("test_tool"));
}
#[test]
fn test_mcp_handle_tools_call() {
let mut client = DcpClient::with_defaults();
client.connect("tcp://localhost:9000").unwrap();
let tool_id = client.register_tool("echo", "Echo tool", [0; 32], 0);
client.register_tool_mcp("echo", tool_id);
let request = r#"{"jsonrpc":"2.0","method":"tools/call","params":{"name":"echo","arguments":{"text":"hello"}},"id":3}"#;
let response = client.handle_mcp_request(request).unwrap();
assert!(response.contains("content"));
assert!(response.contains("echo"));
}
#[test]
fn test_mcp_handle_unknown_method() {
let mut client = DcpClient::with_defaults();
client.connect("tcp://localhost:9000").unwrap();
let request = r#"{"jsonrpc":"2.0","method":"unknown/method","id":4}"#;
let response = client.handle_mcp_request(request).unwrap();
assert!(response.contains("error"));
assert!(response.contains("-32601")); }
#[test]
fn test_mcp_handle_resources_list() {
let mut client = DcpClient::with_defaults();
client.connect("tcp://localhost:9000").unwrap();
let request = r#"{"jsonrpc":"2.0","method":"resources/list","id":5}"#;
let response = client.handle_mcp_request(request).unwrap();
assert!(response.contains("resources"));
assert!(response.contains("[]")); }
#[test]
fn test_mcp_handle_prompts_list() {
let mut client = DcpClient::with_defaults();
client.connect("tcp://localhost:9000").unwrap();
let request = r#"{"jsonrpc":"2.0","method":"prompts/list","id":6}"#;
let response = client.handle_mcp_request(request).unwrap();
assert!(response.contains("prompts"));
assert!(response.contains("[]")); }
#[test]
fn test_dcp_to_mcp_response() {
let client = DcpClient::with_defaults();
let payload = serde_json::to_vec(&serde_json::json!({"result": "success"})).unwrap();
let message = DcpMessage::response(payload);
let mcp_json = client.dcp_to_mcp(&message).unwrap();
assert!(mcp_json.contains("jsonrpc"));
assert!(mcp_json.contains("result"));
}
#[test]
fn test_dcp_to_mcp_error() {
let client = DcpClient::with_defaults();
let message = client.create_error_message(404, "Not found");
let mcp_json = client.dcp_to_mcp(&message).unwrap();
assert!(mcp_json.contains("error"));
assert!(mcp_json.contains("404"));
assert!(mcp_json.contains("Not found"));
}
#[test]
fn test_mcp_to_dcp() {
let mut client = DcpClient::with_defaults();
client.connect("tcp://localhost:9000").unwrap();
let mcp_request =
r#"{"jsonrpc":"2.0","method":"tools/call","params":{"name":"test"},"id":1}"#;
let dcp_message = client.mcp_to_dcp(mcp_request).unwrap();
assert_eq!(dcp_message.message_type, MessageType::Tool);
assert!(!dcp_message.payload.is_empty());
}
fn arb_message_type() -> impl Strategy<Value = MessageType> {
prop_oneof![
Just(MessageType::Tool),
Just(MessageType::Resource),
Just(MessageType::Prompt),
Just(MessageType::Response),
Just(MessageType::Error),
Just(MessageType::Stream),
]
}
fn arb_flags() -> impl Strategy<Value = u8> {
prop_oneof![
Just(0u8),
Just(Flags::STREAMING),
Just(Flags::COMPRESSED),
Just(Flags::SIGNED),
Just(Flags::STREAMING | Flags::COMPRESSED),
Just(Flags::STREAMING | Flags::SIGNED),
Just(Flags::COMPRESSED | Flags::SIGNED),
Just(Flags::STREAMING | Flags::COMPRESSED | Flags::SIGNED),
]
}
fn arb_payload() -> impl Strategy<Value = Vec<u8>> {
prop::collection::vec(any::<u8>(), 0..1024)
}
fn arb_dcp_message() -> impl Strategy<Value = DcpMessage> {
(arb_message_type(), arb_flags(), arb_payload()).prop_map(
|(message_type, flags, payload)| DcpMessage::new(message_type, flags, payload),
)
}
fn arb_chunk_size() -> impl Strategy<Value = usize> {
1usize..256
}
fn arb_stream_data() -> impl Strategy<Value = Vec<u8>> {
prop::collection::vec(any::<u8>(), 1..2048)
}
proptest! {
#![proptest_config(ProptestConfig::with_cases(100))]
#[test]
fn prop_bme_roundtrip(message in arb_dcp_message()) {
let encoded = message.encode();
let decoded = DcpMessage::decode(&encoded).expect("Decoding should succeed");
prop_assert_eq!(decoded.message_type, message.message_type);
prop_assert_eq!(decoded.flags, message.flags);
prop_assert_eq!(decoded.payload, message.payload);
}
#[test]
fn prop_bme_envelope_header(
message_type in arb_message_type(),
flags in arb_flags(),
payload_len in 0u32..65536
) {
let envelope = BinaryMessageEnvelope::new(message_type, flags, payload_len);
let bytes = envelope.as_bytes();
let parsed = BinaryMessageEnvelope::from_bytes(bytes).expect("Parsing should succeed");
let parsed_type = parsed.message_type;
let parsed_flags = parsed.flags;
let parsed_len = parsed.payload_len;
prop_assert_eq!(parsed_type, message_type as u8);
prop_assert_eq!(parsed_flags, flags);
prop_assert_eq!(parsed_len, payload_len);
}
#[test]
fn prop_stream_roundtrip(
data in arb_stream_data(),
chunk_size in arb_chunk_size()
) {
let mut builder = StreamBuilder::new(chunk_size);
let chunks = builder.build_chunks(&data);
let mut assembler = StreamAssembler::new();
for (chunk, payload) in chunks {
assembler.add_chunk(chunk.as_bytes(), &payload).expect("Adding chunk should succeed");
}
prop_assert!(assembler.is_complete());
prop_assert!(!assembler.has_error());
prop_assert_eq!(assembler.take_data(), Some(data));
}
#[test]
fn prop_all_message_types_roundtrip(
message_type in arb_message_type(),
payload in arb_payload()
) {
let message = DcpMessage::new(message_type, 0, payload.clone());
let encoded = message.encode();
let decoded = DcpMessage::decode(&encoded).expect("Decoding should succeed");
prop_assert_eq!(decoded.message_type, message_type);
prop_assert_eq!(decoded.payload, payload);
}
#[test]
fn prop_flags_preserved(flags in arb_flags()) {
let message = DcpMessage::new(MessageType::Tool, flags, vec![1, 2, 3]);
let encoded = message.encode();
let decoded = DcpMessage::decode(&encoded).expect("Decoding should succeed");
prop_assert_eq!(decoded.flags, flags);
prop_assert_eq!(decoded.is_streaming(), flags & Flags::STREAMING != 0);
prop_assert_eq!(decoded.is_compressed(), flags & Flags::COMPRESSED != 0);
prop_assert_eq!(decoded.is_signed(), flags & Flags::SIGNED != 0);
}
#[test]
fn prop_mcp_initialize_response_valid(id in 1i64..1000) {
let request = format!(r#"{{"jsonrpc":"2.0","method":"initialize","id":{}}}"#, id);
let mut client = DcpClient::with_defaults();
client.connect("tcp://localhost:9000").expect("Connection should succeed");
let response = client.handle_mcp_request(&request).expect("Should handle initialize");
let parsed: serde_json::Value = serde_json::from_str(&response).expect("Should be valid JSON");
prop_assert_eq!(parsed["jsonrpc"].as_str(), Some("2.0"));
prop_assert!(parsed.get("result").is_some() || parsed.get("error").is_some());
if let Some(result) = parsed.get("result") {
prop_assert!(result.get("capabilities").is_some());
prop_assert!(result.get("protocolVersion").is_some());
}
}
#[test]
fn prop_mcp_tools_list_contains_registered(
tool_names in prop::collection::vec("[a-z_]+", 1..5)
) {
let mut client = DcpClient::with_defaults();
client.connect("tcp://localhost:9000").expect("Connection should succeed");
for name in &tool_names {
let tool_id = client.register_tool(name, "Test tool", [0; 32], 0);
client.register_tool_mcp(name, tool_id);
}
let request = r#"{"jsonrpc":"2.0","method":"tools/list","id":1}"#;
let response = client.handle_mcp_request(request).expect("Should handle tools/list");
let parsed: serde_json::Value = serde_json::from_str(&response).expect("Should be valid JSON");
let tools = parsed["result"]["tools"].as_array().expect("Should have tools array");
for name in &tool_names {
let found = tools.iter().any(|t| t["name"].as_str() == Some(name.as_str()));
prop_assert!(found, "Tool {} should be in response", name);
}
}
#[test]
fn prop_capability_intersection(
tools1 in prop::collection::vec(0u16..1000, 0..20),
tools2 in prop::collection::vec(0u16..1000, 0..20),
resources1 in prop::collection::vec(0u16..500, 0..10),
resources2 in prop::collection::vec(0u16..500, 0..10),
prompts1 in prop::collection::vec(0u16..200, 0..5),
prompts2 in prop::collection::vec(0u16..200, 0..5),
) {
let mut m1 = CapabilityManifest::new(1);
let mut m2 = CapabilityManifest::new(1);
for &t in &tools1 { m1.set_tool(t); }
for &r in &resources1 { m1.set_resource(r); }
for &p in &prompts1 { m1.set_prompt(p); }
for &t in &tools2 { m2.set_tool(t); }
for &r in &resources2 { m2.set_resource(r); }
for &p in &prompts2 { m2.set_prompt(p); }
let intersection = m1.intersect(&m2);
for t in 0u16..1000 {
let in_both = m1.has_tool(t) && m2.has_tool(t);
prop_assert_eq!(
intersection.has_tool(t), in_both,
"Tool {} should be in intersection iff in both manifests", t
);
}
for r in 0u16..500 {
let in_both = m1.has_resource(r) && m2.has_resource(r);
prop_assert_eq!(
intersection.has_resource(r), in_both,
"Resource {} should be in intersection iff in both manifests", r
);
}
for p in 0u16..200 {
let in_both = m1.has_prompt(p) && m2.has_prompt(p);
prop_assert_eq!(
intersection.has_prompt(p), in_both,
"Prompt {} should be in intersection iff in both manifests", p
);
}
}
#[test]
fn prop_capability_manifest_roundtrip(
tools in prop::collection::vec(0u16..8000, 0..50),
resources in prop::collection::vec(0u16..1000, 0..20),
prompts in prop::collection::vec(0u16..500, 0..10),
extensions in 0u64..u64::MAX,
) {
let mut manifest = CapabilityManifest::new(1);
for &t in &tools { manifest.set_tool(t); }
for &r in &resources { manifest.set_resource(r); }
for &p in &prompts { manifest.set_prompt(p); }
manifest.extensions = extensions;
let bytes = manifest.as_bytes();
let parsed = CapabilityManifest::from_bytes(bytes).expect("Should parse");
for &t in &tools {
prop_assert!(parsed.has_tool(t), "Tool {} should be preserved", t);
}
for &r in &resources {
prop_assert!(parsed.has_resource(r), "Resource {} should be preserved", r);
}
for &p in &prompts {
prop_assert!(parsed.has_prompt(p), "Prompt {} should be preserved", p);
}
prop_assert_eq!(parsed.extensions, extensions);
}
#[test]
fn prop_signature_verification_tool_def(
seed in prop::array::uniform32(any::<u8>()),
tool_id in 1u32..1000,
schema_hash in prop::array::uniform32(any::<u8>()),
capabilities in any::<u64>(),
) {
let signer = Signer::from_seed(&seed);
let signed_def = signer.sign_tool_def(tool_id, schema_hash, capabilities);
prop_assert!(Verifier::verify_tool_def(&signed_def).is_ok());
let mut tampered = signed_def.clone();
tampered.tool_id = tool_id.wrapping_add(1);
prop_assert!(Verifier::verify_tool_def(&tampered).is_err());
let mut tampered = signed_def.clone();
tampered.schema_hash[0] ^= 0xFF;
prop_assert!(Verifier::verify_tool_def(&tampered).is_err());
let mut tampered = signed_def.clone();
tampered.capabilities ^= 0xFFFFFFFF;
prop_assert!(Verifier::verify_tool_def(&tampered).is_err());
}
#[test]
fn prop_signature_verification_invocation(
seed in prop::array::uniform32(any::<u8>()),
wrong_seed in prop::array::uniform32(any::<u8>()),
tool_id in 1u32..1000,
nonce in any::<u64>(),
timestamp in any::<u64>(),
args in prop::collection::vec(any::<u8>(), 0..256),
) {
let signer = Signer::from_seed(&seed);
let public_key = signer.public_key_bytes();
let signed_inv = signer.sign_invocation(tool_id, nonce, timestamp, &args);
prop_assert!(Verifier::verify_invocation(&signed_inv, &public_key).is_ok());
prop_assert!(Verifier::verify_args_hash(&signed_inv, &args));
if seed != wrong_seed {
let wrong_signer = Signer::from_seed(&wrong_seed);
let wrong_key = wrong_signer.public_key_bytes();
prop_assert!(Verifier::verify_invocation(&signed_inv, &wrong_key).is_err());
}
let mut tampered = signed_inv.clone();
tampered.tool_id = tool_id.wrapping_add(1);
prop_assert!(Verifier::verify_invocation(&tampered, &public_key).is_err());
let mut tampered = signed_inv.clone();
tampered.nonce ^= 0xFFFFFFFF;
prop_assert!(Verifier::verify_invocation(&tampered, &public_key).is_err());
}
}
#[test]
fn test_capability_manifest_parsing() {
let mut manifest = CapabilityManifest::new(1);
manifest.set_tool(42);
manifest.set_resource(10);
manifest.set_prompt(5);
let bytes = manifest.as_bytes();
let parsed = DcpClient::parse_capability_manifest(bytes).unwrap();
assert!(parsed.has_tool(42));
assert!(parsed.has_resource(10));
assert!(parsed.has_prompt(5));
}
#[test]
fn test_capability_manifest_serialization() {
let mut client = DcpClient::with_defaults();
client.register_tool("tool1", "Tool 1", [0; 32], 0);
client.register_resource(10);
client.register_prompt(5);
let bytes = client.serialize_capabilities();
let parsed = CapabilityManifest::from_bytes(&bytes).unwrap();
assert!(parsed.has_tool(1)); assert!(parsed.has_resource(10));
assert!(parsed.has_prompt(5));
}
#[test]
fn test_resource_registration() {
let mut client = DcpClient::with_defaults();
assert!(!client.has_resource(42));
client.register_resource(42);
assert!(client.has_resource(42));
client.unregister_resource(42);
assert!(!client.has_resource(42));
}
#[test]
fn test_prompt_registration() {
let mut client = DcpClient::with_defaults();
assert!(!client.has_prompt(10));
client.register_prompt(10);
assert!(client.has_prompt(10));
client.unregister_prompt(10);
assert!(!client.has_prompt(10));
}
#[test]
fn test_extension_flags() {
let mut client = DcpClient::with_defaults();
assert!(!client.has_extension(5));
client.set_extension(5);
assert!(client.has_extension(5));
client.clear_extension(5);
assert!(!client.has_extension(5));
}
#[test]
fn test_enforce_tool_success() {
let mut client = DcpClient::with_defaults();
let tool_id = client.register_tool("test", "Test", [0; 32], 0);
assert!(client.enforce_tool(tool_id).is_ok());
}
#[test]
fn test_enforce_tool_failure() {
let client = DcpClient::with_defaults();
assert!(client.enforce_tool(999).is_err());
}
#[test]
fn test_enforce_resource_success() {
let mut client = DcpClient::with_defaults();
client.register_resource(42);
assert!(client.enforce_resource(42).is_ok());
}
#[test]
fn test_enforce_resource_failure() {
let client = DcpClient::with_defaults();
assert!(client.enforce_resource(999).is_err());
}
#[test]
fn test_enforce_prompt_success() {
let mut client = DcpClient::with_defaults();
client.register_prompt(10);
assert!(client.enforce_prompt(10).is_ok());
}
#[test]
fn test_enforce_prompt_failure() {
let client = DcpClient::with_defaults();
assert!(client.enforce_prompt(999).is_err());
}
#[test]
fn test_negotiate_capabilities() {
let mut client = DcpClient::with_defaults();
client.register_tool("tool1", "Tool 1", [0; 32], 0);
client.register_tool("tool2", "Tool 2", [0; 32], 0);
client.register_tool("tool3", "Tool 3", [0; 32], 0);
client.register_resource(1);
client.register_resource(2);
let mut server_manifest = CapabilityManifest::new(1);
server_manifest.set_tool(2);
server_manifest.set_tool(3);
server_manifest.set_tool(4);
server_manifest.set_resource(2);
server_manifest.set_resource(3);
let negotiated = client.negotiate_capabilities(&server_manifest);
assert!(!negotiated.has_tool(1)); assert!(negotiated.has_tool(2)); assert!(negotiated.has_tool(3)); assert!(!negotiated.has_tool(4));
assert!(!negotiated.has_resource(1)); assert!(negotiated.has_resource(2)); assert!(!negotiated.has_resource(3)); }
#[test]
fn test_capability_stats() {
let mut client = DcpClient::with_defaults();
client.register_tool("tool1", "Tool 1", [0; 32], 0);
client.register_tool("tool2", "Tool 2", [0; 32], 0);
client.register_resource(1);
client.register_prompt(1);
client.register_prompt(2);
client.set_extension(0);
let stats = client.capability_stats();
assert_eq!(stats.tool_count, 2);
assert_eq!(stats.resource_count, 1);
assert_eq!(stats.prompt_count, 2);
assert_eq!(stats.extension_count, 1);
assert_eq!(stats.version, 1);
}
#[test]
fn test_capability_iterators() {
let mut client = DcpClient::with_defaults();
client.register_tool("tool1", "Tool 1", [0; 32], 0);
client.register_tool("tool2", "Tool 2", [0; 32], 0);
client.register_resource(10);
client.register_resource(20);
client.register_prompt(5);
let tool_ids: Vec<_> = client.tool_ids().collect();
assert_eq!(tool_ids, vec![1, 2]);
let resource_ids: Vec<_> = client.resource_ids().collect();
assert_eq!(resource_ids, vec![10, 20]);
let prompt_ids: Vec<_> = client.prompt_ids().collect();
assert_eq!(prompt_ids, vec![5]);
}
}