1use std::sync::Arc;
10
11use tokio::io::{AsyncReadExt, AsyncWriteExt};
12use tokio::net::TcpStream;
13
14use mnemo_core::query::MnemoEngine;
15
16use crate::PgWireConfig;
17use crate::parser::{self, ParsedStatement};
18
19pub async fn handle_connection(
21 mut stream: TcpStream,
22 engine: Arc<MnemoEngine>,
23 config: &PgWireConfig,
24) -> Result<(), Box<dyn std::error::Error + Send + Sync>> {
25 let startup_len_raw = stream.read_i32().await?;
27 let startup_len = usize::try_from(startup_len_raw)
28 .map_err(|_| format!("negative startup message length: {startup_len_raw}"))?;
29 if !(8..=10240).contains(&startup_len) {
30 return Err("invalid startup message length".into());
31 }
32
33 let mut startup_buf = vec![0u8; startup_len - 4];
34 stream.read_exact(&mut startup_buf).await?;
35
36 let protocol_version = i32::from_be_bytes([
37 startup_buf[0],
38 startup_buf[1],
39 startup_buf[2],
40 startup_buf[3],
41 ]);
42
43 if protocol_version == 80877103 {
45 stream.write_all(b"N").await?;
46 let startup_len_raw = stream.read_i32().await?;
48 let startup_len = usize::try_from(startup_len_raw)
49 .map_err(|_| format!("negative startup message length: {startup_len_raw}"))?;
50 if !(8..=10240).contains(&startup_len) {
51 return Err("invalid startup message length after SSL".into());
52 }
53 startup_buf = vec![0u8; startup_len - 4];
54 stream.read_exact(&mut startup_buf).await?;
55 }
56
57 if let Some(ref expected_password) = config.password {
59 stream.write_all(&[b'R', 0, 0, 0, 8, 0, 0, 0, 3]).await?;
61 stream.flush().await?;
62
63 let pw_type = stream.read_u8().await?;
65 if pw_type != b'p' {
66 send_error(&mut stream, "expected password message").await?;
67 return Err("expected password message".into());
68 }
69 let pw_len_raw = stream.read_i32().await?;
70 let pw_len = usize::try_from(pw_len_raw)
71 .map_err(|_| format!("negative password message length: {pw_len_raw}"))?;
72 if !(5..=10240).contains(&pw_len) {
73 return Err("invalid password message length".into());
74 }
75 let mut pw_buf = vec![0u8; pw_len - 4];
76 stream.read_exact(&mut pw_buf).await?;
77 let client_password = String::from_utf8_lossy(&pw_buf)
78 .trim_end_matches('\0')
79 .to_string();
80
81 if client_password != *expected_password {
82 send_error(&mut stream, "password authentication failed").await?;
83 return Err("authentication failed".into());
84 }
85 }
86
87 stream.write_all(&[b'R', 0, 0, 0, 8, 0, 0, 0, 0]).await?;
89
90 send_parameter_status(&mut stream, "server_version", "16.0").await?;
92 send_parameter_status(&mut stream, "server_encoding", "UTF8").await?;
93 send_parameter_status(&mut stream, "client_encoding", "UTF8").await?;
94 send_parameter_status(&mut stream, "application_name", "mnemo-pgwire").await?;
95
96 send_ready_for_query(&mut stream).await?;
98
99 while let Ok(msg_type) = stream.read_u8().await {
101 let msg_len_raw = stream.read_i32().await?;
102 let msg_len = usize::try_from(msg_len_raw)
103 .map_err(|_| format!("negative message length: {msg_len_raw}"))?;
104 if !(4..=1_048_576).contains(&msg_len) {
105 break;
106 }
107
108 let mut msg_buf = vec![0u8; msg_len - 4];
109 if !msg_buf.is_empty() {
110 stream.read_exact(&mut msg_buf).await?;
111 }
112
113 match msg_type {
114 b'Q' => {
115 let sql = String::from_utf8_lossy(&msg_buf)
117 .trim_end_matches('\0')
118 .to_string();
119
120 tracing::debug!("pgwire query: {sql}");
121
122 match handle_query(&sql, &engine, config).await {
123 Ok(response) => {
124 send_query_response(&mut stream, &response).await?;
125 }
126 Err(e) => {
127 send_error(&mut stream, &e.to_string()).await?;
128 }
129 }
130
131 send_ready_for_query(&mut stream).await?;
132 }
133 b'X' => {
134 tracing::debug!("pgwire client terminated");
136 break;
137 }
138 _ => {
139 send_error(
141 &mut stream,
142 &format!("unsupported message type: {}", msg_type as char),
143 )
144 .await?;
145 send_ready_for_query(&mut stream).await?;
146 }
147 }
148 }
149
150 Ok(())
151}
152
153struct QueryResponse {
155 columns: Vec<String>,
156 rows: Vec<Vec<String>>,
157 command_tag: String,
158}
159
160async fn handle_query(
161 sql: &str,
162 engine: &MnemoEngine,
163 config: &PgWireConfig,
164) -> Result<QueryResponse, Box<dyn std::error::Error + Send + Sync>> {
165 let stmt = parser::parse_sql(sql);
166
167 match stmt {
168 ParsedStatement::Select(q) => {
169 let agent_id = q
170 .agent_id
171 .unwrap_or_else(|| config.default_agent_id.clone());
172
173 let orientation_cache_cfg = if q.orientation_cache {
174 Some(mnemo_core::query::orientation_cache::OrientationCacheConfig::new())
175 } else {
176 None
177 };
178 let strategy = if q.reconstruct {
182 "reconstruct"
183 } else {
184 "exact"
185 };
186 let request = mnemo_core::query::recall::RecallRequest {
187 agent_id: Some(agent_id),
188 query: q.query_text.unwrap_or_default(),
189 limit: Some(q.limit),
190 memory_type: None,
191 memory_types: None,
192 scope: None,
193 strategy: Some(strategy.to_string()),
194 min_importance: None,
195 tags: None,
196 org_id: None,
197 temporal_range: None,
198 recency_half_life_hours: None,
199 hybrid_weights: None,
200 rrf_k: None,
201 as_of: None,
202 explain: None,
203 with_provenance: None,
204 mode: None,
205 current_fact_resolver: None,
206 orientation_cache: orientation_cache_cfg,
207 evidence_budget: None,
208 retained_token_budget: None,
209 domain_scope: None,
210 reasoning_trust: None,
211 };
212
213 let response = engine.recall(request).await?;
214
215 let columns = vec![
216 "id".to_string(),
217 "agent_id".to_string(),
218 "content".to_string(),
219 "memory_type".to_string(),
220 "importance".to_string(),
221 "created_at".to_string(),
222 ];
223
224 let rows: Vec<Vec<String>> = response
225 .memories
226 .iter()
227 .skip(q.offset)
228 .map(|m| {
229 vec![
230 m.id.to_string(),
231 m.agent_id.clone(),
232 m.content.clone(),
233 m.memory_type.to_string(),
234 m.importance.to_string(),
235 m.created_at.clone(),
236 ]
237 })
238 .collect();
239
240 let count = rows.len();
241 Ok(QueryResponse {
242 columns,
243 rows,
244 command_tag: format!("SELECT {count}"),
245 })
246 }
247
248 ParsedStatement::Insert(q) => {
249 let agent_id = q
250 .agent_id
251 .unwrap_or_else(|| config.default_agent_id.clone());
252
253 let request = mnemo_core::query::remember::RememberRequest {
254 content: q.content,
255 agent_id: Some(agent_id),
256 memory_type: q.memory_type.as_deref().and_then(parse_memory_type),
257 scope: None,
258 importance: q.importance,
259 tags: if q.tags.is_empty() {
260 None
261 } else {
262 Some(q.tags)
263 },
264 metadata: None,
265 source_type: None,
266 source_id: None,
267 org_id: None,
268 thread_id: None,
269 ttl_seconds: None,
270 related_to: None,
271 decay_rate: None,
272 created_by: None,
273 };
274
275 let response = engine.remember(request).await?;
276
277 Ok(QueryResponse {
278 columns: vec!["id".to_string(), "content_hash".to_string()],
279 rows: vec![vec![response.id.to_string(), response.content_hash.clone()]],
280 command_tag: "INSERT 0 1".to_string(),
281 })
282 }
283
284 ParsedStatement::Delete(q) => {
285 if let Some(memory_id_str) = q.memory_id {
286 let memory_id: uuid::Uuid = memory_id_str
287 .parse()
288 .map_err(|e| format!("invalid UUID in DELETE WHERE id = '...': {e}"))?;
289
290 let agent_id = q
291 .agent_id
292 .unwrap_or_else(|| config.default_agent_id.clone());
293
294 let request = mnemo_core::query::forget::ForgetRequest {
295 memory_ids: vec![memory_id],
296 agent_id: Some(agent_id),
297 strategy: Some(mnemo_core::query::forget::ForgetStrategy::SoftDelete),
298 criteria: None,
299 };
300
301 let response = engine.forget(request).await?;
302 let count = response.forgotten.len();
303
304 Ok(QueryResponse {
305 columns: vec![],
306 rows: vec![],
307 command_tag: format!("DELETE {count}"),
308 })
309 } else {
310 Err("DELETE requires WHERE id = '...' clause".into())
311 }
312 }
313
314 ParsedStatement::Unsupported(s) => Err(format!("unsupported SQL: {s}").into()),
315 }
316}
317
318fn parse_memory_type(s: &str) -> Option<mnemo_core::model::memory::MemoryType> {
319 match s.to_lowercase().as_str() {
320 "episodic" => Some(mnemo_core::model::memory::MemoryType::Episodic),
321 "semantic" => Some(mnemo_core::model::memory::MemoryType::Semantic),
322 "procedural" => Some(mnemo_core::model::memory::MemoryType::Procedural),
323 "working" => Some(mnemo_core::model::memory::MemoryType::Working),
324 _ => None,
325 }
326}
327
328async fn send_parameter_status(
329 stream: &mut TcpStream,
330 name: &str,
331 value: &str,
332) -> Result<(), Box<dyn std::error::Error + Send + Sync>> {
333 let mut buf = Vec::new();
334 buf.push(b'S'); let name_bytes = name.as_bytes();
337 let value_bytes = value.as_bytes();
338 let len = 4 + name_bytes.len() + 1 + value_bytes.len() + 1;
339 buf.extend_from_slice(&(len as i32).to_be_bytes());
340 buf.extend_from_slice(name_bytes);
341 buf.push(0);
342 buf.extend_from_slice(value_bytes);
343 buf.push(0);
344
345 stream.write_all(&buf).await?;
346 Ok(())
347}
348
349async fn send_ready_for_query(
350 stream: &mut TcpStream,
351) -> Result<(), Box<dyn std::error::Error + Send + Sync>> {
352 stream.write_all(&[b'Z', 0, 0, 0, 5, b'I']).await?;
354 Ok(())
355}
356
357async fn send_error(
358 stream: &mut TcpStream,
359 message: &str,
360) -> Result<(), Box<dyn std::error::Error + Send + Sync>> {
361 let mut buf = Vec::new();
362 buf.push(b'E'); let mut fields = Vec::new();
365 fields.push(b'S');
367 fields.extend_from_slice(b"ERROR\0");
368 fields.push(b'C');
370 fields.extend_from_slice(b"42000\0");
371 fields.push(b'M');
373 fields.extend_from_slice(message.as_bytes());
374 fields.push(0);
375 fields.push(0);
377
378 let len = 4 + fields.len();
379 buf.extend_from_slice(&(len as i32).to_be_bytes());
380 buf.extend_from_slice(&fields);
381
382 stream.write_all(&buf).await?;
383 Ok(())
384}
385
386async fn send_query_response(
387 stream: &mut TcpStream,
388 response: &QueryResponse,
389) -> Result<(), Box<dyn std::error::Error + Send + Sync>> {
390 if !response.columns.is_empty() {
391 let mut desc_buf = Vec::new();
393 desc_buf.extend_from_slice(&(response.columns.len() as i16).to_be_bytes());
394
395 for col in &response.columns {
396 desc_buf.extend_from_slice(col.as_bytes());
397 desc_buf.push(0); desc_buf.extend_from_slice(&0i32.to_be_bytes()); desc_buf.extend_from_slice(&0i16.to_be_bytes()); desc_buf.extend_from_slice(&25i32.to_be_bytes()); desc_buf.extend_from_slice(&(-1i16).to_be_bytes()); desc_buf.extend_from_slice(&(-1i32).to_be_bytes()); desc_buf.extend_from_slice(&0i16.to_be_bytes()); }
405
406 let mut msg = Vec::new();
407 msg.push(b'T'); let len = 4 + desc_buf.len();
409 msg.extend_from_slice(&(len as i32).to_be_bytes());
410 msg.extend_from_slice(&desc_buf);
411 stream.write_all(&msg).await?;
412
413 for row in &response.rows {
415 let mut row_buf = Vec::new();
416 row_buf.extend_from_slice(&(row.len() as i16).to_be_bytes());
417
418 for val in row {
419 let bytes = val.as_bytes();
420 row_buf.extend_from_slice(&(bytes.len() as i32).to_be_bytes());
421 row_buf.extend_from_slice(bytes);
422 }
423
424 let mut msg = Vec::new();
425 msg.push(b'D'); let len = 4 + row_buf.len();
427 msg.extend_from_slice(&(len as i32).to_be_bytes());
428 msg.extend_from_slice(&row_buf);
429 stream.write_all(&msg).await?;
430 }
431 }
432
433 let tag = response.command_tag.as_bytes();
435 let mut msg = Vec::new();
436 msg.push(b'C'); let len = 4 + tag.len() + 1;
438 msg.extend_from_slice(&(len as i32).to_be_bytes());
439 msg.extend_from_slice(tag);
440 msg.push(0);
441 stream.write_all(&msg).await?;
442
443 Ok(())
444}