1use mermaid_domain::ProgressEvent;
11use std::collections::VecDeque;
12use std::sync::{Arc, Mutex, OnceLock};
13
14use async_trait::async_trait;
15use futures::{StreamExt, stream};
16
17use mermaid_domain::{FetchBackend, SearchBackend, WebConfig};
18use mermaid_domain::{ToolDefinition, ToolMetadata, ToolOutcome, ToolRunMetadata};
19
20use super::super::ctx::ExecContext;
21use super::ToolExecutor;
22use super::web_client::{
23 FetchProvider, ManagedSearxngBackend, NativeFetchClient, OllamaWebClient, SearchProvider,
24 SearxngClient, ValidatedWebUrl, WebFetchError, WebFetchResult, format_results,
25};
26
27#[derive(Debug, Clone, Copy, PartialEq, Eq)]
32pub enum Egress {
33 OnMachine,
35 OffMachine,
39}
40
41#[derive(Debug, Clone, PartialEq, Eq)]
44pub struct WebCapabilityStatus {
45 pub available: bool,
46 pub backend: &'static str,
47 pub trust_destination: &'static str,
48 pub egress: Egress,
49 pub reason: Option<String>,
50}
51
52impl WebCapabilityStatus {
53 #[must_use]
59 pub fn absence_reason(&self, tool: &str) -> String {
60 let reason = self.reason.as_deref().map_or_else(
61 || "backend initialization failed".to_string(),
62 mermaid_model::utils::redact_secrets,
63 );
64 let remedy = match self.backend {
65 "managed_searxng" => {
66 "The user can set [web] allow_ollama_search_fallback = true (needs \
67 OLLAMA_API_KEY), set [web] search_backend = \"ollama\", or point \
68 [web] searxng_url at a SearXNG instance, then restart mermaid"
69 },
70 "ollama_cloud" => {
71 "The user can set OLLAMA_API_KEY (or switch the [web] backend), \
72 then restart mermaid"
73 },
74 "searxng" => "The user can fix [web] searxng_url, then restart mermaid",
75 _ => "The user can check the [web] config, then restart mermaid",
76 };
77 format!(
78 "{tool} is configured to use the {} backend, which is unavailable: {reason}. {remedy}.",
79 self.backend
80 )
81 }
82}
83
84pub struct WebCapabilities {
88 pub fetch: WebCapabilityStatus,
89 pub search: WebCapabilityStatus,
90 fetch_backend: Option<Arc<dyn FetchProvider>>,
91 search_backend: Option<Arc<dyn SearchProvider>>,
92}
93
94impl WebCapabilities {
95 #[must_use]
96 pub fn resolve(web: &WebConfig) -> Self {
97 let needs_ollama_key = web.fetch_backend == FetchBackend::Ollama
98 || web.search_backend == SearchBackend::Ollama
99 || (web.search_backend == SearchBackend::Auto && web.allow_ollama_search_fallback);
100 let ollama_key = needs_ollama_key
101 .then(|| mermaid_model::utils::resolve_provider_key("ollama", "OLLAMA_API_KEY", None))
102 .flatten();
103
104 let (fetch, fetch_backend): (_, Option<Arc<dyn FetchProvider>>) = match web.fetch_backend {
105 FetchBackend::Native => match NativeFetchClient::new() {
106 Ok(client) => (
107 available("native", "direct from this machine", Egress::OnMachine),
108 Some(Arc::new(client)),
109 ),
110 Err(error) => (
111 unavailable(
112 "native",
113 "direct from this machine",
114 Egress::OnMachine,
115 error.to_string(),
116 ),
117 None,
118 ),
119 },
120 FetchBackend::Ollama => {
121 let (status, client) = ollama_cloud_backend(
122 ollama_key.clone(),
123 "Ollama Cloud (target redirects are provider-managed; final URL is not disclosed)",
124 );
125 (status, client.map(|c| c as Arc<dyn FetchProvider>))
126 },
127 };
128
129 let (search, search_backend): (_, Option<Arc<dyn SearchProvider>>) =
130 match web.search_backend {
131 SearchBackend::Auto => match crate::searxng::managed_backend_viability() {
138 Ok(_) => (
139 available(
140 "managed_searxng",
141 "local managed process",
142 Egress::OnMachine,
143 ),
144 Some(Arc::new(ManagedSearxngBackend)),
145 ),
146 Err(viability) if web.allow_ollama_search_fallback => {
147 let (status, client) = ollama_cloud_backend(ollama_key, "Ollama Cloud");
152 (
153 fallback_status(status, &viability),
154 client.map(|c| c as Arc<dyn SearchProvider>),
155 )
156 },
157 Err(reason) => (
158 unavailable(
159 "managed_searxng",
160 "local managed process",
161 Egress::OnMachine,
162 reason,
163 ),
164 None,
165 ),
166 },
167 SearchBackend::Ollama => {
168 let (status, client) = ollama_cloud_backend(ollama_key, "Ollama Cloud");
169 (status, client.map(|c| c as Arc<dyn SearchProvider>))
170 },
171 SearchBackend::Searxng => match SearxngClient::new(web.searxng_url.clone()) {
174 Ok(client) => (
175 available("searxng", "configured SearXNG instance", Egress::OffMachine),
176 Some(Arc::new(client)),
177 ),
178 Err(error) => (
179 unavailable(
180 "searxng",
181 "configured SearXNG instance",
182 Egress::OffMachine,
183 error.to_string(),
184 ),
185 None,
186 ),
187 },
188 };
189
190 Self {
191 fetch,
192 search,
193 fetch_backend,
194 search_backend,
195 }
196 }
197
198 #[cfg(test)]
202 #[must_use]
203 pub fn from_statuses_for_test(fetch: WebCapabilityStatus, search: WebCapabilityStatus) -> Self {
204 Self {
205 fetch,
206 search,
207 fetch_backend: None,
208 search_backend: None,
209 }
210 }
211
212 #[must_use]
213 pub fn fetch_tool(&self) -> Option<WebFetchTool> {
214 self.fetch_backend
215 .clone()
216 .map(|backend| WebFetchTool::new(backend, self.fetch.backend))
217 }
218
219 #[must_use]
220 pub fn search_tool(&self) -> Option<WebSearchTool> {
221 self.search_backend.clone().map(|backend| WebSearchTool {
222 backend,
223 backend_name: self.search.backend,
224 })
225 }
226}
227
228fn ollama_cloud_backend(
232 key: Option<String>,
233 trust_destination: &'static str,
234) -> (WebCapabilityStatus, Option<Arc<OllamaWebClient>>) {
235 match key {
236 Some(key) => match OllamaWebClient::new(key) {
237 Ok(client) => (
238 available("ollama_cloud", trust_destination, Egress::OffMachine),
239 Some(Arc::new(client)),
240 ),
241 Err(error) => (
242 unavailable(
243 "ollama_cloud",
244 "Ollama Cloud",
245 Egress::OffMachine,
246 error.to_string(),
247 ),
248 None,
249 ),
250 },
251 None => (
252 unavailable(
253 "ollama_cloud",
254 "Ollama Cloud",
255 Egress::OffMachine,
256 "OLLAMA_API_KEY is not configured",
257 ),
258 None,
259 ),
260 }
261}
262
263fn fallback_status(mut status: WebCapabilityStatus, viability: &str) -> WebCapabilityStatus {
269 if let Some(reason) = status.reason.take() {
270 status.reason = Some(format!(
271 "the managed bundle is unavailable ({viability}) and the configured \
272 Ollama Cloud fallback is too: {reason}"
273 ));
274 }
275 status
276}
277
278fn available(
279 backend: &'static str,
280 trust_destination: &'static str,
281 egress: Egress,
282) -> WebCapabilityStatus {
283 WebCapabilityStatus {
284 available: true,
285 backend,
286 trust_destination,
287 egress,
288 reason: None,
289 }
290}
291
292fn unavailable(
293 backend: &'static str,
294 trust_destination: &'static str,
295 egress: Egress,
296 reason: impl Into<String>,
297) -> WebCapabilityStatus {
298 WebCapabilityStatus {
299 available: false,
300 backend,
301 trust_destination,
302 egress,
303 reason: Some(reason.into()),
304 }
305}
306
307pub struct WebSearchTool {
311 backend: Arc<dyn SearchProvider>,
312 backend_name: &'static str,
313}
314
315const MAX_WEB_SEARCH_FAILURE_BYTES: usize = 1024;
316
317#[async_trait]
318impl ToolExecutor for WebSearchTool {
319 fn name(&self) -> &'static str {
320 "web_search"
321 }
322
323 fn schema(&self) -> ToolDefinition {
324 ToolDefinition {
325 name: "web_search".to_string(),
326 description:
327 "Search the web. Takes either a single `query` + `max_results`, or an array of `queries` for parallel fan-out."
328 .to_string(),
329 input_schema: serde_json::json!({
330 "type": "object",
331 "properties": {
332 "query": { "type": "string", "minLength": 1, "maxLength": 2048 },
333 "max_results": { "type": "integer", "minimum": 1, "maximum": 10, "default": 5 },
334 "queries": {
335 "type": "array",
336 "minItems": 1,
337 "maxItems": mermaid_model::constants::MAX_BATCH_TOOL_ITEMS,
338 "items": {
339 "type": "object",
340 "properties": {
341 "query": { "type": "string", "minLength": 1, "maxLength": 2048 },
342 "max_results": { "type": "integer", "minimum": 1, "maximum": 10, "default": 5 }
343 },
344 "required": ["query"],
345 "additionalProperties": false
346 }
347 }
348 },
349 "oneOf": [
350 { "required": ["query"], "not": { "required": ["queries"] } },
351 { "required": ["queries"], "not": { "required": ["query"] } }
352 ],
353 "additionalProperties": false
354 }),
355 }
356 }
357
358 #[expect(
359 clippy::too_many_lines,
360 reason = "the fan-out search: gate, run the queries concurrently, fold each result or \
361 failure into the combined text, then the all-failed error and the success both build the \
362 same ten-field WebSearch metadata from the query list and the failure list; a helper for \
363 the fold would return that pair of lists and the sources, and the outcome builders would \
364 still need everything else"
365 )]
366 async fn execute(&self, args: serde_json::Value, ctx: ExecContext) -> ToolOutcome {
367 let queries = match parse_queries(&args) {
368 Ok(q) => q,
369 Err(e) => return ToolOutcome::error(e, 0.0),
370 };
371 if queries.is_empty() {
372 return ToolOutcome::error("web_search requires at least one query", 0.0);
373 }
374 if let Some(blocked) = super::policy_gate::gate_external(
375 &ctx,
376 "web_search",
377 mermaid_runtime::ToolCategory::Web,
378 format!("web_search ({} queries)", queries.len()),
379 &args,
380 )
381 .await
382 {
383 return blocked;
384 }
385
386 let start = std::time::Instant::now();
387 let jobs = stream::iter(queries.iter().cloned().enumerate())
388 .map(|(idx, (query, count))| {
389 let backend = self.backend.clone();
390 let progress = ctx.progress.clone();
391 let budget = ctx.web_budget();
392 let total = queries.len();
393 async move {
394 let display_query = mermaid_model::utils::redact_secrets(&query);
395 let _ = progress
396 .send(ProgressEvent::Status(format!(
397 "searching {}/{}: {}",
398 idx + 1,
399 total,
400 display_query
401 )))
402 .await;
403 let result = backend.search(&query, count, budget).await;
404 (idx, query, result)
405 }
406 })
407 .buffer_unordered(mermaid_model::constants::MAX_WEB_SEARCH_CONCURRENCY)
408 .collect::<Vec<_>>();
409 let mut completed = tokio::select! {
410 biased;
411 _ = ctx.token.cancelled() => return ToolOutcome::cancelled(),
412 completed = jobs => completed,
413 };
414 completed.sort_by_key(|(idx, _, _)| *idx);
415
416 let mut combined = String::new();
417 let mut result_count = 0usize;
418 let mut sources = Vec::new();
419 let mut errors: Vec<mermaid_domain::WebSearchFailure> = Vec::new();
420 for (idx, query, result) in completed {
421 let display_query = mermaid_model::utils::redact_secrets(&query);
422 let section =
426 match result {
427 Ok(results) => {
428 result_count += results.len();
429 sources.extend(results.iter().map(|result| {
430 mermaid_model::utils::sanitize_url_for_display(&result.url)
431 }));
432 if results.is_empty() {
433 "[SEARCH_RESULTS]\n(no results found)\n[/SEARCH_RESULTS]\n".to_string()
434 } else {
435 format_results(&results)
436 }
437 },
438 Err(e) => {
439 let safe_error = mermaid_model::utils::truncate_middle_bytes(
440 &mermaid_model::utils::redact_secrets(&format!("{e:#}")),
441 MAX_WEB_SEARCH_FAILURE_BYTES,
442 );
443 errors.push(mermaid_domain::WebSearchFailure {
444 query_index: idx,
445 error: safe_error.clone(),
446 });
447 format!("(search failed: {safe_error})\n")
448 },
449 };
450 if queries.len() > 1 {
451 combined.push_str(&format!("=== query: {display_query} ===\n{section}\n\n"));
452 } else {
453 combined = section;
454 }
455 }
456
457 if errors.len() == queries.len() {
461 let summary = errors
462 .iter()
463 .map(|failure| format!("query {}: {}", failure.query_index + 1, failure.error))
464 .collect::<Vec<_>>()
465 .join("; ");
466 let message = format!("web_search via {} failed: {summary}", self.backend_name);
467 let message = mermaid_model::utils::truncate_middle_bytes(
468 &message,
469 mermaid_model::constants::WEB_SEARCH_AGGREGATE_MAX_BYTES
470 .saturating_sub("Error: ".len()),
471 );
472 return ToolOutcome::error(message, start.elapsed().as_secs_f64()).with_metadata(
473 ToolRunMetadata {
474 detail: ToolMetadata::WebSearch {
475 queries: queries.iter().map(|(query, _)| query.clone()).collect(),
476 requested_count: queries.iter().map(|(_, count)| *count).sum(),
477 result_count: 0,
478 sources: Vec::new(),
479 backend: self.backend_name.to_string(),
480 succeeded_queries: 0,
481 failed_queries: errors.len(),
482 partial: false,
483 truncated: false,
484 failures: errors,
485 },
486 result_count: Some(0),
487 ..ToolRunMetadata::default()
488 },
489 );
490 }
491
492 let truncated = combined.len() > mermaid_model::constants::WEB_SEARCH_AGGREGATE_MAX_BYTES;
496 let combined = mermaid_model::utils::truncate_middle_bytes(
497 &combined,
498 mermaid_model::constants::WEB_SEARCH_AGGREGATE_MAX_BYTES,
499 );
500
501 let duration_secs = start.elapsed().as_secs_f64();
502 let requested_count = queries.iter().map(|(_, count)| *count).sum();
503 let query_texts = queries.iter().map(|(query, _)| query.clone()).collect();
504 ToolOutcome::success(
505 combined,
506 format!(
507 "{} {} returned",
508 result_count,
509 if result_count == 1 {
510 "result"
511 } else {
512 "results"
513 }
514 ),
515 duration_secs,
516 )
517 .with_metadata(ToolRunMetadata {
518 detail: ToolMetadata::WebSearch {
519 queries: query_texts,
520 requested_count,
521 result_count,
522 sources,
523 backend: self.backend_name.to_string(),
524 succeeded_queries: queries.len() - errors.len(),
525 failed_queries: errors.len(),
526 partial: !errors.is_empty(),
527 truncated,
528 failures: errors,
529 },
530 result_count: Some(result_count),
531 ..ToolRunMetadata::default()
532 })
533 }
534}
535
536pub struct WebFetchTool {
540 backend: Arc<dyn FetchProvider>,
541 backend_name: &'static str,
542 snapshots: Arc<Mutex<FetchSnapshotStore>>,
543}
544
545impl WebFetchTool {
546 fn new(backend: Arc<dyn FetchProvider>, backend_name: &'static str) -> Self {
547 Self {
548 backend,
549 backend_name,
550 snapshots: global_fetch_snapshot_store(),
551 }
552 }
553
554 #[cfg(test)]
555 fn new_with_test_snapshots(
556 backend: Arc<dyn FetchProvider>,
557 backend_name: &'static str,
558 ) -> Self {
559 Self {
560 backend,
561 backend_name,
562 snapshots: Arc::new(Mutex::new(FetchSnapshotStore::default())),
563 }
564 }
565}
566
567fn fetch_failure_outcome(
568 error: &WebFetchError,
569 requested_url: &str,
570 backend: &str,
571 duration_secs: f64,
572 pattern: Option<String>,
573 context_lines: usize,
574) -> ToolOutcome {
575 let requested_url = mermaid_model::utils::sanitize_url_for_display(requested_url);
576 let message = mermaid_model::utils::redact_secrets(&format!(
577 "web_fetch({requested_url}) via {backend}: {error}"
578 ));
579 let pattern_context = pattern.as_ref().map(|_| context_lines);
580 ToolOutcome::error(message, duration_secs).with_metadata(ToolRunMetadata {
581 detail: ToolMetadata::WebFetch {
582 url: requested_url,
583 final_url: None,
584 status: error.status(),
585 error_kind: Some(error.kind().to_string()),
586 media_type: None,
587 charset: None,
588 backend: backend.to_string(),
589 extraction: String::new(),
590 title: None,
591 line_count: 0,
592 byte_count: 0,
593 source_byte_count: 0,
594 output_byte_count: 0,
595 truncated: false,
596 pattern,
597 context_lines: pattern_context,
598 match_count: None,
599 snapshot_id: None,
600 },
601 line_count: Some(0),
602 byte_count: Some(0),
603 ..ToolRunMetadata::default()
604 })
605}
606
607const MAX_FETCH_SNAPSHOTS: usize = 4;
608const MAX_FETCH_SNAPSHOT_BYTES: usize = 32 * 1024 * 1024;
609const MAX_SNAPSHOT_TITLE_BYTES: usize = 300;
610const MAX_SNAPSHOT_URL_BYTES: usize = 8 * 1024;
611const MAX_SNAPSHOT_MEDIA_TYPE_BYTES: usize = 256;
612const MAX_SNAPSHOT_CHARSET_BYTES: usize = 64;
613
614static FETCH_SNAPSHOT_STORE: OnceLock<Arc<Mutex<FetchSnapshotStore>>> = OnceLock::new();
615
616#[derive(Clone, Debug, PartialEq, Eq)]
617struct FetchSnapshotScope {
618 session_id: Option<String>,
619 task_id: Option<String>,
620 fallback_turn: Option<u64>,
621}
622
623impl FetchSnapshotScope {
624 fn from_context(ctx: &ExecContext) -> Self {
625 let has_owner = ctx.session_id.is_some() || ctx.task_id.is_some();
626 Self {
627 session_id: ctx.session_id.as_deref().map(compact_string),
628 task_id: ctx.task_id.as_deref().map(compact_string),
629 fallback_turn: (!has_owner).then_some(ctx.turn.0),
632 }
633 }
634
635 fn retained_string_bytes(&self) -> usize {
636 option_string_capacity(&self.session_id)
637 .saturating_add(option_string_capacity(&self.task_id))
638 }
639}
640
641#[derive(Clone)]
642struct FetchSnapshot {
643 id: String,
644 scope: FetchSnapshotScope,
645 page: Arc<WebFetchResult>,
646 retained_bytes: usize,
647}
648
649#[derive(Default)]
650struct FetchSnapshotStore {
651 entries: VecDeque<FetchSnapshot>,
652 bytes: usize,
653 next_id: u64,
654}
655
656impl FetchSnapshotStore {
657 fn insert(
658 &mut self,
659 scope: FetchSnapshotScope,
660 page: WebFetchResult,
661 ) -> Result<(String, Arc<WebFetchResult>), String> {
662 self.next_id = self.next_id.wrapping_add(1).max(1);
663 let mut id = format!("web-{}", self.next_id);
664 id.shrink_to_fit();
665
666 let fixed_bytes = id.capacity().saturating_add(scope.retained_string_bytes());
667 if fixed_bytes >= MAX_FETCH_SNAPSHOT_BYTES {
668 return Err("web_fetch: snapshot owner identity exceeds the cache budget".to_string());
669 }
670 let page = Arc::new(bound_snapshot_page(
671 page,
672 MAX_FETCH_SNAPSHOT_BYTES - fixed_bytes,
673 ));
674 let retained_bytes = fixed_bytes.saturating_add(page_retained_string_bytes(&page));
675 if retained_bytes > MAX_FETCH_SNAPSHOT_BYTES {
676 return Err("web_fetch: snapshot metadata exceeds the cache budget".to_string());
677 }
678
679 while !self.entries.is_empty()
680 && (self.entries.len() >= MAX_FETCH_SNAPSHOTS
681 || self.bytes.saturating_add(retained_bytes) > MAX_FETCH_SNAPSHOT_BYTES)
682 {
683 if let Some(removed) = self.entries.pop_front() {
684 self.bytes = self.bytes.saturating_sub(removed.retained_bytes);
685 }
686 }
687 self.bytes = self.bytes.saturating_add(retained_bytes);
688 self.entries.push_back(FetchSnapshot {
689 id: id.clone(),
690 scope,
691 page: page.clone(),
692 retained_bytes,
693 });
694 Ok((id, page))
695 }
696
697 fn get(&self, scope: &FetchSnapshotScope, id: &str) -> Option<Arc<WebFetchResult>> {
698 self.entries
699 .iter()
700 .find(|entry| entry.id == id && &entry.scope == scope)
701 .map(|entry| Arc::clone(&entry.page))
702 }
703}
704
705fn global_fetch_snapshot_store() -> Arc<Mutex<FetchSnapshotStore>> {
706 FETCH_SNAPSHOT_STORE
707 .get_or_init(|| Arc::new(Mutex::new(FetchSnapshotStore::default())))
708 .clone()
709}
710
711fn bound_snapshot_page(mut page: WebFetchResult, max_retained_bytes: usize) -> WebFetchResult {
712 page.title = bounded_title(&page.title);
713 page.title.shrink_to_fit();
714 bound_owned_string(&mut page.requested_url, MAX_SNAPSHOT_URL_BYTES);
715 if let Some(final_url) = page.final_url.as_mut() {
716 bound_owned_string(final_url, MAX_SNAPSHOT_URL_BYTES);
717 }
718 bound_optional_string(&mut page.media_type, MAX_SNAPSHOT_MEDIA_TYPE_BYTES);
719 bound_optional_string(&mut page.charset, MAX_SNAPSHOT_CHARSET_BYTES);
720
721 let metadata_bytes = page_retained_string_bytes_without_content(&page);
722 let content_budget = max_retained_bytes.saturating_sub(metadata_bytes);
723 if page.content.len() > content_budget {
724 let cut = page.content.floor_char_boundary(content_budget);
725 page.content.truncate(cut);
726 page.truncated = true;
727 }
728 page.content.shrink_to_fit();
729 page
730}
731
732fn compact_string(value: &str) -> String {
733 let mut value = value.to_string();
734 value.shrink_to_fit();
735 value
736}
737
738fn bound_owned_string(value: &mut String, max_bytes: usize) {
739 if value.len() > max_bytes {
740 value.truncate(value.floor_char_boundary(max_bytes));
741 }
742 value.shrink_to_fit();
743}
744
745fn bound_optional_string(value: &mut Option<String>, max_bytes: usize) {
746 if let Some(value) = value {
747 bound_owned_string(value, max_bytes);
748 }
749}
750
751fn option_string_capacity(value: &Option<String>) -> usize {
752 value.as_ref().map_or(0, String::capacity)
753}
754
755fn page_retained_string_bytes_without_content(page: &WebFetchResult) -> usize {
756 page.requested_url
757 .capacity()
758 .saturating_add(option_string_capacity(&page.final_url))
759 .saturating_add(option_string_capacity(&page.media_type))
760 .saturating_add(option_string_capacity(&page.charset))
761 .saturating_add(page.title.capacity())
762}
763
764fn page_retained_string_bytes(page: &WebFetchResult) -> usize {
765 page_retained_string_bytes_without_content(page).saturating_add(page.content.capacity())
766}
767
768async fn run_snapshot_blocking<T, F>(work: F) -> Result<T, String>
769where
770 T: Send + 'static,
771 F: FnOnce() -> T + Send + 'static,
772{
773 run_snapshot_blocking_with(super::web_client::extraction_semaphore(), work).await
774}
775
776async fn run_snapshot_blocking_with<T, F>(
777 limiter: Arc<tokio::sync::Semaphore>,
778 work: F,
779) -> Result<T, String>
780where
781 T: Send + 'static,
782 F: FnOnce() -> T + Send + 'static,
783{
784 let permit = limiter
785 .acquire_owned()
786 .await
787 .map_err(|_| "web snapshot renderer is closed".to_string())?;
788 tokio::task::spawn_blocking(move || {
789 let _permit = permit;
792 work()
793 })
794 .await
795 .map_err(|error| format!("web snapshot renderer failed: {error}"))
796}
797
798#[async_trait]
799impl ToolExecutor for WebFetchTool {
800 fn name(&self) -> &'static str {
801 "web_fetch"
802 }
803
804 fn schema(&self) -> ToolDefinition {
805 ToolDefinition {
806 name: "web_fetch".to_string(),
807 description: "Fetch a public HTTP(S) URL into a bounded session snapshot, or inspect \
808 a prior snapshot without refetching. Use pattern for case-insensitive \
809 matching, or start_line + line_count for stable continuation."
810 .to_string(),
811 input_schema: serde_json::json!({
812 "type": "object",
813 "properties": {
814 "url": {
815 "type": "string",
816 "format": "uri",
817 "maxLength": 8192,
818 "description": "Public HTTP(S) URL to fetch"
819 },
820 "snapshot_id": {
821 "type": "string",
822 "pattern": "^web-[0-9]+$",
823 "description": "Snapshot returned by an earlier web_fetch call"
824 },
825 "pattern": {
826 "type": "string",
827 "minLength": 1,
828 "maxLength": 1024,
829 "description": "Case-insensitive substring to find in the page (not a regex)"
830 },
831 "context_lines": {
832 "type": "integer",
833 "minimum": 0,
834 "maximum": 10,
835 "default": 2,
836 "description": "Context lines around each match (default 2, max 10)"
837 },
838 "start_line": {
839 "type": "integer",
840 "minimum": 1,
841 "description": "First 1-based snapshot line to return"
842 },
843 "line_count": {
844 "type": "integer",
845 "minimum": 1,
846 "maximum": 500,
847 "default": 200,
848 "description": "Maximum snapshot lines to return"
849 }
850 },
851 "oneOf": [
852 { "required": ["url"], "not": { "required": ["snapshot_id"] } },
853 { "required": ["snapshot_id"], "not": { "required": ["url"] } }
854 ],
855 "additionalProperties": false
856 }),
857 }
858 }
859
860 #[expect(
861 clippy::too_many_lines,
862 reason = "fetch-then-format: resolve a snapshot or gate and fetch a URL, then render the \
863 page and build the WebFetch metadata; both halves share the parsed request, the timer \
864 and the backend name, and the metadata literal alone is thirty lines of fields taken \
865 from the page and the format result"
866 )]
867 async fn execute(&self, args: serde_json::Value, ctx: ExecContext) -> ToolOutcome {
868 let request = match parse_fetch_args(&args) {
869 Ok(request) => request,
870 Err(error) => return ToolOutcome::error(error, 0.0),
871 };
872 let start = std::time::Instant::now();
873 let snapshot_scope = FetchSnapshotScope::from_context(&ctx);
874 let (page, snapshot_id) = match &request.target {
875 FetchTarget::Snapshot(snapshot_id) => {
876 let page = self
877 .snapshots
878 .lock()
879 .unwrap_or_else(std::sync::PoisonError::into_inner)
880 .get(&snapshot_scope, snapshot_id);
881 let Some(page) = page else {
882 return ToolOutcome::error(
883 format!(
884 "web_fetch: snapshot '{snapshot_id}' is unavailable or was evicted"
885 ),
886 start.elapsed().as_secs_f64(),
887 );
888 };
889 (page, snapshot_id.to_string())
890 },
891 FetchTarget::Url(url) => {
892 let safe_url = mermaid_model::utils::sanitize_url_for_display(url.as_str());
893 if let Some(blocked) = super::policy_gate::gate_external(
894 &ctx,
895 "web_fetch",
896 mermaid_runtime::ToolCategory::Web,
897 format!("web_fetch via {} {safe_url}", self.backend_name),
898 &args,
899 )
900 .await
901 {
902 return blocked;
903 }
904 let fetch = self.backend.fetch(url.as_str(), ctx.web_budget());
905 let page = tokio::select! {
906 biased;
907 _ = ctx.token.cancelled() => return ToolOutcome::cancelled(),
908 result = fetch => match result {
909 Ok(page) => page,
910 Err(error) => {
911 return fetch_failure_outcome(
912 &error,
913 url.as_str(),
914 self.backend_name,
915 start.elapsed().as_secs_f64(),
916 request.pattern.clone(),
917 request.context_lines,
918 );
919 },
920 },
921 };
922 let inserted = self
923 .snapshots
924 .lock()
925 .unwrap_or_else(std::sync::PoisonError::into_inner)
926 .insert(snapshot_scope, page);
927 match inserted {
928 Ok((snapshot_id, page)) => (page, snapshot_id),
929 Err(error) => {
930 return ToolOutcome::error(error, start.elapsed().as_secs_f64());
931 },
932 }
933 },
934 };
935
936 let render_page = Arc::clone(&page);
937 let render_snapshot_id = snapshot_id.clone();
938 let render_pattern = request.pattern.clone();
939 let render_context_lines = request.context_lines;
940 let render_start_line = request.start_line;
941 let render_line_count = request.line_count;
942 let render = run_snapshot_blocking(move || {
943 format_fetch(
944 &render_page,
945 &render_snapshot_id,
946 render_pattern.as_deref(),
947 render_context_lines,
948 render_start_line,
949 render_line_count,
950 )
951 });
952 let formatted = tokio::select! {
953 biased;
954 _ = ctx.token.cancelled() => return ToolOutcome::cancelled(),
955 result = render => match result {
956 Ok(formatted) => formatted,
957 Err(error) => {
958 return ToolOutcome::error(error, start.elapsed().as_secs_f64());
959 },
960 },
961 };
962 let duration_secs = start.elapsed().as_secs_f64();
963 let line_count = formatted.output.lines().count();
964 let byte_count = formatted.output.len();
965 let title = (!page.title.is_empty()).then(|| bounded_title(&page.title));
966 let requested_url = mermaid_model::utils::sanitize_url_for_display(&page.requested_url);
967 let final_url = page
968 .final_url
969 .as_deref()
970 .map(mermaid_model::utils::sanitize_url_for_display);
971 let pattern_context = request.pattern.as_ref().map(|_| request.context_lines);
972 ToolOutcome::success(
973 formatted.output,
974 format!(
975 "{} {} fetched via {}",
976 line_count,
977 if line_count == 1 { "line" } else { "lines" },
978 page.backend.as_str()
979 ),
980 duration_secs,
981 )
982 .with_metadata(ToolRunMetadata {
983 detail: ToolMetadata::WebFetch {
984 url: requested_url,
985 final_url,
986 status: page.status,
987 error_kind: None,
988 media_type: page.media_type.clone(),
989 charset: page.charset.clone(),
990 backend: page.backend.as_str().to_string(),
991 extraction: page.extraction.as_str().to_string(),
992 title,
993 line_count,
994 byte_count,
995 source_byte_count: page.source_bytes,
996 output_byte_count: page.output_bytes,
997 truncated: page.truncated || formatted.truncated,
998 pattern: request.pattern,
999 context_lines: pattern_context,
1000 match_count: formatted.match_count,
1001 snapshot_id: Some(snapshot_id),
1002 },
1003 line_count: Some(line_count),
1004 byte_count: Some(byte_count),
1005 ..ToolRunMetadata::default()
1006 })
1007 }
1008}
1009
1010enum FetchTarget {
1014 Url(ValidatedWebUrl),
1015 Snapshot(String),
1016}
1017
1018struct ParsedFetchArgs {
1019 target: FetchTarget,
1020 pattern: Option<String>,
1021 context_lines: usize,
1022 start_line: Option<usize>,
1023 line_count: usize,
1024}
1025
1026fn parse_fetch_args(args: &serde_json::Value) -> Result<ParsedFetchArgs, String> {
1027 let obj = args
1028 .as_object()
1029 .ok_or_else(|| "web_fetch arguments must be an object".to_string())?;
1030 for key in obj.keys() {
1031 if !matches!(
1032 key.as_str(),
1033 "url" | "snapshot_id" | "pattern" | "context_lines" | "start_line" | "line_count"
1034 ) {
1035 return Err(format!("web_fetch: unknown argument '{key}'"));
1036 }
1037 }
1038
1039 let url = match obj.get("url") {
1040 None => None,
1041 Some(value) => {
1042 let raw = value
1043 .as_str()
1044 .ok_or_else(|| "web_fetch: 'url' must be a string".to_string())?
1045 .trim();
1046 if raw.len() > 8192 {
1047 return Err("web_fetch: URL exceeds 8192 bytes".to_string());
1048 }
1049 Some(ValidatedWebUrl::parse(raw).map_err(|error| format!("web_fetch: {error}"))?)
1050 },
1051 };
1052 let snapshot_id = match obj.get("snapshot_id") {
1053 None => None,
1054 Some(value) => {
1055 let id = value
1056 .as_str()
1057 .ok_or_else(|| "web_fetch: 'snapshot_id' must be a string".to_string())?;
1058 let valid = id.strip_prefix("web-").is_some_and(|suffix| {
1059 !suffix.is_empty() && suffix.bytes().all(|b| b.is_ascii_digit())
1060 });
1061 if !valid {
1062 return Err("web_fetch: invalid snapshot id".to_string());
1063 }
1064 Some(id.to_string())
1065 },
1066 };
1067 let target = match (url, snapshot_id) {
1068 (Some(url), None) => FetchTarget::Url(url),
1069 (None, Some(id)) => FetchTarget::Snapshot(id),
1070 _ => {
1071 return Err("web_fetch requires exactly one of 'url' or 'snapshot_id'".to_string());
1072 },
1073 };
1074
1075 let pattern = match obj.get("pattern") {
1076 None => None,
1077 Some(value) => {
1078 let pattern = value
1079 .as_str()
1080 .ok_or_else(|| "web_fetch: 'pattern' must be a string".to_string())?
1081 .trim();
1082 if pattern.is_empty() {
1083 return Err("web_fetch: 'pattern' must not be empty".to_string());
1084 }
1085 if pattern.contains(['\r', '\n']) {
1086 return Err("web_fetch: 'pattern' must be a single line".to_string());
1087 }
1088 if pattern.chars().count() > 1024 {
1089 return Err("web_fetch: 'pattern' exceeds 1024 characters".to_string());
1090 }
1091 Some(pattern.to_string())
1092 },
1093 };
1094 let context_lines = parse_bounded_usize(obj, "context_lines", 2, 0, 10)?;
1095 if pattern.is_none() && obj.contains_key("context_lines") {
1096 return Err("web_fetch: 'context_lines' requires 'pattern'".to_string());
1097 }
1098 let has_range = obj.contains_key("start_line") || obj.contains_key("line_count");
1099 if pattern.is_some() && has_range {
1100 return Err("web_fetch: use either 'pattern' or a line range, not both".to_string());
1101 }
1102 let start_line = has_range
1103 .then(|| parse_bounded_usize(obj, "start_line", 1, 1, usize::MAX))
1104 .transpose()?;
1105 let line_count = parse_bounded_usize(obj, "line_count", 200, 1, 500)?;
1106
1107 Ok(ParsedFetchArgs {
1108 target,
1109 pattern,
1110 context_lines,
1111 start_line,
1112 line_count,
1113 })
1114}
1115
1116fn parse_bounded_usize(
1117 obj: &serde_json::Map<String, serde_json::Value>,
1118 key: &str,
1119 default: usize,
1120 min: usize,
1121 max: usize,
1122) -> Result<usize, String> {
1123 let Some(value) = obj.get(key) else {
1124 return Ok(default);
1125 };
1126 let value = value
1127 .as_u64()
1128 .and_then(|value| usize::try_from(value).ok())
1129 .ok_or_else(|| format!("web_fetch: '{key}' must be an integer"))?;
1130 if value < min || value > max {
1131 return Err(format!("web_fetch: '{key}' must be from {min} to {max}"));
1132 }
1133 Ok(value)
1134}
1135
1136const WEB_FETCH_MAX_BYTES: usize = mermaid_model::constants::WEB_SEARCH_AGGREGATE_MAX_BYTES;
1139const FETCH_TRUNCATION_SUFFIX: &str = "\n\n...[content truncated]\n[/WEB_FETCH]";
1140
1141struct FormattedFetch {
1142 output: String,
1143 truncated: bool,
1144 match_count: Option<usize>,
1145}
1146
1147const MAX_PATTERN_MATCHES: usize = 20;
1151
1152fn format_fetch(
1153 page: &WebFetchResult,
1154 snapshot_id: &str,
1155 pattern: Option<&str>,
1156 ctx_lines: usize,
1157 start_line: Option<usize>,
1158 line_count: usize,
1159) -> FormattedFetch {
1160 let title = if page.title.trim().is_empty() {
1161 "(no title)".to_string()
1162 } else {
1163 bounded_title(&page.title)
1164 };
1165 let requested_url = bounded_url(&page.requested_url);
1166 let final_url = page
1167 .final_url
1168 .as_deref()
1169 .map(bounded_url)
1170 .unwrap_or_else(|| "(not disclosed by backend)".to_string());
1171 let status = page
1172 .status
1173 .map(|status| status.to_string())
1174 .unwrap_or_else(|| "unknown".to_string());
1175 let media = page.media_type.as_deref().unwrap_or("unknown");
1176 let charset = page.charset.as_deref().unwrap_or("unknown");
1177
1178 let (body, match_count) = if let Some(pattern) = pattern {
1179 match extract_matches(&page.content, pattern, ctx_lines, MAX_PATTERN_MATCHES) {
1180 Some((report, count)) => (report, Some(count)),
1181 None => (format!("No matching lines for \"{pattern}\"."), Some(0)),
1182 }
1183 } else if let Some(start_line) = start_line {
1184 (
1185 format_line_range(&page.content, start_line, line_count),
1186 None,
1187 )
1188 } else {
1189 (page.content.clone(), None)
1190 };
1191
1192 let mut output = format!(
1193 "[WEB_FETCH]\nTitle: {title}\nRequested URL: {requested_url}\nFinal URL: {final_url}\nStatus: {status}\nMedia-Type: {media}\nCharset: {charset}\nBackend: {}\nExtraction: {}\nSnapshot: {snapshot_id}\nSource bytes: {}\nExtracted bytes: {}\nSnapshot bytes: {}\nSource lines: {}\n\nContent:\n{body}\n[/WEB_FETCH]",
1194 page.backend.as_str(),
1195 page.extraction.as_str(),
1196 page.source_bytes,
1197 page.output_bytes,
1198 page.content.len(),
1199 page.content.lines().count(),
1200 );
1201 let truncated = output.len() > WEB_FETCH_MAX_BYTES;
1202 if truncated {
1203 let budget = WEB_FETCH_MAX_BYTES.saturating_sub(FETCH_TRUNCATION_SUFFIX.len());
1204 let cut = output.floor_char_boundary(budget);
1205 output.truncate(cut);
1206 output.push_str(FETCH_TRUNCATION_SUFFIX);
1207 }
1208 FormattedFetch {
1209 output,
1210 truncated,
1211 match_count,
1212 }
1213}
1214
1215fn bounded_title(title: &str) -> String {
1216 let mut bounded = String::with_capacity(title.len().min(MAX_SNAPSHOT_TITLE_BYTES));
1217 let content_budget = MAX_SNAPSHOT_TITLE_BYTES.saturating_sub(3);
1218 let mut truncated = false;
1219
1220 for word in title.split_whitespace() {
1221 let separator_bytes = usize::from(!bounded.is_empty());
1222 if bounded
1223 .len()
1224 .saturating_add(separator_bytes)
1225 .saturating_add(word.len())
1226 <= MAX_SNAPSHOT_TITLE_BYTES
1227 {
1228 if separator_bytes != 0 {
1229 bounded.push(' ');
1230 }
1231 bounded.push_str(word);
1232 continue;
1233 }
1234
1235 if bounded.len() > content_budget {
1236 bounded.truncate(bounded.floor_char_boundary(content_budget));
1237 }
1238 if bounded.len() < content_budget {
1239 if separator_bytes != 0 && bounded.len() < content_budget {
1240 bounded.push(' ');
1241 }
1242 let remaining = content_budget.saturating_sub(bounded.len());
1243 let cut = word.floor_char_boundary(remaining);
1244 bounded.push_str(&word[..cut]);
1245 }
1246 truncated = true;
1247 break;
1248 }
1249
1250 if truncated {
1251 bounded.push_str("...");
1252 }
1253 bounded
1254}
1255
1256fn bounded_url(url: &str) -> String {
1257 const MAX_DISPLAY_URL_BYTES: usize = 2048;
1258 let url = mermaid_model::utils::sanitize_url_for_display(url);
1259 if url.len() <= MAX_DISPLAY_URL_BYTES {
1260 return url;
1261 }
1262 let cut = url.floor_char_boundary(MAX_DISPLAY_URL_BYTES.saturating_sub(3));
1263 format!("{}...", &url[..cut])
1264}
1265
1266fn format_line_range(content: &str, start_line: usize, line_count: usize) -> String {
1267 let total = content.lines().count();
1268 if start_line > total {
1269 return format!("Requested line {start_line}, but the snapshot contains {total} lines.");
1270 }
1271 let mut output = format!(
1272 "Lines {start_line}-{} of {total}:\n",
1273 start_line
1274 .saturating_add(line_count)
1275 .saturating_sub(1)
1276 .min(total)
1277 );
1278 for (offset, line) in content
1279 .lines()
1280 .skip(start_line.saturating_sub(1))
1281 .take(line_count)
1282 .enumerate()
1283 {
1284 output.push_str(&format!("L{}: {line}\n", start_line + offset));
1285 }
1286 output
1287}
1288
1289fn normalized_case_fold(value: &str) -> String {
1294 use caseless::Caseless;
1295 use unicode_normalization::UnicodeNormalization;
1296
1297 value.nfd().default_case_fold().nfd().collect()
1298}
1299
1300fn extract_matches(
1307 content: &str,
1308 pattern: &str,
1309 context_lines: usize,
1310 max_blocks: usize,
1311) -> Option<(String, usize)> {
1312 let needle = normalized_case_fold(pattern);
1313 let lines: Vec<&str> = content.lines().collect();
1314 let matched: Vec<usize> = lines
1315 .iter()
1316 .enumerate()
1317 .filter(|(_, line)| normalized_case_fold(line).contains(&needle))
1318 .map(|(i, _)| i)
1319 .collect();
1320 if matched.is_empty() {
1321 return None;
1322 }
1323
1324 let mut blocks: Vec<(usize, usize)> = Vec::new();
1327 for &i in &matched {
1328 let start = i.saturating_sub(context_lines);
1329 let end = (i + context_lines).min(lines.len() - 1);
1330 match blocks.last_mut() {
1331 Some((_, last_end)) if start <= *last_end + 1 => *last_end = (*last_end).max(end),
1332 _ => blocks.push((start, end)),
1333 }
1334 }
1335 let included = &blocks[..blocks.len().min(max_blocks)];
1336 let cutoff = included.last().map(|&(_, end)| end).unwrap_or(0);
1337 let dropped = matched.iter().filter(|&&i| i > cutoff).count();
1338
1339 let mut out = format!(
1340 "{} match{} for \"{}\":\n",
1341 matched.len(),
1342 if matched.len() == 1 { "" } else { "es" },
1343 pattern
1344 );
1345 for (bi, &(start, end)) in included.iter().enumerate() {
1346 if bi > 0 {
1347 out.push_str("---\n");
1348 }
1349 for (offset, line) in lines[start..=end].iter().enumerate() {
1350 out.push_str(&format!("L{}: {}\n", start + offset + 1, line));
1352 }
1353 }
1354 if dropped > 0 {
1355 out.push_str(&format!(
1356 "(+{dropped} more match{})\n",
1357 if dropped == 1 { "" } else { "es" }
1358 ));
1359 }
1360 let match_count = matched.len();
1361 Some((out, match_count))
1362}
1363
1364fn parse_queries(args: &serde_json::Value) -> Result<Vec<(String, usize)>, String> {
1365 let obj = args
1366 .as_object()
1367 .ok_or_else(|| "web_search arguments must be an object".to_string())?;
1368 for key in obj.keys() {
1369 if !matches!(key.as_str(), "query" | "max_results" | "queries") {
1370 return Err(format!("web_search: unknown argument '{key}'"));
1371 }
1372 }
1373 if obj.contains_key("query") && obj.contains_key("queries") {
1374 return Err("web_search accepts either 'query' or 'queries', not both".to_string());
1375 }
1376
1377 if let Some(value) = obj.get("queries") {
1378 let Some(arr) = value.as_array() else {
1379 return Err("web_search: 'queries' must be an array".to_string());
1380 };
1381 if arr.is_empty() {
1382 return Err("web_search: 'queries' must contain at least one entry".to_string());
1383 }
1384 if arr.len() > mermaid_model::constants::MAX_BATCH_TOOL_ITEMS {
1385 return Err(format!(
1386 "web_search: too many queries ({}); cap is {} per call — split the request",
1387 arr.len(),
1388 mermaid_model::constants::MAX_BATCH_TOOL_ITEMS
1389 ));
1390 }
1391 let mut out = Vec::with_capacity(arr.len());
1392 for v in arr {
1393 let Some(obj) = v.as_object() else {
1394 return Err(
1395 "web_search: 'queries' must be an array of {query, max_results}".to_string(),
1396 );
1397 };
1398 for key in obj.keys() {
1399 if !matches!(key.as_str(), "query" | "max_results") {
1400 return Err(format!("web_search: unknown query argument '{key}'"));
1401 }
1402 }
1403 out.push(parse_query_entry(obj)?);
1404 }
1405 return Ok(out);
1406 }
1407 if obj.contains_key("query") {
1408 return Ok(vec![parse_query_entry(obj)?]);
1409 }
1410 Err("web_search requires 'query' (string) or 'queries' (array)".to_string())
1411}
1412
1413fn parse_query_entry(
1414 obj: &serde_json::Map<String, serde_json::Value>,
1415) -> Result<(String, usize), String> {
1416 let query = obj
1417 .get("query")
1418 .and_then(|value| value.as_str())
1419 .ok_or_else(|| "web_search: each query needs 'query' (string)".to_string())?
1420 .trim();
1421 if query.is_empty() {
1422 return Err("web_search: query must not be empty".to_string());
1423 }
1424 if query.contains(['\r', '\n', '\0']) {
1425 return Err("web_search: query must be a single text line".to_string());
1426 }
1427 if query.chars().count() > 2048 {
1428 return Err("web_search: query exceeds 2048 characters".to_string());
1429 }
1430 let count = match obj.get("max_results") {
1431 None => 5,
1432 Some(value) => {
1433 let count = value.as_u64().ok_or_else(|| {
1434 "web_search: 'max_results' must be an integer from 1 to 10".to_string()
1435 })?;
1436 if !(1..=10).contains(&count) {
1437 return Err("web_search: 'max_results' must be from 1 to 10".to_string());
1438 }
1439 count as usize
1440 },
1441 };
1442 Ok((query.to_string(), count))
1443}
1444
1445pub(crate) fn require_http_scheme(url: &str) -> Result<reqwest::Url, String> {
1457 let parsed = reqwest::Url::parse(url).map_err(|e| format!("invalid URL: {e}"))?;
1458 match parsed.scheme() {
1459 "http" | "https" => Ok(parsed),
1460 other => Err(format!(
1461 "unsupported URL scheme '{other}' (only http/https allowed)"
1462 )),
1463 }
1464}
1465
1466#[cfg(test)]
1467mod tests {
1468 use super::*;
1469
1470 fn page(content: impl Into<String>) -> WebFetchResult {
1471 let content = content.into();
1472 WebFetchResult {
1473 requested_url: "https://example.com/start".to_string(),
1474 final_url: Some("https://example.com/final".to_string()),
1475 status: Some(200),
1476 media_type: Some("text/html".to_string()),
1477 charset: Some("utf-8".to_string()),
1478 backend: super::super::web_client::FetchBackend::Native,
1479 extraction: super::super::web_client::ExtractionMode::Readability,
1480 source_bytes: content.len(),
1481 output_bytes: content.len(),
1482 truncated: false,
1483 title: "T".to_string(),
1484 content,
1485 }
1486 }
1487
1488 fn scope(session_id: &str) -> FetchSnapshotScope {
1489 FetchSnapshotScope {
1490 session_id: Some(compact_string(session_id)),
1491 task_id: None,
1492 fallback_turn: None,
1493 }
1494 }
1495
1496 #[test]
1497 fn require_http_scheme_accepts_http_rejects_exotic() {
1498 for good in [
1501 "http://example.com",
1502 "https://example.com/path?a=1&b=2",
1503 "http://localhost:3000",
1504 "http://127.0.0.1:8080",
1505 ] {
1506 assert!(require_http_scheme(good).is_ok(), "{good} should pass");
1507 }
1508 for bad in [
1510 "file:///etc/passwd",
1511 "javascript:alert(1)",
1512 "data:text/html,<script>",
1513 "ftp://example.com",
1514 "not a url",
1515 ] {
1516 assert!(
1517 require_http_scheme(bad).is_err(),
1518 "{bad} should be rejected"
1519 );
1520 }
1521 }
1522
1523 #[test]
1524 fn format_fetch_caps_long_content() {
1525 let big = "z".repeat(WEB_FETCH_MAX_BYTES * 2);
1527 let big_page = page(big);
1528 let out = format_fetch(&big_page, "web-1", None, 2, None, 200);
1529 assert!(
1530 out.output.len() <= WEB_FETCH_MAX_BYTES,
1531 "content must be capped, got {} bytes",
1532 out.output.len()
1533 );
1534 assert!(
1535 out.output.contains("truncated"),
1536 "expected truncation marker"
1537 );
1538 assert!(out.truncated);
1539
1540 let small = page("hello world");
1542 let out = format_fetch(&small, "web-1", None, 2, None, 200);
1543 assert!(out.output.contains("hello world"));
1544 assert!(!out.output.contains("truncated"));
1545 }
1546
1547 #[test]
1548 fn format_fetch_caps_the_complete_envelope_and_sanitizes_provenance() {
1549 let mut page = page("body");
1550 page.title = format!(" {}\n{} ", "title ".repeat(100), "tail");
1551 page.requested_url = format!(
1552 "https://alice:hunter2@example.com/page?token=opaque-secret&q={}",
1553 "x".repeat(10_000)
1554 );
1555 page.final_url = Some(page.requested_url.clone());
1556
1557 let out = format_fetch(&page, "web-1", None, 2, None, 200);
1558 assert!(out.output.len() <= WEB_FETCH_MAX_BYTES);
1559 assert!(
1560 !out.output.contains("alice"),
1561 "userinfo leaked: {}",
1562 out.output
1563 );
1564 assert!(
1565 !out.output.contains("hunter2"),
1566 "password leaked: {}",
1567 out.output
1568 );
1569 assert!(!out.output.contains("opaque-secret"), "query secret leaked");
1570 let title = out
1571 .output
1572 .lines()
1573 .find_map(|line| line.strip_prefix("Title: "))
1574 .expect("title header");
1575 assert!(title.len() <= 300);
1576 }
1577
1578 #[test]
1579 fn complete_output_budget_holds_for_multibyte_boundary_sizes() {
1580 for unit in ["a", "é", "界"] {
1581 for units in [0, 1, 14_900, 15_000, 15_100, 40_000] {
1582 let mut candidate = page(unit.repeat(units));
1583 candidate.title = unit.repeat(1_000);
1584 let formatted = format_fetch(&candidate, "web-99", None, 2, None, 200);
1585 assert!(
1586 formatted.output.len() <= WEB_FETCH_MAX_BYTES,
1587 "{} bytes escaped the complete-result cap",
1588 formatted.output.len()
1589 );
1590 assert!(std::str::from_utf8(formatted.output.as_bytes()).is_ok());
1591 assert!(formatted.output.ends_with("[/WEB_FETCH]"));
1592 }
1593 }
1594 }
1595
1596 #[test]
1597 fn snapshot_store_accounts_for_and_bounds_every_retained_string() {
1598 let original_content = "é".repeat(MAX_FETCH_SNAPSHOT_BYTES / 2 + 1_000);
1599 let original_output_bytes = original_content.len();
1600 let mut oversized = page(original_content);
1601 oversized.output_bytes = original_output_bytes;
1602 oversized.title = "title ".repeat(10_000);
1603 oversized.requested_url = format!("https://example.com/{}", "r".repeat(20_000));
1604 oversized.final_url = Some(format!("https://example.com/{}", "f".repeat(20_000)));
1605 oversized.media_type = Some("m".repeat(1_000));
1606 oversized.charset = Some("c".repeat(1_000));
1607
1608 let owner = scope("session-retained-size");
1609 let mut store = FetchSnapshotStore::default();
1610 let (id, bounded) = store.insert(owner.clone(), oversized).unwrap();
1611 let entry = store.entries.back().expect("snapshot entry");
1612 let expected = entry
1613 .id
1614 .capacity()
1615 .saturating_add(entry.scope.retained_string_bytes())
1616 .saturating_add(page_retained_string_bytes(&entry.page));
1617
1618 assert_eq!(entry.id, id);
1619 assert_eq!(entry.retained_bytes, expected);
1620 assert_eq!(store.bytes, expected);
1621 assert!(store.bytes <= MAX_FETCH_SNAPSHOT_BYTES);
1622 assert_eq!(bounded.output_bytes, original_output_bytes);
1623 assert!(bounded.truncated);
1624 assert!(bounded.title.len() <= MAX_SNAPSHOT_TITLE_BYTES);
1625 assert!(bounded.requested_url.len() <= MAX_SNAPSHOT_URL_BYTES);
1626 assert!(bounded.final_url.as_ref().unwrap().len() <= MAX_SNAPSHOT_URL_BYTES);
1627 assert!(bounded.media_type.as_ref().unwrap().len() <= MAX_SNAPSHOT_MEDIA_TYPE_BYTES);
1628 assert!(bounded.charset.as_ref().unwrap().len() <= MAX_SNAPSHOT_CHARSET_BYTES);
1629 assert!(std::str::from_utf8(bounded.content.as_bytes()).is_ok());
1630 }
1631
1632 #[test]
1633 fn snapshot_store_isolates_session_and_task_owners() {
1634 let owner = FetchSnapshotScope {
1635 session_id: Some(compact_string("session-a")),
1636 task_id: Some(compact_string("task-a")),
1637 fallback_turn: None,
1638 };
1639 let mut store = FetchSnapshotStore::default();
1640 let (id, _) = store.insert(owner.clone(), page("private page")).unwrap();
1641
1642 assert!(store.get(&owner, &id).is_some());
1643 for outsider in [
1644 FetchSnapshotScope {
1645 session_id: Some(compact_string("session-b")),
1646 task_id: Some(compact_string("task-a")),
1647 fallback_turn: None,
1648 },
1649 FetchSnapshotScope {
1650 session_id: Some(compact_string("session-a")),
1651 task_id: Some(compact_string("task-b")),
1652 fallback_turn: None,
1653 },
1654 ] {
1655 assert!(store.get(&outsider, &id).is_none());
1656 }
1657 }
1658
1659 #[test]
1660 fn snapshot_store_evicts_the_oldest_entry_at_the_count_limit() {
1661 let mut store = FetchSnapshotStore::default();
1662 let owner = scope("session-eviction");
1663 let mut ids = Vec::new();
1664 for index in 0..=MAX_FETCH_SNAPSHOTS {
1665 ids.push(
1666 store
1667 .insert(owner.clone(), page(format!("page {index}")))
1668 .unwrap()
1669 .0,
1670 );
1671 }
1672 assert!(
1673 store.get(&owner, &ids[0]).is_none(),
1674 "oldest snapshot was not evicted"
1675 );
1676 assert!(store.get(&owner, ids.last().unwrap()).is_some());
1677 assert_eq!(store.entries.len(), MAX_FETCH_SNAPSHOTS);
1678 }
1679
1680 #[test]
1681 fn extract_matches_finds_case_insensitive_with_context() {
1682 let content = "line one\nline two\nTARGET here\nline four\nline five";
1683 let (out, count) = extract_matches(content, "target", 1, 20).unwrap();
1684 assert_eq!(count, 1);
1685 assert!(out.starts_with("1 match for \"target\":"));
1686 assert!(out.contains("L2: line two"));
1687 assert!(out.contains("L3: TARGET here"));
1688 assert!(out.contains("L4: line four"));
1689 assert!(!out.contains("L1:"), "context clipped to 1 line: {out}");
1690 assert!(!out.contains("L5:"));
1691 }
1692
1693 #[test]
1694 fn extract_matches_merges_overlapping_windows() {
1695 let content = "a\nhit one\nhit two\nb\nc\nd\ne\nf\ng\nhit three\nz";
1697 let (out, count) = extract_matches(content, "hit", 1, 20).unwrap();
1698 assert_eq!(count, 3);
1699 assert!(out.starts_with("3 matches"));
1700 assert_eq!(out.matches("---").count(), 1, "two blocks: {out}");
1701 assert_eq!(out.matches("hit one").count(), 1);
1703 }
1704
1705 #[test]
1706 fn extract_matches_caps_blocks_and_reports_tail() {
1707 let content = (0..25)
1709 .map(|i| format!("match {i}\nx\nx\nx\nx\nx"))
1710 .collect::<Vec<_>>()
1711 .join("\n");
1712 let (out, count) = extract_matches(&content, "match", 0, 20).unwrap();
1713 assert_eq!(count, 25);
1714 assert!(out.starts_with("25 matches"));
1715 assert_eq!(out.matches("---").count(), 19, "20 blocks: {out}");
1716 assert!(out.contains("(+5 more matches)"), "tail note: {out}");
1717 }
1718
1719 #[test]
1720 fn extract_matches_none_and_multibyte() {
1721 assert!(extract_matches("nothing here", "absent", 2, 20).is_none());
1722 let content = "voil\u{e0} un r\u{e9}sultat\nplain line";
1724 let (out, count) = extract_matches(content, "R\u{c9}SULTAT", 0, 20).unwrap();
1725 assert_eq!(count, 1);
1726 assert!(out.contains("L1: voil\u{e0} un r\u{e9}sultat"));
1727 assert!(!out.contains("plain line"));
1729 }
1730
1731 #[test]
1732 fn extract_matches_uses_full_unicode_case_folding() {
1733 let content = "Die Straße ist lang\nSTRASSE in capitals\nother";
1734 let (out, count) = extract_matches(content, "strasse", 0, 20).unwrap();
1735 assert!(out.starts_with("2 matches for \"strasse\":"), "{out}");
1736 assert!(out.contains("L1: Die Straße ist lang"), "{out}");
1737 assert!(out.contains("L2: STRASSE in capitals"), "{out}");
1738 assert_eq!(count, 2);
1739 }
1740
1741 #[test]
1742 fn extract_matches_normalizes_composed_and_decomposed_text() {
1743 let content = "Café noir\nCafe\u{301} blanc\nplain";
1744 let decomposed_pattern = "CAFE\u{301}";
1745 let (out, count) = extract_matches(content, decomposed_pattern, 0, 20).unwrap();
1746 assert!(out.starts_with("2 matches"), "{out}");
1747 assert!(out.contains("L1: Café noir"), "{out}");
1748 assert!(out.contains("L2: Cafe\u{301} blanc"), "{out}");
1749 assert_eq!(count, 2);
1750 }
1751
1752 #[test]
1753 fn format_fetch_pattern_paths() {
1754 let page = page("alpha\nbeta\ngamma");
1755 let out = format_fetch(&page, "web-1", Some("beta"), 1, None, 200);
1757 assert!(out.output.contains("1 match for \"beta\""));
1758 assert!(out.output.contains("L2: beta"));
1759 let out = format_fetch(&page, "web-1", Some("nope"), 1, None, 200);
1761 assert!(out.output.contains("No matching lines for \"nope\"."));
1762 assert!(!out.output.contains("alpha"));
1763 }
1764
1765 #[test]
1766 fn find_in_page_runs_before_the_cap() {
1767 let mut content = "x\n".repeat(WEB_FETCH_MAX_BYTES / 2);
1770 content.push_str("needle in the tail\n");
1771 let page = page(content);
1772 let out = format_fetch(&page, "web-1", Some("needle"), 1, None, 200);
1773 assert!(
1774 out.output.contains("1 match for \"needle\""),
1775 "tail match found"
1776 );
1777 assert!(out.output.contains("needle in the tail"));
1778 }
1779
1780 #[test]
1781 fn parse_queries_single_form() {
1782 let args = serde_json::json!({"query": "rust async", "max_results": 3});
1783 let q = parse_queries(&args).unwrap();
1784 assert_eq!(q.len(), 1);
1785 assert_eq!(q[0].0, "rust async");
1786 assert_eq!(q[0].1, 3);
1787 }
1788
1789 #[test]
1790 fn parse_queries_array_form() {
1791 let args = serde_json::json!({"queries": [
1792 {"query": "a", "max_results": 2},
1793 {"query": "b", "max_results": 5},
1794 ]});
1795 let q = parse_queries(&args).unwrap();
1796 assert_eq!(q.len(), 2);
1797 assert_eq!(q[1].1, 5);
1798 }
1799
1800 #[test]
1801 fn parse_queries_missing_errors() {
1802 let args = serde_json::json!({});
1803 assert!(parse_queries(&args).is_err());
1804 }
1805
1806 #[test]
1807 fn parse_queries_rejects_out_of_range_count() {
1808 let args = serde_json::json!({"query": "q", "max_results": 999});
1809 assert!(parse_queries(&args).is_err());
1810 let args = serde_json::json!({"query": "q", "max_results": 0});
1811 assert!(parse_queries(&args).is_err());
1812 let args = serde_json::json!({"query": "q", "max_results": "5"});
1813 assert!(parse_queries(&args).is_err());
1814 }
1815
1816 #[test]
1817 fn parse_queries_rejects_ambiguous_and_unknown_arguments() {
1818 assert!(
1819 parse_queries(&serde_json::json!({"query":"a", "queries":[{"query":"b"}]})).is_err()
1820 );
1821 assert!(parse_queries(&serde_json::json!({"query":"a", "extra":true})).is_err());
1822 assert!(parse_queries(&serde_json::json!({"query":" "})).is_err());
1823 assert!(parse_queries(&serde_json::json!({"query":"safe\n=== injected ==="})).is_err());
1824 }
1825
1826 #[test]
1827 fn parse_queries_rejects_excess_fan_out() {
1828 let many: Vec<_> = (0..mermaid_model::constants::MAX_BATCH_TOOL_ITEMS + 1)
1830 .map(|i| serde_json::json!({"query": format!("q{i}")}))
1831 .collect();
1832 let args = serde_json::json!({ "queries": many });
1833 assert!(parse_queries(&args).is_err());
1834
1835 let at_cap: Vec<_> = (0..mermaid_model::constants::MAX_BATCH_TOOL_ITEMS)
1837 .map(|i| serde_json::json!({"query": format!("q{i}")}))
1838 .collect();
1839 let args = serde_json::json!({ "queries": at_cap });
1840 assert_eq!(
1841 parse_queries(&args).unwrap().len(),
1842 mermaid_model::constants::MAX_BATCH_TOOL_ITEMS
1843 );
1844 }
1845
1846 #[test]
1847 fn parse_fetch_args_is_strict_and_rejects_credentialed_urls() {
1848 for invalid in [
1849 serde_json::json!({}),
1850 serde_json::json!({"url": "https://example.com", "snapshot_id": "web-1"}),
1851 serde_json::json!({"url": "https://user:password@example.com"}),
1852 serde_json::json!({"url": "http://127.0.0.1/private"}),
1853 serde_json::json!({"snapshot_id": "bad"}),
1854 serde_json::json!({"snapshot_id": "web-1", "context_lines": 2}),
1855 serde_json::json!({"snapshot_id": "web-1", "pattern": "x", "start_line": 1}),
1856 serde_json::json!({"snapshot_id": "web-1", "unknown": true}),
1857 ] {
1858 assert!(parse_fetch_args(&invalid).is_err(), "accepted {invalid}");
1859 }
1860
1861 let parsed = parse_fetch_args(&serde_json::json!({
1862 "url": "https://example.com/page#fragment",
1863 "start_line": 4,
1864 "line_count": 2
1865 }))
1866 .unwrap();
1867 let FetchTarget::Url(url) = &parsed.target else {
1868 panic!("expected a URL target");
1869 };
1870 assert_eq!(url.as_str(), "https://example.com/page");
1871 assert_eq!(parsed.start_line, Some(4));
1872 assert_eq!(parsed.line_count, 2);
1873 }
1874
1875 #[tokio::test]
1876 async fn web_fetch_failure_retains_typed_backend_provenance() {
1877 use crate::providers::ctx::test_exec_context;
1878 use mermaid_domain::{ToolCallId, ToolStatus, TurnId};
1879
1880 struct FailingFetch;
1881
1882 #[async_trait]
1883 impl FetchProvider for FailingFetch {
1884 async fn fetch(
1885 &self,
1886 url: &str,
1887 _budget: crate::providers::ctx::WebByteBudget,
1888 ) -> Result<WebFetchResult, WebFetchError> {
1889 Err(WebFetchError::HttpStatus {
1890 status: 503,
1891 url: url.to_string(),
1892 })
1893 }
1894 }
1895
1896 let tool = WebFetchTool::new_with_test_snapshots(Arc::new(FailingFetch), "mock");
1897 let (ctx, _rx) =
1898 test_exec_context(TurnId(9), ToolCallId(9), std::path::PathBuf::from("/tmp"));
1899 let outcome = tool
1900 .execute(serde_json::json!({"url": "https://example.com/page"}), ctx)
1901 .await;
1902
1903 assert_eq!(outcome.status, ToolStatus::Error);
1904 match &outcome.metadata.detail {
1905 ToolMetadata::WebFetch {
1906 status,
1907 error_kind,
1908 backend,
1909 final_url,
1910 ..
1911 } => {
1912 assert_eq!(*status, Some(503));
1913 assert_eq!(error_kind.as_deref(), Some("http_status"));
1914 assert_eq!(backend, "mock");
1915 assert!(final_url.is_none());
1916 },
1917 other => panic!("expected web metadata, got {other:?}"),
1918 }
1919 }
1920
1921 #[tokio::test]
1922 async fn web_search_progress_redacts_without_changing_transport_query() {
1923 use crate::providers::ctx::test_exec_context;
1924 use mermaid_domain::{ToolCallId, ToolStatus, TurnId};
1925
1926 struct RecordingSearch {
1927 seen: Arc<Mutex<Option<String>>>,
1928 }
1929
1930 #[async_trait]
1931 impl SearchProvider for RecordingSearch {
1932 async fn search(
1933 &self,
1934 query: &str,
1935 _count: usize,
1936 _budget: crate::providers::ctx::WebByteBudget,
1937 ) -> anyhow::Result<Vec<crate::providers::tool::web_client::SearchResult>> {
1938 *self
1939 .seen
1940 .lock()
1941 .unwrap_or_else(std::sync::PoisonError::into_inner) = Some(query.to_string());
1942 Ok(Vec::new())
1943 }
1944 }
1945
1946 let seen = Arc::new(Mutex::new(None));
1947 let tool = WebSearchTool {
1948 backend: Arc::new(RecordingSearch { seen: seen.clone() }),
1949 backend_name: "mock",
1950 };
1951 let (ctx, mut progress) =
1952 test_exec_context(TurnId(91), ToolCallId(91), std::path::PathBuf::from("/tmp"));
1953 let query = "research OPENAI_API_KEY=abc";
1954 let outcome = tool.execute(serde_json::json!({"query": query}), ctx).await;
1955 assert_eq!(outcome.status, ToolStatus::Success);
1956 assert_eq!(
1957 seen.lock()
1958 .unwrap_or_else(std::sync::PoisonError::into_inner)
1959 .as_deref(),
1960 Some(query),
1961 "redaction must not alter the transport query"
1962 );
1963 let ProgressEvent::Status(status) = progress.recv().await.expect("search progress") else {
1964 panic!("expected search status progress");
1965 };
1966 assert!(
1967 !status.contains("abc"),
1968 "progress leaked query secret: {status}"
1969 );
1970 assert!(status.contains("OPENAI_API_KEY=[REDACTED]"));
1971 }
1972
1973 #[tokio::test]
1974 async fn snapshot_line_ranges_do_not_refetch_mutable_pages() {
1975 use crate::providers::ctx::test_exec_context;
1976 use mermaid_domain::{ToolCallId, ToolStatus, TurnId};
1977 use std::sync::atomic::{AtomicUsize, Ordering};
1978
1979 struct MockFetch {
1980 calls: Arc<AtomicUsize>,
1981 }
1982
1983 #[async_trait]
1984 impl FetchProvider for MockFetch {
1985 async fn fetch(
1986 &self,
1987 url: &str,
1988 _budget: crate::providers::ctx::WebByteBudget,
1989 ) -> Result<WebFetchResult, WebFetchError> {
1990 self.calls.fetch_add(1, Ordering::SeqCst);
1991 let mut result = page("one\ntwo\nthree");
1992 result.requested_url = url.to_string();
1993 result.final_url = Some(url.to_string());
1994 Ok(result)
1995 }
1996 }
1997
1998 let calls = Arc::new(AtomicUsize::new(0));
1999 let tool = WebFetchTool::new_with_test_snapshots(
2000 Arc::new(MockFetch {
2001 calls: calls.clone(),
2002 }),
2003 "mock",
2004 );
2005 let (mut ctx, _rx) =
2006 test_exec_context(TurnId(10), ToolCallId(10), std::path::PathBuf::from("/tmp"));
2007 ctx.session_id = Some("session-a".to_string());
2008 let first = tool
2009 .execute(serde_json::json!({"url": "https://example.com/page"}), ctx)
2010 .await;
2011 assert_eq!(first.status, ToolStatus::Success);
2012 let snapshot_id = match &first.metadata.detail {
2013 ToolMetadata::WebFetch { snapshot_id, .. } => snapshot_id.clone().expect("snapshot id"),
2014 other => panic!("expected web metadata, got {other:?}"),
2015 };
2016
2017 let (mut foreign_ctx, _rx) =
2018 test_exec_context(TurnId(11), ToolCallId(11), std::path::PathBuf::from("/tmp"));
2019 foreign_ctx.session_id = Some("session-b".to_string());
2020 let foreign = tool
2021 .execute(
2022 serde_json::json!({
2023 "snapshot_id": snapshot_id.clone(),
2024 "start_line": 2,
2025 "line_count": 1
2026 }),
2027 foreign_ctx,
2028 )
2029 .await;
2030 assert_eq!(foreign.status, ToolStatus::Error);
2031 assert!(foreign.output().contains("unavailable or was evicted"));
2032
2033 let (mut ctx, _rx) =
2034 test_exec_context(TurnId(12), ToolCallId(12), std::path::PathBuf::from("/tmp"));
2035 ctx.session_id = Some("session-a".to_string());
2036 let continuation = tool
2037 .execute(
2038 serde_json::json!({
2039 "snapshot_id": snapshot_id,
2040 "start_line": 2,
2041 "line_count": 1
2042 }),
2043 ctx,
2044 )
2045 .await;
2046 assert_eq!(continuation.status, ToolStatus::Success);
2047 assert!(continuation.output().contains("L2: two"));
2048 assert!(!continuation.output().contains("L1: one"));
2049 assert_eq!(
2050 calls.load(Ordering::SeqCst),
2051 1,
2052 "snapshot triggered a refetch"
2053 );
2054 }
2055
2056 #[tokio::test(flavor = "multi_thread", worker_threads = 4)]
2057 async fn snapshot_blocking_work_respects_the_global_extractor_limit() {
2058 use std::sync::atomic::{AtomicUsize, Ordering};
2059
2060 let active = Arc::new(AtomicUsize::new(0));
2061 let peak = Arc::new(AtomicUsize::new(0));
2062 let mut jobs = Vec::new();
2063 for _ in 0..8 {
2064 let active = Arc::clone(&active);
2065 let peak = Arc::clone(&peak);
2066 jobs.push(tokio::spawn(async move {
2067 run_snapshot_blocking(move || {
2068 let now = active.fetch_add(1, Ordering::SeqCst) + 1;
2069 peak.fetch_max(now, Ordering::SeqCst);
2070 std::thread::sleep(std::time::Duration::from_millis(20));
2071 active.fetch_sub(1, Ordering::SeqCst);
2072 })
2073 .await
2074 .unwrap();
2075 }));
2076 }
2077 for job in jobs {
2078 job.await.unwrap();
2079 }
2080
2081 assert_eq!(active.load(Ordering::SeqCst), 0);
2082 assert!(
2083 peak.load(Ordering::SeqCst) <= mermaid_model::constants::MAX_WEB_EXTRACTION_CONCURRENCY
2084 );
2085 }
2086
2087 #[tokio::test(flavor = "multi_thread", worker_threads = 2)]
2088 async fn cancelled_snapshot_waiter_keeps_its_permit_until_blocking_work_finishes() {
2089 let limiter = Arc::new(tokio::sync::Semaphore::new(1));
2090 let (started_tx, started_rx) = tokio::sync::oneshot::channel();
2091 let (release_tx, release_rx) = std::sync::mpsc::sync_channel(0);
2092 let worker_limiter = limiter.clone();
2093 let worker = tokio::spawn(async move {
2094 run_snapshot_blocking_with(worker_limiter, move || {
2095 let _ = started_tx.send(());
2096 release_rx.recv().expect("test releases blocking worker");
2097 })
2098 .await
2099 });
2100 started_rx.await.expect("blocking worker started");
2101 worker.abort();
2102 let _ = worker.await;
2103 assert_eq!(
2104 limiter.available_permits(),
2105 0,
2106 "cancelling the async waiter released a still-running blocking job"
2107 );
2108
2109 release_tx.send(()).expect("release blocking worker");
2110 tokio::time::timeout(std::time::Duration::from_secs(1), async {
2111 while limiter.available_permits() == 0 {
2112 tokio::task::yield_now().await;
2113 }
2114 })
2115 .await
2116 .expect("blocking worker did not release its permit");
2117 }
2118
2119 #[tokio::test]
2120 async fn web_search_batch_survives_empty_and_failed_queries() {
2121 use crate::providers::ctx::test_exec_context;
2122 use crate::providers::tool::web_client::SearchResult;
2123 use async_trait::async_trait;
2124 use mermaid_domain::{ToolCallId, ToolStatus, TurnId};
2125 use std::sync::Arc;
2126
2127 struct Mock;
2128 #[async_trait]
2129 impl SearchProvider for Mock {
2130 async fn search(
2131 &self,
2132 query: &str,
2133 _count: usize,
2134 _budget: crate::providers::ctx::WebByteBudget,
2135 ) -> anyhow::Result<Vec<SearchResult>> {
2136 match query {
2137 "boom" => Err(anyhow::anyhow!("backend down")),
2138 "empty" => Ok(Vec::new()),
2139 _ => Ok(vec![SearchResult {
2140 title: "Title".to_string(),
2141 url: "https://example.com".to_string(),
2142 snippet: "snip".to_string(),
2143 full_content: "content".to_string(),
2144 }]),
2145 }
2146 }
2147 }
2148
2149 let mk = || WebSearchTool {
2150 backend: Arc::new(Mock),
2151 backend_name: "mock",
2152 };
2153 let tmp = std::path::PathBuf::from("/tmp");
2154
2155 let (ctx, _rx) = test_exec_context(TurnId(1), ToolCallId(1), tmp.clone());
2157 let out = mk()
2158 .execute(
2159 serde_json::json!({"queries": [{"query":"good"},{"query":"empty"},{"query":"boom"}]}),
2160 ctx,
2161 )
2162 .await;
2163 assert_eq!(
2164 out.status,
2165 ToolStatus::Success,
2166 "a partial batch must not abort"
2167 );
2168 assert!(
2169 out.output().contains("https://example.com"),
2170 "keeps the good result"
2171 );
2172 match &out.metadata.detail {
2173 ToolMetadata::WebSearch {
2174 partial, failures, ..
2175 } => {
2176 assert!(*partial);
2177 assert_eq!(failures.len(), 1);
2178 assert_eq!(failures[0].query_index, 2);
2179 assert!(failures[0].error.contains("backend down"));
2180 },
2181 other => panic!("expected web search metadata, got {other:?}"),
2182 }
2183
2184 let (ctx, _rx) = test_exec_context(TurnId(2), ToolCallId(2), tmp.clone());
2186 let out = mk()
2187 .execute(serde_json::json!({"query": "empty"}), ctx)
2188 .await;
2189 assert_eq!(out.status, ToolStatus::Success, "empty is not an error");
2190 assert!(out.output().contains("no results"));
2191
2192 let (ctx, _rx) = test_exec_context(TurnId(3), ToolCallId(3), tmp);
2194 let out = mk()
2195 .execute(
2196 serde_json::json!({"queries": [{"query":"boom"},{"query":"boom"}]}),
2197 ctx,
2198 )
2199 .await;
2200 assert_eq!(out.status, ToolStatus::Error, "total failure is an error");
2201 match &out.metadata.detail {
2202 ToolMetadata::WebSearch {
2203 failed_queries,
2204 failures,
2205 ..
2206 } => {
2207 assert_eq!(*failed_queries, 2);
2208 assert_eq!(failures.len(), 2);
2209 },
2210 other => panic!("expected web search metadata, got {other:?}"),
2211 }
2212 }
2213
2214 #[tokio::test]
2215 async fn web_search_batch_caps_concurrency_and_restores_input_order() {
2216 use crate::providers::ctx::test_exec_context;
2217 use crate::providers::tool::web_client::SearchResult;
2218 use mermaid_domain::{ToolCallId, ToolStatus, TurnId};
2219 use std::sync::atomic::{AtomicUsize, Ordering};
2220
2221 struct ConcurrencyMock {
2222 active: Arc<AtomicUsize>,
2223 peak: Arc<AtomicUsize>,
2224 }
2225
2226 #[async_trait]
2227 impl SearchProvider for ConcurrencyMock {
2228 async fn search(
2229 &self,
2230 query: &str,
2231 _count: usize,
2232 _budget: crate::providers::ctx::WebByteBudget,
2233 ) -> anyhow::Result<Vec<SearchResult>> {
2234 let active = self.active.fetch_add(1, Ordering::SeqCst) + 1;
2235 self.peak.fetch_max(active, Ordering::SeqCst);
2236 let delay = 5 + (6 - query.parse::<u64>().unwrap_or(0)) * 5;
2237 tokio::time::sleep(std::time::Duration::from_millis(delay)).await;
2238 self.active.fetch_sub(1, Ordering::SeqCst);
2239 Ok(vec![SearchResult {
2240 title: format!("result {query}"),
2241 url: format!("https://example.com/{query}"),
2242 snippet: String::new(),
2243 full_content: format!("content {query}"),
2244 }])
2245 }
2246 }
2247
2248 let active = Arc::new(AtomicUsize::new(0));
2249 let peak = Arc::new(AtomicUsize::new(0));
2250 let tool = WebSearchTool {
2251 backend: Arc::new(ConcurrencyMock {
2252 active: active.clone(),
2253 peak: peak.clone(),
2254 }),
2255 backend_name: "mock",
2256 };
2257 let (ctx, _rx) =
2258 test_exec_context(TurnId(20), ToolCallId(20), std::path::PathBuf::from("/tmp"));
2259 let queries: Vec<_> = (0..6)
2260 .map(|index| serde_json::json!({"query": index.to_string()}))
2261 .collect();
2262 let outcome = tool
2263 .execute(serde_json::json!({"queries": queries}), ctx)
2264 .await;
2265 assert_eq!(outcome.status, ToolStatus::Success);
2266 assert_eq!(active.load(Ordering::SeqCst), 0);
2267 assert_eq!(
2268 peak.load(Ordering::SeqCst),
2269 mermaid_model::constants::MAX_WEB_SEARCH_CONCURRENCY
2270 );
2271 let mut cursor = 0;
2272 for index in 0..6 {
2273 let marker = format!("=== query: {index} ===");
2274 let position = outcome.output()[cursor..]
2275 .find(&marker)
2276 .map(|offset| cursor + offset)
2277 .expect("ordered query section");
2278 assert!(position >= cursor);
2279 cursor = position + marker.len();
2280 }
2281 }
2282
2283 #[tokio::test]
2284 async fn web_search_complete_output_budget_is_byte_exact_for_multibyte_text() {
2285 use crate::providers::ctx::test_exec_context;
2286 use crate::providers::tool::web_client::SearchResult;
2287 use mermaid_domain::{ToolCallId, ToolStatus, TurnId};
2288
2289 struct MultibyteMock;
2290
2291 #[async_trait]
2292 impl SearchProvider for MultibyteMock {
2293 async fn search(
2294 &self,
2295 _query: &str,
2296 _count: usize,
2297 _budget: crate::providers::ctx::WebByteBudget,
2298 ) -> anyhow::Result<Vec<SearchResult>> {
2299 Ok(vec![SearchResult {
2300 title: "界".repeat(5_000),
2301 url: "https://example.com/result".to_string(),
2302 snippet: String::new(),
2303 full_content: "界".repeat(20_000),
2304 }])
2305 }
2306 }
2307
2308 let tool = WebSearchTool {
2309 backend: Arc::new(MultibyteMock),
2310 backend_name: "mock",
2311 };
2312 let (ctx, _rx) =
2313 test_exec_context(TurnId(21), ToolCallId(21), std::path::PathBuf::from("/tmp"));
2314 let outcome = tool
2315 .execute(serde_json::json!({"query": "multibyte"}), ctx)
2316 .await;
2317
2318 assert_eq!(outcome.status, ToolStatus::Success);
2319 assert!(
2320 outcome.output().len() <= mermaid_model::constants::WEB_SEARCH_AGGREGATE_MAX_BYTES,
2321 "{} bytes escaped the complete search-result cap",
2322 outcome.output().len()
2323 );
2324 assert!(std::str::from_utf8(outcome.output().as_bytes()).is_ok());
2325 }
2326
2327 #[tokio::test]
2328 async fn web_search_total_failure_and_structured_errors_are_byte_bounded() {
2329 use crate::providers::ctx::test_exec_context;
2330 use mermaid_domain::{ToolCallId, ToolStatus, TurnId};
2331
2332 struct LargeFailure;
2333
2334 #[async_trait]
2335 impl SearchProvider for LargeFailure {
2336 async fn search(
2337 &self,
2338 _query: &str,
2339 _count: usize,
2340 _budget: crate::providers::ctx::WebByteBudget,
2341 ) -> anyhow::Result<Vec<crate::providers::tool::web_client::SearchResult>> {
2342 Err(anyhow::anyhow!("{}", "界".repeat(20_000)))
2343 }
2344 }
2345
2346 let tool = WebSearchTool {
2347 backend: Arc::new(LargeFailure),
2348 backend_name: "mock",
2349 };
2350 let (ctx, _rx) =
2351 test_exec_context(TurnId(22), ToolCallId(22), std::path::PathBuf::from("/tmp"));
2352 let outcome = tool
2353 .execute(
2354 serde_json::json!({"queries": [{"query": "one"}, {"query": "two"}]}),
2355 ctx,
2356 )
2357 .await;
2358
2359 assert_eq!(outcome.status, ToolStatus::Error);
2360 assert!(
2361 outcome.output().len() <= mermaid_model::constants::WEB_SEARCH_AGGREGATE_MAX_BYTES,
2362 "{} bytes escaped the complete search error cap",
2363 outcome.output().len()
2364 );
2365 let ToolMetadata::WebSearch { failures, .. } = &outcome.metadata.detail else {
2366 panic!("expected web search metadata");
2367 };
2368 assert_eq!(failures.len(), 2);
2369 assert!(
2370 failures
2371 .iter()
2372 .all(|failure| failure.error.len() <= MAX_WEB_SEARCH_FAILURE_BYTES)
2373 );
2374 }
2375}