use crate::DCPError;
#[repr(C)]
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct ToolInvocation {
pub tool_id: u32,
pub arg_layout: u64,
pub args_offset: u32,
pub args_len: u32,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
#[repr(u8)]
pub enum ArgType {
Null = 0,
Bool = 1,
I32 = 2,
I64 = 3,
F64 = 4,
String = 5,
Bytes = 6,
Array = 7,
Object = 8,
}
impl ArgType {
pub fn from_u8(value: u8) -> Option<Self> {
match value {
0 => Some(Self::Null),
1 => Some(Self::Bool),
2 => Some(Self::I32),
3 => Some(Self::I64),
4 => Some(Self::F64),
5 => Some(Self::String),
6 => Some(Self::Bytes),
7 => Some(Self::Array),
8 => Some(Self::Object),
_ => None,
}
}
}
impl ToolInvocation {
pub const SIZE: usize = 24;
pub const MAX_ARGS: usize = 16;
pub fn new(tool_id: u32, arg_layout: u64, args_offset: u32, args_len: u32) -> Self {
Self {
tool_id,
arg_layout,
args_offset,
args_len,
}
}
#[inline(always)]
pub fn from_bytes(bytes: impl AsRef<[u8]>) -> Result<Self, DCPError> {
let bytes = bytes.as_ref();
if bytes.len() < Self::SIZE {
return Err(DCPError::InsufficientData);
}
if bytes.len() != Self::SIZE {
return Err(DCPError::ValidationFailed);
}
if bytes[4..8].iter().any(|&byte| byte != 0) {
return Err(DCPError::ValidationFailed);
}
Ok(Self {
tool_id: u32::from_le_bytes(bytes[0..4].try_into().unwrap()),
arg_layout: u64::from_le_bytes(bytes[8..16].try_into().unwrap()),
args_offset: u32::from_le_bytes(bytes[16..20].try_into().unwrap()),
args_len: u32::from_le_bytes(bytes[20..24].try_into().unwrap()),
})
}
#[inline(always)]
pub fn as_bytes(&self) -> [u8; Self::SIZE] {
let mut bytes = [0u8; Self::SIZE];
bytes[0..4].copy_from_slice(&self.tool_id.to_le_bytes());
bytes[8..16].copy_from_slice(&self.arg_layout.to_le_bytes());
bytes[16..20].copy_from_slice(&self.args_offset.to_le_bytes());
bytes[20..24].copy_from_slice(&self.args_len.to_le_bytes());
bytes
}
pub fn get_arg_type(&self, index: usize) -> Option<ArgType> {
if index >= Self::MAX_ARGS {
return None;
}
let shift = index * 4;
let type_bits = ((self.arg_layout >> shift) & 0xF) as u8;
ArgType::from_u8(type_bits)
}
pub fn set_arg_type(&mut self, index: usize, arg_type: ArgType) {
if index >= Self::MAX_ARGS {
return;
}
let shift = index * 4;
self.arg_layout &= !(0xF << shift);
self.arg_layout |= (arg_type as u64) << shift;
}
pub fn arg_count(&self) -> usize {
let mut count = 0;
for i in 0..Self::MAX_ARGS {
if let Some(arg_type) = self.get_arg_type(i) {
if arg_type != ArgType::Null {
count += 1;
} else {
break; }
}
}
count
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_invocation_size() {
assert_eq!(std::mem::size_of::<ToolInvocation>(), ToolInvocation::SIZE);
}
#[test]
fn test_invocation_round_trip() {
let inv = ToolInvocation::new(42, 0x12345678, 100, 200);
let bytes = inv.as_bytes();
let parsed = ToolInvocation::from_bytes(bytes).unwrap();
assert_eq!(parsed.tool_id, 42);
assert_eq!(parsed.arg_layout, 0x12345678);
assert_eq!(parsed.args_offset, 100);
assert_eq!(parsed.args_len, 200);
}
#[test]
fn test_arg_type_encoding() {
let mut inv = ToolInvocation::new(1, 0, 0, 0);
inv.set_arg_type(0, ArgType::String);
inv.set_arg_type(1, ArgType::I32);
inv.set_arg_type(2, ArgType::Bool);
assert_eq!(inv.get_arg_type(0), Some(ArgType::String));
assert_eq!(inv.get_arg_type(1), Some(ArgType::I32));
assert_eq!(inv.get_arg_type(2), Some(ArgType::Bool));
assert_eq!(inv.get_arg_type(3), Some(ArgType::Null));
}
#[test]
fn test_arg_count() {
let mut inv = ToolInvocation::new(1, 0, 0, 0);
assert_eq!(inv.arg_count(), 0);
inv.set_arg_type(0, ArgType::String);
inv.set_arg_type(1, ArgType::I32);
assert_eq!(inv.arg_count(), 2);
}
}