Skip to main content

codex_api/endpoint/
memories.rs

1use crate::auth::SharedAuthProvider;
2use crate::common::MemorySummarizeInput;
3use crate::common::MemorySummarizeOutput;
4use crate::endpoint::session::EndpointSession;
5use crate::error::ApiError;
6use crate::provider::Provider;
7use codex_client::HttpTransport;
8use codex_client::RequestTelemetry;
9use http::HeaderMap;
10use http::Method;
11use serde::Deserialize;
12use serde_json::to_value;
13use std::sync::Arc;
14
15pub struct MemoriesClient<T: HttpTransport> {
16    session: EndpointSession<T>,
17}
18
19impl<T: HttpTransport> MemoriesClient<T> {
20    pub fn new(transport: T, provider: Provider, auth: SharedAuthProvider) -> Self {
21        Self {
22            session: EndpointSession::new(transport, provider, auth),
23        }
24    }
25
26    pub fn with_telemetry(self, request: Option<Arc<dyn RequestTelemetry>>) -> Self {
27        Self {
28            session: self.session.with_request_telemetry(request),
29        }
30    }
31
32    fn path() -> &'static str {
33        "memories/trace_summarize"
34    }
35
36    pub async fn summarize(
37        &self,
38        body: serde_json::Value,
39        extra_headers: HeaderMap,
40    ) -> Result<Vec<MemorySummarizeOutput>, ApiError> {
41        let resp = self
42            .session
43            .execute(Method::POST, Self::path(), extra_headers, Some(body))
44            .await?;
45        let parsed: SummarizeResponse =
46            serde_json::from_slice(&resp.body).map_err(|e| ApiError::Stream(e.to_string()))?;
47        Ok(parsed.output)
48    }
49
50    pub async fn summarize_input(
51        &self,
52        input: &MemorySummarizeInput,
53        extra_headers: HeaderMap,
54    ) -> Result<Vec<MemorySummarizeOutput>, ApiError> {
55        let body = to_value(input).map_err(|e| {
56            ApiError::Stream(format!("failed to encode memory summarize input: {e}"))
57        })?;
58        self.summarize(body, extra_headers).await
59    }
60}
61
62#[derive(Debug, Deserialize)]
63struct SummarizeResponse {
64    output: Vec<MemorySummarizeOutput>,
65}
66
67#[cfg(test)]
68mod tests {
69    use super::*;
70    use crate::auth::AuthProvider;
71    use crate::common::RawMemory;
72    use crate::common::RawMemoryMetadata;
73    use crate::provider::RetryConfig;
74    use codex_client::Request;
75    use codex_client::RequestBody;
76    use codex_client::Response;
77    use codex_client::StreamResponse;
78    use codex_client::TransportError;
79    use http::HeaderMap;
80    use http::Method;
81    use http::StatusCode;
82    use pretty_assertions::assert_eq;
83    use serde_json::json;
84    use std::sync::Arc;
85    use std::sync::Mutex;
86    use std::time::Duration;
87
88    #[derive(Clone, Default)]
89    struct DummyTransport;
90
91    impl HttpTransport for DummyTransport {
92        async fn execute(&self, _req: Request) -> Result<Response, TransportError> {
93            Err(TransportError::Build("execute should not run".to_string()))
94        }
95
96        async fn stream(&self, _req: Request) -> Result<StreamResponse, TransportError> {
97            Err(TransportError::Build("stream should not run".to_string()))
98        }
99    }
100
101    #[derive(Clone, Default)]
102    struct DummyAuth;
103
104    impl AuthProvider for DummyAuth {
105        fn add_auth_headers(&self, _headers: &mut HeaderMap) {}
106    }
107
108    #[derive(Clone)]
109    struct CapturingTransport {
110        last_request: Arc<Mutex<Option<Request>>>,
111        response_body: Arc<Vec<u8>>,
112    }
113
114    impl CapturingTransport {
115        fn new(response_body: Vec<u8>) -> Self {
116            Self {
117                last_request: Arc::new(Mutex::new(None)),
118                response_body: Arc::new(response_body),
119            }
120        }
121    }
122
123    impl HttpTransport for CapturingTransport {
124        async fn execute(&self, req: Request) -> Result<Response, TransportError> {
125            *self.last_request.lock().expect("lock request store") = Some(req);
126            Ok(Response {
127                status: StatusCode::OK,
128                headers: HeaderMap::new(),
129                body: self.response_body.as_ref().clone().into(),
130            })
131        }
132
133        async fn stream(&self, _req: Request) -> Result<StreamResponse, TransportError> {
134            Err(TransportError::Build("stream should not run".to_string()))
135        }
136    }
137
138    fn provider(base_url: &str) -> Provider {
139        Provider {
140            name: "test".to_string(),
141            base_url: base_url.to_string(),
142            query_params: None,
143            headers: HeaderMap::new(),
144            retry: RetryConfig {
145                max_attempts: 1,
146                base_delay: Duration::from_millis(1),
147                retry_429: false,
148                retry_5xx: true,
149                retry_transport: true,
150            },
151            stream_idle_timeout: Duration::from_secs(1),
152        }
153    }
154
155    #[test]
156    fn path_is_memories_trace_summarize_for_wire_compatibility() {
157        assert_eq!(
158            MemoriesClient::<DummyTransport>::path(),
159            "memories/trace_summarize"
160        );
161    }
162
163    #[tokio::test]
164    async fn summarize_input_posts_expected_payload_and_parses_output() {
165        let transport = CapturingTransport::new(
166            serde_json::to_vec(&json!({
167                "output": [
168                    {
169                        "trace_summary": "raw summary",
170                        "memory_summary": "memory summary"
171                    }
172                ]
173            }))
174            .expect("serialize response"),
175        );
176        let client = MemoriesClient::new(
177            transport.clone(),
178            provider("https://example.com/api/codex"),
179            Arc::new(DummyAuth),
180        );
181
182        let input = MemorySummarizeInput {
183            model: "gpt-test".to_string(),
184            raw_memories: vec![RawMemory {
185                id: "trace-1".to_string(),
186                metadata: RawMemoryMetadata {
187                    source_path: "/tmp/trace.json".to_string(),
188                },
189                items: vec![json!({"type": "message", "role": "user", "content": []})],
190            }],
191            reasoning: None,
192        };
193
194        let output = client
195            .summarize_input(&input, HeaderMap::new())
196            .await
197            .expect("summarize input request should succeed");
198        assert_eq!(output.len(), 1);
199        assert_eq!(output[0].raw_memory, "raw summary");
200        assert_eq!(output[0].memory_summary, "memory summary");
201
202        let request = transport
203            .last_request
204            .lock()
205            .expect("lock request store")
206            .clone()
207            .expect("request should be captured");
208        assert_eq!(request.method, Method::POST);
209        assert_eq!(
210            request.url,
211            "https://example.com/api/codex/memories/trace_summarize"
212        );
213        let body = request
214            .body
215            .as_ref()
216            .and_then(RequestBody::json)
217            .expect("request body should be JSON");
218        assert_eq!(body["model"], "gpt-test");
219        assert_eq!(body["traces"][0]["id"], "trace-1");
220        assert_eq!(
221            body["traces"][0]["metadata"]["source_path"],
222            "/tmp/trace.json"
223        );
224    }
225}