1use sqlx::PgPool;
2use std::sync::Arc;
3use systemprompt_identifiers::{ContextId, MessageId, SessionId, TaskId, TraceId, UserId};
4use systemprompt_traits::{DynFileUploadProvider, FileUploadInput, RepositoryError};
5
6use crate::models::a2a::Part;
7
8#[derive(Clone)]
9pub struct FileUploadContext<'a> {
10 pub upload_provider: &'a DynFileUploadProvider,
11 pub context_id: &'a ContextId,
12 pub user_id: Option<&'a UserId>,
13 pub session_id: Option<&'a SessionId>,
14 pub trace_id: Option<&'a TraceId>,
15}
16
17impl std::fmt::Debug for FileUploadContext<'_> {
18 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
19 f.debug_struct("FileUploadContext")
20 .field("upload_provider", &"<DynFileUploadProvider>")
21 .field("context_id", &self.context_id)
22 .field("user_id", &self.user_id)
23 .field("session_id", &self.session_id)
24 .field("trace_id", &self.trace_id)
25 .finish()
26 }
27}
28
29pub async fn get_message_parts(
30 pool: &Arc<PgPool>,
31 message_id: &MessageId,
32) -> Result<Vec<Part>, RepositoryError> {
33 let part_rows: Vec<crate::models::MessagePart> = sqlx::query_as!(
34 crate::models::MessagePart,
35 r#"SELECT
36 id as "id!",
37 message_id as "message_id!: MessageId",
38 task_id as "task_id!: TaskId",
39 part_kind as "part_kind!",
40 sequence_number as "sequence_number!",
41 text_content,
42 file_name,
43 file_mime_type,
44 file_uri,
45 file_bytes,
46 data_content,
47 metadata
48 FROM message_parts WHERE message_id = $1 ORDER BY sequence_number ASC"#,
49 message_id.as_str()
50 )
51 .fetch_all(pool.as_ref())
52 .await
53 .map_err(RepositoryError::database)?;
54
55 let mut parts = Vec::new();
56
57 for row in part_rows {
58 let part = match row.part_kind.as_str() {
59 "text" => {
60 let text = row
61 .text_content
62 .ok_or_else(|| RepositoryError::InvalidData("Missing text_content".into()))?;
63 Part::Text(crate::models::a2a::TextPart { text })
64 },
65 "file" => Part::File(crate::models::a2a::FilePart {
66 file: crate::models::a2a::FileContent {
67 name: row.file_name,
68 mime_type: row.file_mime_type,
69 bytes: row.file_bytes,
70 url: row.file_uri,
71 },
72 }),
73 "data" => {
74 let data_value = row
75 .data_content
76 .ok_or_else(|| RepositoryError::InvalidData("Missing data_content".into()))?;
77 let serde_json::Value::Object(data) = data_value else {
78 return Err(RepositoryError::InvalidData(
79 "Data content must be a JSON object".into(),
80 ));
81 };
82 Part::Data(crate::models::a2a::DataPart { data })
83 },
84 _ => {
85 return Err(RepositoryError::InvalidData(format!(
86 "Unknown part kind: {}",
87 row.part_kind
88 )));
89 },
90 };
91
92 parts.push(part);
93 }
94
95 Ok(parts)
96}
97
98#[expect(
99 missing_debug_implementations,
100 reason = "params struct holds non-Debug references"
101)]
102pub struct PersistPartSqlxParams<'a> {
103 pub tx: &'a mut sqlx::Transaction<'static, sqlx::Postgres>,
104 pub part: &'a Part,
105 pub message_id: &'a MessageId,
106 pub task_id: &'a TaskId,
107 pub sequence_number: i32,
108 pub upload_ctx: Option<&'a FileUploadContext<'a>>,
109}
110
111pub(super) async fn persist_part_sqlx(
112 params: PersistPartSqlxParams<'_>,
113) -> Result<(), RepositoryError> {
114 let PersistPartSqlxParams {
115 tx,
116 part,
117 message_id,
118 task_id,
119 sequence_number,
120 upload_ctx,
121 } = params;
122 match part {
123 Part::Text(text_part) => {
124 sqlx::query!(
125 r#"INSERT INTO message_parts (message_id, task_id, part_kind, sequence_number, text_content)
126 VALUES ($1, $2, 'text', $3, $4)"#,
127 message_id.as_str(),
128 task_id.as_str(),
129 sequence_number,
130 text_part.text
131 )
132 .execute(&mut **tx)
133 .await
134 .map_err(RepositoryError::database)?;
135 },
136 Part::File(file_part) => {
137 let upload_result = try_upload_file(file_part, upload_ctx).await;
138
139 let (file_id, file_uri) = match upload_result {
140 Some((id, uri)) => (Some(id), Some(uri)),
141 None => (None, None),
142 };
143
144 sqlx::query!(
145 r#"INSERT INTO message_parts (message_id, task_id, part_kind, sequence_number, file_name, file_mime_type, file_uri, file_bytes, file_id)
146 VALUES ($1, $2, 'file', $3, $4, $5, $6, $7, $8)"#,
147 message_id.as_str(),
148 task_id.as_str(),
149 sequence_number,
150 file_part.file.name,
151 file_part.file.mime_type,
152 file_uri,
153 file_part.file.bytes.as_deref(),
154 file_id
155 )
156 .execute(&mut **tx)
157 .await
158 .map_err(RepositoryError::database)?;
159 },
160 Part::Data(data_part) => {
161 let data_json =
162 serde_json::to_value(&data_part.data).map_err(RepositoryError::Serialization)?;
163 sqlx::query!(
164 r#"INSERT INTO message_parts (message_id, task_id, part_kind, sequence_number, data_content)
165 VALUES ($1, $2, 'data', $3, $4)"#,
166 message_id.as_str(),
167 task_id.as_str(),
168 sequence_number,
169 data_json
170 )
171 .execute(&mut **tx)
172 .await
173 .map_err(RepositoryError::database)?;
174 },
175 }
176
177 Ok(())
178}
179
180async fn try_upload_file(
181 file_part: &crate::models::a2a::FilePart,
182 upload_ctx: Option<&FileUploadContext<'_>>,
183) -> Option<(uuid::Uuid, String)> {
184 let ctx = upload_ctx?;
185
186 if !ctx.upload_provider.is_enabled() {
187 return None;
188 }
189
190 let mime_type = file_part
191 .file
192 .mime_type
193 .as_deref()
194 .unwrap_or("application/octet-stream");
195
196 let bytes = file_part.file.bytes.as_deref()?;
197 let mut input = FileUploadInput::new(mime_type, bytes, Some(ctx.context_id.clone()));
198
199 if let Some(name) = &file_part.file.name {
200 input = input.with_name(name);
201 }
202
203 if let Some(user_id) = ctx.user_id {
204 input = input.with_user_id(user_id.clone());
205 }
206
207 if let Some(session_id) = ctx.session_id {
208 input = input.with_session_id(session_id.clone());
209 }
210
211 if let Some(trace_id) = ctx.trace_id {
212 input = input.with_trace_id(trace_id.clone());
213 }
214
215 match ctx.upload_provider.upload_file(input).await {
216 Ok(uploaded) => {
217 let file_uuid = uuid::Uuid::parse_str(uploaded.file_id.as_str())
218 .map_err(|e| {
219 tracing::warn!(file_id = %uploaded.file_id, error = %e, "Invalid UUID from file service");
220 e
221 })
222 .ok()?;
223 Some((file_uuid, uploaded.public_url))
224 },
225 Err(e) => {
226 tracing::warn!(error = %e, "File upload failed, continuing with base64 only");
227 None
228 },
229 }
230}
231
232pub(super) async fn persist_part_with_tx(
233 tx: &mut dyn systemprompt_database::DatabaseTransaction,
234 part: &Part,
235 message_id: &MessageId,
236 task_id: &TaskId,
237 sequence_number: i32,
238) -> Result<(), RepositoryError> {
239 let message_id_str = message_id.as_str();
240 let task_id_str = task_id.as_str();
241 match part {
242 Part::Text(text_part) => {
243 let query: &str = "INSERT INTO message_parts (message_id, task_id, part_kind, \
244 sequence_number, text_content) VALUES ($1, $2, 'text', $3, $4)";
245 tx.execute(
246 &query,
247 &[
248 &message_id_str,
249 &task_id_str,
250 &sequence_number,
251 &text_part.text,
252 ],
253 )
254 .await?;
255 },
256 Part::File(file_part) => {
257 let uri_opt: Option<&str> = None;
258 let query: &str = "INSERT INTO message_parts (message_id, task_id, part_kind, \
259 sequence_number, file_name, file_mime_type, file_uri, file_bytes) \
260 VALUES ($1, $2, 'file', $3, $4, $5, $6, $7)";
261 tx.execute(
262 &query,
263 &[
264 &message_id_str,
265 &task_id_str,
266 &sequence_number,
267 &file_part.file.name,
268 &file_part.file.mime_type,
269 &uri_opt,
270 &file_part.file.bytes.as_deref(),
271 ],
272 )
273 .await?;
274 },
275 Part::Data(data_part) => {
276 let data_json = serde_json::to_string(&data_part.data)?;
277 let query: &str = "INSERT INTO message_parts (message_id, task_id, part_kind, \
280 sequence_number, data_content) VALUES ($1, $2, 'data', $3, \
281 $4::jsonb)";
282 tx.execute(
283 &query,
284 &[&message_id_str, &task_id_str, &sequence_number, &data_json],
285 )
286 .await?;
287 },
288 }
289
290 Ok(())
291}