Skip to main content

gate4agent_node_protocol/
correlation.rs

1//! Opaque correlation ids and tool-class labels the node mints for a
2//! session's provider events.
3//!
4//! The ids are part of the wire contract: `NodeRequest::ResolveInteraction`
5//! answers an `AgentStreamChunkKindV1::InteractionPrompt::correlation_id`, and
6//! a client that derives its own telemetry from `NodeEvent::Control` must mint
7//! the very same id to line the two up. They live here, next to the request
8//! that consumes them, so the node and every client compute them through one
9//! function.
10//!
11//! Every id is a short prefix plus the first eight bytes of a SHA-256 over the
12//! session identity and the provider's own id, so the provider's raw request or
13//! tool id never reaches the wire.
14
15use gate4agent_types::{AgentInstanceId, ControlEvent, ProviderSource, SessionGeneration};
16use ring::digest::{digest, SHA256};
17use std::fmt::Write as _;
18
19fn hex_suffix_correlation(prefix: &str, material: &[u8]) -> String {
20    let digest = digest(&SHA256, material);
21    let mut correlation = String::with_capacity(prefix.len() + 16);
22    correlation.push_str(prefix);
23    for byte in &digest.as_ref()[..8] {
24        // Writing to a `String` never fails; the `Result` only exists because
25        // `fmt::Write` is shared with fallible sinks.
26        let _ = write!(&mut correlation, "{byte:02x}");
27    }
28    correlation
29}
30
31fn provider_source_material(source: &ProviderSource) -> Vec<u8> {
32    // `ProviderSource` is a plain serde tree of strings and enums, so its JSON
33    // form cannot fail. An empty fallback keeps the id deterministic rather
34    // than panicking in library code.
35    serde_json::to_vec(source).unwrap_or_default()
36}
37
38/// Correlation id of one provider-side subagent.
39pub fn subagent_correlation(
40    instance_id: AgentInstanceId,
41    generation: SessionGeneration,
42    source: &ProviderSource,
43    provider_agent_id: &str,
44) -> String {
45    let mut material = Vec::with_capacity(16 + provider_agent_id.len());
46    material.extend_from_slice(&instance_id.0.to_le_bytes());
47    material.extend_from_slice(&generation.0.to_le_bytes());
48    material.extend_from_slice(&provider_source_material(source));
49    material.extend_from_slice(provider_agent_id.as_bytes());
50    hex_suffix_correlation("sub-", &material)
51}
52
53/// Correlation id of one provider tool call carried by `event`.
54pub fn tool_correlation(
55    event: &ControlEvent,
56    source: &ProviderSource,
57    provider_tool_id: &str,
58) -> String {
59    let mut material = Vec::new();
60    material.extend_from_slice(&event.instance_id.0.to_le_bytes());
61    material.extend_from_slice(&event.generation.0.to_le_bytes());
62    material.extend_from_slice(&provider_source_material(source));
63    material.extend_from_slice(provider_tool_id.as_bytes());
64    hex_suffix_correlation("tool-", &material)
65}
66
67/// Correlation id of one interaction (approval or question) of a session.
68///
69/// Takes the `(instance_id, generation, interaction_id)` triple directly so the
70/// reverse lookup in a request handler can replay the same digest for each
71/// interaction a session's live snapshot still remembers, without a
72/// `ControlEvent` to borrow the identity from.
73pub fn interaction_correlation(
74    instance_id: AgentInstanceId,
75    generation: SessionGeneration,
76    interaction_id: u64,
77) -> String {
78    let mut material = Vec::with_capacity(24);
79    material.extend_from_slice(b"interaction");
80    material.extend_from_slice(&instance_id.0.to_le_bytes());
81    material.extend_from_slice(&generation.0.to_le_bytes());
82    material.extend_from_slice(&interaction_id.to_le_bytes());
83    hex_suffix_correlation("int-", &material)
84}
85
86/// Correlation id of the provider process a session owns.
87pub fn process_correlation(instance_id: AgentInstanceId, generation: SessionGeneration) -> String {
88    let mut material = Vec::with_capacity(24);
89    material.extend_from_slice(b"provider-session");
90    material.extend_from_slice(&instance_id.0.to_le_bytes());
91    material.extend_from_slice(&generation.0.to_le_bytes());
92    hex_suffix_correlation("proc-", &material)
93}
94
95/// Collapses a provider tool name into a small class label (`Read`, `Write`,
96/// `Shell`, ...). Returns `None` for an empty name; sets `truncated` whenever
97/// the label differs from the trimmed input, so a reader can tell a
98/// generalised label from a verbatim one.
99pub fn sanitize_progress_tool_label(value: &str, truncated: &mut bool) -> Option<String> {
100    let trimmed = value.trim();
101    if trimmed.is_empty() {
102        *truncated = true;
103        return None;
104    }
105    let normalized = trimmed.to_ascii_lowercase();
106    let class = if ["read", "view", "open", "get"].iter().any(|term| normalized.contains(term)) {
107        "Read"
108    } else if ["write", "create", "save"].iter().any(|term| normalized.contains(term)) {
109        "Write"
110    } else if ["edit", "patch", "replace"].iter().any(|term| normalized.contains(term)) {
111        "Edit"
112    } else if ["shell", "bash", "powershell", "terminal", "exec", "command"]
113        .iter()
114        .any(|term| normalized.contains(term))
115    {
116        "Shell"
117    } else if ["search", "find", "grep", "query"].iter().any(|term| normalized.contains(term)) {
118        "Search"
119    } else if ["browser", "web", "http", "fetch"].iter().any(|term| normalized.contains(term)) {
120        "Browse"
121    } else if ["git", "commit", "diff"].iter().any(|term| normalized.contains(term)) {
122        "Git"
123    } else if ["ask", "question", "approval", "input"].iter().any(|term| normalized.contains(term)) {
124        "Ask"
125    } else if ["task", "agent", "spawn"].iter().any(|term| normalized.contains(term)) {
126        "Task"
127    } else {
128        "Tool"
129    };
130    if trimmed != class {
131        *truncated = true;
132    }
133    Some(class.to_owned())
134}
135
136/// The class label for a tool or request method, `"Tool"` when the name is empty.
137pub fn tool_class_label(name: &str) -> String {
138    let mut truncated = false;
139    sanitize_progress_tool_label(name, &mut truncated).unwrap_or_else(|| "Tool".to_owned())
140}
141
142/// Cuts `value` to at most `max_bytes` on a character boundary. The returned
143/// flag says whether a cut happened; a caller that cuts is responsible for
144/// reporting it.
145pub fn truncate_text(value: &str, max_bytes: usize) -> (String, bool) {
146    if value.len() <= max_bytes {
147        return (value.to_owned(), false);
148    }
149    let mut end = max_bytes;
150    while end > 0 && !value.is_char_boundary(end) {
151        end -= 1;
152    }
153    (value[..end].to_owned(), true)
154}
155
156#[cfg(test)]
157mod tests {
158    use super::*;
159
160    #[test]
161    fn interaction_correlation_is_stable_and_identity_bound() {
162        let a = interaction_correlation(AgentInstanceId(7), SessionGeneration(3), 9);
163        assert_eq!(a, interaction_correlation(AgentInstanceId(7), SessionGeneration(3), 9));
164        assert!(a.starts_with("int-"));
165        assert_eq!(a.len(), 4 + 16);
166        assert_ne!(a, interaction_correlation(AgentInstanceId(7), SessionGeneration(4), 9));
167        assert_ne!(a, interaction_correlation(AgentInstanceId(8), SessionGeneration(3), 9));
168        assert_ne!(a, interaction_correlation(AgentInstanceId(7), SessionGeneration(3), 10));
169    }
170
171    #[test]
172    fn process_correlation_has_its_own_prefix() {
173        let id = process_correlation(AgentInstanceId(1), SessionGeneration(1));
174        assert!(id.starts_with("proc-"));
175        assert_eq!(id.len(), 5 + 16);
176    }
177
178    #[test]
179    fn tool_class_label_generalises_and_defaults() {
180        assert_eq!(tool_class_label("Bash"), "Shell");
181        assert_eq!(tool_class_label("  "), "Tool");
182        assert_eq!(tool_class_label("zzz"), "Tool");
183    }
184
185    #[test]
186    fn truncate_text_cuts_on_a_character_boundary() {
187        assert_eq!(truncate_text("abc", 8), ("abc".to_owned(), false));
188        assert_eq!(truncate_text("ééé", 3), ("é".to_owned(), true));
189    }
190}