1use std::collections::{BTreeMap, HashSet};
2use std::fmt::Display;
3use std::hash;
4
5use common::fmt::{EscapeKwFreeIdent, EscapeObjectKey, Float, Fmt, QuoteStr, SqlDuration};
6use rust_decimal::Decimal;
7use surrealdb_strand::Strand;
8use surrealdb_types::{SqlFormat, ToSql, write_sql};
9
10use crate::TableName;
11
12#[derive(Clone, Debug, Eq, PartialEq, Hash)]
13#[cfg_attr(feature = "arbitrary", derive(arbitrary::Arbitrary))]
14pub enum GeometryKind {
15 Point,
16 Line,
17 Polygon,
18 MultiPoint,
19 MultiLine,
20 MultiPolygon,
21 Collection,
22}
23
24impl ToSql for GeometryKind {
25 fn fmt_sql(&self, f: &mut String, _fmt: SqlFormat) {
26 match self {
27 GeometryKind::Point => f.push_str("point"),
28 GeometryKind::Line => f.push_str("line"),
29 GeometryKind::Polygon => f.push_str("polygon"),
30 GeometryKind::MultiPoint => f.push_str("multipoint"),
31 GeometryKind::MultiLine => f.push_str("multiline"),
32 GeometryKind::MultiPolygon => f.push_str("multipolygon"),
33 GeometryKind::Collection => f.push_str("collection"),
34 }
35 }
36}
37
38#[derive(Clone, Debug, Default, PartialEq, Eq, Hash)]
40#[cfg_attr(feature = "arbitrary", derive(arbitrary::Arbitrary))]
41pub enum Kind {
42 #[default]
44 Any,
45 None,
47 Null,
49 Bool,
51 Bytes,
53 Datetime,
55 Decimal,
57 Duration,
59 Float,
61 Int,
63 Number,
66 Object,
68 String,
70 Uuid,
72 Regex,
74 Table(Vec<TableName>),
76 Record(Vec<TableName>),
78 Geometry(Vec<GeometryKind>),
80 Either(
83 #[cfg_attr(feature = "arbitrary", arbitrary(with = crate::arbitrary::either_kind))]
84 Vec<Kind>,
85 ),
86 Set(Box<Kind>, Option<u64>),
88 Array(Box<Kind>, Option<u64>),
90 Function(Option<Vec<Kind>>, Option<Box<Kind>>),
94 Range,
96 Literal(KindLiteral),
102 File(Vec<String>),
106}
107
108impl Kind {
109 pub fn flatten(self) -> Vec<Kind> {
110 match self {
111 Kind::Either(x) => x.into_iter().flat_map(|k| k.flatten()).collect(),
112 _ => vec![self],
113 }
114 }
115
116 pub fn either(kinds: Vec<Kind>) -> Kind {
117 let mut seen = HashSet::new();
118 let mut kinds = kinds
119 .into_iter()
120 .flat_map(|k| k.flatten())
121 .filter(|k| seen.insert(k.clone()))
122 .collect::<Vec<_>>();
123 match kinds.len() {
124 0 => Kind::None,
125 1 => kinds.remove(0),
126 _ => Kind::Either(kinds),
127 }
128 }
129}
130
131impl ToSql for Kind {
132 fn fmt_sql(&self, f: &mut String, fmt: SqlFormat) {
133 match self {
134 Kind::Any => f.push_str("any"),
135 Kind::None => f.push_str("none"),
136 Kind::Null => f.push_str("null"),
137 Kind::Bool => f.push_str("bool"),
138 Kind::Bytes => f.push_str("bytes"),
139 Kind::Datetime => f.push_str("datetime"),
140 Kind::Decimal => f.push_str("decimal"),
141 Kind::Duration => f.push_str("duration"),
142 Kind::Float => f.push_str("float"),
143 Kind::Int => f.push_str("int"),
144 Kind::Number => f.push_str("number"),
145 Kind::Object => f.push_str("object"),
146 Kind::String => f.push_str("string"),
147 Kind::Uuid => f.push_str("uuid"),
148 Kind::Regex => f.push_str("regex"),
149 Kind::Function(_, _) => f.push_str("function"),
150 Kind::Table(k) => {
151 if k.is_empty() {
152 f.push_str("table");
153 } else {
154 write_sql!(
155 f,
156 fmt,
157 "table<{}>",
158 Fmt::verbar_separated(k.iter().map(|x| EscapeKwFreeIdent(x.as_str())))
159 );
160 }
161 }
162 Kind::Record(k) => {
163 if k.is_empty() {
164 f.push_str("record");
165 } else {
166 write_sql!(
167 f,
168 fmt,
169 "record<{}>",
170 Fmt::verbar_separated(k.iter().map(|x| EscapeKwFreeIdent(x.as_str())))
171 );
172 }
173 }
174 Kind::Geometry(k) => {
175 if k.is_empty() {
176 f.push_str("geometry");
177 } else {
178 write_sql!(f, fmt, "geometry<{}>", Fmt::verbar_separated(k));
179 }
180 }
181 Kind::Set(k, l) => match (k, l) {
182 (k, None) if matches!(**k, Kind::Any) => f.push_str("set"),
183 (k, None) => write_sql!(f, fmt, "set<{k}>"),
184 (k, Some(l)) => write_sql!(f, fmt, "set<{k}, {l}>"),
185 },
186 Kind::Array(k, l) => match (k, l) {
187 (k, None) if matches!(**k, Kind::Any) => f.push_str("array"),
188 (k, None) => write_sql!(f, fmt, "array<{k}>"),
189 (k, Some(l)) => write_sql!(f, fmt, "array<{k}, {l}>"),
190 },
191 Kind::Either(k) => write_sql!(f, fmt, "{}", Fmt::verbar_separated(k)),
192 Kind::Range => f.push_str("range"),
193 Kind::Literal(l) => l.fmt_sql(f, fmt),
194 Kind::File(k) => {
195 if k.is_empty() {
196 f.push_str("file");
197 } else {
198 write_sql!(
199 f,
200 fmt,
201 "file<{}>",
202 Fmt::verbar_separated(k.iter().map(|x| EscapeKwFreeIdent(x)))
203 );
204 }
205 }
206 }
207 }
208}
209
210#[derive(Clone, Debug)]
211#[cfg_attr(feature = "arbitrary", derive(arbitrary::Arbitrary))]
212pub enum KindLiteral {
213 String(Strand),
214 Integer(i64),
215 Float(f64),
216 Decimal(Decimal),
217 Duration(std::time::Duration),
218 Array(Vec<Kind>),
219 Object(BTreeMap<Strand, Kind>),
220 Bool(bool),
221}
222
223impl hash::Hash for KindLiteral {
224 fn hash<H: hash::Hasher>(&self, state: &mut H) {
225 match self {
226 Self::String(v) => v.hash(state),
227 Self::Integer(v) => v.hash(state),
228 Self::Float(v) => v.to_bits().hash(state),
229 Self::Decimal(v) => v.hash(state),
230 Self::Duration(v) => v.hash(state),
231 Self::Array(v) => v.hash(state),
232 Self::Object(v) => v.hash(state),
233 Self::Bool(v) => v.hash(state),
234 }
235 }
236}
237
238impl PartialEq for KindLiteral {
239 fn eq(&self, other: &Self) -> bool {
240 match self {
241 KindLiteral::String(a) => {
242 if let KindLiteral::String(b) = other {
243 a == b
244 } else {
245 false
246 }
247 }
248 KindLiteral::Integer(a) => {
249 if let KindLiteral::Integer(b) = other {
250 a == b
251 } else {
252 false
253 }
254 }
255 KindLiteral::Float(a) => {
256 if let KindLiteral::Float(b) = other {
257 a.to_bits() == b.to_bits()
259 } else {
260 false
261 }
262 }
263 KindLiteral::Decimal(a) => {
264 if let KindLiteral::Decimal(b) = other {
265 a == b
266 } else {
267 false
268 }
269 }
270 KindLiteral::Duration(a) => {
271 if let KindLiteral::Duration(b) = other {
272 a == b
273 } else {
274 false
275 }
276 }
277 KindLiteral::Array(a) => {
278 if let KindLiteral::Array(b) = other {
279 a == b
280 } else {
281 false
282 }
283 }
284 KindLiteral::Object(a) => {
285 if let KindLiteral::Object(b) = other {
286 a == b
287 } else {
288 false
289 }
290 }
291 KindLiteral::Bool(a) => {
292 if let KindLiteral::Bool(b) = other {
293 a == b
294 } else {
295 false
296 }
297 }
298 }
299 }
300}
301impl Eq for KindLiteral {}
302
303impl ToSql for KindLiteral {
304 fn fmt_sql(&self, f: &mut String, fmt: SqlFormat) {
305 match self {
306 KindLiteral::String(s) => write_sql!(f, fmt, "{}", QuoteStr(s)),
307 KindLiteral::Integer(n) => write_sql!(f, fmt, "{}", n),
308 KindLiteral::Float(n) => write_sql!(f, fmt, " {}", Float(*n)),
309 KindLiteral::Decimal(n) => write_sql!(f, fmt, " {}", n),
310 KindLiteral::Duration(d) => write_sql!(f, fmt, "{}", SqlDuration(*d)),
311 KindLiteral::Bool(b) => write_sql!(f, fmt, "{}", b),
312 KindLiteral::Array(a) => {
313 f.push('[');
314 if !a.is_empty() {
315 let fmt = fmt.increment();
316 write_sql!(f, fmt, "{}", Fmt::pretty_comma_separated(a.as_slice()));
317 }
318 f.push(']');
319 }
320 KindLiteral::Object(o) => {
321 if fmt.is_pretty() {
322 f.push('{');
323 } else {
324 f.push_str("{ ");
325 }
326 if !o.is_empty() {
327 let fmt = fmt.increment();
328 write_sql!(
329 f,
330 fmt,
331 "{}",
332 Fmt::pretty_comma_separated(o.iter().map(|args| Fmt::new(
333 args,
334 |(k, v), f, fmt| {
335 write_sql!(f, fmt, "{}: {}", EscapeObjectKey(k), v)
336 }
337 )),)
338 );
339 }
340 if fmt.is_pretty() {
341 f.push('}');
342 } else {
343 f.push_str(" }");
344 }
345 }
346 }
347 }
348}
349
350impl Display for Kind {
351 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
352 write!(f, "{}", self.to_sql())
353 }
354}
355
356pub fn kind_contains_object(kind: &Kind) -> bool {
362 match kind {
363 Kind::Object => true,
364 Kind::Either(kinds) => kinds.iter().any(kind_contains_object),
365 Kind::Array(inner, _) | Kind::Set(inner, _) => kind_contains_object(inner),
366 Kind::Literal(KindLiteral::Object(_)) => true,
367 Kind::Literal(KindLiteral::Array(kinds)) => kinds.iter().any(kind_contains_object),
368 _ => false,
369 }
370}
371
372impl From<GeometryKind> for surrealdb_types::GeometryKind {
373 fn from(v: GeometryKind) -> Self {
374 match v {
375 GeometryKind::Point => surrealdb_types::GeometryKind::Point,
376 GeometryKind::Line => surrealdb_types::GeometryKind::Line,
377 GeometryKind::Polygon => surrealdb_types::GeometryKind::Polygon,
378 GeometryKind::MultiPoint => surrealdb_types::GeometryKind::MultiPoint,
379 GeometryKind::MultiLine => surrealdb_types::GeometryKind::MultiLine,
380 GeometryKind::MultiPolygon => surrealdb_types::GeometryKind::MultiPolygon,
381 GeometryKind::Collection => surrealdb_types::GeometryKind::Collection,
382 }
383 }
384}
385
386impl From<surrealdb_types::GeometryKind> for GeometryKind {
387 fn from(v: surrealdb_types::GeometryKind) -> Self {
388 match v {
389 surrealdb_types::GeometryKind::Point => GeometryKind::Point,
390 surrealdb_types::GeometryKind::Line => GeometryKind::Line,
391 surrealdb_types::GeometryKind::Polygon => GeometryKind::Polygon,
392 surrealdb_types::GeometryKind::MultiPoint => GeometryKind::MultiPoint,
393 surrealdb_types::GeometryKind::MultiLine => GeometryKind::MultiLine,
394 surrealdb_types::GeometryKind::MultiPolygon => GeometryKind::MultiPolygon,
395 surrealdb_types::GeometryKind::Collection => GeometryKind::Collection,
396 }
397 }
398}
399
400impl From<Kind> for surrealdb_types::Kind {
401 fn from(v: Kind) -> Self {
402 match v {
403 Kind::Any => surrealdb_types::Kind::Any,
404 Kind::None => surrealdb_types::Kind::None,
405 Kind::Null => surrealdb_types::Kind::Null,
406 Kind::Bool => surrealdb_types::Kind::Bool,
407 Kind::Bytes => surrealdb_types::Kind::Bytes,
408 Kind::Datetime => surrealdb_types::Kind::Datetime,
409 Kind::Decimal => surrealdb_types::Kind::Decimal,
410 Kind::Duration => surrealdb_types::Kind::Duration,
411 Kind::Float => surrealdb_types::Kind::Float,
412 Kind::Int => surrealdb_types::Kind::Int,
413 Kind::Number => surrealdb_types::Kind::Number,
414 Kind::Object => surrealdb_types::Kind::Object,
415 Kind::String => surrealdb_types::Kind::String,
416 Kind::Uuid => surrealdb_types::Kind::Uuid,
417 Kind::Regex => surrealdb_types::Kind::Regex,
418 Kind::Table(k) => surrealdb_types::Kind::Table(k.into_iter().map(Into::into).collect()),
419 Kind::Record(k) => {
420 surrealdb_types::Kind::Record(k.into_iter().map(Into::into).collect())
421 }
422 Kind::Geometry(k) => {
423 surrealdb_types::Kind::Geometry(k.into_iter().map(Into::into).collect())
424 }
425 Kind::Either(k) => {
426 surrealdb_types::Kind::Either(k.into_iter().map(Into::into).collect())
427 }
428 Kind::Set(k, l) => surrealdb_types::Kind::Set(Box::new((*k).into()), l),
429 Kind::Array(k, l) => surrealdb_types::Kind::Array(Box::new((*k).into()), l),
430 Kind::Function(args, ret) => surrealdb_types::Kind::Function(
431 args.map(|args| args.into_iter().map(Into::into).collect()),
432 ret.map(|ret| Box::new((*ret).into())),
433 ),
434 Kind::Range => surrealdb_types::Kind::Range,
435 Kind::Literal(l) => surrealdb_types::Kind::Literal(l.into()),
436 Kind::File(k) => surrealdb_types::Kind::File(k),
437 }
438 }
439}
440
441impl From<surrealdb_types::Kind> for Kind {
442 fn from(v: surrealdb_types::Kind) -> Self {
443 match v {
444 surrealdb_types::Kind::None => Kind::None,
445 surrealdb_types::Kind::Null => Kind::Null,
446 surrealdb_types::Kind::Any => Kind::Any,
447 surrealdb_types::Kind::Bool => Kind::Bool,
448 surrealdb_types::Kind::Bytes => Kind::Bytes,
449 surrealdb_types::Kind::Datetime => Kind::Datetime,
450 surrealdb_types::Kind::Decimal => Kind::Decimal,
451 surrealdb_types::Kind::Duration => Kind::Duration,
452 surrealdb_types::Kind::Float => Kind::Float,
453 surrealdb_types::Kind::Int => Kind::Int,
454 surrealdb_types::Kind::Number => Kind::Number,
455 surrealdb_types::Kind::Object => Kind::Object,
456 surrealdb_types::Kind::String => Kind::String,
457 surrealdb_types::Kind::Uuid => Kind::Uuid,
458 surrealdb_types::Kind::Regex => Kind::Regex,
459 surrealdb_types::Kind::Table(k) => Kind::Table(k.into_iter().map(Into::into).collect()),
460 surrealdb_types::Kind::Record(k) => {
461 Kind::Record(k.into_iter().map(Into::into).collect())
462 }
463 surrealdb_types::Kind::Geometry(k) => {
464 Kind::Geometry(k.into_iter().map(Into::into).collect())
465 }
466 surrealdb_types::Kind::Either(k) => {
467 Kind::Either(k.into_iter().map(Into::into).collect())
468 }
469 surrealdb_types::Kind::Set(k, l) => Kind::Set(Box::new((*k).into()), l),
470 surrealdb_types::Kind::Array(k, l) => Kind::Array(Box::new((*k).into()), l),
471 surrealdb_types::Kind::Function(args, ret) => Kind::Function(
472 args.map(|args| args.into_iter().map(Into::into).collect()),
473 ret.map(|ret| Box::new((*ret).into())),
474 ),
475 surrealdb_types::Kind::Range => Kind::Range,
476 surrealdb_types::Kind::Literal(l) => Kind::Literal(l.into()),
477 surrealdb_types::Kind::File(k) => Kind::File(k),
478 }
479 }
480}
481
482impl From<KindLiteral> for surrealdb_types::KindLiteral {
483 fn from(v: KindLiteral) -> Self {
484 match v {
485 KindLiteral::Bool(b) => surrealdb_types::KindLiteral::Bool(b),
486 KindLiteral::Integer(i) => surrealdb_types::KindLiteral::Integer(i),
487 KindLiteral::Float(f) => surrealdb_types::KindLiteral::Float(f),
488 KindLiteral::Decimal(d) => surrealdb_types::KindLiteral::Decimal(d),
489 KindLiteral::String(s) => surrealdb_types::KindLiteral::String(s.into_string()),
490 KindLiteral::Duration(d) => {
491 surrealdb_types::KindLiteral::Duration(surrealdb_types::Duration::from(d))
492 }
493 KindLiteral::Array(a) => {
494 surrealdb_types::KindLiteral::Array(a.into_iter().map(Into::into).collect())
495 }
496 KindLiteral::Object(o) => surrealdb_types::KindLiteral::Object(
497 o.into_iter().map(|(k, v)| (k.into_string(), v.into())).collect(),
498 ),
499 }
500 }
501}
502
503impl From<surrealdb_types::KindLiteral> for KindLiteral {
504 fn from(v: surrealdb_types::KindLiteral) -> Self {
505 match v {
506 surrealdb_types::KindLiteral::Bool(b) => Self::Bool(b),
507 surrealdb_types::KindLiteral::Integer(i) => Self::Integer(i),
508 surrealdb_types::KindLiteral::Float(f) => Self::Float(f),
509 surrealdb_types::KindLiteral::Decimal(d) => Self::Decimal(d),
510 surrealdb_types::KindLiteral::String(s) => Self::String(s.into()),
511 surrealdb_types::KindLiteral::Duration(d) => Self::Duration(d.into_inner()),
512 surrealdb_types::KindLiteral::Array(a) => {
513 Self::Array(a.into_iter().map(Into::into).collect())
514 }
515 surrealdb_types::KindLiteral::Object(o) => {
516 Self::Object(o.into_iter().map(|(k, v)| (k.into(), v.into())).collect())
517 }
518 }
519 }
520}
521
522#[cfg(test)]
523mod tests {
524 use rstest::rstest;
525
526 use super::*;
527
528 #[rstest]
529 #[case::any(Kind::Any, "any")]
530 #[case::none(Kind::None, "none")]
531 #[case::null(Kind::Null, "null")]
532 #[case::bool(Kind::Bool, "bool")]
533 #[case::bytes(Kind::Bytes, "bytes")]
534 #[case::datetime(Kind::Datetime, "datetime")]
535 #[case::decimal(Kind::Decimal, "decimal")]
536 #[case::duration(Kind::Duration, "duration")]
537 #[case::float(Kind::Float, "float")]
538 #[case::int(Kind::Int, "int")]
539 #[case::number(Kind::Number, "number")]
540 #[case::object(Kind::Object, "object")]
541 #[case::string(Kind::String, "string")]
542 #[case::uuid(Kind::Uuid, "uuid")]
543 #[case::regex(Kind::Regex, "regex")]
544 #[case::range(Kind::Range, "range")]
545 #[case::function(Kind::Function(None, None), "function")]
546 #[case::table_empty(Kind::Table(vec![]), "table")]
547 #[case::table_single(Kind::Table(vec!["users".into()]), "table<users>")]
548 #[case::table_multiple(Kind::Table(vec!["users".into(), "posts".into()]), "table<users | posts>")]
549 #[case::record_empty(Kind::Record(vec![]), "record")]
550 #[case::record_single(Kind::Record(vec!["users".into()]), "record<users>")]
551 #[case::geometry_empty(Kind::Geometry(vec![]), "geometry")]
552 #[case::geometry_single(Kind::Geometry(vec![GeometryKind::Point]), "geometry<point>")]
553 #[case::set_any(Kind::Set(Box::new(Kind::Any), None), "set")]
554 #[case::set_typed(Kind::Set(Box::new(Kind::String), None), "set<string>")]
555 #[case::array_any(Kind::Array(Box::new(Kind::Any), None), "array")]
556 #[case::array_typed(Kind::Array(Box::new(Kind::String), Some(5)), "array<string, 5>")]
557 #[case::either(Kind::Either(vec![Kind::String, Kind::Int]), "string | int")]
558 #[case::file_empty(Kind::File(vec![]), "file")]
559 #[case::file_single(Kind::File(vec!["bucket".to_string()]), "file<bucket>")]
560 fn test_kind_to_sql(#[case] kind: Kind, #[case] expected: &str) {
561 assert_eq!(kind.to_sql(), expected);
562 assert_eq!(kind.to_string(), expected);
563 }
564
565 #[rstest]
566 #[case::any(Kind::Any)]
567 #[case::none(Kind::None)]
568 #[case::null(Kind::Null)]
569 #[case::bool(Kind::Bool)]
570 #[case::bytes(Kind::Bytes)]
571 #[case::datetime(Kind::Datetime)]
572 #[case::decimal(Kind::Decimal)]
573 #[case::duration(Kind::Duration)]
574 #[case::float(Kind::Float)]
575 #[case::int(Kind::Int)]
576 #[case::number(Kind::Number)]
577 #[case::object(Kind::Object)]
578 #[case::string(Kind::String)]
579 #[case::uuid(Kind::Uuid)]
580 #[case::regex(Kind::Regex)]
581 #[case::range(Kind::Range)]
582 #[case::table(Kind::Table(vec!["users".into()]))]
583 #[case::record(Kind::Record(vec!["users".into()]))]
584 #[case::geometry(Kind::Geometry(vec![GeometryKind::Point]))]
585 #[case::set(Kind::Set(Box::new(Kind::String), None))]
586 #[case::array(Kind::Array(Box::new(Kind::String), None))]
587 #[case::either(Kind::Either(vec![Kind::String, Kind::Int]))]
588 #[case::file(Kind::File(vec!["bucket".to_string()]))]
589 fn test_kind_conversions_public(#[case] sql_kind: Kind) {
590 let public_kind: surrealdb_types::Kind = sql_kind.clone().into();
591 let back_to_sql: Kind = public_kind.into();
592 assert_eq!(sql_kind, back_to_sql);
593 }
594
595 #[rstest]
596 #[case::any(Kind::Any)]
597 #[case::table(Kind::Table(vec!["users".into()]))]
598 #[case::record(Kind::Record(vec!["users".into()]))]
599 #[case::geometry(Kind::Geometry(vec![GeometryKind::Point]))]
600 fn test_kind_flatten(#[case] kind: Kind) {
601 let flattened = kind.clone().flatten();
602 assert_eq!(flattened.len(), 1);
603 assert_eq!(flattened[0], kind);
604 }
605
606 #[test]
607 fn test_kind_either() {
608 let kinds = vec![Kind::Table(vec!["users".into()]), Kind::Table(vec!["posts".into()])];
609 let either = Kind::either(kinds);
610 assert!(matches!(either, Kind::Either(_)));
611 if let Kind::Either(inner) = either {
612 assert_eq!(inner.len(), 2);
613 }
614 }
615}