File size: 3,387 Bytes
e5034c3
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
use std::collections::{BTreeMap, HashSet};
use std::fmt::Display;

use serde::Serialize;
use serde_json::Value;

use crate::ToolDefinition;

pub struct ToolUsagePrompt<'a> {
    tools: &'a Vec<ToolDefinition>,
}

impl<'a> From<&'a Vec<ToolDefinition>> for ToolUsagePrompt<'a> {
    fn from(value: &'a Vec<ToolDefinition>) -> Self {
        Self { tools: value }
    }
}

impl Display for ToolUsagePrompt<'_> {
    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
        for tool in self.tools.iter() {
            let schema_value = tool.input_schema.as_value();

            // Extract required fields
            let required = schema_value
                .as_object()
                .and_then(|obj| obj.get("required"))
                .and_then(|req| req.as_array())
                .map(|arr| {
                    arr.iter()
                        .filter_map(|v| v.as_str().map(String::from))
                        .collect::<HashSet<_>>()
                })
                .unwrap_or_default();

            // Extract properties
            let parameters = schema_value
                .as_object()
                .and_then(|obj| obj.get("properties"))
                .and_then(|props| props.as_object())
                .map(|props| {
                    props
                        .iter()
                        .map(|(name, prop)| {
                            let description = prop
                                .as_object()
                                .and_then(|p| p.get("description"))
                                .and_then(|d| d.as_str())
                                .unwrap_or("")
                                .to_string();

                            let type_of = prop.as_object().and_then(|p| p.get("type")).cloned();

                            let parameter = Parameter {
                                description,
                                type_of,
                                is_required: required.contains(name),
                            };

                            (name.clone(), parameter)
                        })
                        .collect::<BTreeMap<_, _>>()
                })
                .unwrap_or_default();

            let schema = Schema {
                name: tool.name.to_string(),
                arguments: parameters,
                description: tool.description.clone(),
            };

            writeln!(f, "<tool>{schema}</tool>")?;
        }

        Ok(())
    }
}

#[derive(Serialize)]
struct Schema {
    name: String,
    description: String,
    arguments: BTreeMap<String, Parameter>,
}

#[derive(Serialize)]
struct Parameter {
    description: String,
    #[serde(rename = "type")]
    #[serde(skip_serializing_if = "Option::is_none")]
    type_of: Option<Value>,
    is_required: bool,
}

impl Display for Schema {
    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
        write!(f, "{}", serde_json::to_string(self).unwrap())
    }
}

#[cfg(test)]
mod tests {

    use insta::assert_snapshot;
    use strum::IntoEnumIterator;

    use super::*;
    use crate::ToolCatalog;

    #[test]
    fn test_tool_usage() {
        let tools = ToolCatalog::iter()
            .map(|v| v.definition())
            .collect::<Vec<_>>();
        let prompt = ToolUsagePrompt::from(&tools);
        assert_snapshot!(prompt);
    }
}