rig_core/providers/internal/
wire_ids.rs1use std::collections::{BTreeMap, HashMap, HashSet};
19
20use crate::message::{AssistantContent, CallId, LocalCallId, Message, UserContent};
21
22#[derive(Debug)]
25pub struct WireIds {
26 ids: BTreeMap<(usize, usize), String>,
27}
28
29#[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 pub fn new(history: &[Message]) -> Self {
42 Self::with_reserved(history, std::iter::empty())
43 }
44
45 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 pub fn get(&self, message: usize, content: usize) -> Option<&str> {
107 self.ids.get(&(message, content)).map(String::as_str)
108 }
109
110 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;