1#![allow(warnings)]
2use std::collections::BTreeMap;
3use std::future::Future;
4use std::pin::Pin;
5
6use chrono::{DateTime, NaiveDate, Utc};
7use deadpool_postgres::Pool;
8use rust_decimal::Decimal;
9use std::sync::Arc;
10use teaql_core::{
11 BinaryOp, DataType, EntityDescriptor, Expr, InsertCommand, PropertyDescriptor, Record,
12 SelectQuery, UpdateCommand, Value,
13};
14use teaql_runtime::{GraphNode, InternalIdGenerator, RuntimeError, SchemaProvider, UserContext};
15use teaql_sql::{
16 CompiledQuery, DatabaseKind, SqlCompileError, SqlDialect, SqlTransport,
17 quote_identifier_if_needed,
18};
19use tokio::sync::Mutex;
20
21pub const DEFAULT_ID_SPACE_TABLE: &str = "teaql_id_space";
22
23#[derive(Debug, Default, Clone, Copy)]
24pub struct PostgresDialect;
25
26impl SqlDialect for PostgresDialect {
27 fn kind(&self) -> DatabaseKind {
28 DatabaseKind::PostgreSql
29 }
30
31 fn quote_ident(&self, ident: &str) -> String {
32 quote_ident(ident)
33 }
34
35 fn placeholder(&self, index: usize) -> String {
36 format!("${index}")
37 }
38
39 fn schema_setup_sqls(&self) -> &'static [&'static str] {
40 &[CREATE_SOUNDEX_FUNCTION]
41 }
42
43 fn schema_type_sql(
44 &self,
45 data_type: DataType,
46 _property: &PropertyDescriptor,
47 ) -> Result<&'static str, SqlCompileError> {
48 match data_type {
49 DataType::Bool => Ok("BOOLEAN"),
50 DataType::I64 | DataType::U64 => Ok("BIGINT"),
51 DataType::F64 => Ok("DOUBLE PRECISION"),
52 DataType::Decimal => Ok("NUMERIC"),
53 DataType::Text => Ok("VARCHAR(255)"),
54 DataType::LargeText => Ok("TEXT"),
55 DataType::Json => Ok("JSONB"),
56 DataType::Date => Ok("DATE"),
57 DataType::Timestamp => Ok("TIMESTAMPTZ"),
58 }
59 }
60
61 fn compile_in(
62 &self,
63 entity: &EntityDescriptor,
64 left: &Expr,
65 op: BinaryOp,
66 right: &Expr,
67 params: &mut Vec<Value>,
68 ) -> Result<String, SqlCompileError> {
69 match op {
70 BinaryOp::InLarge | BinaryOp::NotInLarge => {
71 let Expr::Value(Value::List(values)) = right else {
72 let lhs = self.compile_expr(entity, left, params)?;
73 let rhs = self.compile_expr(entity, right, params)?;
74 let operator = match op {
75 BinaryOp::InLarge => "= ANY",
76 BinaryOp::NotInLarge => "<> ALL",
77 _ => unreachable!(),
78 };
79 return Ok(format!("({lhs} {operator} ({rhs}))"));
80 };
81 if values.is_empty() {
82 return Err(SqlCompileError::EmptyInList);
83 }
84 let lhs = self.compile_expr(entity, left, params)?;
85 params.push(Value::List(values.clone()));
86 let placeholder = self.placeholder(params.len());
87 let operator = match op {
88 BinaryOp::InLarge => "= ANY",
89 BinaryOp::NotInLarge => "<> ALL",
90 _ => unreachable!(),
91 };
92 Ok(format!("({lhs} {operator}({placeholder}))"))
93 }
94 _ => {
95 let lhs = self.compile_expr(entity, left, params)?;
96 let operator = match op {
97 BinaryOp::In => "IN",
98 BinaryOp::NotIn => "NOT IN",
99 _ => unreachable!(),
100 };
101 match right {
102 Expr::Value(Value::List(values)) => {
103 if values.is_empty() {
104 return Err(SqlCompileError::EmptyInList);
105 }
106 let mut placeholders = Vec::with_capacity(values.len());
107 for value in values {
108 params.push(value.clone());
109 placeholders.push(self.placeholder(params.len()));
110 }
111 Ok(format!("({lhs} {operator} ({}))", placeholders.join(", ")))
112 }
113 _ => {
114 let rhs = self.compile_expr(entity, right, params)?;
115 Ok(format!("({lhs} {operator} ({rhs}))"))
116 }
117 }
118 }
119 }
120 }
121}
122
123const CREATE_SOUNDEX_FUNCTION: &str = r#"
124CREATE OR REPLACE FUNCTION soundex(input text)
125RETURNS text
126LANGUAGE plpgsql
127IMMUTABLE
128STRICT
129AS $$
130DECLARE
131 normalized text := upper(regexp_replace(input, '[^A-Za-z]', '', 'g'));
132 first_char text;
133 output text;
134 previous_code text;
135 code text;
136 ch text;
137 i integer;
138BEGIN
139 IF normalized = '' THEN
140 RETURN '0000';
141 END IF;
142
143 first_char := substr(normalized, 1, 1);
144 output := first_char;
145 previous_code := CASE
146 WHEN first_char IN ('B', 'F', 'P', 'V') THEN '1'
147 WHEN first_char IN ('C', 'G', 'J', 'K', 'Q', 'S', 'X', 'Z') THEN '2'
148 WHEN first_char IN ('D', 'T') THEN '3'
149 WHEN first_char = 'L' THEN '4'
150 WHEN first_char IN ('M', 'N') THEN '5'
151 WHEN first_char = 'R' THEN '6'
152 ELSE '0'
153 END;
154
155 FOR i IN 2..char_length(normalized) LOOP
156 ch := substr(normalized, i, 1);
157 code := CASE
158 WHEN ch IN ('B', 'F', 'P', 'V') THEN '1'
159 WHEN ch IN ('C', 'G', 'J', 'K', 'Q', 'S', 'X', 'Z') THEN '2'
160 WHEN ch IN ('D', 'T') THEN '3'
161 WHEN ch = 'L' THEN '4'
162 WHEN ch IN ('M', 'N') THEN '5'
163 WHEN ch = 'R' THEN '6'
164 ELSE '0'
165 END;
166
167 IF code <> '0' AND code <> previous_code THEN
168 output := output || code;
169 IF char_length(output) = 4 THEN
170 RETURN output;
171 END IF;
172 END IF;
173 previous_code := code;
174 END LOOP;
175
176 RETURN rpad(output, 4, '0');
177END;
178$$
179"#;
180
181#[derive(Debug)]
182pub enum MutationExecutorError {
183 Driver(tokio_postgres::Error),
184 Pool(String),
185 SqlCompile(SqlCompileError),
186 UnsupportedValue(&'static str),
187 UnsupportedColumnType(String),
188 Bind(String),
189}
190
191impl std::fmt::Display for MutationExecutorError {
192 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
193 match self {
194 Self::Driver(err) => err.fmt(f),
195 Self::Pool(err) => write!(f, "postgres pool error: {err}"),
196 Self::SqlCompile(err) => err.fmt(f),
197 Self::UnsupportedValue(kind) => {
198 write!(f, "unsupported bind value for mutation executor: {kind}")
199 }
200 Self::UnsupportedColumnType(kind) => {
201 write!(f, "unsupported column type for record decoding: {kind}")
202 }
203 Self::Bind(message) => write!(f, "bind error: {message}"),
204 }
205 }
206}
207
208impl std::error::Error for MutationExecutorError {}
209
210impl From<tokio_postgres::Error> for MutationExecutorError {
211 fn from(value: tokio_postgres::Error) -> Self {
212 Self::Driver(value)
213 }
214}
215
216impl From<SqlCompileError> for MutationExecutorError {
217 fn from(value: SqlCompileError) -> Self {
218 Self::SqlCompile(value)
219 }
220}
221
222#[derive(Clone)]
223pub struct PgMutationExecutor {
224 pool: Pool,
225}
226
227impl SqlTransport for PgMutationExecutor {
228 type Error = MutationExecutorError;
229
230 async fn fetch_all_sql(&self, query: &CompiledQuery) -> Result<Vec<Record>, Self::Error> {
231 self.fetch_all(query).await
232 }
233
234 async fn execute_sql(&self, query: &CompiledQuery) -> Result<u64, Self::Error> {
235 self.execute(query).await
236 }
237}
238
239impl teaql_sql::SqlTransaction for PgMutationExecutor {
240 type Error = MutationExecutorError;
241
242 async fn commit_sql(self) -> Result<(), Self::Error> {
243 Err(MutationExecutorError::Bind(
244 "Transactions not supported yet".to_string(),
245 ))
246 }
247
248 async fn rollback_sql(self) -> Result<(), Self::Error> {
249 Err(MutationExecutorError::Bind(
250 "Transactions not supported yet".to_string(),
251 ))
252 }
253}
254
255impl teaql_sql::SqlTransactionTransport for PgMutationExecutor {
256 type Tx<'a>
257 = Self
258 where
259 Self: 'a;
260
261 async fn begin_sql(&self) -> Result<Self::Tx<'_>, Self::Error> {
262 Err(MutationExecutorError::Bind(
263 "Transactions not supported yet".to_string(),
264 ))
265 }
266}
267
268impl PgMutationExecutor {
269 pub fn new(pool: Pool) -> Self {
270 Self { pool }
271 }
272
273 pub fn pool(&self) -> Pool {
274 self.pool.clone()
275 }
276
277 pub async fn ensure_schema(
278 &self,
279 dialect: &PostgresDialect,
280 entities: &[&EntityDescriptor],
281 ) -> Result<(), MutationExecutorError> {
282 let client = self
283 .pool
284 .get()
285 .await
286 .map_err(|e| MutationExecutorError::Pool(e.to_string()))?;
287 for sql in dialect.schema_setup_sqls() {
288 client.execute(*sql, &[]).await?;
289 }
290 self.ensure_id_space_table(DEFAULT_ID_SPACE_TABLE).await?;
291
292 for entity in entities {
293 if !self.table_exists(&entity.table_name).await? {
294 let sql = dialect.compile_create_table(entity)?;
295 client.execute(&sql, &[]).await?;
296 continue;
297 }
298
299 let existing_columns = self.table_columns(&entity.table_name).await?;
300 for property in &entity.properties {
301 let bare_column = strip_identifier_quotes(&property.column_name).to_lowercase();
302 if existing_columns.contains(&bare_column) {
303 continue;
304 }
305 let sql = dialect.compile_add_column(entity, property)?;
306 client.execute(&sql, &[]).await?;
307 }
308
309 for sql in dialect.schema_indexes_sqls(entity)? {
310 client.execute(&sql, &[]).await?;
311 }
312 }
313 Ok(())
314 }
315
316 pub async fn ensure_id_space_table(
317 &self,
318 table_name: &str,
319 ) -> Result<(), MutationExecutorError> {
320 let sql = format!(
321 "CREATE TABLE IF NOT EXISTS {} (type_name VARCHAR(100) PRIMARY KEY, current_level BIGINT NOT NULL)",
322 quote_ident(table_name)
323 );
324 let client = self
325 .pool
326 .get()
327 .await
328 .map_err(|e| MutationExecutorError::Pool(e.to_string()))?;
329 client.execute(&sql, &[]).await?;
330 Ok(())
331 }
332
333 pub async fn execute(&self, query: &CompiledQuery) -> Result<u64, MutationExecutorError> {
334 let mut args = PgArgs { values: Vec::new() };
335 for value in &query.params {
336 bind_pg(&mut args, value)?;
337 }
338 let client = self
339 .pool
340 .get()
341 .await
342 .map_err(|e| MutationExecutorError::Pool(e.to_string()))?;
343 let result = client.execute(&query.sql, &args.as_refs()).await?;
344 Ok(result)
345 }
346
347 pub async fn fetch_all(
348 &self,
349 query: &CompiledQuery,
350 ) -> Result<Vec<Record>, MutationExecutorError> {
351 let mut args = PgArgs { values: Vec::new() };
352 for value in &query.params {
353 bind_pg(&mut args, value)?;
354 }
355 let client = self
356 .pool
357 .get()
358 .await
359 .map_err(|e| MutationExecutorError::Pool(e.to_string()))?;
360 let rows = client.query(&query.sql, &args.as_refs()).await?;
361 rows.iter().map(decode_pg_row).collect()
362 }
363
364 async fn table_exists(&self, table_name: &str) -> Result<bool, MutationExecutorError> {
365 let client = self
366 .pool
367 .get()
368 .await
369 .map_err(|e| MutationExecutorError::Pool(e.to_string()))?;
370 let row = client
371 .query_one(
372 "SELECT COUNT(1)
373 FROM information_schema.tables
374 WHERE table_schema = current_schema()
375 AND table_name = $1",
376 &[&table_name],
377 )
378 .await?;
379 let exists: i64 = row.try_get(0)?;
380 Ok(exists > 0)
381 }
382
383 async fn table_columns(
384 &self,
385 table_name: &str,
386 ) -> Result<std::collections::BTreeSet<String>, MutationExecutorError> {
387 let client = self
388 .pool
389 .get()
390 .await
391 .map_err(|e| MutationExecutorError::Pool(e.to_string()))?;
392 let rows = client
393 .query(
394 "SELECT column_name
395 FROM information_schema.columns
396 WHERE table_schema = current_schema()
397 AND table_name = $1",
398 &[&table_name],
399 )
400 .await?;
401 let mut columns = std::collections::BTreeSet::new();
402 for row in rows {
403 let name: String = row.try_get("column_name")?;
404 columns.insert(name.to_lowercase());
405 }
406 Ok(columns)
407 }
408}
409
410async fn ensure_initial_graphs_postgres(
411 executor: &PgMutationExecutor,
412 dialect: &PostgresDialect,
413 ctx: &UserContext,
414) -> Result<(), MutationExecutorError> {
415 for graph in ctx.initial_graphs() {
416 let entity = ctx.entity(&graph.entity).ok_or_else(|| {
417 MutationExecutorError::Bind(format!("missing entity: {}", graph.entity))
418 })?;
419 if initial_graph_exists_postgres(executor, dialect, entity, graph).await? {
420 if let Some(query) = compile_initial_graph_update(dialect, entity, graph)? {
421 executor.execute(&query).await?;
422 }
423 continue;
424 }
425 let query = compile_initial_graph_insert(dialect, entity, graph)?;
426 executor.execute(&query).await?;
427 }
428 Ok(())
429}
430
431async fn initial_graph_exists_postgres(
432 executor: &PgMutationExecutor,
433 dialect: &PostgresDialect,
434 entity: &EntityDescriptor,
435 graph: &GraphNode,
436) -> Result<bool, MutationExecutorError> {
437 let Some(id) = graph.values.get("id") else {
438 return Ok(false);
439 };
440 let query = dialect.compile_select(
441 entity,
442 &SelectQuery::new(&graph.entity)
443 .project("id")
444 .filter(Expr::eq("id", id.clone()))
445 .limit(1),
446 )?;
447 Ok(!executor.fetch_all(&query).await?.is_empty())
448}
449
450fn compile_initial_graph_insert(
451 dialect: &impl SqlDialect,
452 entity: &EntityDescriptor,
453 graph: &GraphNode,
454) -> Result<CompiledQuery, MutationExecutorError> {
455 let mut command = InsertCommand::new(&graph.entity);
456 for (field, value) in &graph.values {
457 command = command.value(field.clone(), value.clone());
458 }
459 dialect.compile_insert(entity, &command).map_err(Into::into)
460}
461
462fn compile_initial_graph_update(
463 dialect: &impl SqlDialect,
464 entity: &EntityDescriptor,
465 graph: &crate::GraphNode,
466) -> Result<Option<CompiledQuery>, MutationExecutorError> {
467 let Some(id) = graph.values.get("id") else {
468 return Ok(None);
469 };
470 let mut command = UpdateCommand::new(&graph.entity, id.clone());
471 for (field, value) in &graph.values {
472 if field == "id" {
473 continue;
474 }
475 command = command.value(field.clone(), value.clone());
476 }
477 match dialect.compile_update(entity, &command) {
478 Ok(query) => Ok(Some(query)),
479 Err(SqlCompileError::EmptyMutation(_)) => Ok(None),
480 Err(err) => Err(err.into()),
481 }
482}
483
484pub trait PostgresSchemaExt {
485 fn ensure_postgres_schema(
486 &self,
487 ) -> Pin<Box<dyn Future<Output = Result<(), MutationExecutorError>> + '_>>;
488}
489
490pub async fn ensure_postgres_schema_for(ctx: &UserContext) -> Result<(), MutationExecutorError> {
491 let dialect = ctx.get_resource::<PostgresDialect>().ok_or_else(|| {
492 MutationExecutorError::Bind("missing typed resource: PostgresDialect".to_owned())
493 })?;
494 let executor = ctx.get_resource::<PgMutationExecutor>().ok_or_else(|| {
495 MutationExecutorError::Bind("missing typed resource: PgMutationExecutor".to_owned())
496 })?;
497
498 let entities = ctx.all_entities();
499
500 executor.ensure_schema(dialect, &entities).await?;
501 ensure_initial_graphs_postgres(executor, dialect, ctx).await
502}
503
504impl PostgresSchemaExt for UserContext {
505 fn ensure_postgres_schema(
506 &self,
507 ) -> Pin<Box<dyn Future<Output = Result<(), MutationExecutorError>> + '_>> {
508 Box::pin(ensure_postgres_schema_for(self))
509 }
510}
511
512#[derive(Debug, Default, Clone, Copy)]
513pub struct PostgresSchemaProvider;
514
515impl SchemaProvider for PostgresSchemaProvider {
516 fn ensure_schema<'a>(
517 &'a self,
518 ctx: &'a UserContext,
519 ) -> Pin<Box<dyn Future<Output = Result<(), RuntimeError>> + Send + 'a>> {
520 Box::pin(async move {
521 ensure_postgres_schema_for(ctx)
522 .await
523 .map_err(|err| RuntimeError::Schema(err.to_string()))
524 })
525 }
526}
527
528pub trait PostgresProviderExt {
529 fn use_postgres_provider(&mut self, executor: PgMutationExecutor) -> &mut Self;
530}
531
532impl PostgresProviderExt for UserContext {
533 fn use_postgres_provider(&mut self, executor: PgMutationExecutor) -> &mut Self {
534 self.insert_resource(PostgresDialect);
535 self.insert_resource(executor);
536 self.set_schema_provider(PostgresSchemaProvider);
537 self
538 }
539}
540
541#[derive(Clone)]
542pub struct PgIdSpaceGenerator {
543 pool: Pool,
544 table_name: String,
545}
546
547impl PgIdSpaceGenerator {
548 pub fn new(pool: Pool) -> Self {
549 Self {
550 pool,
551 table_name: DEFAULT_ID_SPACE_TABLE.to_owned(),
552 }
553 }
554
555 pub fn from_executor(executor: PgMutationExecutor) -> Self {
556 Self::new(executor.pool())
557 }
558
559 pub fn with_table_name(mut self, table_name: impl Into<String>) -> Self {
560 self.table_name = table_name.into();
561 self
562 }
563
564 pub async fn ensure_table(&self) -> Result<(), MutationExecutorError> {
565 PgMutationExecutor::new(self.pool.clone())
566 .ensure_id_space_table(&self.table_name)
567 .await
568 }
569
570 pub async fn next_id(&self, entity: &str) -> Result<u64, MutationExecutorError> {
571 self.ensure_table().await?;
572 let update_sql = format!(
573 "UPDATE {} SET current_level = current_level + 1 WHERE type_name = $1 RETURNING current_level",
574 quote_ident(&self.table_name)
575 );
576 let client = self
577 .pool
578 .get()
579 .await
580 .map_err(|e| MutationExecutorError::Pool(e.to_string()))?;
581 let row = client.query_opt(&update_sql, &[&entity]).await?;
582
583 let id = match row {
584 Some(r) => {
585 let level: i64 = r.try_get(0)?;
586 level
587 }
588 None => {
589 let insert_sql = format!(
590 "INSERT INTO {} (type_name, current_level) VALUES ($1, 1) RETURNING current_level",
591 quote_ident(&self.table_name)
592 );
593 let insert_res = client.query_one(&insert_sql, &[&entity]).await;
594 match insert_res {
595 Ok(r) => {
596 let level: i64 = r.try_get(0)?;
597 level
598 }
599 Err(_) => {
600 let row = client.query_one(&update_sql, &[&entity]).await?;
601 let level: i64 = row.try_get(0)?;
602 level
603 }
604 }
605 }
606 };
607
608 u64::try_from(id).map_err(|_| {
609 MutationExecutorError::Bind(format!("generated id {id} cannot be represented as u64"))
610 })
611 }
612}
613
614impl InternalIdGenerator for PgIdSpaceGenerator {
615 fn generate_id(&self, entity: &str) -> Result<u64, RuntimeError> {
616 let generator = self.clone();
617 let entity = entity.to_owned();
618 block_on_id_generation(async move { generator.next_id(&entity).await })
619 }
620}
621
622fn block_on_id_generation<F>(future: F) -> Result<u64, RuntimeError>
623where
624 F: Future<Output = Result<u64, MutationExecutorError>> + Send + 'static,
625{
626 let result = match tokio::runtime::Handle::try_current() {
627 Ok(handle) => tokio::task::block_in_place(|| handle.block_on(future)),
628 Err(_) => tokio::runtime::Builder::new_current_thread()
629 .enable_all()
630 .build()
631 .map_err(|err| RuntimeError::IdGeneration(err.to_string()))?
632 .block_on(future),
633 };
634 result.map_err(|err| RuntimeError::IdGeneration(err.to_string()))
635}
636
637fn quote_ident(ident: &str) -> String {
638 quote_identifier_if_needed(ident, '"')
639}
640
641fn strip_identifier_quotes(ident: &str) -> &str {
645 let bytes = ident.as_bytes();
646 if bytes.len() >= 2 {
647 let (first, last) = (bytes[0], bytes[bytes.len() - 1]);
648 if (first == b'"' && last == b'"')
649 || (first == b'`' && last == b'`')
650 || (first == b'[' && last == b']')
651 {
652 return &ident[1..ident.len() - 1];
653 }
654 }
655 ident
656}
657
658fn try_parse_datetime_from_str(s: &str) -> Option<chrono::DateTime<chrono::Utc>> {
659 if let Ok(dt) = chrono::DateTime::parse_from_rfc3339(s) {
660 return Some(dt.with_timezone(&chrono::Utc));
661 }
662 if let Ok(ndt) = chrono::NaiveDateTime::parse_from_str(s, "%Y-%m-%d %H:%M:%S") {
663 return Some(chrono::DateTime::from_naive_utc_and_offset(
664 ndt,
665 chrono::Utc,
666 ));
667 }
668 if let Ok(nd) = chrono::NaiveDate::parse_from_str(s, "%Y-%m-%d") {
669 let ndt = nd.and_hms_opt(0, 0, 0)?;
670 return Some(chrono::DateTime::from_naive_utc_and_offset(
671 ndt,
672 chrono::Utc,
673 ));
674 }
675 None
676}
677
678#[derive(Debug, Clone, Copy)]
679struct PgNull;
680
681impl tokio_postgres::types::ToSql for PgNull {
682 fn to_sql(
683 &self,
684 ty: &tokio_postgres::types::Type,
685 out: &mut bytes::BytesMut,
686 ) -> Result<tokio_postgres::types::IsNull, Box<dyn std::error::Error + Sync + Send>> {
687 Ok(tokio_postgres::types::IsNull::Yes)
688 }
689
690 fn accepts(ty: &tokio_postgres::types::Type) -> bool {
691 true
692 }
693
694 fn to_sql_checked(
695 &self,
696 ty: &tokio_postgres::types::Type,
697 out: &mut bytes::BytesMut,
698 ) -> Result<tokio_postgres::types::IsNull, Box<dyn std::error::Error + Sync + Send>> {
699 Ok(tokio_postgres::types::IsNull::Yes)
700 }
701}
702
703struct PgArgs {
704 values: Vec<Box<dyn tokio_postgres::types::ToSql + Sync + Send>>,
705}
706impl PgArgs {
707 fn add<T: tokio_postgres::types::ToSql + Sync + Send + 'static>(&mut self, v: T) {
708 self.values.push(Box::new(v));
709 }
710 fn as_refs(&self) -> Vec<&(dyn tokio_postgres::types::ToSql + Sync)> {
711 self.values.iter().map(|b| b.as_ref() as _).collect()
712 }
713}
714
715fn bind_pg(args: &mut PgArgs, value: &Value) -> Result<(), MutationExecutorError> {
716 match value {
717 Value::Null => {
718 args.add(PgNull);
719 }
720 Value::Bool(v) => args.add(*v),
721 Value::I64(v) => args.add(*v),
722 Value::U64(v) => {
723 let v = i64::try_from(*v).map_err(|_| {
724 MutationExecutorError::Bind(format!("u64 value {v} exceeds i64 range"))
725 })?;
726 args.add(v);
727 }
728 Value::F64(v) => args.add(*v),
729 Value::Decimal(v) => args.add(*v),
730 Value::Text(v) => match try_parse_datetime_from_str(v) {
731 Some(dt) => args.add(dt),
732 None => args.add(v.clone()),
733 },
734 Value::Json(v) => {
735 let j_val: serde_json::Value =
736 serde_json::to_value(v).map_err(|e| MutationExecutorError::Bind(e.to_string()))?;
737 args.add(j_val);
738 }
739 Value::Date(v) => args.add(*v),
740 Value::Timestamp(v) => args.add(*v),
741 Value::Object(_) => return Err(MutationExecutorError::UnsupportedValue("object")),
742 Value::List(values) => bind_pg_list(args, values)?,
743 Value::TypedNull(dt) => match dt {
744 DataType::Bool => args.add(Option::<bool>::None),
745 DataType::I64 | DataType::U64 => args.add(Option::<i64>::None),
746 DataType::F64 => args.add(Option::<f64>::None),
747 DataType::Decimal => args.add(Option::<Decimal>::None),
748 DataType::Text | DataType::LargeText => args.add(Option::<String>::None),
749 DataType::Json => args.add(Option::<serde_json::Value>::None),
750 DataType::Date => args.add(Option::<NaiveDate>::None),
751 DataType::Timestamp => args.add(Option::<DateTime<Utc>>::None),
752 },
753 }
754 Ok(())
755}
756
757fn bind_pg_list(args: &mut PgArgs, values: &[Value]) -> Result<(), MutationExecutorError> {
758 let Some(first) = values.first() else {
759 return Err(MutationExecutorError::UnsupportedValue("empty list"));
760 };
761 match first {
762 Value::Bool(_) => {
763 let values = values
764 .iter()
765 .map(|value| match value {
766 Value::Bool(value) => Ok(*value),
767 _ => Err(MutationExecutorError::UnsupportedValue("mixed bool list")),
768 })
769 .collect::<Result<Vec<_>, _>>()?;
770 args.add(values);
771 }
772 Value::I64(_) => {
773 let values = values
774 .iter()
775 .map(|value| match value {
776 Value::I64(value) => Ok(*value),
777 _ => Err(MutationExecutorError::UnsupportedValue("mixed i64 list")),
778 })
779 .collect::<Result<Vec<_>, _>>()?;
780 args.add(values);
781 }
782 Value::U64(_) => {
783 let values = values
784 .iter()
785 .map(|value| match value {
786 Value::U64(value) => i64::try_from(*value).map_err(|_| {
787 MutationExecutorError::Bind(format!("u64 value {value} exceeds i64 range"))
788 }),
789 _ => Err(MutationExecutorError::UnsupportedValue("mixed u64 list")),
790 })
791 .collect::<Result<Vec<_>, _>>()?;
792 args.add(values);
793 }
794 Value::F64(_) => {
795 let values = values
796 .iter()
797 .map(|value| match value {
798 Value::F64(value) => Ok(*value),
799 _ => Err(MutationExecutorError::UnsupportedValue("mixed f64 list")),
800 })
801 .collect::<Result<Vec<_>, _>>()?;
802 args.add(values);
803 }
804 Value::Decimal(_) => {
805 let values = values
806 .iter()
807 .map(|value| match value {
808 Value::Decimal(value) => Ok(*value),
809 _ => Err(MutationExecutorError::UnsupportedValue(
810 "mixed decimal list",
811 )),
812 })
813 .collect::<Result<Vec<_>, _>>()?;
814 args.add(values);
815 }
816 Value::Text(_) => {
817 let values = values
818 .iter()
819 .map(|value| match value {
820 Value::Text(value) => Ok(value.clone()),
821 _ => Err(MutationExecutorError::UnsupportedValue("mixed text list")),
822 })
823 .collect::<Result<Vec<_>, _>>()?;
824 args.add(values);
825 }
826 Value::Date(_) => {
827 let values = values
828 .iter()
829 .map(|value| match value {
830 Value::Date(value) => Ok(*value),
831 _ => Err(MutationExecutorError::UnsupportedValue("mixed date list")),
832 })
833 .collect::<Result<Vec<_>, _>>()?;
834 args.add(values);
835 }
836 Value::Timestamp(_) => {
837 let values = values
838 .iter()
839 .map(|value| match value {
840 Value::Timestamp(value) => Ok(*value),
841 _ => Err(MutationExecutorError::UnsupportedValue(
842 "mixed timestamp list",
843 )),
844 })
845 .collect::<Result<Vec<_>, _>>()?;
846 args.add(values);
847 }
848 Value::Null => return Err(MutationExecutorError::UnsupportedValue("null list")),
849 Value::Json(_) => return Err(MutationExecutorError::UnsupportedValue("json list")),
850 Value::Object(_) => return Err(MutationExecutorError::UnsupportedValue("object list")),
851 Value::List(_) => return Err(MutationExecutorError::UnsupportedValue("nested list")),
852 Value::TypedNull(_) => return Err(MutationExecutorError::UnsupportedValue("null list")),
853 }
854 Ok(())
855}
856
857fn decode_pg_row(row: &tokio_postgres::Row) -> Result<Record, MutationExecutorError> {
858 let mut record = BTreeMap::new();
859 for (index, column) in row.columns().iter().enumerate() {
860 let name = column.name().to_owned();
861 let type_name = column.type_().name().to_ascii_uppercase();
862
863 let value = match type_name.as_str() {
864 "BOOL" | "BOOLEAN" => {
865 let v: Option<bool> = row.try_get(index)?;
866 match v {
867 Some(v) => Value::Bool(v),
868 None => Value::Null,
869 }
870 }
871 "INT2" => {
872 let v: Option<i16> = row.try_get(index)?;
873 match v {
874 Some(v) => Value::I64(v as i64),
875 None => Value::Null,
876 }
877 }
878 "INT4" => {
879 let v: Option<i32> = row.try_get(index)?;
880 match v {
881 Some(v) => Value::I64(v as i64),
882 None => Value::Null,
883 }
884 }
885 "INT8" => {
886 let v: Option<i64> = row.try_get(index)?;
887 match v {
888 Some(v) => Value::I64(v),
889 None => Value::Null,
890 }
891 }
892 "FLOAT4" => {
893 let v: Option<f32> = row.try_get(index)?;
894 match v {
895 Some(v) => Value::F64(v as f64),
896 None => Value::Null,
897 }
898 }
899 "FLOAT8" => {
900 let v: Option<f64> = row.try_get(index)?;
901 match v {
902 Some(v) => Value::F64(v),
903 None => Value::Null,
904 }
905 }
906 "NUMERIC" => {
907 let v: Option<Decimal> = row.try_get(index)?;
908 match v {
909 Some(v) => Value::Decimal(v),
910 None => Value::Null,
911 }
912 }
913 "JSON" | "JSONB" => {
914 let v: Option<serde_json::Value> = row.try_get(index)?;
915 match v {
916 Some(j) => Value::Json(j.into()),
917 None => Value::Null,
918 }
919 }
920 "DATE" => {
921 let v: Option<NaiveDate> = row.try_get(index)?;
922 match v {
923 Some(v) => Value::Date(v),
924 None => Value::Null,
925 }
926 }
927 "TIMESTAMP" | "TIMESTAMPTZ" => {
928 let v: Option<DateTime<Utc>> = row.try_get(index)?;
929 match v {
930 Some(v) => Value::Timestamp(v),
931 None => Value::Null,
932 }
933 }
934 "TEXT" | "VARCHAR" | "BPCHAR" | "NAME" | "UUID" => {
935 let v: Option<String> = row.try_get(index)?;
936 match v {
937 Some(v) => Value::Text(v),
938 None => Value::Null,
939 }
940 }
941 other => {
942 return Err(MutationExecutorError::UnsupportedColumnType(
943 other.to_owned(),
944 ));
945 }
946 };
947 record.insert(name, value);
948 }
949 Ok(record)
950}
951
952#[cfg(test)]
953mod tests {
954 use super::*;
955 use teaql_core::{DeleteCommand, RecoverCommand};
956
957 fn entity() -> EntityDescriptor {
958 EntityDescriptor::new("Order")
959 .table_name("orders")
960 .property(
961 PropertyDescriptor::new("id", DataType::U64)
962 .column_name("id")
963 .id()
964 .not_null(),
965 )
966 .property(
967 PropertyDescriptor::new("version", DataType::I64)
968 .column_name("version")
969 .version()
970 .not_null(),
971 )
972 .property(PropertyDescriptor::new("name", DataType::Text).column_name("name"))
973 }
974
975 #[test]
976 fn postgres_dialect_compiles_mutations_with_numbered_placeholders() {
977 let insert = PostgresDialect
978 .compile_insert(
979 &entity(),
980 &InsertCommand::new("Order")
981 .value("id", 1_u64)
982 .value("name", "A"),
983 )
984 .unwrap();
985 assert_eq!(insert.sql, "INSERT INTO orders (id, name) VALUES ($1, $2)");
986
987 let update = PostgresDialect
988 .compile_update(
989 &entity(),
990 &UpdateCommand::new("Order", 1_u64)
991 .expected_version(3)
992 .value("name", "B"),
993 )
994 .unwrap();
995 assert_eq!(
996 update.sql,
997 "UPDATE orders SET name = $1, version = $2 WHERE id = $3 AND version = $4"
998 );
999
1000 let delete = PostgresDialect
1001 .compile_delete(
1002 &entity(),
1003 &DeleteCommand::new("Order", 1_u64).expected_version(3),
1004 )
1005 .unwrap();
1006 let recover = PostgresDialect
1007 .compile_recover(&entity(), &RecoverCommand::new("Order", 1_u64, -4))
1008 .unwrap();
1009 assert_eq!(
1010 delete.sql,
1011 "UPDATE orders SET version = $1 WHERE id = $2 AND version = $3"
1012 );
1013 assert_eq!(
1014 recover.sql,
1015 "UPDATE orders SET version = $1 WHERE id = $2 AND version = $3"
1016 );
1017 }
1018
1019 #[test]
1020 fn postgres_dialect_compiles_schema_and_large_in_array_binds() {
1021 let create = PostgresDialect.compile_create_table(&entity()).unwrap();
1022 assert_eq!(
1023 create,
1024 "CREATE TABLE IF NOT EXISTS orders (id BIGINT PRIMARY KEY NOT NULL, version BIGINT NOT NULL, name VARCHAR(255))"
1025 );
1026 assert!(
1027 PostgresDialect
1028 .schema_setup_sqls()
1029 .iter()
1030 .any(|sql| sql.contains("CREATE OR REPLACE FUNCTION soundex"))
1031 );
1032
1033 let query = PostgresDialect
1034 .compile_select(
1035 &entity(),
1036 &SelectQuery::new("Order")
1037 .filter(Expr::in_large(
1038 "id",
1039 vec![Value::from(1_u64), Value::from(2_u64)],
1040 ))
1041 .order_asc("id"),
1042 )
1043 .unwrap();
1044 assert_eq!(
1045 query.sql,
1046 "SELECT id, version, name FROM orders WHERE (id = ANY($1)) ORDER BY id ASC"
1047 );
1048 assert_eq!(
1049 query.params,
1050 vec![Value::List(vec![Value::U64(1), Value::U64(2)])]
1051 );
1052 }
1053}