rskit-llm-common 0.2.0-alpha.2

Shared LLM provider parsing and error utilities for rskit
Documentation
use rskit_ai::ToolUseBlock;
use rskit_errors::{AppError, AppResult, ErrorCode};
use serde_json::{Map, Value};

use crate::{StreamChunk, StreamToolCall};

/// Merge a streamed tool-call delta into the accumulated call list.
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);
}

/// Parse a provider tool-input JSON fragment into a JSON object map.
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)
}

/// Convert a JSON value into a tool-input object map.
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}"),
        )),
    }
}

/// Reconstruct tool-use blocks from a sequence of streamed chunks.
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");
    }
}