1use std::sync::atomic::{AtomicU64, Ordering};
20use std::sync::{Arc, Mutex, Weak};
21
22use crate::protocol::message::{backend::Message, frontend};
23use crate::types::Oid;
24use tracing::{trace, warn};
25
26use super::connection::RawConnection;
27use super::error::{Error, Result};
28use super::row::Row;
29use super::statement::{Column, ParamFormat, bind_format_codes};
30use super::sync_stream::SyncStream;
31pub trait SqlParam {
40 fn encode(&self) -> Vec<u8>;
42}
43
44impl SqlParam for i16 {
45 #[inline]
46 fn encode(&self) -> Vec<u8> {
47 self.to_le_bytes().to_vec()
48 }
49}
50
51impl SqlParam for i32 {
52 #[inline]
53 fn encode(&self) -> Vec<u8> {
54 self.to_le_bytes().to_vec()
55 }
56}
57
58impl SqlParam for i64 {
59 #[inline]
60 fn encode(&self) -> Vec<u8> {
61 self.to_le_bytes().to_vec()
62 }
63}
64
65impl SqlParam for f32 {
66 #[inline]
67 fn encode(&self) -> Vec<u8> {
68 self.to_le_bytes().to_vec()
69 }
70}
71
72impl SqlParam for f64 {
73 #[inline]
74 fn encode(&self) -> Vec<u8> {
75 self.to_le_bytes().to_vec()
76 }
77}
78
79impl SqlParam for bool {
80 #[inline]
81 fn encode(&self) -> Vec<u8> {
82 vec![u8::from(*self)]
83 }
84}
85
86impl SqlParam for &str {
87 #[inline]
88 fn encode(&self) -> Vec<u8> {
89 self.as_bytes().to_vec()
90 }
91}
92
93impl SqlParam for String {
94 #[inline]
95 fn encode(&self) -> Vec<u8> {
96 self.as_bytes().to_vec()
97 }
98}
99
100impl SqlParam for &String {
101 #[inline]
102 fn encode(&self) -> Vec<u8> {
103 self.as_bytes().to_vec()
104 }
105}
106
107impl SqlParam for Vec<u8> {
108 #[inline]
109 fn encode(&self) -> Vec<u8> {
110 self.clone()
111 }
112}
113
114impl SqlParam for &[u8] {
115 #[inline]
116 fn encode(&self) -> Vec<u8> {
117 self.to_vec()
118 }
119}
120
121#[macro_export]
139macro_rules! params {
140 () => {
141 &[] as &[Option<Vec<u8>>]
142 };
143 ($($val:expr),+ $(,)?) => {{
144 use $crate::client::prepare::SqlParam;
145 vec![$(Some($val.encode())),+]
146 }};
147}
148
149static STATEMENT_COUNTER: AtomicU64 = AtomicU64::new(0);
151
152fn generate_statement_name() -> String {
154 let id = STATEMENT_COUNTER.fetch_add(1, Ordering::Relaxed);
155 format!("__hyper_stmt_{id}")
156}
157
158#[derive(Debug)]
178pub struct PreparedStatement {
179 name: String,
181 query: String,
183 param_types: Vec<Oid>,
185 columns: Vec<Column>,
187}
188
189#[derive(Debug)]
208pub struct OwnedPreparedStatement {
209 statement: PreparedStatement,
211 connection: Weak<Mutex<RawConnection<SyncStream>>>,
213}
214
215impl OwnedPreparedStatement {
216 pub(crate) fn new(
218 statement: PreparedStatement,
219 connection: &Arc<Mutex<RawConnection<SyncStream>>>,
220 ) -> Self {
221 OwnedPreparedStatement {
222 statement,
223 connection: Arc::downgrade(connection),
224 }
225 }
226
227 #[must_use]
229 pub fn name(&self) -> &str {
230 self.statement.name()
231 }
232
233 #[must_use]
235 pub fn query(&self) -> &str {
236 self.statement.query()
237 }
238
239 #[must_use]
241 pub fn param_types(&self) -> &[Oid] {
242 self.statement.param_types()
243 }
244
245 #[must_use]
247 pub fn param_count(&self) -> usize {
248 self.statement.param_count()
249 }
250
251 #[must_use]
253 pub fn columns(&self) -> &[Column] {
254 self.statement.columns()
255 }
256
257 #[must_use]
259 pub fn column_count(&self) -> usize {
260 self.statement.column_count()
261 }
262
263 #[must_use]
265 pub fn statement(&self) -> &PreparedStatement {
266 &self.statement
267 }
268
269 pub fn close(self) -> Result<()> {
282 if let Some(conn) = self.connection.upgrade() {
283 close_statement(&conn, &self.statement)?;
284 }
285 std::mem::forget(self);
287 Ok(())
288 }
289}
290
291impl Drop for OwnedPreparedStatement {
292 fn drop(&mut self) {
293 if let Some(conn) = self.connection.upgrade()
295 && let Err(e) = close_statement_internal(&conn, &self.statement)
296 {
297 warn!(
298 target: "hyperdb_api",
299 statement_name = %self.statement.name,
300 error = %e,
301 "failed-to-close-prepared-statement-during-drop"
302 );
303 }
304 }
307}
308
309impl PreparedStatement {
310 #[must_use]
312 pub fn name(&self) -> &str {
313 &self.name
314 }
315
316 #[must_use]
318 pub fn query(&self) -> &str {
319 &self.query
320 }
321
322 #[must_use]
324 pub fn param_types(&self) -> &[Oid] {
325 &self.param_types
326 }
327
328 #[must_use]
330 pub fn param_count(&self) -> usize {
331 self.param_types.len()
332 }
333
334 #[must_use]
336 pub fn columns(&self) -> &[Column] {
337 &self.columns
338 }
339
340 #[must_use]
342 pub fn column_count(&self) -> usize {
343 self.columns.len()
344 }
345}
346
347pub fn prepare(
357 connection: &Arc<Mutex<RawConnection<SyncStream>>>,
358 query: &str,
359 param_types: &[Oid],
360) -> Result<PreparedStatement> {
361 let name = generate_statement_name();
362 let mut conn = connection
363 .lock()
364 .map_err(|_| Error::connection("connection mutex poisoned"))?;
365
366 frontend::parse(&name, query, param_types, conn.write_buf())?;
368
369 frontend::describe(b'S', &name, conn.write_buf())?;
371
372 frontend::sync(conn.write_buf());
374 conn.flush()?;
375
376 let mut parsed_params = Vec::new();
378 let mut parsed_columns = Vec::new();
379
380 loop {
381 let msg = conn.read_message()?;
382 match msg {
383 Message::ParseComplete => {
384 }
386 Message::ParameterDescription(desc) => {
387 for oid in desc.parameters().filter_map(|r| {
388 r.map_err(|e| trace!(target: "hyperdb_api_core::client", error = %e, "dropped error parsing parameter OID")).ok()
389 }) {
390 parsed_params.push(oid);
391 }
392 }
393 Message::RowDescription(desc) => {
394 for f in desc.fields().filter_map(|r| {
395 r.map_err(|e| trace!(target: "hyperdb_api_core::client", error = %e, "dropped error parsing row description field")).ok()
396 }) {
397 parsed_columns.push(Column::new(
398 f.name().to_string(),
399 f.type_oid(),
400 f.type_modifier(),
401 super::statement::ColumnFormat::from_code(f.format()),
402 ));
403 }
404 }
405 Message::NoData => {
406 }
408 Message::ReadyForQuery(_) => {
409 break;
410 }
411 Message::ErrorResponse(body) => {
412 return Err(conn.consume_error(&body));
413 }
414 _ => {}
415 }
416 }
417
418 Ok(PreparedStatement {
419 name,
420 query: query.to_string(),
421 param_types: parsed_params,
422 columns: parsed_columns,
423 })
424}
425
426pub fn execute_prepared(
447 connection: &Arc<Mutex<RawConnection<SyncStream>>>,
448 statement: &PreparedStatement,
449 params: &[Option<&[u8]>],
450) -> Result<Vec<Row>> {
451 let param_formats = bind_format_codes(&[], params.len())?;
453 let result_formats: Vec<i16> = vec![1; statement.columns.len()]; let mut conn = connection
456 .lock()
457 .map_err(|_| Error::connection("connection mutex poisoned"))?;
458
459 frontend::bind(
460 "", &statement.name,
462 ¶m_formats,
463 params,
464 &result_formats,
465 conn.write_buf(),
466 )?;
467
468 frontend::execute("", 0, conn.write_buf())?; frontend::sync(conn.write_buf());
473 conn.flush()?;
474
475 let mut rows = Vec::new();
477 let columns = Arc::new(statement.columns.clone());
478
479 loop {
480 let msg = conn.read_message()?;
481 match msg {
482 Message::BindComplete => {
483 }
485 Message::DataRow(data) => {
486 rows.push(Row::new(Arc::clone(&columns), data)?);
487 }
488 Message::CommandComplete(_) => {
489 }
491 Message::EmptyQueryResponse => {
492 }
494 Message::ReadyForQuery(_) => {
495 break;
496 }
497 Message::ErrorResponse(body) => {
498 return Err(conn.consume_error(&body));
499 }
500 _ => {}
501 }
502 }
503
504 Ok(rows)
505}
506
507pub fn execute_prepared_no_result(
514 connection: &Arc<Mutex<RawConnection<SyncStream>>>,
515 statement: &PreparedStatement,
516 params: &[Option<&[u8]>],
517) -> Result<u64> {
518 execute_prepared_no_result_with_formats(connection, statement, params, &[])
519}
520
521pub fn execute_prepared_no_result_with_formats(
533 connection: &Arc<Mutex<RawConnection<SyncStream>>>,
534 statement: &PreparedStatement,
535 params: &[Option<&[u8]>],
536 param_formats: &[ParamFormat],
537) -> Result<u64> {
538 let param_format_codes = bind_format_codes(param_formats, params.len())?;
539
540 let mut conn = connection
541 .lock()
542 .map_err(|_| Error::connection("connection mutex poisoned"))?;
543
544 let result_formats: Vec<i16> = vec![];
545
546 frontend::bind(
547 "",
548 &statement.name,
549 ¶m_format_codes,
550 params,
551 &result_formats,
552 conn.write_buf(),
553 )?;
554
555 frontend::execute("", 0, conn.write_buf())?;
557
558 frontend::sync(conn.write_buf());
560 conn.flush()?;
561
562 let mut affected_rows = 0u64;
564
565 loop {
566 let msg = conn.read_message()?;
567 match msg {
568 Message::BindComplete => {}
569 Message::CommandComplete(body) => {
570 if let Ok(tag) = body.tag() {
571 affected_rows = parse_affected_rows(tag);
572 }
573 }
574 Message::EmptyQueryResponse => {}
575 Message::ReadyForQuery(_) => {
576 break;
577 }
578 Message::ErrorResponse(body) => {
579 return Err(conn.consume_error(&body));
580 }
581 _ => {}
582 }
583 }
584
585 Ok(affected_rows)
586}
587
588pub fn close_statement(
598 connection: &Arc<Mutex<RawConnection<SyncStream>>>,
599 statement: &PreparedStatement,
600) -> Result<()> {
601 close_statement_internal(connection, statement)
602}
603
604fn close_statement_internal(
606 connection: &Arc<Mutex<RawConnection<SyncStream>>>,
607 statement: &PreparedStatement,
608) -> Result<()> {
609 let mut conn = connection
610 .lock()
611 .map_err(|_| Error::connection("connection mutex poisoned"))?;
612
613 frontend::close(b'S', &statement.name, conn.write_buf())?;
615
616 frontend::sync(conn.write_buf());
618 conn.flush()?;
619
620 loop {
622 let msg = conn.read_message()?;
623 match msg {
624 Message::CloseComplete => {}
625 Message::ReadyForQuery(_) => {
626 break;
627 }
628 Message::ErrorResponse(body) => {
629 return Err(conn.consume_error(&body));
630 }
631 _ => {}
632 }
633 }
634
635 Ok(())
636}
637
638pub fn prepare_owned(
644 connection: &Arc<Mutex<RawConnection<SyncStream>>>,
645 query: &str,
646 param_types: &[Oid],
647) -> Result<OwnedPreparedStatement> {
648 let statement = prepare(connection, query, param_types)?;
649 Ok(OwnedPreparedStatement::new(statement, connection))
650}
651
652fn parse_affected_rows(tag: &str) -> u64 {
654 let parts: Vec<&str> = tag.split_whitespace().collect();
655
656 match parts.first() {
657 Some(&"INSERT") => parts.get(2).and_then(|s| s.parse().ok()).unwrap_or(0),
658 Some(&"UPDATE" | &"DELETE" | &"SELECT" | &"COPY") => {
659 parts.get(1).and_then(|s| s.parse().ok()).unwrap_or(0)
660 }
661 _ => 0,
662 }
663}
664
665#[cfg(test)]
666mod tests {
667 use super::*;
668
669 #[test]
670 fn test_sql_param_i16() {
671 assert_eq!(0_i16.encode(), vec![0, 0]);
672 assert_eq!(1_i16.encode(), vec![1, 0]);
673 assert_eq!((-1_i16).encode(), vec![255, 255]);
674 }
675
676 #[test]
677 fn test_sql_param_i32() {
678 assert_eq!(0_i32.encode(), vec![0, 0, 0, 0]);
679 assert_eq!(1_i32.encode(), vec![1, 0, 0, 0]);
680 assert_eq!((-1_i32).encode(), vec![255, 255, 255, 255]);
681 assert_eq!(256_i32.encode(), vec![0, 1, 0, 0]);
682 }
683
684 #[test]
685 fn test_sql_param_i64() {
686 assert_eq!(0_i64.encode(), vec![0, 0, 0, 0, 0, 0, 0, 0]);
687 assert_eq!(1_i64.encode(), vec![1, 0, 0, 0, 0, 0, 0, 0]);
688 assert_eq!(
689 (-1_i64).encode(),
690 vec![255, 255, 255, 255, 255, 255, 255, 255]
691 );
692 }
693
694 #[test]
695 #[expect(
696 clippy::float_cmp,
697 reason = "1.5 is exactly representable; encode/decode must round-trip bit-for-bit"
698 )]
699 fn test_sql_param_f32() {
700 let encoded = 1.5_f32.encode();
701 assert_eq!(encoded.len(), 4);
702 let decoded = f32::from_le_bytes([encoded[0], encoded[1], encoded[2], encoded[3]]);
703 assert_eq!(decoded, 1.5);
704 }
705
706 #[test]
707 #[expect(
708 clippy::float_cmp,
709 reason = "1.5 is exactly representable; encode/decode must round-trip bit-for-bit"
710 )]
711 fn test_sql_param_f64() {
712 let encoded = 1.5_f64.encode();
713 assert_eq!(encoded.len(), 8);
714 let decoded = f64::from_le_bytes([
715 encoded[0], encoded[1], encoded[2], encoded[3], encoded[4], encoded[5], encoded[6],
716 encoded[7],
717 ]);
718 assert_eq!(decoded, 1.5);
719 }
720
721 #[test]
722 fn test_sql_param_bool() {
723 assert_eq!(true.encode(), vec![1]);
724 assert_eq!(false.encode(), vec![0]);
725 }
726
727 #[test]
728 fn test_sql_param_str() {
729 assert_eq!("hello".encode(), b"hello".to_vec());
730 assert_eq!("".encode(), Vec::<u8>::new());
731 assert_eq!("héllo".encode(), "héllo".as_bytes().to_vec());
732 }
733
734 #[test]
735 fn test_sql_param_string() {
736 let s = String::from("hello");
737 assert_eq!(s.encode(), b"hello".to_vec());
738 assert_eq!(s.encode(), b"hello".to_vec());
739 }
740
741 #[test]
742 fn test_sql_param_bytes() {
743 let bytes: Vec<u8> = vec![1, 2, 3, 4];
744 assert_eq!(bytes.encode(), vec![1, 2, 3, 4]);
745 assert_eq!(bytes.as_slice().encode(), vec![1, 2, 3, 4]);
746 }
747
748 #[test]
749 fn test_params_macro_empty() {
750 let p = params![];
751 assert!(p.is_empty());
752 }
753
754 #[test]
755 fn test_params_macro_single() {
756 let p = params![42_i32];
757 assert_eq!(p.len(), 1);
758 assert_eq!(p[0], Some(vec![42, 0, 0, 0]));
759 }
760
761 #[test]
762 fn test_params_macro_multiple() {
763 let p = params![42_i32, "hello", true];
764 assert_eq!(p.len(), 3);
765 assert_eq!(p[0], Some(vec![42, 0, 0, 0]));
766 assert_eq!(p[1], Some(b"hello".to_vec()));
767 assert_eq!(p[2], Some(vec![1]));
768 }
769}