Skip to main content

rig_core/providers/internal/
wire_ids.rs

1//! How a request spells its tool calls' ids on a wire that requires one. A
2//! provider's id is sent as it is; an id rig issued is spelled as a
3//! request-local alias, `tool-<n>` in order of first appearance, that no
4//! provider id of the request uses. The spelling is deterministic, so the
5//! same history always encodes to the same bytes.
6//!
7//! ```
8//! use rig_core::message::{CallId, Message, ToolName};
9//! use rig_core::providers::internal::wire_ids::WireIds;
10//!
11//! let call = CallId::from_wire("");
12//! let history = vec![Message::tool_result(call, ToolName::new("add")?, "5")];
13//! let ids = WireIds::new(&history);
14//! assert_eq!(ids.get(0, 0), Some("tool-0"));
15//! # Ok::<(), rig_core::message::EmptyToolName>(())
16//! ```
17
18use std::collections::{BTreeMap, HashMap, HashSet};
19
20use crate::message::{AssistantContent, CallId, LocalCallId, Message, UserContent};
21
22/// The wire spelling of every tool call and result of one request's
23/// history, by message and content position.
24#[derive(Debug)]
25pub struct WireIds {
26    ids: BTreeMap<(usize, usize), String>,
27}
28
29/// Converted messages carried a different number of tool calls and results
30/// than their source.
31#[derive(Debug, thiserror::Error, PartialEq, Eq)]
32#[error("tool id slot count at message {message}: expected {expected}, got {actual}")]
33pub struct WireSlotCount {
34    message: usize,
35    expected: usize,
36    actual: usize,
37}
38
39impl WireIds {
40    /// The spelling of every call and result in `history`.
41    pub fn new(history: &[Message]) -> Self {
42        Self::with_reserved(history, std::iter::empty())
43    }
44
45    /// [`Self::new`], with provider handles carried outside tool calls (such
46    /// as Anthropic server tools) that no alias may take.
47    pub fn with_reserved(history: &[Message], reserved: impl IntoIterator<Item = String>) -> Self {
48        let occurrences: Vec<((usize, usize), &CallId)> = history
49            .iter()
50            .enumerate()
51            .flat_map(|(message, entry)| {
52                let ids: Vec<(usize, &CallId)> = match entry {
53                    Message::Assistant { content, .. } => content
54                        .iter()
55                        .enumerate()
56                        .filter_map(|(index, part)| match part {
57                            AssistantContent::ToolCall(call) => Some((index, &call.id)),
58                            _ => None,
59                        })
60                        .collect(),
61                    Message::User { content } => content
62                        .iter()
63                        .enumerate()
64                        .filter_map(|(index, part)| match part {
65                            UserContent::ToolResult(result) => Some((index, &result.call)),
66                            _ => None,
67                        })
68                        .collect(),
69                    Message::System { .. } => Vec::new(),
70                };
71                ids.into_iter()
72                    .map(move |(content, id)| ((message, content), id))
73            })
74            .collect();
75        let mut used: HashSet<String> = occurrences
76            .iter()
77            .filter_map(|(_, id)| id.provider().map(|provider| provider.call_id.clone()))
78            .chain(reserved)
79            .collect();
80        let mut aliases: HashMap<&LocalCallId, String> = HashMap::new();
81        let mut next = 0usize;
82        let mut ids = BTreeMap::new();
83        for (position, id) in occurrences {
84            let spelled = match id {
85                CallId::Provider(provider) => provider.call_id.clone(),
86                CallId::Local(local) => aliases
87                    .entry(local)
88                    .or_insert_with(|| {
89                        loop {
90                            let candidate = format!("tool-{next}");
91                            next += 1;
92                            if used.insert(candidate.clone()) {
93                                break candidate;
94                            }
95                        }
96                    })
97                    .clone(),
98            };
99            ids.insert(position, spelled);
100        }
101        Self { ids }
102    }
103
104    /// The spelling for the call or result at this position of the history,
105    /// `None` for other content.
106    pub fn get(&self, message: usize, content: usize) -> Option<&str> {
107        self.ids.get(&(message, content)).map(String::as_str)
108    }
109
110    /// Assign one converted message's id slots, in the source message's
111    /// tool-content order. Conversion may split or drop text and reasoning
112    /// but keeps every call and result, so the counts must agree.
113    pub fn apply<'a>(
114        &self,
115        message: usize,
116        slots: impl IntoIterator<Item = &'a mut String>,
117    ) -> Result<(), WireSlotCount> {
118        let planned: Vec<_> = self
119            .ids
120            .range((message, 0)..=(message, usize::MAX))
121            .map(|(_, id)| id)
122            .collect();
123        let slots: Vec<_> = slots.into_iter().collect();
124        if planned.len() != slots.len() {
125            return Err(WireSlotCount {
126                message,
127                expected: planned.len(),
128                actual: slots.len(),
129            });
130        }
131        for (slot, id) in slots.into_iter().zip(planned) {
132            slot.clone_from(id);
133        }
134        Ok(())
135    }
136}
137
138#[cfg(test)]
139#[allow(clippy::expect_used)]
140pub(crate) mod tests;