1use std::collections::HashSet;
4
5use arrow_schema::ArrowError;
6use geo_traits::{
7 GeometryCollectionTrait, GeometryTrait, MultiLineStringTrait, MultiPointTrait,
8 MultiPolygonTrait,
9};
10use geoarrow_array::cast::AsGeoArrowArray;
11use geoarrow_array::{GeoArrowArray, GeoArrowArrayAccessor};
12use geoarrow_schema::error::{GeoArrowError, GeoArrowResult};
13use geoarrow_schema::{Dimension, GeoArrowType};
14
15pub fn infer_downcast_type<'a>(
83 arrays: impl Iterator<Item = &'a dyn GeoArrowArray>,
84) -> GeoArrowResult<Option<(NativeType, Dimension)>> {
85 let mut type_ids = HashSet::new();
86 for array in arrays {
87 let type_id = get_type_ids(array)?;
88 type_ids.extend(type_id);
89 }
90
91 if type_ids.is_empty() {
92 return Err(ArrowError::CastError(
93 "Empty iterator of arrays passed to infer_downcast_type".to_string(),
94 )
95 .into());
96 }
97
98 infer_from_native_type_and_dimension(type_ids)
99}
100
101fn get_type_ids(array: &dyn GeoArrowArray) -> GeoArrowResult<HashSet<NativeTypeAndDimension>> {
103 use GeoArrowType::*;
104 let type_ids: HashSet<NativeTypeAndDimension> = match array.data_type() {
105 Point(typ) => [NativeTypeAndDimension::new(
106 NativeType::Point,
107 typ.dimension(),
108 )]
109 .into_iter()
110 .collect(),
111 LineString(typ) => [NativeTypeAndDimension::new(
112 NativeType::LineString,
113 typ.dimension(),
114 )]
115 .into_iter()
116 .collect(),
117 Polygon(typ) => [NativeTypeAndDimension::new(
118 NativeType::Polygon,
119 typ.dimension(),
120 )]
121 .into_iter()
122 .collect(),
123 MultiPoint(typ) => {
124 let dim = typ.dimension();
125 let array = array.as_multi_point();
126 array
127 .iter()
128 .flatten()
129 .map(|multi_point| {
130 let geom_type = if multi_point?.num_points() >= 2 {
131 NativeTypeAndDimension::new(NativeType::MultiPoint, dim)
132 } else {
133 NativeTypeAndDimension::new(NativeType::Point, dim)
134 };
135 Ok::<_, GeoArrowError>(geom_type)
136 })
137 .collect::<GeoArrowResult<HashSet<NativeTypeAndDimension>>>()?
138 }
139 MultiLineString(typ) => {
140 let dim = typ.dimension();
141 let array = array.as_multi_line_string();
142 array
143 .iter()
144 .flatten()
145 .map(|multi_line_string| {
146 let geom_type = if multi_line_string?.num_line_strings() >= 2 {
147 NativeTypeAndDimension::new(NativeType::MultiLineString, dim)
148 } else {
149 NativeTypeAndDimension::new(NativeType::LineString, dim)
150 };
151 Ok::<_, GeoArrowError>(geom_type)
152 })
153 .collect::<GeoArrowResult<HashSet<NativeTypeAndDimension>>>()?
154 }
155 MultiPolygon(typ) => {
156 let dim = typ.dimension();
157 let array = array.as_multi_polygon();
158 array
159 .iter()
160 .flatten()
161 .map(|multi_polygon| {
162 let geom_type = if multi_polygon?.num_polygons() >= 2 {
163 NativeTypeAndDimension::new(NativeType::MultiPolygon, dim)
164 } else {
165 NativeTypeAndDimension::new(NativeType::Polygon, dim)
166 };
167 Ok::<_, GeoArrowError>(geom_type)
168 })
169 .collect::<GeoArrowResult<HashSet<NativeTypeAndDimension>>>()?
170 }
171 GeometryCollection(typ) => {
172 let dim = typ.dimension();
173 let array = array.as_geometry_collection();
174 array
175 .iter()
176 .flatten()
177 .map(|geometry_collection| {
178 let geometry_collection = geometry_collection?;
179 let geom_type = if geometry_collection.num_geometries() == 1 {
180 let geom_type = NativeType::from_geometry_trait(
181 &geometry_collection.geometry(0).unwrap(),
182 );
183 NativeTypeAndDimension::new(geom_type, dim)
184 } else {
185 NativeTypeAndDimension::new(NativeType::GeometryCollection, dim)
186 };
187 Ok::<_, GeoArrowError>(geom_type)
188 })
189 .collect::<GeoArrowResult<HashSet<NativeTypeAndDimension>>>()?
190 }
191 Rect(typ) => [NativeTypeAndDimension::new(
192 NativeType::Rect,
193 typ.dimension(),
194 )]
195 .into_iter()
196 .collect(),
197 Geometry(_) => {
198 let type_ids: HashSet<i8> =
199 HashSet::from_iter(array.as_geometry().type_ids().iter().copied());
200 type_ids
201 .into_iter()
202 .map(NativeTypeAndDimension::from_type_id)
203 .collect()
204 }
205 Wkb(_) => array
206 .as_wkb::<i32>()
207 .iter()
208 .flatten()
209 .map(|wkb| {
210 let wkb = wkb?;
211 let dim = wkb.dim().try_into()?;
212 let geom_type = NativeType::from_geometry_trait(&wkb);
213 Ok(NativeTypeAndDimension::new(geom_type, dim))
214 })
215 .collect::<GeoArrowResult<HashSet<NativeTypeAndDimension>>>()?,
216 LargeWkb(_) => array
217 .as_wkb::<i64>()
218 .iter()
219 .flatten()
220 .map(|wkb| {
221 let wkb = wkb?;
222 let dim = wkb.dim().try_into()?;
223 let geom_type = NativeType::from_geometry_trait(&wkb);
224 Ok(NativeTypeAndDimension::new(geom_type, dim))
225 })
226 .collect::<GeoArrowResult<HashSet<NativeTypeAndDimension>>>()?,
227 WkbView(_) => array
228 .as_wkb_view()
229 .iter()
230 .flatten()
231 .map(|wkb| {
232 let wkb = wkb?;
233 let dim = wkb.dim().try_into()?;
234 let geom_type = NativeType::from_geometry_trait(&wkb);
235 Ok(NativeTypeAndDimension::new(geom_type, dim))
236 })
237 .collect::<GeoArrowResult<HashSet<NativeTypeAndDimension>>>()?,
238 Wkt(_) => array
239 .as_wkt::<i32>()
240 .inner()
241 .iter()
242 .flatten()
243 .map(|s| {
244 let (wkt_type, wkt_dim) = wkt::infer_type(s).map_err(ArrowError::CastError)?;
245 let geom_type =
246 NativeTypeAndDimension::new(wkt_type.into(), wkt_dim_to_geoarrow_dim(wkt_dim));
247 Ok(geom_type)
248 })
249 .collect::<GeoArrowResult<HashSet<NativeTypeAndDimension>>>()?,
250 LargeWkt(_) => array
251 .as_wkt::<i64>()
252 .inner()
253 .iter()
254 .flatten()
255 .map(|s| {
256 let (wkt_type, wkt_dim) = wkt::infer_type(s).map_err(ArrowError::CastError)?;
257 let geom_type =
258 NativeTypeAndDimension::new(wkt_type.into(), wkt_dim_to_geoarrow_dim(wkt_dim));
259 Ok(geom_type)
260 })
261 .collect::<GeoArrowResult<HashSet<NativeTypeAndDimension>>>()?,
262 WktView(_) => array
263 .as_wkt_view()
264 .inner()
265 .iter()
266 .flatten()
267 .map(|s| {
268 let (wkt_type, wkt_dim) = wkt::infer_type(s).map_err(ArrowError::CastError)?;
269 let geom_type =
270 NativeTypeAndDimension::new(wkt_type.into(), wkt_dim_to_geoarrow_dim(wkt_dim));
271 Ok(geom_type)
272 })
273 .collect::<GeoArrowResult<HashSet<NativeTypeAndDimension>>>()?,
274 };
275 Ok(type_ids)
276}
277
278fn wkt_dim_to_geoarrow_dim(wkt_dim: wkt::types::Dimension) -> Dimension {
279 match wkt_dim {
280 wkt::types::Dimension::XY => Dimension::XY,
281 wkt::types::Dimension::XYZ => Dimension::XYZ,
282 wkt::types::Dimension::XYM => Dimension::XYM,
283 wkt::types::Dimension::XYZM => Dimension::XYZM,
284 }
285}
286
287fn infer_from_native_type_and_dimension(
288 type_ids: HashSet<NativeTypeAndDimension>,
289) -> GeoArrowResult<Option<(NativeType, Dimension)>> {
290 if type_ids.len() == 1 {
292 let type_id = type_ids.into_iter().next().unwrap();
293 return Ok(Some((type_id.geometry_type, type_id.dim)));
294 }
295
296 let (dims, native_types): (HashSet<_>, HashSet<_>) = type_ids
298 .iter()
299 .map(|type_id| (type_id.dim, type_id.geometry_type))
300 .unzip();
301 if dims.len() > 1 {
302 return Ok(None);
303 }
304 let dim = dims.into_iter().next().unwrap();
305
306 if native_types.len() == 2 {
307 if native_types.contains(&NativeType::Point)
308 && native_types.contains(&NativeType::MultiPoint)
309 {
310 return Ok(Some((NativeType::MultiPoint, dim)));
311 }
312
313 if native_types.contains(&NativeType::LineString)
314 && native_types.contains(&NativeType::MultiLineString)
315 {
316 return Ok(Some((NativeType::MultiLineString, dim)));
317 }
318
319 if native_types.contains(&NativeType::Polygon)
320 && native_types.contains(&NativeType::MultiPolygon)
321 {
322 return Ok(Some((NativeType::MultiPolygon, dim)));
323 }
324 }
325
326 Ok(None)
327}
328
329#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
331pub enum NativeType {
332 #[allow(missing_docs)]
333 Point,
334 #[allow(missing_docs)]
335 LineString,
336 #[allow(missing_docs)]
337 Polygon,
338 #[allow(missing_docs)]
339 MultiPoint,
340 #[allow(missing_docs)]
341 MultiLineString,
342 #[allow(missing_docs)]
343 MultiPolygon,
344 #[allow(missing_docs)]
345 GeometryCollection,
346 #[allow(missing_docs)]
347 Rect,
348}
349
350impl NativeType {
351 fn from_geometry_trait(geometry: &impl GeometryTrait) -> Self {
352 match geometry.as_type() {
353 geo_traits::GeometryType::Point(_) => Self::Point,
354 geo_traits::GeometryType::LineString(_) => Self::LineString,
355 geo_traits::GeometryType::Polygon(_) => Self::Polygon,
356 geo_traits::GeometryType::MultiPoint(_) => Self::MultiPoint,
357 geo_traits::GeometryType::MultiLineString(_) => Self::MultiLineString,
358 geo_traits::GeometryType::MultiPolygon(_) => Self::MultiPolygon,
359 geo_traits::GeometryType::GeometryCollection(_) => Self::GeometryCollection,
360 _ => panic!("Unsupported geometry type"),
361 }
362 }
363}
364
365impl From<wkt::types::GeometryType> for NativeType {
366 fn from(value: wkt::types::GeometryType) -> Self {
367 match value {
368 wkt::types::GeometryType::Point => Self::Point,
369 wkt::types::GeometryType::LineString => Self::LineString,
370 wkt::types::GeometryType::Polygon => Self::Polygon,
371 wkt::types::GeometryType::MultiPoint => Self::MultiPoint,
372 wkt::types::GeometryType::MultiLineString => Self::MultiLineString,
373 wkt::types::GeometryType::MultiPolygon => Self::MultiPolygon,
374 wkt::types::GeometryType::GeometryCollection => Self::GeometryCollection,
375 }
376 }
377}
378
379#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
380struct NativeTypeAndDimension {
381 geometry_type: NativeType,
382 dim: Dimension,
383}
384
385impl NativeTypeAndDimension {
386 fn new(geometry_type: NativeType, dim: Dimension) -> Self {
387 Self { geometry_type, dim }
388 }
389
390 fn from_type_id(type_id: i8) -> Self {
391 let dim = match type_id / 10 {
392 0 => Dimension::XY,
393 1 => Dimension::XYZ,
394 2 => Dimension::XYM,
395 3 => Dimension::XYZM,
396 _ => panic!("unsupported type_id: {type_id}"),
397 };
398 let geometry_type = match type_id % 10 {
399 1 => NativeType::Point,
400 2 => NativeType::LineString,
401 3 => NativeType::Polygon,
402 4 => NativeType::MultiPoint,
403 5 => NativeType::MultiLineString,
404 6 => NativeType::MultiPolygon,
405 7 => NativeType::GeometryCollection,
406 _ => panic!("unsupported type id"),
407 };
408 Self { geometry_type, dim }
409 }
410}
411
412impl From<(NativeType, Dimension)> for NativeTypeAndDimension {
413 fn from(value: (NativeType, Dimension)) -> Self {
414 Self::new(value.0, value.1)
415 }
416}
417
418#[cfg(test)]
419mod test {
420 use geoarrow_array::cast::{to_wkb, to_wkt};
421 use geoarrow_array::test;
422 use geoarrow_schema::CoordType;
423
424 use super::*;
425
426 #[test]
427 fn infer_get_type_ids_point() {
428 for dim in [
430 Dimension::XY,
431 Dimension::XYZ,
432 Dimension::XYM,
433 Dimension::XYZM,
434 ] {
435 let array = test::point::array(CoordType::Interleaved, dim);
436 assert_eq!(
437 get_type_ids(&array).unwrap(),
438 HashSet::from_iter([NativeTypeAndDimension::new(NativeType::Point, dim)])
439 );
440 }
441 }
442
443 #[test]
444 fn infer_get_type_ids_linestring() {
445 for dim in [
447 Dimension::XY,
448 Dimension::XYZ,
449 Dimension::XYM,
450 Dimension::XYZM,
451 ] {
452 let array = test::linestring::array(CoordType::Interleaved, dim);
453 assert_eq!(
454 get_type_ids(&array).unwrap(),
455 HashSet::from_iter([NativeTypeAndDimension::new(NativeType::LineString, dim)])
456 );
457 }
458 }
459
460 #[test]
461 fn infer_get_type_ids_polygon() {
462 for dim in [
464 Dimension::XY,
465 Dimension::XYZ,
466 Dimension::XYM,
467 Dimension::XYZM,
468 ] {
469 let array = test::polygon::array(CoordType::Interleaved, dim);
470 assert_eq!(
471 get_type_ids(&array).unwrap(),
472 HashSet::from_iter([NativeTypeAndDimension::new(NativeType::Polygon, dim)])
473 );
474 }
475 }
476
477 #[test]
478 fn infer_get_type_ids_multipoint() {
479 for dim in [
481 Dimension::XY,
482 Dimension::XYZ,
483 Dimension::XYM,
484 Dimension::XYZM,
485 ] {
486 let array = test::multipoint::array(CoordType::Interleaved, dim);
487 assert_eq!(
488 get_type_ids(&array).unwrap(),
489 HashSet::from_iter([
490 NativeTypeAndDimension::new(NativeType::Point, dim),
491 NativeTypeAndDimension::new(NativeType::MultiPoint, dim),
492 ])
493 );
494 }
495 }
496
497 #[test]
498 fn infer_get_type_ids_multilinestring() {
499 for dim in [
501 Dimension::XY,
502 Dimension::XYZ,
503 Dimension::XYM,
504 Dimension::XYZM,
505 ] {
506 let array = test::multilinestring::array(CoordType::Interleaved, dim);
507 assert_eq!(
508 get_type_ids(&array).unwrap(),
509 HashSet::from_iter([
510 NativeTypeAndDimension::new(NativeType::LineString, dim),
511 NativeTypeAndDimension::new(NativeType::MultiLineString, dim),
512 ])
513 );
514 }
515 }
516
517 #[test]
518 fn infer_get_type_ids_multipolygon() {
519 for dim in [
521 Dimension::XY,
522 Dimension::XYZ,
523 Dimension::XYM,
524 Dimension::XYZM,
525 ] {
526 let array = test::multipolygon::array(CoordType::Interleaved, dim);
527 assert_eq!(
528 get_type_ids(&array).unwrap(),
529 HashSet::from_iter([
530 NativeTypeAndDimension::new(NativeType::Polygon, dim),
531 NativeTypeAndDimension::new(NativeType::MultiPolygon, dim),
532 ])
533 );
534 }
535 }
536
537 #[test]
538 fn infer_get_type_ids_geometrycollection() {
539 for dim in [
541 Dimension::XY,
542 Dimension::XYZ,
543 Dimension::XYM,
544 Dimension::XYZM,
545 ] {
546 let array = test::geometrycollection::array(CoordType::Interleaved, dim, false);
547 assert_eq!(
548 get_type_ids(&array).unwrap(),
549 HashSet::from_iter([
550 NativeTypeAndDimension::new(NativeType::Point, dim),
551 NativeTypeAndDimension::new(NativeType::LineString, dim),
552 NativeTypeAndDimension::new(NativeType::Polygon, dim),
553 NativeTypeAndDimension::new(NativeType::MultiPoint, dim),
554 NativeTypeAndDimension::new(NativeType::MultiLineString, dim),
555 NativeTypeAndDimension::new(NativeType::MultiPolygon, dim),
556 NativeTypeAndDimension::new(NativeType::GeometryCollection, dim),
557 ])
558 );
559 }
560 }
561
562 #[test]
563 fn infer_get_type_ids_geometry_wkb_wkt() {
564 let array = test::geometry::array(CoordType::Interleaved, false);
565 let wkb_array = to_wkb::<i32>(&array).unwrap();
566 let large_wkb_array = to_wkb::<i64>(&array).unwrap();
567 let wkt_array = to_wkt::<i32>(&array).unwrap();
568 let large_wkt_array = to_wkt::<i64>(&array).unwrap();
569
570 let mut expected_types = HashSet::new();
571 for dim in [
572 Dimension::XY,
573 Dimension::XYZ,
574 Dimension::XYM,
575 Dimension::XYZM,
576 ] {
577 expected_types.insert(NativeTypeAndDimension::new(NativeType::Point, dim));
578 expected_types.insert(NativeTypeAndDimension::new(NativeType::LineString, dim));
579 expected_types.insert(NativeTypeAndDimension::new(NativeType::Polygon, dim));
580 expected_types.insert(NativeTypeAndDimension::new(NativeType::MultiPoint, dim));
581 expected_types.insert(NativeTypeAndDimension::new(
582 NativeType::MultiLineString,
583 dim,
584 ));
585 expected_types.insert(NativeTypeAndDimension::new(NativeType::MultiPolygon, dim));
586 expected_types.insert(NativeTypeAndDimension::new(
587 NativeType::GeometryCollection,
588 dim,
589 ));
590 }
591
592 assert_eq!(get_type_ids(&array).unwrap(), expected_types);
593 assert_eq!(get_type_ids(&wkb_array).unwrap(), expected_types);
594 assert_eq!(get_type_ids(&large_wkb_array).unwrap(), expected_types);
595 assert_eq!(get_type_ids(&wkt_array).unwrap(), expected_types);
596 assert_eq!(get_type_ids(&large_wkt_array).unwrap(), expected_types);
597 }
598
599 #[test]
600 fn infer_from_one_type() {
601 let input_type = NativeTypeAndDimension::new(NativeType::Point, Dimension::XY);
602 let type_ids = [input_type].into_iter().collect::<HashSet<_>>();
603 let resolved_type = infer_from_native_type_and_dimension(type_ids)
604 .unwrap()
605 .unwrap();
606 assert_eq!(input_type, resolved_type.into());
607 }
608
609 #[test]
610 fn cant_infer_from_two_dims() {
611 let input_types = [
612 NativeTypeAndDimension::new(NativeType::Point, Dimension::XY),
613 NativeTypeAndDimension::new(NativeType::Point, Dimension::XYZ),
614 ];
615 let type_ids = input_types.into_iter().collect::<HashSet<_>>();
616 assert!(
617 infer_from_native_type_and_dimension(type_ids)
618 .unwrap()
619 .is_none()
620 );
621 }
622
623 #[test]
624 fn infer_point_multi_point() {
625 let input_types = [
626 NativeTypeAndDimension::new(NativeType::Point, Dimension::XYZ),
627 NativeTypeAndDimension::new(NativeType::MultiPoint, Dimension::XYZ),
628 ];
629 let type_ids = input_types.into_iter().collect::<HashSet<_>>();
630 let resolved_type = infer_from_native_type_and_dimension(type_ids)
631 .unwrap()
632 .unwrap();
633 assert_eq!(
634 NativeTypeAndDimension::new(NativeType::MultiPoint, Dimension::XYZ),
635 resolved_type.into()
636 );
637 }
638
639 #[test]
640 fn infer_linestring_multilinestring() {
641 let input_types = [
642 NativeTypeAndDimension::new(NativeType::LineString, Dimension::XYM),
643 NativeTypeAndDimension::new(NativeType::MultiLineString, Dimension::XYM),
644 ];
645 let type_ids = input_types.into_iter().collect::<HashSet<_>>();
646 let resolved_type = infer_from_native_type_and_dimension(type_ids)
647 .unwrap()
648 .unwrap();
649 assert_eq!(
650 NativeTypeAndDimension::new(NativeType::MultiLineString, Dimension::XYM),
651 resolved_type.into()
652 );
653 }
654
655 #[test]
656 fn infer_polygon_multipolygon() {
657 let input_types = [
658 NativeTypeAndDimension::new(NativeType::Polygon, Dimension::XYZM),
659 NativeTypeAndDimension::new(NativeType::MultiPolygon, Dimension::XYZM),
660 ];
661 let type_ids = input_types.into_iter().collect::<HashSet<_>>();
662 let resolved_type = infer_from_native_type_and_dimension(type_ids)
663 .unwrap()
664 .unwrap();
665 assert_eq!(
666 NativeTypeAndDimension::new(NativeType::MultiPolygon, Dimension::XYZM),
667 resolved_type.into()
668 );
669 }
670
671 #[test]
672 fn unable_to_infer() {
673 let input_types = [
674 NativeTypeAndDimension::new(NativeType::Point, Dimension::XY),
675 NativeTypeAndDimension::new(NativeType::LineString, Dimension::XY),
676 ];
677 let type_ids = input_types.into_iter().collect::<HashSet<_>>();
678 assert!(
679 infer_from_native_type_and_dimension(type_ids)
680 .unwrap()
681 .is_none()
682 );
683 }
684}