Skip to main content

mnemo_pgwire/
server.rs

1//! PostgreSQL wire protocol connection handler.
2//!
3//! Implements the subset of the PostgreSQL wire protocol needed for
4//! simple query execution. Handles startup, authentication (trust mode),
5//! and the simple query flow.
6//!
7//! Reference: <https://www.postgresql.org/docs/current/protocol.html>
8
9use 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
19/// Handle a single PostgreSQL wire protocol connection.
20pub 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    // Phase 1: Startup message
26    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    // SSL request (80877103) — respond with 'N' (no SSL)
44    if protocol_version == 80877103 {
45        stream.write_all(b"N").await?;
46        // Client will retry with normal startup
47        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    // Phase 2: Authentication
58    if let Some(ref expected_password) = config.password {
59        // Send AuthenticationCleartextPassword (type 3)
60        stream.write_all(&[b'R', 0, 0, 0, 8, 0, 0, 0, 3]).await?;
61        stream.flush().await?;
62
63        // Read password message (type 'p')
64        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    // Send AuthenticationOk
88    stream.write_all(&[b'R', 0, 0, 0, 8, 0, 0, 0, 0]).await?;
89
90    // Send ParameterStatus messages
91    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 ReadyForQuery
97    send_ready_for_query(&mut stream).await?;
98
99    // Phase 3: Query loop
100    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                // Simple Query
116                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                // Terminate
135                tracing::debug!("pgwire client terminated");
136                break;
137            }
138            _ => {
139                // Unsupported message type — send error and continue
140                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
153/// Query response rows.
154struct 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            // v0.5.1 — the `/*+ reconstruct */` hint selects the
179            // active-reconstruction strategy (MRAgent, arXiv:2606.06036);
180            // otherwise the default pgwire read is filter-based "exact".
181            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'); // ParameterStatus type
335
336    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    // ReadyForQuery: type 'Z', length 5, transaction status 'I' (idle)
353    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'); // ErrorResponse type
363
364    let mut fields = Vec::new();
365    // Severity
366    fields.push(b'S');
367    fields.extend_from_slice(b"ERROR\0");
368    // SQLSTATE (42000 = syntax error)
369    fields.push(b'C');
370    fields.extend_from_slice(b"42000\0");
371    // Message
372    fields.push(b'M');
373    fields.extend_from_slice(message.as_bytes());
374    fields.push(0);
375    // Terminator
376    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        // RowDescription
392        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); // null terminator
398            desc_buf.extend_from_slice(&0i32.to_be_bytes()); // table OID
399            desc_buf.extend_from_slice(&0i16.to_be_bytes()); // column attr number
400            desc_buf.extend_from_slice(&25i32.to_be_bytes()); // type OID (text = 25)
401            desc_buf.extend_from_slice(&(-1i16).to_be_bytes()); // type size (-1 = variable)
402            desc_buf.extend_from_slice(&(-1i32).to_be_bytes()); // type modifier
403            desc_buf.extend_from_slice(&0i16.to_be_bytes()); // format code (text = 0)
404        }
405
406        let mut msg = Vec::new();
407        msg.push(b'T'); // RowDescription type
408        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        // DataRow for each row
414        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'); // DataRow type
426            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    // CommandComplete
434    let tag = response.command_tag.as_bytes();
435    let mut msg = Vec::new();
436    msg.push(b'C'); // CommandComplete type
437    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}