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;