agent_sdk_core/application/
stream.rs1use std::collections::{BTreeMap, BTreeSet};
8
9use serde::{Deserialize, Serialize};
10
11use crate::{
12 domain::AgentError,
13 stream_records::{
14 RedactedMatch, RepeatPolicy, StreamDelta, StreamIntervention, StreamMatcher, StreamRule,
15 StreamRuleRepeatStateSnapshot,
16 },
17};
18
19#[derive(Clone, Debug, Default, Deserialize, Eq, PartialEq, Serialize)]
20pub struct StreamRuleEngineState {
23 #[serde(default, skip_serializing_if = "Vec::is_empty")]
24 pub seen_match_keys: Vec<String>,
28}
29
30#[derive(Clone, Debug)]
31pub struct StreamRuleEngine {
34 rules: Vec<StreamRule>,
35 buffers: BTreeMap<String, String>,
36 seen_match_keys: BTreeSet<String>,
37}
38
39impl StreamRuleEngine {
40 pub fn new(rules: Vec<StreamRule>) -> Result<Self, AgentError> {
44 for rule in &rules {
45 rule.validate()?;
46 }
47 Ok(Self {
48 rules,
49 buffers: BTreeMap::new(),
50 seen_match_keys: BTreeSet::new(),
51 })
52 }
53
54 pub fn restore(
58 rules: Vec<StreamRule>,
59 state: StreamRuleEngineState,
60 ) -> Result<Self, AgentError> {
61 let mut engine = Self::new(rules)?;
62 engine.seen_match_keys = state.seen_match_keys.into_iter().collect();
63 Ok(engine)
64 }
65
66 pub fn rules(&self) -> &[StreamRule] {
70 &self.rules
71 }
72
73 pub fn snapshot_state(&self) -> StreamRuleEngineState {
77 StreamRuleEngineState {
78 seen_match_keys: self.seen_match_keys.iter().cloned().collect(),
79 }
80 }
81
82 pub fn repeat_state_for(&self, rule: &StreamRule) -> StreamRuleRepeatStateSnapshot {
85 let prefix = format!("{}:", rule.id.as_str());
86 StreamRuleRepeatStateSnapshot {
87 seen_match_keys: self
88 .seen_match_keys
89 .iter()
90 .filter(|key| key.starts_with(&prefix))
91 .cloned()
92 .collect(),
93 }
94 }
95
96 pub fn observe_delta(
100 &mut self,
101 delta: StreamDelta,
102 ) -> Result<Vec<StreamIntervention>, AgentError> {
103 if !delta.channel.is_policy_visible() {
104 return Ok(Vec::new());
105 }
106
107 let mut interventions = Vec::new();
108 let rules = self.rules.clone();
109 for rule in &rules {
110 if !rule
111 .channels
112 .iter()
113 .any(|selector| selector.matches(&delta))
114 {
115 continue;
116 }
117
118 let match_text = match &rule.matcher {
119 StreamMatcher::Marker { marker_id, .. } => {
120 if delta.marker_id.as_ref() == Some(marker_id) {
121 Some(delta.redacted_summary.as_str())
122 } else {
123 None
124 }
125 }
126 StreamMatcher::Literal { .. } | StreamMatcher::Regex { .. } => delta.matcher_text(),
127 StreamMatcher::HostMatcher { .. } => None,
128 };
129
130 let Some(chunk_text) = match_text else {
131 continue;
132 };
133
134 let buffer_key = buffer_key(rule, &delta);
135 let buffer = self.buffers.entry(buffer_key).or_default();
136 buffer.push_str(chunk_text);
137 truncate_utf8_suffix(buffer, rule.matcher.window_bytes() as usize);
138
139 let Some((start, end)) = find_match(&rule.matcher, buffer)? else {
140 continue;
141 };
142 let matched_text = &buffer[start..end];
143 let redacted = RedactedMatch::from_text(rule, &delta, matched_text);
144 let repeat_key = repeat_key(rule, &delta, &redacted);
145 if !matches!(rule.repeat, RepeatPolicy::Always)
146 && !self.seen_match_keys.insert(repeat_key)
147 {
148 continue;
149 }
150
151 interventions.push(StreamIntervention::proposed(rule, redacted));
152 }
153 Ok(interventions)
154 }
155}
156
157fn buffer_key(rule: &StreamRule, delta: &StreamDelta) -> String {
158 format!(
159 "{}:{:?}:{:?}:{:?}:{:?}",
160 rule.id.as_str(),
161 delta.channel,
162 delta.direction,
163 delta.attempt_id.as_ref().map(|id| id.as_str().to_string()),
164 delta
165 .realtime_session_id
166 .as_ref()
167 .map(|id| id.as_str().to_string())
168 )
169}
170
171fn repeat_key(rule: &StreamRule, delta: &StreamDelta, redacted: &RedactedMatch) -> String {
172 match rule.repeat {
173 RepeatPolicy::Always => format!(
174 "{}:always:{}:{}",
175 rule.id.as_str(),
176 redacted.text_hash,
177 delta.cursor.chunk_sequence
178 ),
179 RepeatPolicy::OncePerRun => format!("{}:run:{}", rule.id.as_str(), delta.run_id.as_str()),
180 RepeatPolicy::OncePerTurn => format!(
181 "{}:turn:{}",
182 rule.id.as_str(),
183 delta
184 .turn_id
185 .as_ref()
186 .map(|id| id.as_str())
187 .unwrap_or(delta.run_id.as_str())
188 ),
189 RepeatPolicy::OncePerAttemptAndSpan => format!(
190 "{}:attempt:{:?}:{:?}:{}:{}:{}",
191 rule.id.as_str(),
192 delta.attempt_id.as_ref().map(|id| id.as_str().to_string()),
193 delta
194 .realtime_session_id
195 .as_ref()
196 .map(|id| id.as_str().to_string()),
197 delta.channel.as_contract_name(),
198 redacted.text_hash,
199 redacted.cursor.chunk_sequence
200 ),
201 }
202}
203
204fn find_match(matcher: &StreamMatcher, buffer: &str) -> Result<Option<(usize, usize)>, AgentError> {
205 match matcher {
206 StreamMatcher::Literal {
207 text,
208 case_sensitive,
209 ..
210 } => {
211 if *case_sensitive {
212 Ok(buffer.find(text).map(|start| (start, start + text.len())))
213 } else {
214 let haystack = buffer.to_lowercase();
215 let needle = text.to_lowercase();
216 Ok(haystack
217 .find(&needle)
218 .map(|start| (start, start + needle.len())))
219 }
220 }
221 StreamMatcher::Regex { pattern, .. } => safe_regex_find(pattern, buffer),
222 StreamMatcher::Marker { .. } => Ok(Some((0, buffer.len()))),
223 StreamMatcher::HostMatcher { .. } => Ok(None),
224 }
225}
226
227fn safe_regex_find(pattern: &str, buffer: &str) -> Result<Option<(usize, usize)>, AgentError> {
228 crate::stream_records::validate_safe_regex(pattern)?;
229
230 if let Some(match_range) = find_char_class_repetition(pattern, buffer) {
231 return Ok(Some(match_range));
232 }
233 if let Some(match_range) = find_digit_repetition(pattern, buffer) {
234 return Ok(Some(match_range));
235 }
236 if pattern.contains(".*") {
237 return Ok(find_ordered_parts(pattern, buffer));
238 }
239
240 let literal = unescape_regex_literal(pattern);
241 Ok(buffer
242 .find(&literal)
243 .map(|start| (start, start + literal.len())))
244}
245
246fn find_char_class_repetition(pattern: &str, buffer: &str) -> Option<(usize, usize)> {
247 let class_start = pattern.find('[')?;
248 let class_end = pattern[class_start..].find(']')? + class_start;
249 let quantifier = &pattern[class_end + 1..];
250 let min = if let Some(open) = quantifier.find('{') {
251 let close = quantifier[open + 1..].find('}')? + open + 1;
252 quantifier[open + 1..close]
253 .trim_end_matches(',')
254 .parse::<usize>()
255 .ok()?
256 } else {
257 return None;
258 };
259 let prefix = unescape_regex_literal(&pattern[..class_start]);
260 let suffix = "";
261 let start = buffer.find(&prefix)?;
262 let mut index = start + prefix.len();
263 let mut count = 0;
264 for character in buffer[index..].chars() {
265 if character.is_ascii_alphanumeric() {
266 index += character.len_utf8();
267 count += 1;
268 } else {
269 break;
270 }
271 }
272 if count >= min && buffer[index..].starts_with(suffix) {
273 Some((start, index + suffix.len()))
274 } else {
275 None
276 }
277}
278
279fn find_digit_repetition(pattern: &str, buffer: &str) -> Option<(usize, usize)> {
280 let marker = "\\d+";
281 let digit_start = pattern.find(marker)?;
282 let prefix = unescape_regex_literal(&pattern[..digit_start]);
283 let suffix = unescape_regex_literal(&pattern[digit_start + marker.len()..]);
284 let start = buffer.find(&prefix)?;
285 let mut index = start + prefix.len();
286 let mut count = 0;
287 for character in buffer[index..].chars() {
288 if character.is_ascii_digit() {
289 index += character.len_utf8();
290 count += 1;
291 } else {
292 break;
293 }
294 }
295 if count > 0 && buffer[index..].starts_with(&suffix) {
296 Some((start, index + suffix.len()))
297 } else {
298 None
299 }
300}
301
302fn find_ordered_parts(pattern: &str, buffer: &str) -> Option<(usize, usize)> {
303 let parts = pattern
304 .split(".*")
305 .map(unescape_regex_literal)
306 .collect::<Vec<_>>();
307 let first = parts.first()?;
308 let mut start = buffer.find(first)?;
309 let mut cursor = start;
310 for part in &parts {
311 if part.is_empty() {
312 continue;
313 }
314 let relative = buffer[cursor..].find(part)?;
315 cursor += relative + part.len();
316 }
317 if first.is_empty() {
318 start = 0;
319 }
320 Some((start, cursor))
321}
322
323fn unescape_regex_literal(pattern: &str) -> String {
324 pattern
325 .replace("\\.", ".")
326 .replace("\\-", "-")
327 .replace("\\_", "_")
328}
329
330fn truncate_utf8_suffix(buffer: &mut String, max_bytes: usize) {
331 if max_bytes == 0 || buffer.len() <= max_bytes {
332 return;
333 }
334 let mut start = buffer.len() - max_bytes;
335 while !buffer.is_char_boundary(start) {
336 start += 1;
337 }
338 buffer.replace_range(..start, "");
339}