Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
50 changes: 12 additions & 38 deletions rust/sedona-functions/src/st_setsrid.rs
Original file line number Diff line number Diff line change
Expand Up @@ -14,10 +14,7 @@
// KIND, either express or implied. See the License for the
// specific language governing permissions and limitations
// under the License.
use std::{
collections::{HashMap, HashSet},
sync::{Arc, OnceLock},
};
use std::sync::{Arc, OnceLock};

use arrow_array::{
builder::{BinaryBuilder, NullBufferBuilder},
Expand All @@ -43,7 +40,11 @@ use sedona_expr::{
scalar_udf::{ScalarKernelRef, SedonaScalarKernel, SedonaScalarUDF},
};
use sedona_geometry::transform::CrsEngine;
use sedona_schema::{crs::deserialize_crs, datatypes::SedonaType, matchers::ArgMatcher};
use sedona_schema::{
crs::{deserialize_crs, CachedCrsNormalization, CachedSRIDToCrs},
datatypes::SedonaType,
matchers::ArgMatcher,
};

/// ST_SetSRID() scalar UDF implementation
///
Expand Down Expand Up @@ -475,8 +476,7 @@ fn normalize_crs_array(
| DataType::UInt16
| DataType::UInt32
| DataType::UInt64 => {
// Local cache to avoid re-validating inputs
let mut known_valid = HashSet::new();
let mut srid_to_crs = CachedSRIDToCrs::new();

let int_value = crs_value.cast_to(&DataType::Int64, None)?;
let int_array_ref = ColumnarValue::values_to_arrays(&[int_value])?;
Expand All @@ -485,18 +485,10 @@ fn normalize_crs_array(
.iter()
.map(|maybe_srid| -> Result<Option<String>> {
if let Some(srid) = maybe_srid {
if srid == 0 {
let Some(auth_code) = srid_to_crs.get_crs(srid)? else {
return Ok(None);
} else if srid == 4326 {
return Ok(Some("OGC:CRS84".to_string()));
}

let auth_code = format!("EPSG:{srid}");
if !known_valid.contains(&srid) {
validate_crs(&auth_code, maybe_engine)?;
known_valid.insert(srid);
}

};
validate_crs(&auth_code, maybe_engine)?;
Ok(Some(auth_code))
} else {
Ok(None)
Expand All @@ -507,7 +499,7 @@ fn normalize_crs_array(
Ok(Arc::new(utf8_view_array))
}
_ => {
let mut known_abbreviated = HashMap::<String, String>::new();
let mut crs_norm = CachedCrsNormalization::new();

let string_value = crs_value.cast_to(&DataType::Utf8View, None)?;
let string_array_ref = ColumnarValue::values_to_arrays(&[string_value])?;
Expand All @@ -516,25 +508,7 @@ fn normalize_crs_array(
.iter()
.map(|maybe_crs| -> Result<Option<String>> {
if let Some(crs_str) = maybe_crs {
if crs_str == "0" {
return Ok(None);
}

if let Some(abbreviated_crs) = known_abbreviated.get(crs_str) {
Ok(Some(abbreviated_crs.clone()))
} else if let Some(crs) = deserialize_crs(crs_str)? {
let abbreviated_crs =
if let Some(auth_code) = crs.to_authority_code()? {
auth_code
} else {
crs_str.to_string()
};

known_abbreviated.insert(crs.to_string(), abbreviated_crs.clone());
Ok(Some(abbreviated_crs))
} else {
Ok(None)
}
crs_norm.normalize(crs_str)
} else {
Ok(None)
}
Expand Down
14 changes: 14 additions & 0 deletions rust/sedona-raster-functions/benches/native-raster-functions.rs
Original file line number Diff line number Diff line change
Expand Up @@ -55,6 +55,20 @@ fn criterion_benchmark(c: &mut Criterion) {
BenchmarkArgs::ArrayScalarScalar(Raster(64, 64), Int32(0, 63), Int32(0, 63)),
);
benchmark::scalar(c, &f, "native-raster", "rs_rotation", Raster(64, 64));
benchmark::scalar(
c,
&f,
"native-raster",
"rs_setcrs",
BenchmarkArgs::ArrayScalar(Raster(64, 64), String("EPSG:3857".to_string())),
);
benchmark::scalar(
c,
&f,
"native-raster",
"rs_setsrid",
BenchmarkArgs::ArrayScalar(Raster(64, 64), Int32(3857, 3858)),
);
benchmark::scalar(c, &f, "native-raster", "rs_scalex", Raster(64, 64));
benchmark::scalar(c, &f, "native-raster", "rs_scaley", Raster(64, 64));
benchmark::scalar(c, &f, "native-raster", "rs_skewx", Raster(64, 64));
Expand Down
1 change: 1 addition & 0 deletions rust/sedona-raster-functions/src/lib.rs
Original file line number Diff line number Diff line change
Expand Up @@ -24,6 +24,7 @@ pub mod rs_georeference;
pub mod rs_geotransform;
pub mod rs_numbands;
pub mod rs_rastercoordinate;
pub mod rs_setsrid;
pub mod rs_size;
pub mod rs_srid;
pub mod rs_worldcoordinate;
2 changes: 2 additions & 0 deletions rust/sedona-raster-functions/src/register.rs
Original file line number Diff line number Diff line change
Expand Up @@ -55,6 +55,8 @@ pub fn default_function_set() -> FunctionSet {
crate::rs_rastercoordinate::rs_worldtorastercoordy_udf,
crate::rs_size::rs_height_udf,
crate::rs_size::rs_width_udf,
crate::rs_setsrid::rs_set_crs_udf,
crate::rs_setsrid::rs_set_srid_udf,
crate::rs_srid::rs_crs_udf,
crate::rs_srid::rs_srid_udf,
crate::rs_worldcoordinate::rs_rastertoworldcoord_udf,
Expand Down
Loading