Skip to main content

agent_first_psql/
types.rs

1use serde::{Deserialize, Serialize};
2use serde_json::Value;
3use std::collections::HashMap;
4
5#[derive(Debug, Deserialize)]
6#[serde(tag = "code", deny_unknown_fields)]
7pub enum Input {
8    #[serde(rename = "query")]
9    Query {
10        id: String,
11        #[serde(default)]
12        session: Option<String>,
13        sql: String,
14        #[serde(default)]
15        params: Vec<Value>,
16        #[serde(default)]
17        options: QueryOptions,
18    },
19    #[serde(rename = "config")]
20    Config(ConfigPatch),
21    #[serde(rename = "cancel")]
22    Cancel { id: String },
23    #[serde(rename = "ping")]
24    Ping,
25    #[serde(rename = "close")]
26    Close,
27}
28
29#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
30pub enum Permission {
31    #[serde(rename = "read")]
32    Read,
33    #[serde(rename = "write")]
34    Write,
35    #[serde(rename = "ssh-read")]
36    SshRead,
37    #[serde(rename = "ssh-write")]
38    SshWrite,
39}
40
41impl Permission {
42    pub fn as_str(self) -> &'static str {
43        match self {
44            Self::Read => "read",
45            Self::Write => "write",
46            Self::SshRead => "ssh-read",
47            Self::SshWrite => "ssh-write",
48        }
49    }
50
51    pub fn is_read_only(self) -> bool {
52        matches!(self, Self::Read | Self::SshRead)
53    }
54
55    pub fn allows_ssh(self) -> bool {
56        matches!(self, Self::SshRead | Self::SshWrite)
57    }
58}
59
60impl std::str::FromStr for Permission {
61    type Err = String;
62
63    fn from_str(value: &str) -> Result<Self, Self::Err> {
64        match value {
65            "read" => Ok(Self::Read),
66            "write" => Ok(Self::Write),
67            "ssh-read" => Ok(Self::SshRead),
68            "ssh-write" => Ok(Self::SshWrite),
69            _ => Err(format!(
70                "invalid permission `{value}`; expected read, write, ssh-read, or ssh-write"
71            )),
72        }
73    }
74}
75
76#[derive(Debug, Deserialize, Default, Clone)]
77#[serde(deny_unknown_fields)]
78#[allow(dead_code)]
79pub struct QueryOptions {
80    #[serde(default)]
81    pub stream_rows: bool,
82    pub batch_rows: Option<usize>,
83    pub batch_bytes: Option<usize>,
84    pub statement_timeout_ms: Option<u64>,
85    pub lock_timeout_ms: Option<u64>,
86    pub permission: Option<Permission>,
87    pub inline_max_rows: Option<usize>,
88    pub inline_max_bytes: Option<usize>,
89}
90
91#[derive(Debug, Serialize)]
92#[serde(tag = "code")]
93pub enum Output {
94    #[serde(rename = "result")]
95    Result {
96        #[serde(skip_serializing_if = "Option::is_none")]
97        id: Option<String>,
98        #[serde(skip_serializing_if = "Option::is_none")]
99        session: Option<String>,
100        command_tag: String,
101        columns: Vec<ColumnInfo>,
102        rows: Vec<Value>,
103        row_count: usize,
104        trace: Trace,
105    },
106    #[serde(rename = "result_start")]
107    ResultStart {
108        id: String,
109        #[serde(skip_serializing_if = "Option::is_none")]
110        session: Option<String>,
111        columns: Vec<ColumnInfo>,
112    },
113    #[serde(rename = "result_rows")]
114    ResultRows {
115        id: String,
116        rows: Vec<Value>,
117        rows_batch_count: usize,
118    },
119    #[serde(rename = "result_end")]
120    ResultEnd {
121        id: String,
122        #[serde(skip_serializing_if = "Option::is_none")]
123        session: Option<String>,
124        command_tag: String,
125        trace: Trace,
126    },
127    #[serde(rename = "sql_error")]
128    SqlError {
129        #[serde(skip_serializing_if = "Option::is_none")]
130        id: Option<String>,
131        #[serde(skip_serializing_if = "Option::is_none")]
132        session: Option<String>,
133        sqlstate: String,
134        message: String,
135        #[serde(skip_serializing_if = "Option::is_none")]
136        detail: Option<String>,
137        #[serde(skip_serializing_if = "Option::is_none")]
138        hint: Option<String>,
139        #[serde(skip_serializing_if = "Option::is_none")]
140        position: Option<String>,
141        trace: Trace,
142    },
143    #[serde(rename = "error")]
144    Error {
145        #[serde(skip_serializing_if = "Option::is_none")]
146        id: Option<String>,
147        error_code: String,
148        error: String,
149        #[serde(skip_serializing_if = "Option::is_none")]
150        hint: Option<String>,
151        retryable: bool,
152        trace: Trace,
153    },
154    #[serde(rename = "error")]
155    ConnectError {
156        #[serde(skip_serializing_if = "Option::is_none")]
157        id: Option<String>,
158        error_code: String,
159        error: String,
160        #[serde(skip_serializing_if = "Option::is_none")]
161        sqlstate: Option<String>,
162        #[serde(skip_serializing_if = "Option::is_none")]
163        message: Option<String>,
164        #[serde(skip_serializing_if = "Option::is_none")]
165        detail: Option<String>,
166        #[serde(skip_serializing_if = "Option::is_none")]
167        hint: Option<String>,
168        retryable: bool,
169        trace: Trace,
170    },
171    #[serde(rename = "dry_run")]
172    DryRun {
173        #[serde(skip_serializing_if = "Option::is_none")]
174        id: Option<String>,
175        sql: String,
176        params: Vec<String>,
177        #[serde(skip_serializing_if = "Option::is_none")]
178        session: Option<String>,
179        trace: Trace,
180    },
181    #[serde(rename = "config")]
182    Config(RuntimeConfig),
183    #[serde(rename = "pong")]
184    Pong { trace: PongTrace },
185    #[serde(rename = "close")]
186    Close { message: String, trace: CloseTrace },
187    #[serde(rename = "log")]
188    Log {
189        event: String,
190        #[serde(skip_serializing_if = "Option::is_none")]
191        request_id: Option<String>,
192        #[serde(skip_serializing_if = "Option::is_none")]
193        session: Option<String>,
194        #[serde(skip_serializing_if = "Option::is_none")]
195        error_code: Option<String>,
196        #[serde(skip_serializing_if = "Option::is_none")]
197        command_tag: Option<String>,
198        #[serde(skip_serializing_if = "Option::is_none")]
199        version: Option<String>,
200        #[serde(skip_serializing_if = "Option::is_none")]
201        config: Option<Value>,
202        #[serde(skip_serializing_if = "Option::is_none")]
203        args: Option<Value>,
204        #[serde(skip_serializing_if = "Option::is_none")]
205        env: Option<Value>,
206        trace: Trace,
207    },
208}
209
210#[derive(Debug, Serialize, Clone)]
211pub struct ColumnInfo {
212    pub name: String,
213    #[serde(rename = "type")]
214    pub type_name: String,
215}
216
217#[derive(Debug, Serialize, Clone)]
218pub struct Trace {
219    pub duration_ms: u64,
220    #[serde(skip_serializing_if = "Option::is_none")]
221    pub row_count: Option<usize>,
222    #[serde(skip_serializing_if = "Option::is_none")]
223    pub payload_bytes: Option<usize>,
224}
225
226impl Trace {
227    pub fn only_duration(duration_ms: u64) -> Self {
228        Self {
229            duration_ms,
230            row_count: None,
231            payload_bytes: None,
232        }
233    }
234}
235
236#[derive(Debug, Serialize)]
237pub struct PongTrace {
238    pub uptime_s: u64,
239    pub requests_total: u64,
240    pub in_flight: usize,
241}
242
243#[derive(Debug, Serialize)]
244pub struct CloseTrace {
245    pub uptime_s: u64,
246    pub requests_total: u64,
247}
248
249#[derive(Debug, Serialize, Deserialize, Clone, Default)]
250#[serde(deny_unknown_fields)]
251pub struct SessionConfig {
252    #[serde(skip_serializing_if = "Option::is_none")]
253    pub dsn_secret: Option<String>,
254    #[serde(skip_serializing_if = "Option::is_none")]
255    pub conninfo_secret: Option<String>,
256    #[serde(skip_serializing_if = "Option::is_none")]
257    pub host: Option<String>,
258    #[serde(skip_serializing_if = "Option::is_none")]
259    pub port: Option<u16>,
260    #[serde(skip_serializing_if = "Option::is_none")]
261    pub user: Option<String>,
262    #[serde(skip_serializing_if = "Option::is_none")]
263    pub dbname: Option<String>,
264    #[serde(skip_serializing_if = "Option::is_none")]
265    pub password_secret: Option<String>,
266    #[serde(skip_serializing_if = "Option::is_none")]
267    pub ssh: Option<String>,
268    #[serde(default, skip_serializing_if = "Vec::is_empty")]
269    pub ssh_options: Vec<String>,
270    #[serde(skip_serializing_if = "Option::is_none")]
271    pub ssh_local_host: Option<String>,
272    #[serde(skip_serializing_if = "Option::is_none")]
273    pub ssh_local_port: Option<u16>,
274    #[serde(skip_serializing_if = "Option::is_none")]
275    pub ssh_remote_socket: Option<String>,
276    #[serde(skip_serializing_if = "Option::is_none")]
277    pub ssh_sudo_user: Option<String>,
278}
279
280impl SessionConfig {
281    pub fn uses_ssh_transport(&self) -> bool {
282        self.ssh.is_some()
283            || !self.ssh_options.is_empty()
284            || self.ssh_local_host.is_some()
285            || self.ssh_local_port.is_some()
286            || self.ssh_remote_socket.is_some()
287            || self.ssh_sudo_user.is_some()
288    }
289}
290
291#[derive(Debug, Serialize, Deserialize, Clone)]
292pub struct RuntimeConfig {
293    pub default_session: String,
294    #[serde(default)]
295    pub sessions: HashMap<String, SessionConfig>,
296    pub inline_max_rows: usize,
297    pub inline_max_bytes: usize,
298    pub statement_timeout_ms: u64,
299    pub lock_timeout_ms: u64,
300    #[serde(default)]
301    pub log: Vec<String>,
302}
303
304impl Default for RuntimeConfig {
305    fn default() -> Self {
306        let mut sessions = HashMap::new();
307        sessions.insert("default".to_string(), SessionConfig::default());
308        Self {
309            default_session: "default".to_string(),
310            sessions,
311            inline_max_rows: 1000,
312            inline_max_bytes: 1_048_576,
313            statement_timeout_ms: 30_000,
314            lock_timeout_ms: 5_000,
315            log: vec![],
316        }
317    }
318}
319
320#[derive(Debug, Deserialize, Default)]
321#[serde(deny_unknown_fields)]
322pub struct ConfigPatch {
323    pub default_session: Option<String>,
324    pub sessions: Option<HashMap<String, SessionConfigPatch>>,
325    pub inline_max_rows: Option<usize>,
326    pub inline_max_bytes: Option<usize>,
327    pub statement_timeout_ms: Option<u64>,
328    pub lock_timeout_ms: Option<u64>,
329    pub log: Option<Vec<String>>,
330}
331
332#[derive(Debug, Default)]
333pub enum PatchField<T> {
334    #[default]
335    Missing,
336    Null,
337    Value(T),
338}
339
340impl<'de, T> Deserialize<'de> for PatchField<T>
341where
342    T: Deserialize<'de>,
343{
344    fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
345    where
346        D: serde::Deserializer<'de>,
347    {
348        let value = Option::<T>::deserialize(deserializer)?;
349        match value {
350            Some(value) => Ok(Self::Value(value)),
351            None => Ok(Self::Null),
352        }
353    }
354}
355
356impl<T> PatchField<T> {
357    pub fn into_update(self) -> Option<Option<T>> {
358        match self {
359            Self::Missing => None,
360            Self::Null => Some(None),
361            Self::Value(value) => Some(Some(value)),
362        }
363    }
364}
365
366#[derive(Debug, Deserialize, Default)]
367#[serde(deny_unknown_fields)]
368pub struct SessionConfigPatch {
369    #[serde(default)]
370    pub dsn_secret: PatchField<String>,
371    #[serde(default)]
372    pub conninfo_secret: PatchField<String>,
373    #[serde(default)]
374    pub host: PatchField<String>,
375    #[serde(default)]
376    pub port: PatchField<u16>,
377    #[serde(default)]
378    pub user: PatchField<String>,
379    #[serde(default)]
380    pub dbname: PatchField<String>,
381    #[serde(default)]
382    pub password_secret: PatchField<String>,
383    #[serde(default)]
384    pub ssh: PatchField<String>,
385    #[serde(default)]
386    pub ssh_options: PatchField<Vec<String>>,
387    #[serde(default)]
388    pub ssh_local_host: PatchField<String>,
389    #[serde(default)]
390    pub ssh_local_port: PatchField<u16>,
391    #[serde(default)]
392    pub ssh_remote_socket: PatchField<String>,
393    #[serde(default)]
394    pub ssh_sudo_user: PatchField<String>,
395}
396
397#[derive(Debug, Clone)]
398#[allow(dead_code)]
399pub struct ResolvedOptions {
400    pub stream_rows: bool,
401    pub batch_rows: usize,
402    pub batch_bytes: usize,
403    pub statement_timeout_ms: u64,
404    pub lock_timeout_ms: u64,
405    pub read_only: bool,
406    pub inline_max_rows: usize,
407    pub inline_max_bytes: usize,
408}
409
410#[cfg(test)]
411#[path = "../tests/support/unit_types.rs"]
412mod tests;