1use agent_client_protocol::schema::v2 as acp;
2use base64::Engine;
3use base64::engine::general_purpose::STANDARD as BASE64;
4use std::io::Read;
5use std::path::{Path, PathBuf};
6use url::Url;
7
8#[derive(Debug, Clone, PartialEq, Eq)]
9pub struct PromptAttachment {
10 pub path: PathBuf,
11 pub display_name: String,
12}
13
14#[derive(Debug, Clone, Copy, PartialEq, Eq)]
15pub enum AttachmentKind {
16 Text,
17 Image,
18 Audio,
19 Unsupported,
20}
21
22pub struct AttachmentOutcome {
23 pub blocks: Vec<acp::ContentBlock>,
24 pub warnings: Vec<String>,
25}
26
27pub fn classify_attachment(path: &Path) -> AttachmentKind {
28 let mime = mime_guess::from_path(path).first_or_octet_stream().to_string();
29 if IMAGE_MIME_TYPES.contains(&mime.as_str()) {
30 AttachmentKind::Image
31 } else if AUDIO_MIME_TYPES.contains(&mime.as_str()) {
32 AttachmentKind::Audio
33 } else if mime.starts_with("text/") {
34 AttachmentKind::Text
35 } else {
36 AttachmentKind::Unsupported
37 }
38}
39
40pub fn build_attachments(attachments: &[PromptAttachment]) -> AttachmentOutcome {
41 build_attachments_with(attachments, read_capped)
42}
43
44pub(crate) fn build_attachments_with(
48 attachments: &[PromptAttachment],
49 mut read: impl FnMut(&Path, &str) -> Result<Vec<u8>, String>,
50) -> AttachmentOutcome {
51 let mut outcome = AttachmentOutcome { blocks: Vec::new(), warnings: Vec::new() };
52 for attachment in attachments {
53 let encoded = read(&attachment.path, &attachment.display_name)
54 .and_then(|bytes| encode_attachment(&attachment.path, &attachment.display_name, bytes));
55 match encoded {
56 Ok((block, warning)) => {
57 outcome.blocks.push(block);
58 if let Some(warning) = warning {
59 outcome.warnings.push(warning);
60 }
61 }
62 Err(warning) => outcome.warnings.push(warning),
63 }
64 }
65 outcome
66}
67
68pub(crate) fn read_capped(path: &Path, display_name: &str) -> Result<Vec<u8>, String> {
71 let mut bytes = Vec::new();
72 std::fs::File::open(path)
73 .map_err(|error| format!("Failed to read {display_name}: {error}"))?
74 .take((MAX_MEDIA_BYTES + 1) as u64)
75 .read_to_end(&mut bytes)
76 .map_err(|error| format!("Failed to read {display_name}: {error}"))?;
77 Ok(bytes)
78}
79
80const MAX_EMBED_TEXT_BYTES: usize = 1024 * 1024;
81const MAX_MEDIA_BYTES: usize = 10 * 1024 * 1024;
82const IMAGE_MIME_TYPES: &[&str] = &["image/png", "image/jpeg", "image/gif", "image/webp"];
83const AUDIO_MIME_TYPES: &[&str] = &["audio/wav", "audio/mpeg", "audio/mp3", "audio/ogg"];
84
85pub(crate) fn encode_attachment(
87 path: &Path,
88 display_name: &str,
89 bytes: Vec<u8>,
90) -> Result<(acp::ContentBlock, Option<String>), String> {
91 let mime_type = mime_guess::from_path(path).first_or_octet_stream().to_string();
92 match classify_attachment(path) {
93 AttachmentKind::Image | AttachmentKind::Audio => {
94 encode_media_block(&bytes, display_name, &mime_type).map(|block| (block, None))
95 }
96 AttachmentKind::Text | AttachmentKind::Unsupported => encode_text_block(bytes, path, display_name, &mime_type),
97 }
98}
99
100fn encode_media_block(bytes: &[u8], display_name: &str, mime_type: &str) -> Result<acp::ContentBlock, String> {
101 if bytes.len() > MAX_MEDIA_BYTES {
102 return Err(format!("Skipped {display_name}: file too large (max {MAX_MEDIA_BYTES})"));
103 }
104 let data = BASE64.encode(bytes);
105 Ok(if IMAGE_MIME_TYPES.contains(&mime_type) {
106 acp::ContentBlock::Image(acp::ImageContent::new(data, mime_type))
107 } else {
108 acp::ContentBlock::Audio(acp::AudioContent::new(data, mime_type))
109 })
110}
111
112fn encode_text_block(
113 mut bytes: Vec<u8>,
114 path: &Path,
115 display_name: &str,
116 mime_type: &str,
117) -> Result<(acp::ContentBlock, Option<String>), String> {
118 let truncated = bytes.len() > MAX_EMBED_TEXT_BYTES;
119 if truncated {
120 bytes.truncate(MAX_EMBED_TEXT_BYTES);
121 }
122
123 let text = match std::str::from_utf8(&bytes) {
124 Ok(text) => text.to_string(),
125 Err(error) if truncated && error.error_len().is_none() => {
129 std::str::from_utf8(&bytes[..error.valid_up_to()]).expect("valid_up_to marks a UTF-8 boundary").to_string()
130 }
131 Err(_) => return Err(format!("Skipped binary or non-UTF8 file: {display_name}")),
132 };
133
134 let uri = attachment_uri(path, display_name)?;
135 let warning = truncated.then(|| format!("Truncated {display_name} to {MAX_EMBED_TEXT_BYTES} bytes"));
136 Ok((
137 acp::ContentBlock::Resource(acp::EmbeddedResource::new(acp::EmbeddedResourceResource::TextResourceContents(
138 acp::TextResourceContents::new(text, uri).mime_type(mime_type),
139 ))),
140 warning,
141 ))
142}
143
144fn attachment_uri(path: &Path, display_name: &str) -> Result<String, String> {
145 let uri_path = std::fs::canonicalize(path).unwrap_or_else(|_| path.to_path_buf());
146 Url::from_file_path(uri_path)
147 .map(|url| url.to_string())
148 .map_err(|()| format!("Failed to build file URI for {display_name}"))
149}