diff --git a/core/wren-core-base/.gitignore b/core/wren-core-base/.gitignore new file mode 100644 index 000000000..7358fa77e --- /dev/null +++ b/core/wren-core-base/.gitignore @@ -0,0 +1,4 @@ +Cargo.lock +target/ +manifest-macro/Cargo.lock +manifest-macro/target/ \ No newline at end of file diff --git a/core/wren-core-base/Cargo.toml b/core/wren-core-base/Cargo.toml new file mode 100644 index 000000000..316e1a5b8 --- /dev/null +++ b/core/wren-core-base/Cargo.toml @@ -0,0 +1,19 @@ +[package] +name = "wren-core-base" +version = "0.1.0" +edition = "2021" + +[features] +python-binding = ["dep:pyo3"] +default = [] + +[dependencies] +pyo3 = { version = "0.23.3", features = ["extension-module"], optional = true } +serde = { version = "1.0.201", features = ["derive", "rc"] } +wren-manifest-macro = { path = "manifest-macro" } +serde_json = { version = "1.0.117" } +serde_with = { version = "3.11.0" } + +[lib] +name = "wren_core_base" +path = "src/lib.rs" diff --git a/core/wren-core-base/README.md b/core/wren-core-base/README.md new file mode 100644 index 000000000..74ac0aa16 --- /dev/null +++ b/core/wren-core-base/README.md @@ -0,0 +1,5 @@ +# Wren Core Base Module +This module is the base module for Wren Core. It contains the common utilities, the base traits and structs for the Wren Core. + +## Crate Features +- `python-binding`: Enable the Python binding to access the Manifest struct in Python. This feature is disabled by default. \ No newline at end of file diff --git a/core/wren-core-base/manifest-macro/Cargo.toml b/core/wren-core-base/manifest-macro/Cargo.toml new file mode 100644 index 000000000..cce385b94 --- /dev/null +++ b/core/wren-core-base/manifest-macro/Cargo.toml @@ -0,0 +1,17 @@ +[package] +name = "wren-manifest-macro" +version = "0.1.0" +edition = "2021" + +[dependencies] +syn = { version = "2.0", default-features = false, features = [ + "printing", + "parsing", + "proc-macro", +] } +quote = "1.0" + +[lib] +proc-macro = true +name = "manifest_macro" +path = "src/lib.rs" \ No newline at end of file diff --git a/core/wren-core-base/manifest-macro/README.md b/core/wren-core-base/manifest-macro/README.md new file mode 100644 index 000000000..b2dac2aca --- /dev/null +++ b/core/wren-core-base/manifest-macro/README.md @@ -0,0 +1,14 @@ +# Wren Core Manifest Macro +This is module to collect the generating macros for the manifest struct of Wren MDL. +They are used to generate the manifest struct for different bindings. +Currently, we have the following bindings: +- Python +- Rust + +## Example +```rust +use wren_core_manifest_macro::manifest; + +manifest!(true); // Generate the manifest struct for Python binding +manifest!(false); // Generate the manifest struct for Rust binding +``` \ No newline at end of file diff --git a/core/wren-core-base/manifest-macro/src/lib.rs b/core/wren-core-base/manifest-macro/src/lib.rs new file mode 100644 index 000000000..6e8e8f06d --- /dev/null +++ b/core/wren-core-base/manifest-macro/src/lib.rs @@ -0,0 +1,344 @@ +/* + * Licensed to the Apache Software Foundation (ASF) under one + * or more contributor license agreements. See the NOTICE file + * distributed with this work for additional information + * regarding copyright ownership. The ASF licenses this file + * to you under the Apache License, Version 2.0 (the + * "License"); you may not use this file except in compliance + * with the License. You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, + * software distributed under the License is distributed on an + * "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY + * KIND, either express or implied. See the License for the + * specific language governing permissions and limitations + * under the License. + */ + +use quote::quote; +use syn::{parse_macro_input, LitBool}; + +/// This macro generates a struct for `Manifest` +/// If python_binding is true, it will generate a `pyclass` attribute +#[proc_macro] +pub fn manifest(python_binding: proc_macro::TokenStream) -> proc_macro::TokenStream { + let input = parse_macro_input!(python_binding as LitBool); + let python_binding = if input.value { + quote! { + #[pyclass] + } + } else { + quote! {} + }; + + let expanded = quote! { + #python_binding + #[derive(Serialize, Deserialize, Debug, PartialEq, Eq, Hash, Clone)] + #[serde(rename_all = "camelCase")] + pub struct Manifest { + pub catalog: String, + pub schema: String, + #[serde(default)] + pub models: Vec>, + #[serde(default)] + pub relationships: Vec>, + #[serde(default)] + pub metrics: Vec>, + #[serde(default)] + pub views: Vec>, + #[serde(default)] + pub data_source: Option, + } + }; + proc_macro::TokenStream::from(expanded) +} + +/// This macro generates an enum for `DataSource` +/// If python_binding is true, it will generate a `pyclass` attribute +#[proc_macro] +pub fn data_source(python_binding: proc_macro::TokenStream) -> proc_macro::TokenStream { + let input = parse_macro_input!(python_binding as LitBool); + let python_binding = if input.value { + quote! { + #[pyclass(eq, eq_int)] + } + } else { + quote! {} + }; + + let expanded = quote! { + #python_binding + #[derive(Serialize, Deserialize, Debug, Default, PartialEq, Eq, Hash, Clone, Copy)] + #[serde(rename_all = "UPPERCASE")] + pub enum DataSource { + #[serde(alias = "bigquery")] + BigQuery, + #[serde(alias = "clickhouse")] + Clickhouse, + #[serde(alias = "canner")] + Canner, + #[serde(alias = "trino")] + Trino, + #[serde(alias = "mssql")] + MSSQL, + #[serde(alias = "mysql")] + MySQL, + #[serde(alias = "postgres")] + Postgres, + #[serde(alias = "snowflake")] + Snowflake, + #[default] + #[serde(alias = "datafusion")] + Datafusion, + #[serde(alias = "duckdb")] + DuckDB, + } + }; + proc_macro::TokenStream::from(expanded) +} + +/// This macro generates a struct for `Model` +/// If python_binding is true, it will generate a `pyclass` attribute +#[proc_macro] +pub fn model(python_binding: proc_macro::TokenStream) -> proc_macro::TokenStream { + let input = parse_macro_input!(python_binding as LitBool); + let python_binding = if input.value { + quote! { + #[pyclass] + } + } else { + quote! {} + }; + + let expanded = quote! { + #python_binding + #[serde_as] + #[derive(Serialize, Deserialize, Debug, PartialEq, Eq, Hash, Clone)] + #[serde(rename_all = "camelCase")] + pub struct Model { + pub name: String, + #[serde(default)] + pub ref_sql: Option, + #[serde(default)] + pub base_object: Option, + #[serde(default, with = "table_reference")] + pub table_reference: Option, + pub columns: Vec>, + #[serde(default)] + pub primary_key: Option, + #[serde(default, with = "bool_from_int")] + pub cached: bool, + #[serde(default)] + pub refresh_time: Option, + } + }; + proc_macro::TokenStream::from(expanded) +} + +/// This macro generates a struct for `Column` +/// If python_binding is true, it will generate a `pyclass` attribute +#[proc_macro] +pub fn column(python_binding: proc_macro::TokenStream) -> proc_macro::TokenStream { + let input = parse_macro_input!(python_binding as LitBool); + let python_binding = if input.value { + quote! { + #[pyclass] + } + } else { + quote! {} + }; + + let expanded = quote! { + #python_binding + #[serde_as] + #[derive(Serialize, Deserialize, Debug, PartialEq, Eq, Hash)] + #[serde(rename_all = "camelCase")] + pub struct Column { + pub name: String, + pub r#type: String, + #[serde(default)] + pub relationship: Option, + #[serde(default, with = "bool_from_int")] + pub is_calculated: bool, + #[serde(default, with = "bool_from_int")] + pub not_null: bool, + #[serde_as(as = "NoneAsEmptyString")] + #[serde(default)] + pub expression: Option, + #[serde(default, with = "bool_from_int")] + pub is_hidden: bool, + } + }; + proc_macro::TokenStream::from(expanded) +} + +/// This macro generates a struct for `Relationship` +/// If python_binding is true, it will generate a `pyclass` attribute +#[proc_macro] +pub fn relationship(python_binding: proc_macro::TokenStream) -> proc_macro::TokenStream { + let input = parse_macro_input!(python_binding as LitBool); + let python_binding = if input.value { + quote! { + #[pyclass] + } + } else { + quote! {} + }; + + let expanded = quote! { + #python_binding + #[serde_as] + #[derive(Serialize, Deserialize, Debug, Hash, PartialEq, Eq)] + #[serde(rename_all = "camelCase")] + pub struct Relationship { + pub name: String, + pub models: Vec, + pub join_type: JoinType, + pub condition: String, + } + }; + proc_macro::TokenStream::from(expanded) +} + +/// This macro generates an enum for `JoinType` +/// If python_binding is true, it will generate a `pyclass` attribute +#[proc_macro] +pub fn join_type(python_binding: proc_macro::TokenStream) -> proc_macro::TokenStream { + let input = parse_macro_input!(python_binding as LitBool); + let python_binding = if input.value { + quote! { + #[pyclass(eq, eq_int)] + } + } else { + quote! {} + }; + + let expanded = quote! { + #python_binding + #[derive(Serialize, Deserialize, Debug, PartialEq, Eq, Hash, Clone, Copy)] + #[serde(rename_all = "SCREAMING_SNAKE_CASE")] + pub enum JoinType { + #[serde(alias = "one_to_one")] + OneToOne, + #[serde(alias = "one_to_many")] + OneToMany, + #[serde(alias = "many_to_one")] + ManyToOne, + #[serde(alias = "many_to_many")] + ManyToMany, + } + }; + proc_macro::TokenStream::from(expanded) +} + +/// This macro generates a struct for `Metric` +/// If python_binding is true, it will generate a `pyclass` attribute +#[proc_macro] +pub fn metric(python_binding: proc_macro::TokenStream) -> proc_macro::TokenStream { + let input = parse_macro_input!(python_binding as LitBool); + let python_binding = if input.value { + quote! { + #[pyclass] + } + } else { + quote! {} + }; + + let expanded = quote! { + #python_binding + #[serde_as] + #[derive(Serialize, Deserialize, Debug, PartialEq, Eq, Hash)] + #[serde(rename_all = "camelCase")] + pub struct Metric { + pub name: String, + pub base_object: String, + pub dimension: Vec>, + pub measure: Vec>, + pub time_grain: Vec, + #[serde(default, with = "bool_from_int")] + pub cached: bool, + pub refresh_time: Option, + } + }; + proc_macro::TokenStream::from(expanded) +} + +/// This macro generates a struct for `TimeGrain` +/// If python_binding is true, it will generate a `pyclass` attribute +#[proc_macro] +pub fn time_grain(python_binding: proc_macro::TokenStream) -> proc_macro::TokenStream { + let input = parse_macro_input!(python_binding as LitBool); + let python_binding = if input.value { + quote! { + #[pyclass] + } + } else { + quote! {} + }; + + let expanded = quote! { + #python_binding + #[derive(Serialize, Deserialize, Debug, PartialEq, Eq, Hash, Clone)] + #[serde(rename_all = "camelCase")] + pub struct TimeGrain { + pub name: String, + pub ref_column: String, + pub date_parts: Vec, + } + }; + proc_macro::TokenStream::from(expanded) +} + +/// This macro generates an enum for `TimeUnit` +/// If python_binding is true, it will generate a `pyclass` attribute +#[proc_macro] +pub fn time_unit(python_binding: proc_macro::TokenStream) -> proc_macro::TokenStream { + let input = parse_macro_input!(python_binding as LitBool); + let python_binding = if input.value { + quote! { + #[pyclass(eq, eq_int)] + } + } else { + quote! {} + }; + + let expanded = quote! { + #python_binding + #[derive(Serialize, Deserialize, Debug, PartialEq, Eq, Hash, Clone)] + pub enum TimeUnit { + Year, + Month, + Day, + Hour, + Minute, + Second, + } + }; + proc_macro::TokenStream::from(expanded) +} + +/// This macro generates a struct for `View` +/// If python_binding is true, it will generate a `pyclass` attribute +#[proc_macro] +pub fn view(python_binding: proc_macro::TokenStream) -> proc_macro::TokenStream { + let input = parse_macro_input!(python_binding as LitBool); + let python_binding = if input.value { + quote! { + #[pyclass] + } + } else { + quote! {} + }; + + let expanded = quote! { + #python_binding + #[derive(Serialize, Deserialize, Debug, PartialEq, Eq, Hash)] + pub struct View { + pub name: String, + pub statement: String, + } + }; + proc_macro::TokenStream::from(expanded) +} diff --git a/core/wren-core-base/src/lib.rs b/core/wren-core-base/src/lib.rs new file mode 100644 index 000000000..35c771dd0 --- /dev/null +++ b/core/wren-core-base/src/lib.rs @@ -0,0 +1 @@ +pub mod mdl; diff --git a/core/wren-core/core/src/mdl/builder.rs b/core/wren-core-base/src/mdl/builder.rs similarity index 94% rename from core/wren-core/core/src/mdl/builder.rs rename to core/wren-core-base/src/mdl/builder.rs index 59dfd5f37..d69adf52b 100644 --- a/core/wren-core/core/src/mdl/builder.rs +++ b/core/wren-core-base/src/mdl/builder.rs @@ -1,8 +1,26 @@ +/* + * Licensed to the Apache Software Foundation (ASF) under one + * or more contributor license agreements. See the NOTICE file + * distributed with this work for additional information + * regarding copyright ownership. The ASF licenses this file + * to you under the Apache License, Version 2.0 (the + * "License"); you may not use this file except in compliance + * with the License. You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, + * software distributed under the License is distributed on an + * "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY + * KIND, either express or implied. See the License for the + * specific language governing permissions and limitations + * under the License. + */ + #![allow(dead_code)] use crate::mdl::manifest::{ - Column, DataSource, JoinType, Manifest, Metric, Model, Relationship, TimeGrain, - TimeUnit, View, + Column, DataSource, JoinType, Manifest, Metric, Model, Relationship, TimeGrain, TimeUnit, View, }; use std::sync::Arc; @@ -336,8 +354,7 @@ mod test { }; use crate::mdl::manifest::DataSource::MySQL; use crate::mdl::manifest::{ - Column, DataSource, JoinType, Manifest, Metric, Model, Relationship, TimeUnit, - View, + Column, DataSource, JoinType, Manifest, Metric, Model, Relationship, TimeUnit, View, }; use std::fs; use std::path::PathBuf; @@ -561,17 +578,15 @@ mod test { .build(); let json_str = serde_json::to_string(&expected).unwrap(); - let actual: crate::mdl::manifest::Manifest = - serde_json::from_str(&json_str).unwrap(); + let actual: crate::mdl::manifest::Manifest = serde_json::from_str(&json_str).unwrap(); assert_eq!(actual, expected) } #[test] fn test_json_serde() { - let test_data: PathBuf = - [env!("CARGO_MANIFEST_DIR"), "tests", "data", "mdl.json"] - .iter() - .collect(); + let test_data: PathBuf = [env!("CARGO_MANIFEST_DIR"), "tests", "data", "mdl.json"] + .iter() + .collect(); let mdl_json = fs::read_to_string(test_data.as_path()).unwrap(); let mdl = serde_json::from_str::(&mdl_json).unwrap(); diff --git a/core/wren-core/core/src/mdl/manifest.rs b/core/wren-core-base/src/mdl/manifest.rs similarity index 57% rename from core/wren-core/core/src/mdl/manifest.rs rename to core/wren-core-base/src/mdl/manifest.rs index aa99b0e5b..c9d6bce5b 100644 --- a/core/wren-core/core/src/mdl/manifest.rs +++ b/core/wren-core-base/src/mdl/manifest.rs @@ -1,53 +1,76 @@ +/* + * Licensed to the Apache Software Foundation (ASF) under one + * or more contributor license agreements. See the NOTICE file + * distributed with this work for additional information + * regarding copyright ownership. The ASF licenses this file + * to you under the Apache License, Version 2.0 (the + * "License"); you may not use this file except in compliance + * with the License. You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, + * software distributed under the License is distributed on an + * "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY + * KIND, either express or implied. See the License for the + * specific language governing permissions and limitations + * under the License. + */ use std::fmt::Display; use std::sync::Arc; -use serde::{Deserialize, Serialize}; -use serde_with::serde_as; -use serde_with::NoneAsEmptyString; - -/// This is the main struct that holds all the information about the manifest -#[derive(Serialize, Deserialize, Debug, PartialEq, Eq, Hash, Clone)] -#[serde(rename_all = "camelCase")] -pub struct Manifest { - pub catalog: String, - pub schema: String, - #[serde(default)] - pub models: Vec>, - #[serde(default)] - pub relationships: Vec>, - #[serde(default)] - pub metrics: Vec>, - #[serde(default)] - pub views: Vec>, - pub data_source: Option, +#[cfg(not(feature = "python-binding"))] +mod manifest_impl { + use crate::mdl::manifest::bool_from_int; + use crate::mdl::manifest::table_reference; + use manifest_macro::{ + column, data_source, join_type, manifest, metric, model, relationship, time_grain, + time_unit, view, + }; + use serde::{Deserialize, Serialize}; + use serde_with::serde_as; + use serde_with::NoneAsEmptyString; + use std::sync::Arc; + manifest!(false); + data_source!(false); + model!(false); + column!(false); + relationship!(false); + metric!(false); + view!(false); + join_type!(false); + time_grain!(false); + time_unit!(false); } -#[derive(Serialize, Deserialize, Debug, Default, PartialEq, Eq, Hash, Clone, Copy)] -#[serde(rename_all = "UPPERCASE")] -pub enum DataSource { - #[serde(alias = "bigquery")] - BigQuery, - #[serde(alias = "clickhouse")] - Clickhouse, - #[serde(alias = "canner")] - Canner, - #[serde(alias = "trino")] - Trino, - #[serde(alias = "mssql")] - MSSQL, - #[serde(alias = "mysql")] - MySQL, - #[serde(alias = "postgres")] - Postgres, - #[serde(alias = "snowflake")] - Snowflake, - #[serde(alias = "datafusion")] - #[default] - Datafusion, - #[serde(alias = "duckdb")] - DuckDB, +#[cfg(feature = "python-binding")] +mod manifest_impl { + use crate::mdl::manifest::bool_from_int; + use crate::mdl::manifest::table_reference; + use manifest_macro::{ + column, data_source, join_type, manifest, metric, model, relationship, time_grain, + time_unit, view, + }; + use pyo3::pyclass; + use serde::{Deserialize, Serialize}; + use serde_with::serde_as; + use serde_with::NoneAsEmptyString; + use std::sync::Arc; + + data_source!(true); + model!(true); + column!(true); + relationship!(true); + metric!(true); + view!(true); + join_type!(true); + time_grain!(true); + time_unit!(true); + manifest!(true); } +pub use crate::mdl::manifest::manifest_impl::*; + impl Display for DataSource { fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { match self { @@ -65,32 +88,6 @@ impl Display for DataSource { } } -#[serde_as] -#[derive(Serialize, Deserialize, Debug, PartialEq, Eq, Hash)] -#[serde(rename_all = "camelCase")] -pub struct Model { - pub name: String, - #[serde(default)] - pub ref_sql: Option, - #[serde(default)] - pub base_object: Option, - #[serde(default, with = "table_reference")] - pub table_reference: Option, - pub columns: Vec>, - #[serde(default)] - pub primary_key: Option, - #[serde(default, with = "bool_from_int")] - pub cached: bool, - #[serde(default)] - pub refresh_time: Option, -} - -impl Model { - pub fn table_reference(&self) -> &str { - self.table_reference.as_deref().unwrap_or("") - } -} - mod table_reference { use serde::{self, Deserialize, Deserializer, Serialize, Serializer}; @@ -122,16 +119,12 @@ mod table_reference { .filter(|s| !s.is_empty())) } - pub fn serialize( - table_ref: &Option, - serializer: S, - ) -> Result + pub fn serialize(table_ref: &Option, serializer: S) -> Result where S: Serializer, { if let Some(table_ref) = table_ref { - let parts: Vec<&str> = - table_ref.split('.').filter(|p| !p.is_empty()).collect(); + let parts: Vec<&str> = table_ref.split('.').filter(|p| !p.is_empty()).collect(); if parts.len() > 3 { return Err(serde::ser::Error::custom(format!( "Invalid table reference: {table_ref}" @@ -190,47 +183,6 @@ mod bool_from_int { } } -#[serde_as] -#[derive(Serialize, Deserialize, Debug, PartialEq, Eq, Hash)] -#[serde(rename_all = "camelCase")] -pub struct Column { - pub name: String, - pub r#type: String, - #[serde(default)] - pub relationship: Option, - #[serde(default, with = "bool_from_int")] - pub is_calculated: bool, - #[serde(default, with = "bool_from_int")] - pub not_null: bool, - #[serde_as(as = "NoneAsEmptyString")] - #[serde(default)] - pub expression: Option, - #[serde(default, with = "bool_from_int")] - pub is_hidden: bool, -} - -#[derive(Serialize, Deserialize, Debug, Hash, PartialEq, Eq)] -#[serde(rename_all = "camelCase")] -pub struct Relationship { - pub name: String, - pub models: Vec, - pub join_type: JoinType, - pub condition: String, -} - -#[derive(Serialize, Deserialize, Debug, PartialEq, Eq, Hash, Clone, Copy)] -#[serde(rename_all = "SCREAMING_SNAKE_CASE")] -pub enum JoinType { - #[serde(alias = "one_to_one")] - OneToOne, - #[serde(alias = "one_to_many")] - OneToMany, - #[serde(alias = "many_to_one")] - ManyToOne, - #[serde(alias = "many_to_many")] - ManyToMany, -} - impl JoinType { pub fn is_to_one(&self) -> bool { matches!(self, JoinType::OneToOne | JoinType::ManyToOne) @@ -248,17 +200,55 @@ impl Display for JoinType { } } -#[derive(Serialize, Deserialize, Debug, PartialEq, Eq, Hash)] -#[serde(rename_all = "camelCase")] -pub struct Metric { - pub name: String, - pub base_object: String, - pub dimension: Vec>, - pub measure: Vec>, - pub time_grain: Vec, - #[serde(default, with = "bool_from_int")] - pub cached: bool, - pub refresh_time: Option, +impl Model { + /// Physical columns are columns that can be selected from the model. + /// All physical columns are visible columns, but not all visible columns are physical columns + /// e.g. columns that are not a relationship column + pub fn get_physical_columns(&self) -> Vec> { + self.get_visible_columns() + .filter(|c| c.relationship.is_none()) + .map(|c| Arc::clone(&c)) + .collect() + } + + /// Return the name of the model + pub fn name(&self) -> &str { + &self.name + } + + /// Return the iterator of all visible columns + pub fn get_visible_columns(&self) -> impl Iterator> + '_ { + self.columns.iter().filter(|f| !f.is_hidden).map(Arc::clone) + } + + /// Get the specified visible column by name + pub fn get_column(&self, column_name: &str) -> Option> { + self.get_visible_columns() + .find(|c| c.name == column_name) + .map(|c| Arc::clone(&c)) + } + + /// Return the primary key of the model + pub fn primary_key(&self) -> Option<&str> { + self.primary_key.as_deref() + } + + /// Return the table reference of the model + pub fn table_reference(&self) -> &str { + self.table_reference.as_deref().unwrap_or("") + } +} + +impl Column { + /// Return the name of the column + pub fn name(&self) -> &str { + &self.name + } + + /// Return the expression of the column + pub fn expression(&self) -> Option<&str> { + self.expression.as_deref() + } } impl Metric { @@ -267,30 +257,6 @@ impl Metric { } } -#[derive(Serialize, Deserialize, Debug, PartialEq, Eq, Hash, Clone)] -#[serde(rename_all = "camelCase")] -pub struct TimeGrain { - pub name: String, - pub ref_column: String, - pub date_parts: Vec, -} - -#[derive(Serialize, Deserialize, Debug, PartialEq, Eq, Hash, Clone)] -pub enum TimeUnit { - Year, - Month, - Day, - Hour, - Minute, - Second, -} - -#[derive(Serialize, Deserialize, Debug, PartialEq, Eq, Hash)] -pub struct View { - pub name: String, - pub statement: String, -} - impl View { pub fn name(&self) -> &str { &self.name @@ -322,8 +288,7 @@ mod tests { .iter() .for_each(|(table_ref, expected)| { let mut buf = Vec::new(); - table_reference::serialize(table_ref, &mut Serializer::new(&mut buf)) - .unwrap(); + table_reference::serialize(table_ref, &mut Serializer::new(&mut buf)).unwrap(); assert_eq!(String::from_utf8(buf).unwrap(), *expected); }); } diff --git a/core/wren-core-base/src/mdl/mod.rs b/core/wren-core-base/src/mdl/mod.rs new file mode 100644 index 000000000..bb6c52514 --- /dev/null +++ b/core/wren-core-base/src/mdl/mod.rs @@ -0,0 +1,25 @@ +/* + * Licensed to the Apache Software Foundation (ASF) under one + * or more contributor license agreements. See the NOTICE file + * distributed with this work for additional information + * regarding copyright ownership. The ASF licenses this file + * to you under the Apache License, Version 2.0 (the + * "License"); you may not use this file except in compliance + * with the License. You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, + * software distributed under the License is distributed on an + * "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY + * KIND, either express or implied. See the License for the + * specific language governing permissions and limitations + * under the License. + */ + +pub mod builder; +pub mod manifest; +mod py_method; + +pub use builder::*; +pub use manifest::*; diff --git a/core/wren-core-base/src/mdl/py_method.rs b/core/wren-core-base/src/mdl/py_method.rs new file mode 100644 index 000000000..ef921aaee --- /dev/null +++ b/core/wren-core-base/src/mdl/py_method.rs @@ -0,0 +1,61 @@ +/* + * Licensed to the Apache Software Foundation (ASF) under one + * or more contributor license agreements. See the NOTICE file + * distributed with this work for additional information + * regarding copyright ownership. The ASF licenses this file + * to you under the Apache License, Version 2.0 (the + * "License"); you may not use this file except in compliance + * with the License. You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, + * software distributed under the License is distributed on an + * "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY + * KIND, either express or implied. See the License for the + * specific language governing permissions and limitations + * under the License. + */ + +#[cfg(feature = "python-binding")] +mod manifest_python_impl { + use crate::mdl::manifest::{Manifest, Model}; + use crate::mdl::DataSource; + use pyo3::{pymethods, PyResult}; + use std::sync::Arc; + + #[pymethods] + impl Manifest { + #[getter] + fn catalog(&self) -> PyResult { + Ok(self.catalog.clone()) + } + + #[getter] + fn schema(&self) -> PyResult { + Ok(self.schema.clone()) + } + + #[getter] + fn models(&self) -> PyResult> { + Ok(self + .models + .iter() + .map(|m| Arc::unwrap_or_clone(Arc::clone(m))) + .collect()) + } + + #[getter] + fn data_source(&self) -> PyResult> { + Ok(self.data_source) + } + } + + #[pymethods] + impl Model { + #[getter] + fn get_name(&self) -> PyResult { + Ok(self.name.clone()) + } + } +} diff --git a/core/wren-core-py/Cargo.lock b/core/wren-core-py/Cargo.lock index 975eeb8ac..631274010 100644 --- a/core/wren-core-py/Cargo.lock +++ b/core/wren-core-py/Cargo.lock @@ -524,9 +524,9 @@ dependencies = [ [[package]] name = "cc" -version = "1.2.3" +version = "1.2.5" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "27f657647bcff5394bf56c7317665bbf790a137a50eaaa5c6bfbb9e27a518f2d" +checksum = "c31a0499c1dc64f458ad13872de75c0eb7e3fdb0e67964610c914b034fc5956e" dependencies = [ "jobserver", "libc", @@ -642,9 +642,9 @@ dependencies = [ [[package]] name = "crossbeam-utils" -version = "0.8.20" +version = "0.8.21" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "22ec99545bb0ed0ea7bb9b8e1e9122ea386ff8a48c0922e43f36d45ab09e0e80" +checksum = "d0a5c400df2834b80a4c3327b3aad3a4c4cd4de0629063962b03235697506a28" [[package]] name = "crunchy" @@ -1186,9 +1186,9 @@ checksum = "60b1af1c220855b6ceac025d3f6ecdd2b7c4894bfe9cd9bda4fbb4bc7c0d4cf0" [[package]] name = "env_filter" -version = "0.1.2" +version = "0.1.3" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "4f2c92ceda6ceec50f43169f9ee8424fe2db276791afde7b2cd8bc084cb376ab" +checksum = "186e05a59d4c50738528153b83b0b0194d3a29507dfec16eccd4b342903397d0" dependencies = [ "log", "regex", @@ -1196,9 +1196,9 @@ dependencies = [ [[package]] name = "env_logger" -version = "0.11.5" +version = "0.11.6" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "e13fa619b91fb2381732789fc5de83b45675e882f66623b7d8cb4f643017018d" +checksum = "dcaee3d8e3cfc3fd92428d477bc97fc29ec8716d180c0d74c643bb26166660e0" dependencies = [ "anstream", "anstyle", @@ -1769,9 +1769,9 @@ dependencies = [ [[package]] name = "libc" -version = "0.2.168" +version = "0.2.169" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "5aaeb2981e0606ca11d79718f8bb01164f1d6ed75080182d3abf017e6d244b6d" +checksum = "b5aba8db14291edd000dfcc4d620c7ebfb122c613afb886ca8803fa4e128a20a" [[package]] name = "libm" @@ -1854,9 +1854,9 @@ dependencies = [ [[package]] name = "miniz_oxide" -version = "0.8.0" +version = "0.8.2" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "e2d80299ef12ff69b16a84bb182e3b9df68b5a91574d3d4fa6e41b65deec4df1" +checksum = "4ffbe83022cedc1d264172192511ae958937694cd57ce297164951b8b3568394" dependencies = [ "adler2", ] @@ -1953,9 +1953,9 @@ dependencies = [ [[package]] name = "object" -version = "0.36.5" +version = "0.36.7" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "aedf0a2d09c573ed1d8d85b30c119153926a2b36dce0ab28322c09a117a4683e" +checksum = "62948e14d923ea95ea2c7c86c71013138b66525b86bdc08d2dcc262bdb497b87" dependencies = [ "memchr", ] @@ -2453,9 +2453,9 @@ checksum = "94143f37725109f92c262ed2cf5e59bce7498c01bcc1502d7b9afe439a4e9f49" [[package]] name = "semver" -version = "1.0.23" +version = "1.0.24" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "61697e0a1c7e512e84a621326239844a24d8207b4669b41bc18b32ea5cbf988b" +checksum = "3cb6eb87a131f756572d7fb904f6e7b68633f09cca868c5df1c4b8d1a694bbba" [[package]] name = "seq-macro" @@ -2485,9 +2485,9 @@ dependencies = [ [[package]] name = "serde_json" -version = "1.0.133" +version = "1.0.134" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "c7fceb2473b9166b2294ef05efcb65a3db80803f0b03ef86a5fc88a2b85ee377" +checksum = "d00f4175c42ee48b15416f6193a959ba3a0d67fc699a0db9ad12df9f83991c7d" dependencies = [ "itoa", "memchr", @@ -2670,9 +2670,9 @@ checksum = "13c2bddecc57b384dee18652358fb23172facb8a2c51ccc10d74c157bdea3292" [[package]] name = "syn" -version = "2.0.90" +version = "2.0.91" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "919d3b74a5dd0ccd15aeb8f93e7006bd9e14c295087c9896a110f490752bcf31" +checksum = "d53cbcb5a243bd33b7858b1d7f4aca2153490815872d86d955d6ea29f743c035" dependencies = [ "proc-macro2", "quote", @@ -2711,18 +2711,18 @@ dependencies = [ [[package]] name = "thiserror" -version = "2.0.6" +version = "2.0.9" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "8fec2a1820ebd077e2b90c4df007bebf344cd394098a13c563957d0afc83ea47" +checksum = "f072643fd0190df67a8bab670c20ef5d8737177d6ac6b2e9a236cb096206b2cc" dependencies = [ "thiserror-impl", ] [[package]] name = "thiserror-impl" -version = "2.0.6" +version = "2.0.9" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "d65750cab40f4ff1929fb1ba509e9914eb756131cef4210da8d5d700d26f6312" +checksum = "7b50fa271071aae2e6ee85f842e2e28ba8cd2c5fb67f11fcb1fd70b276f9e7d4" dependencies = [ "proc-macro2", "quote", @@ -3155,6 +3155,18 @@ dependencies = [ "serde_json", "serde_with", "tokio", + "wren-core-base", +] + +[[package]] +name = "wren-core-base" +version = "0.1.0" +dependencies = [ + "pyo3", + "serde", + "serde_json", + "serde_with", + "wren-manifest-macro", ] [[package]] @@ -3173,6 +3185,15 @@ dependencies = [ "thiserror", "tokio", "wren-core", + "wren-core-base", +] + +[[package]] +name = "wren-manifest-macro" +version = "0.1.0" +dependencies = [ + "quote", + "syn", ] [[package]] diff --git a/core/wren-core-py/Cargo.toml b/core/wren-core-py/Cargo.toml index 5fdf2d45c..53605d695 100644 --- a/core/wren-core-py/Cargo.toml +++ b/core/wren-core-py/Cargo.toml @@ -11,6 +11,7 @@ crate-type = ["cdylib"] [dependencies] pyo3 = { version = "0.23.3", features = ["extension-module"] } wren-core = { path = "../wren-core/core" } +wren-core-base = { path = "../wren-core-base", features = ["python-binding"] } base64 = "0.22.1" serde_json = "1.0.117" thiserror = "2.0.3" diff --git a/core/wren-core-py/src/extractor.rs b/core/wren-core-py/src/extractor.rs index 20307dacf..ca66071ea 100644 --- a/core/wren-core-py/src/extractor.rs +++ b/core/wren-core-py/src/extractor.rs @@ -1,11 +1,12 @@ use crate::errors::CoreError; -use crate::manifest::{to_manifest, PyManifest}; +use crate::manifest::to_manifest; use pyo3::{pyclass, pymethods}; use std::collections::hash_map::Entry; use std::collections::{HashMap, HashSet}; use std::sync::Arc; use wren_core::mdl::manifest::{Model, Relationship, View}; use wren_core::mdl::WrenMDL; +use wren_core_base::mdl::Manifest; #[pyclass] #[derive(Clone)] @@ -36,10 +37,7 @@ impl PyManifestExtractor { /// If a model is related to another dataset, both datasets will be kept. /// The relationship between of them will be kept as well. /// A dataset could be model, view. - pub fn extract_by( - &self, - used_datasets: Vec, - ) -> Result { + pub fn extract_by(&self, used_datasets: Vec) -> Result { extract_manifest(&self.mdl, &used_datasets) } } @@ -69,19 +67,19 @@ fn resolve_used_table_names(mdl: &WrenMDL, sql: &str) -> Result, Cor fn extract_manifest( mdl: &WrenMDL, used_datasets: &[String], -) -> Result { +) -> Result { let extracted_models = extract_models(mdl, used_datasets); let (used_views, models_of_views) = extract_views(mdl, used_datasets); let used_models = [extracted_models, models_of_views].concat(); let used_relationships = extract_relationships(mdl, &used_models); - Ok(PyManifest { + Ok(Manifest { catalog: mdl.catalog().to_string(), schema: mdl.schema().to_string(), models: used_models, relationships: used_relationships, metrics: mdl.metrics().to_vec(), views: used_views, - data_source: *mdl.data_source(), + data_source: mdl.data_source(), }) } @@ -157,13 +155,13 @@ fn extract_relationships( #[cfg(test)] mod tests { use crate::extractor::PyManifestExtractor; - use crate::manifest::{to_json_base64, PyManifest}; + use crate::manifest::to_json_base64; use rstest::{fixture, rstest}; use std::iter::Iterator; - use wren_core::mdl::builder::{ + use wren_core::mdl::manifest::{DataSource, JoinType}; + use wren_core_base::mdl::builder::{ ColumnBuilder, ManifestBuilder, ModelBuilder, RelationshipBuilder, ViewBuilder, }; - use wren_core::mdl::manifest::{DataSource, JoinType}; #[fixture] pub fn mdl_base64() -> String { @@ -216,7 +214,7 @@ mod tests { .view(c_view) .data_source(DataSource::BigQuery) .build(); - to_json_base64(PyManifest::from(&manifest)).unwrap() + to_json_base64(manifest).unwrap() } #[fixture] diff --git a/core/wren-core-py/src/lib.rs b/core/wren-core-py/src/lib.rs index 30b81df6d..af7f1fdf1 100644 --- a/core/wren-core-py/src/lib.rs +++ b/core/wren-core-py/src/lib.rs @@ -14,7 +14,7 @@ fn wren_core_wrapper(m: &Bound<'_, PyModule>) -> PyResult<()> { env_logger::init(); m.add_class::()?; m.add_class::()?; - m.add_class::()?; + m.add_class::()?; m.add_class::()?; m.add_function(wrap_pyfunction!(manifest::to_json_base64, m)?)?; Ok(()) diff --git a/core/wren-core-py/src/manifest.rs b/core/wren-core-py/src/manifest.rs index ce9297036..9749b1303 100644 --- a/core/wren-core-py/src/manifest.rs +++ b/core/wren-core-py/src/manifest.rs @@ -1,18 +1,13 @@ use crate::errors::CoreError; use base64::prelude::BASE64_STANDARD; use base64::Engine; -use pyo3::{pyclass, pyfunction, pymethods, PyResult}; -use serde::{Deserialize, Serialize}; -use std::iter::Iterator; -use std::sync::Arc; -use wren_core::mdl::manifest::{ - Column, DataSource, JoinType, Manifest, Metric, Model, Relationship, TimeGrain, - TimeUnit, View, -}; +use pyo3::pyfunction; + +pub use wren_core_base::mdl::*; /// Convert a manifest to a JSON string and then encode it as base64. #[pyfunction] -pub fn to_json_base64(mdl: PyManifest) -> Result { +pub fn to_json_base64(mdl: Manifest) -> Result { let mdl_json = serde_json::to_string(&mdl)?; let mdl_base64 = BASE64_STANDARD.encode(mdl_json.as_bytes()); Ok(mdl_base64) @@ -26,342 +21,16 @@ pub fn to_manifest(mdl_base64: &str) -> Result { Ok(manifest) } -#[pyclass(name = "Manifest")] -#[derive(Serialize, Deserialize, Clone, Debug)] -#[serde(rename_all = "camelCase")] -pub struct PyManifest { - pub catalog: String, - pub schema: String, - pub data_source: Option, - pub models: Vec>, - pub relationships: Vec>, - pub metrics: Vec>, - pub views: Vec>, -} - -#[pymethods] -impl PyManifest { - #[getter] - fn catalog(&self) -> PyResult { - Ok(self.catalog.clone()) - } - - #[getter] - fn schema(&self) -> PyResult { - Ok(self.schema.clone()) - } - - #[getter] - fn models(&self) -> PyResult> { - Ok(self - .models - .iter() - .map(|m| PyModel::from(m.as_ref())) - .collect()) - } - - #[getter] - fn relationships(&self) -> PyResult> { - Ok(self - .relationships - .iter() - .map(|r| PyRelationship::from(r.as_ref())) - .collect()) - } - - #[getter] - fn metrics(&self) -> PyResult> { - Ok(self - .metrics - .iter() - .map(|m| PyMetric::from(m.as_ref())) - .collect()) - } - - #[getter] - fn views(&self) -> PyResult> { - Ok(self - .views - .iter() - .map(|v| PyView::from(v.as_ref())) - .collect()) - } - - #[getter] - fn data_source(&self) -> PyResult> { - Ok(self.data_source.map(PyDataSource::from)) - } -} - -impl From<&Manifest> for PyManifest { - fn from(manifest: &Manifest) -> Self { - Self { - catalog: manifest.catalog.clone(), - schema: manifest.schema.clone(), - models: manifest.models.clone(), - relationships: manifest.relationships.clone(), - metrics: manifest.metrics.clone(), - views: manifest.views.clone(), - data_source: manifest.data_source, - } - } -} - -#[pyclass(name = "Model")] -#[derive(Serialize, Deserialize, Debug)] -pub struct PyModel { - #[pyo3(get)] - pub name: String, - #[pyo3(get)] - pub ref_sql: Option, - #[pyo3(get)] - pub base_object: Option, - #[pyo3(get)] - pub table_reference: Option, - pub columns: Vec>, - #[pyo3(get)] - pub primary_key: Option, - #[pyo3(get)] - pub cached: bool, - #[pyo3(get)] - pub refresh_time: Option, -} - -#[pymethods] -impl PyModel { - #[getter] - fn columns(&self) -> PyResult> { - Ok(self - .columns - .iter() - .map(|c| PyColumn::from(c.as_ref())) - .collect()) - } -} - -impl From<&Model> for PyModel { - fn from(model: &Model) -> Self { - Self { - name: model.name.clone(), - ref_sql: model.ref_sql.clone(), - base_object: model.base_object.clone(), - table_reference: Some(String::from(model.table_reference())), - columns: model.columns.clone(), - primary_key: model.primary_key.clone(), - cached: model.cached, - refresh_time: model.refresh_time.clone(), - } - } -} - -#[pyclass(name = "Column")] -#[derive(Serialize, Deserialize, Debug)] -pub struct PyColumn { - #[pyo3(get)] - pub name: String, - #[pyo3(get)] - pub r#type: String, - #[pyo3(get)] - pub relationship: Option, - #[pyo3(get)] - pub is_calculated: bool, - #[pyo3(get)] - pub not_null: bool, - #[pyo3(get)] - pub expression: Option, - #[pyo3(get)] - pub is_hidden: bool, -} - -impl From<&Column> for PyColumn { - fn from(column: &Column) -> Self { - Self { - name: column.name.clone(), - r#type: column.r#type.clone(), - relationship: column.relationship.clone(), - is_calculated: column.is_calculated, - not_null: column.not_null, - expression: column.expression.clone(), - is_hidden: column.is_hidden, - } - } -} - -#[pyclass(name = "Relationship")] -#[derive(Serialize, Deserialize, Debug)] -pub struct PyRelationship { - #[pyo3(get)] - pub name: String, - #[pyo3(get)] - pub models: Vec, - pub join_type: JoinType, - #[pyo3(get)] - pub condition: String, -} - -#[pymethods] -impl PyRelationship { - #[getter] - fn join_type(&self) -> PyResult { - Ok(PyJoinType::from(&self.join_type)) - } -} - -impl From<&Relationship> for PyRelationship { - fn from(relationship: &Relationship) -> Self { - Self { - name: relationship.name.clone(), - models: relationship.models.clone(), - join_type: relationship.join_type, - condition: relationship.condition.clone(), - } - } -} - -#[pyclass(name = "JoinType", eq)] -#[derive(Serialize, Deserialize, PartialEq, Eq, Debug)] -pub enum PyJoinType { - #[serde(alias = "one_to_one")] - OneToOne, - #[serde(alias = "one_to_many")] - OneToMany, - #[serde(alias = "many_to_one")] - ManyToOne, - #[serde(alias = "many_to_many")] - ManyToMany, -} - -impl From<&JoinType> for PyJoinType { - fn from(join_type: &JoinType) -> Self { - match join_type { - JoinType::OneToOne => PyJoinType::OneToOne, - JoinType::OneToMany => PyJoinType::OneToMany, - JoinType::ManyToOne => PyJoinType::ManyToOne, - JoinType::ManyToMany => PyJoinType::ManyToMany, - } - } -} - -#[pyclass(name = "Metric")] -#[derive(Serialize, Deserialize, Debug)] -pub struct PyMetric { - #[pyo3(get)] - pub name: String, - #[pyo3(get)] - pub base_object: String, - pub dimension: Vec>, - pub measure: Vec>, - pub time_grain: Vec, - #[pyo3(get)] - pub cached: bool, - #[pyo3(get)] - pub refresh_time: Option, -} - -impl From<&Metric> for PyMetric { - fn from(metric: &Metric) -> Self { - Self { - name: metric.name.clone(), - base_object: metric.base_object.clone(), - dimension: metric.dimension.clone(), - measure: metric.measure.clone(), - time_grain: metric.time_grain.clone(), - cached: metric.cached, - refresh_time: metric.refresh_time.clone(), - } - } -} - -#[pyclass(name = "TimeGrain")] -#[derive(Serialize, Deserialize, Debug)] -pub struct PyTimeGrain { - pub name: String, - pub ref_column: String, - pub date_parts: Vec, -} - -#[pyclass(name = "TimeUnit", eq)] -#[derive(Serialize, Deserialize, PartialEq, Eq, Debug)] -pub enum PyTimeUnit { - Year, - Month, - Day, - Hour, - Minute, - Second, -} - -#[pyclass(name = "View")] -#[derive(Serialize, Deserialize, Debug)] -pub struct PyView { - #[pyo3(get)] - pub name: String, - #[pyo3(get)] - pub statement: String, -} - -impl From<&View> for PyView { - fn from(view: &View) -> Self { - Self { - name: view.name.clone(), - statement: view.statement.clone(), - } - } -} - -#[pyclass(name = "DataSource", eq)] -#[derive(Serialize, Deserialize, Debug, PartialEq, Eq, Hash, Clone, Copy)] -#[serde(rename_all = "SCREAMING_SNAKE_CASE")] -pub enum PyDataSource { - #[serde(alias = "bigquery")] - BigQuery, - #[serde(alias = "clickhouse")] - Clickhouse, - #[serde(alias = "canner")] - Canner, - #[serde(alias = "trino")] - Trino, - #[serde(alias = "mssql")] - MsSQL, - #[serde(alias = "mysql")] - MySQL, - #[serde(alias = "postgres")] - Postgres, - #[serde(alias = "snowflake")] - Snowflake, - #[serde(alias = "datafusion")] - Datafusion, - #[serde(alias = "duckdb")] - DuckDB, -} - -impl From for PyDataSource { - fn from(data_source: DataSource) -> Self { - match data_source { - DataSource::BigQuery => PyDataSource::BigQuery, - DataSource::Clickhouse => PyDataSource::Clickhouse, - DataSource::Canner => PyDataSource::Canner, - DataSource::Trino => PyDataSource::Trino, - DataSource::MSSQL => PyDataSource::MsSQL, - DataSource::MySQL => PyDataSource::MySQL, - DataSource::Postgres => PyDataSource::Postgres, - DataSource::Snowflake => PyDataSource::Snowflake, - DataSource::Datafusion => PyDataSource::Datafusion, - DataSource::DuckDB => PyDataSource::DuckDB, - } - } -} - #[cfg(test)] mod tests { - use crate::manifest::{to_json_base64, to_manifest, PyManifest}; + use crate::manifest::{to_json_base64, to_manifest, Manifest}; use std::sync::Arc; use wren_core::mdl::manifest::DataSource::BigQuery; use wren_core::mdl::manifest::Model; #[test] fn test_manifest_to_json_base64() { - let py_manifest = PyManifest { + let py_manifest = Manifest { catalog: "catalog".to_string(), schema: "schema".to_string(), models: vec![ diff --git a/core/wren-core/Cargo.toml b/core/wren-core/Cargo.toml index 5ff81c851..ba8294efb 100644 --- a/core/wren-core/Cargo.toml +++ b/core/wren-core/Cargo.toml @@ -27,3 +27,4 @@ serde_json = { version = "1.0.117" } serde_with = { version = "3.11.0" } tokio = { version = "1.4.0", features = ["rt", "rt-multi-thread", "macros"] } wren-core = { path = "core" } +wren-core-base = { path = "../wren-core-base" } diff --git a/core/wren-core/core/Cargo.toml b/core/wren-core/core/Cargo.toml index 392b05056..3c1336de0 100644 --- a/core/wren-core/core/Cargo.toml +++ b/core/wren-core/core/Cargo.toml @@ -33,3 +33,4 @@ serde = { workspace = true } serde_json = { workspace = true } serde_with = { workspace = true } tokio = { workspace = true, features = ["rt", "rt-multi-thread", "macros"] } +wren-core-base = { workspace = true } diff --git a/core/wren-core/core/src/mdl/dataset.rs b/core/wren-core/core/src/mdl/dataset.rs index d1b5880c9..13e7f1c1f 100644 --- a/core/wren-core/core/src/mdl/dataset.rs +++ b/core/wren-core/core/src/mdl/dataset.rs @@ -1,100 +1,12 @@ -use crate::logical_plan::utils::map_data_type; -use crate::mdl::manifest::{Column, Metric, Model}; -use crate::mdl::utils::quoted; +use crate::mdl::manifest::{Metric, Model}; +use crate::mdl::utils::{quoted, to_field, to_remote_field}; use crate::mdl::{RegisterTables, SessionStateRef}; use datafusion::arrow::datatypes::Field; use datafusion::common::DFSchema; use datafusion::common::Result; -use datafusion::logical_expr::sqlparser::ast::Expr::CompoundIdentifier; -use datafusion::sql::sqlparser::ast::Expr::Identifier; -use datafusion::sql::sqlparser::ast::{visit_expressions, Expr, Ident}; use std::fmt::Display; -use std::ops::ControlFlow; use std::sync::Arc; -impl Model { - /// Physical columns are columns that can be selected from the model. - /// All physical columns are visible columns, but not all visible columns are physical columns - /// e.g. columns that are not a relationship column - pub fn get_physical_columns(&self) -> Vec> { - self.get_visible_columns() - .filter(|c| c.relationship.is_none()) - .map(|c| Arc::clone(&c)) - .collect() - } - - /// Return the name of the model - pub fn name(&self) -> &str { - &self.name - } - - /// Return the iterator of all visible columns - pub fn get_visible_columns(&self) -> impl Iterator> + '_ { - self.columns.iter().filter(|f| !f.is_hidden).map(Arc::clone) - } - - /// Get the specified visible column by name - pub fn get_column(&self, column_name: &str) -> Option> { - self.get_visible_columns() - .find(|c| c.name == column_name) - .map(|c| Arc::clone(&c)) - } - - /// Return the primary key of the model - pub fn primary_key(&self) -> Option<&str> { - self.primary_key.as_deref() - } -} - -impl Column { - /// Return the name of the column - pub fn name(&self) -> &str { - &self.name - } - - /// Return the expression of the column - pub fn expression(&self) -> Option<&str> { - self.expression.as_deref() - } - - /// Transform the column to a datafusion field - pub fn to_field(&self) -> Result { - let data_type = map_data_type(&self.r#type)?; - Ok(Field::new(&self.name, data_type, self.not_null)) - } - - /// Transform the column to a datafusion field for a remote table - pub fn to_remote_field(&self, session_state: SessionStateRef) -> Result> { - if self.expression().is_some() { - let session_state = session_state.read(); - let expr = session_state.sql_to_expr( - self.expression().unwrap(), - session_state.config_options().sql_parser.dialect.as_str(), - )?; - let columns = Self::collect_columns(expr); - columns - .into_iter() - .map(|c| Ok(Field::new(c.value, map_data_type(&self.r#type)?, false))) - .collect::>() - } else { - Ok(vec![self.to_field()?]) - } - } - - fn collect_columns(expr: Expr) -> Vec { - let mut visited = vec![]; - visit_expressions(&expr, |e| { - if let CompoundIdentifier(ids) = e { - ids.iter().cloned().for_each(|id| visited.push(id)); - } else if let Identifier(id) = e { - visited.push(id.clone()); - } - ControlFlow::<()>::Continue(()) - }); - visited - } -} - #[derive(PartialEq, Eq, Hash, Debug, Clone)] pub enum Dataset { Model(Arc), @@ -122,7 +34,7 @@ impl Dataset { let fields: Vec<_> = model .get_physical_columns() .iter() - .map(|c| c.to_field()) + .map(|c| to_field(c)) .collect::>()?; let arrow_schema = datafusion::arrow::datatypes::Schema::new(fields); DFSchema::try_from_qualified_schema(quoted(&model.name), &arrow_schema) @@ -151,7 +63,7 @@ impl Dataset { .get_physical_columns() .iter() .filter(|c| !c.is_calculated) - .map(|c| c.to_remote_field(Arc::clone(&session_state))) + .map(|c| to_remote_field(c, Arc::clone(&session_state))) .collect::>>>()? .iter() .flat_map(|c| c.clone()) diff --git a/core/wren-core/core/src/mdl/mod.rs b/core/wren-core/core/src/mdl/mod.rs index a7f2274ad..baeb657ad 100644 --- a/core/wren-core/core/src/mdl/mod.rs +++ b/core/wren-core/core/src/mdl/mod.rs @@ -6,7 +6,8 @@ use crate::mdl::function::{ ByPassAggregateUDF, ByPassScalarUDF, ByPassWindowFunction, FunctionType, RemoteFunction, }; -use crate::mdl::manifest::{Column, DataSource, Manifest, Metric, Model, View}; +use crate::mdl::manifest::{Column, Manifest, Metric, Model, View}; +use crate::mdl::utils::to_field; use crate::DataFusionError; use datafusion::arrow::datatypes::Field; use datafusion::common::internal_datafusion_err; @@ -26,14 +27,19 @@ use manifest::Relationship; use parking_lot::RwLock; use std::hash::Hash; use std::{collections::HashMap, sync::Arc}; +use wren_core_base::mdl::DataSource; -pub mod builder; +pub mod builder { + pub use wren_core_base::mdl::builder::*; +} pub mod context; pub(crate) mod dataset; mod dialect; pub mod function; pub mod lineage; -pub mod manifest; +pub mod manifest { + pub use wren_core_base::mdl::manifest::*; +} pub mod utils; pub type SessionStateRef = Arc>; @@ -220,7 +226,7 @@ impl WrenMDL { Ok(None) } } else { - Ok(Some(column.to_field()?)) + Ok(Some(to_field(column)?)) } } @@ -281,8 +287,8 @@ impl WrenMDL { &self.manifest.metrics } - pub fn data_source(&self) -> &Option { - &self.manifest.data_source + pub fn data_source(&self) -> Option { + self.manifest.data_source } pub fn get_model(&self, name: &str) -> Option> { diff --git a/core/wren-core/core/src/mdl/utils.rs b/core/wren-core/core/src/mdl/utils.rs index 66013f336..6c205c8b9 100644 --- a/core/wren-core/core/src/mdl/utils.rs +++ b/core/wren-core/core/src/mdl/utils.rs @@ -1,7 +1,4 @@ -use std::collections::{BTreeSet, VecDeque}; -use std::ops::ControlFlow; -use std::sync::Arc; - +use datafusion::arrow::datatypes::Field; use datafusion::common::{plan_err, Column, DFSchema}; use datafusion::error::Result; use datafusion::execution::session_state::SessionState; @@ -12,8 +9,11 @@ use datafusion::sql::sqlparser::dialect::GenericDialect; use datafusion::sql::sqlparser::parser::Parser; use petgraph::algo::is_cyclic_directed; use petgraph::{EdgeType, Graph}; +use std::collections::{BTreeSet, VecDeque}; +use std::ops::ControlFlow; +use std::sync::Arc; -use crate::logical_plan::utils::from_qualified_name; +use crate::logical_plan::utils::{from_qualified_name, map_data_type}; use crate::mdl::manifest::Model; use crate::mdl::{AnalyzedWrenMDL, ColumnReference, Dataset, SessionStateRef}; @@ -210,6 +210,46 @@ pub fn quoted(s: &str) -> String { format!("\"{}\"", s) } +/// Transform the column to a datafusion field +pub fn to_field(column: &wren_core_base::mdl::Column) -> Result { + let data_type = map_data_type(&column.r#type)?; + Ok(Field::new(&column.name, data_type, column.not_null)) +} + +/// Transform the column to a datafusion field for a remote table +pub fn to_remote_field( + column: &wren_core_base::mdl::Column, + session_state: SessionStateRef, +) -> Result> { + if column.expression().is_some() { + let session_state = session_state.read(); + let expr = session_state.sql_to_expr( + column.expression().unwrap(), + session_state.config_options().sql_parser.dialect.as_str(), + )?; + let columns = collect_columns(expr); + columns + .into_iter() + .map(|c| Ok(Field::new(c.value, map_data_type(&column.r#type)?, false))) + .collect::>() + } else { + Ok(vec![to_field(column)?]) + } +} + +fn collect_columns(expr: datafusion::logical_expr::sqlparser::ast::Expr) -> Vec { + let mut visited = vec![]; + visit_expressions(&expr, |e| { + if let CompoundIdentifier(ids) = e { + ids.iter().cloned().for_each(|id| visited.push(id)); + } else if let Identifier(id) = e { + visited.push(id.clone()); + } + ControlFlow::<()>::Continue(()) + }); + visited +} + #[cfg(test)] mod tests { use std::fs;