use rskit_ai::ToolUseBlock;
use rskit_errors::{AppError, AppResult, ErrorCode};
use serde_json::{Map, Value};
use crate::{StreamChunk, StreamToolCall};
pub fn merge_tool_delta(calls: &mut Vec<StreamToolCall>, delta: StreamToolCall) {
if calls.len() <= delta.index {
calls.resize_with(delta.index + 1, StreamToolCall::default);
}
let current = &mut calls[delta.index];
current.index = delta.index;
if !delta.id.is_empty() {
current.id = delta.id;
}
if !delta.name.is_empty() {
current.name = delta.name;
}
current.input_delta.push_str(&delta.input_delta);
}
pub fn parse_input_json(input: &str) -> AppResult<Map<String, Value>> {
if input.trim().is_empty() {
return Ok(Map::new());
}
let value: Value = serde_json::from_str(input).map_err(|error| {
AppError::new(
ErrorCode::InvalidFormat,
format!("failed to parse tool input JSON: {error}"),
)
})?;
value_to_input_map(value)
}
pub fn value_to_input_map(value: Value) -> AppResult<Map<String, Value>> {
match value {
Value::Object(map) => Ok(map),
Value::Null => Ok(Map::new()),
other => Err(AppError::new(
ErrorCode::InvalidFormat,
format!("expected tool input object, got {other}"),
)),
}
}
pub fn accumulate_tool_uses(
chunks: impl IntoIterator<Item = StreamChunk>,
) -> AppResult<Vec<ToolUseBlock>> {
let mut calls = Vec::new();
for chunk in chunks {
for delta in chunk.tool_calls {
merge_tool_delta(&mut calls, delta);
}
}
calls
.into_iter()
.enumerate()
.filter(|(_, call)| !call.name.is_empty() || !call.id.is_empty())
.map(|(index, call)| {
Ok(ToolUseBlock {
id: if call.id.is_empty() {
format!("tool_call_{index}")
} else {
call.id
},
name: call.name,
input: parse_input_json(&call.input_delta)?,
})
})
.collect()
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn parse_input_json_accepts_empty_null_and_objects() {
assert!(parse_input_json("").unwrap().is_empty());
assert!(parse_input_json("null").unwrap().is_empty());
assert_eq!(parse_input_json(r#"{"a":1}"#).unwrap()["a"], 1);
}
#[test]
fn parse_input_json_rejects_invalid_or_non_object_values() {
assert_eq!(
parse_input_json("{").unwrap_err().code(),
ErrorCode::InvalidFormat
);
assert_eq!(
value_to_input_map(Value::String("bad".into()))
.unwrap_err()
.code(),
ErrorCode::InvalidFormat
);
}
#[test]
fn accumulate_tool_uses_merges_fragments_and_defaults_missing_ids() {
let calls = accumulate_tool_uses([
StreamChunk {
tool_calls: vec![StreamToolCall {
index: 1,
name: "lookup".into(),
input_delta: r#"{"q":"#.into(),
..Default::default()
}],
..Default::default()
},
StreamChunk {
tool_calls: vec![StreamToolCall {
index: 1,
input_delta: r#""rust"}"#.into(),
..Default::default()
}],
..Default::default()
},
])
.unwrap();
assert_eq!(calls[0].id, "tool_call_1");
assert_eq!(calls[0].name, "lookup");
assert_eq!(calls[0].input["q"], "rust");
}
}