codex_api/endpoint/
memories.rs1use 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}