Skip to main content

systemprompt_api/services/middleware/context/sources/
payload.rs

1//! Request-context extraction from A2A JSON-RPC bodies.
2//!
3//! Copyright (c) systemprompt.io — Business Source License 1.1.
4//! See <https://systemprompt.io> for licensing details.
5
6use axum::body::Body;
7use axum::extract::Request;
8use serde_json::Value;
9use systemprompt_models::execution::{ContextExtractionError, ContextIdSource};
10
11#[derive(Debug, Clone, Copy)]
12pub struct PayloadSource;
13
14impl PayloadSource {
15    pub fn extract_context_source(
16        body_bytes: &[u8],
17    ) -> Result<ContextIdSource, ContextExtractionError> {
18        // JSON: A2A JSON-RPC envelope is an external protocol boundary; the
19        // method name drives which typed field is read, so the shape is dynamic.
20        let payload: Value = serde_json::from_slice(body_bytes).map_err(|error| {
21            tracing::debug!(%error, "request payload is not valid JSON");
22            ContextExtractionError::InvalidHeaderValue {
23                header: "payload".to_owned(),
24                reason: "body is not valid JSON".to_owned(),
25            }
26        })?;
27
28        let method = payload.get("method").and_then(|m| m.as_str()).unwrap_or("");
29
30        if method.starts_with("tasks/") {
31            let task_id = payload
32                .get("params")
33                .and_then(|p| p.get("id"))
34                .and_then(|id| id.as_str())
35                .map(str::to_owned)
36                .ok_or_else(|| ContextExtractionError::InvalidHeaderValue {
37                    header: "params.id".to_owned(),
38                    reason: "Task ID required for task methods".to_owned(),
39                })?;
40
41            return Ok(ContextIdSource::FromTask {
42                task_id: systemprompt_identifiers::TaskId::new(task_id),
43            });
44        }
45
46        payload
47            .get("params")
48            .and_then(|p| p.get("message"))
49            .and_then(|m| m.get("contextId"))
50            .and_then(|c| c.as_str())
51            .map(|s| ContextIdSource::Direct(s.to_owned()))
52            .ok_or(ContextExtractionError::MissingContextId)
53    }
54
55    pub async fn read_and_reconstruct(
56        request: Request<Body>,
57    ) -> Result<(Vec<u8>, Request<Body>), ContextExtractionError> {
58        let (parts, body) = request.into_parts();
59
60        let body_bytes = axum::body::to_bytes(body, usize::MAX)
61            .await
62            .map_err(|error| {
63                tracing::debug!(%error, "request body could not be read");
64                ContextExtractionError::InvalidHeaderValue {
65                    header: "body".to_owned(),
66                    reason: "request body could not be read".to_owned(),
67                }
68            })?
69            .to_vec();
70
71        let new_body = Body::from(body_bytes.clone());
72        let new_request = Request::from_parts(parts, new_body);
73
74        Ok((body_bytes, new_request))
75    }
76}