Skip to main content

systemprompt_agent/repository/context/message/
parts.rs

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            // `data_content` is a jsonb column; cast the bound text payload
278            // so Postgres accepts it without an OID mismatch.
279            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}