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")]
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, Deserialize, Default, Clone)]
30#[allow(dead_code)]
31pub struct QueryOptions {
32    #[serde(default)]
33    pub stream_rows: bool,
34    pub batch_rows: Option<usize>,
35    pub batch_bytes: Option<usize>,
36    pub statement_timeout_ms: Option<u64>,
37    pub lock_timeout_ms: Option<u64>,
38    pub read_only: Option<bool>,
39    pub inline_max_rows: Option<usize>,
40    pub inline_max_bytes: Option<usize>,
41}
42
43#[derive(Debug, Serialize)]
44#[serde(tag = "code")]
45pub enum Output {
46    #[serde(rename = "result")]
47    Result {
48        #[serde(skip_serializing_if = "Option::is_none")]
49        id: Option<String>,
50        #[serde(skip_serializing_if = "Option::is_none")]
51        session: Option<String>,
52        command_tag: String,
53        columns: Vec<ColumnInfo>,
54        rows: Vec<Value>,
55        row_count: usize,
56        trace: Trace,
57    },
58    #[serde(rename = "result_start")]
59    ResultStart {
60        id: String,
61        #[serde(skip_serializing_if = "Option::is_none")]
62        session: Option<String>,
63        columns: Vec<ColumnInfo>,
64    },
65    #[serde(rename = "result_rows")]
66    ResultRows {
67        id: String,
68        rows: Vec<Value>,
69        rows_batch_count: usize,
70    },
71    #[serde(rename = "result_end")]
72    ResultEnd {
73        id: String,
74        #[serde(skip_serializing_if = "Option::is_none")]
75        session: Option<String>,
76        command_tag: String,
77        trace: Trace,
78    },
79    #[serde(rename = "sql_error")]
80    SqlError {
81        #[serde(skip_serializing_if = "Option::is_none")]
82        id: Option<String>,
83        #[serde(skip_serializing_if = "Option::is_none")]
84        session: Option<String>,
85        sqlstate: String,
86        message: String,
87        #[serde(skip_serializing_if = "Option::is_none")]
88        detail: Option<String>,
89        #[serde(skip_serializing_if = "Option::is_none")]
90        hint: Option<String>,
91        #[serde(skip_serializing_if = "Option::is_none")]
92        position: Option<String>,
93        trace: Trace,
94    },
95    #[serde(rename = "error")]
96    Error {
97        #[serde(skip_serializing_if = "Option::is_none")]
98        id: Option<String>,
99        error_code: String,
100        error: String,
101        #[serde(skip_serializing_if = "Option::is_none")]
102        hint: Option<String>,
103        retryable: bool,
104        trace: Trace,
105    },
106    #[serde(rename = "dry_run")]
107    DryRun {
108        #[serde(skip_serializing_if = "Option::is_none")]
109        id: Option<String>,
110        sql: String,
111        params: Vec<String>,
112        #[serde(skip_serializing_if = "Option::is_none")]
113        session: Option<String>,
114        trace: Trace,
115    },
116    #[serde(rename = "config")]
117    Config(RuntimeConfig),
118    #[serde(rename = "pong")]
119    Pong { trace: PongTrace },
120    #[serde(rename = "close")]
121    Close { message: String, trace: CloseTrace },
122    #[serde(rename = "log")]
123    Log {
124        event: String,
125        #[serde(skip_serializing_if = "Option::is_none")]
126        request_id: Option<String>,
127        #[serde(skip_serializing_if = "Option::is_none")]
128        session: Option<String>,
129        #[serde(skip_serializing_if = "Option::is_none")]
130        error_code: Option<String>,
131        #[serde(skip_serializing_if = "Option::is_none")]
132        command_tag: Option<String>,
133        #[serde(skip_serializing_if = "Option::is_none")]
134        version: Option<String>,
135        #[serde(skip_serializing_if = "Option::is_none")]
136        config: Option<Value>,
137        #[serde(skip_serializing_if = "Option::is_none")]
138        args: Option<Value>,
139        #[serde(skip_serializing_if = "Option::is_none")]
140        env: Option<Value>,
141        trace: Trace,
142    },
143}
144
145#[derive(Debug, Serialize, Clone)]
146pub struct ColumnInfo {
147    pub name: String,
148    #[serde(rename = "type")]
149    pub type_name: String,
150}
151
152#[derive(Debug, Serialize, Clone)]
153pub struct Trace {
154    pub duration_ms: u64,
155    #[serde(skip_serializing_if = "Option::is_none")]
156    pub row_count: Option<usize>,
157    #[serde(skip_serializing_if = "Option::is_none")]
158    pub payload_bytes: Option<usize>,
159}
160
161impl Trace {
162    pub fn only_duration(duration_ms: u64) -> Self {
163        Self {
164            duration_ms,
165            row_count: None,
166            payload_bytes: None,
167        }
168    }
169}
170
171#[derive(Debug, Serialize)]
172pub struct PongTrace {
173    pub uptime_s: u64,
174    pub requests_total: u64,
175    pub in_flight: usize,
176}
177
178#[derive(Debug, Serialize)]
179pub struct CloseTrace {
180    pub uptime_s: u64,
181    pub requests_total: u64,
182}
183
184#[derive(Debug, Serialize, Deserialize, Clone, Default)]
185pub struct SessionConfig {
186    #[serde(skip_serializing_if = "Option::is_none")]
187    pub dsn_secret: Option<String>,
188    #[serde(skip_serializing_if = "Option::is_none")]
189    pub conninfo_secret: Option<String>,
190    #[serde(skip_serializing_if = "Option::is_none")]
191    pub host: Option<String>,
192    #[serde(skip_serializing_if = "Option::is_none")]
193    pub port: Option<u16>,
194    #[serde(skip_serializing_if = "Option::is_none")]
195    pub user: Option<String>,
196    #[serde(skip_serializing_if = "Option::is_none")]
197    pub dbname: Option<String>,
198    #[serde(skip_serializing_if = "Option::is_none")]
199    pub password_secret: Option<String>,
200}
201
202#[derive(Debug, Serialize, Deserialize, Clone)]
203pub struct RuntimeConfig {
204    pub default_session: String,
205    #[serde(default)]
206    pub sessions: HashMap<String, SessionConfig>,
207    pub inline_max_rows: usize,
208    pub inline_max_bytes: usize,
209    pub statement_timeout_ms: u64,
210    pub lock_timeout_ms: u64,
211    #[serde(default)]
212    pub log: Vec<String>,
213}
214
215impl Default for RuntimeConfig {
216    fn default() -> Self {
217        let mut sessions = HashMap::new();
218        sessions.insert("default".to_string(), SessionConfig::default());
219        Self {
220            default_session: "default".to_string(),
221            sessions,
222            inline_max_rows: 1000,
223            inline_max_bytes: 1_048_576,
224            statement_timeout_ms: 30_000,
225            lock_timeout_ms: 5_000,
226            log: vec![],
227        }
228    }
229}
230
231#[derive(Debug, Deserialize, Default)]
232pub struct ConfigPatch {
233    pub default_session: Option<String>,
234    pub sessions: Option<HashMap<String, SessionConfigPatch>>,
235    pub inline_max_rows: Option<usize>,
236    pub inline_max_bytes: Option<usize>,
237    pub statement_timeout_ms: Option<u64>,
238    pub lock_timeout_ms: Option<u64>,
239    pub log: Option<Vec<String>>,
240}
241
242#[derive(Debug, Default)]
243pub enum PatchField<T> {
244    #[default]
245    Missing,
246    Null,
247    Value(T),
248}
249
250impl<'de, T> Deserialize<'de> for PatchField<T>
251where
252    T: Deserialize<'de>,
253{
254    fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
255    where
256        D: serde::Deserializer<'de>,
257    {
258        let value = Option::<T>::deserialize(deserializer)?;
259        match value {
260            Some(value) => Ok(Self::Value(value)),
261            None => Ok(Self::Null),
262        }
263    }
264}
265
266impl<T> PatchField<T> {
267    pub fn into_update(self) -> Option<Option<T>> {
268        match self {
269            Self::Missing => None,
270            Self::Null => Some(None),
271            Self::Value(value) => Some(Some(value)),
272        }
273    }
274}
275
276#[derive(Debug, Deserialize, Default)]
277pub struct SessionConfigPatch {
278    #[serde(default)]
279    pub dsn_secret: PatchField<String>,
280    #[serde(default)]
281    pub conninfo_secret: PatchField<String>,
282    #[serde(default)]
283    pub host: PatchField<String>,
284    #[serde(default)]
285    pub port: PatchField<u16>,
286    #[serde(default)]
287    pub user: PatchField<String>,
288    #[serde(default)]
289    pub dbname: PatchField<String>,
290    #[serde(default)]
291    pub password_secret: PatchField<String>,
292}
293
294#[derive(Debug, Clone)]
295#[allow(dead_code)]
296pub struct ResolvedOptions {
297    pub stream_rows: bool,
298    pub batch_rows: usize,
299    pub batch_bytes: usize,
300    pub statement_timeout_ms: u64,
301    pub lock_timeout_ms: u64,
302    pub read_only: bool,
303    pub inline_max_rows: usize,
304    pub inline_max_bytes: usize,
305}
306
307#[cfg(test)]
308#[path = "../tests/support/unit_types.rs"]
309mod tests;