refactor(core): introduce wren-core-base module to collect the common structs and utilities (#1003)

This commit is contained in:
Jax Liu
2024-12-24 14:56:03 +08:00
committed by GitHub
parent 1449143d62
commit 1c4692291b
21 changed files with 759 additions and 640 deletions
+4
View File
@@ -0,0 +1,4 @@
Cargo.lock
target/
manifest-macro/Cargo.lock
manifest-macro/target/
+19
View File
@@ -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"
+5
View File
@@ -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.
@@ -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"
@@ -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
```
@@ -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<Arc<Model>>,
#[serde(default)]
pub relationships: Vec<Arc<Relationship>>,
#[serde(default)]
pub metrics: Vec<Arc<Metric>>,
#[serde(default)]
pub views: Vec<Arc<View>>,
#[serde(default)]
pub data_source: Option<DataSource>,
}
};
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<String>,
#[serde(default)]
pub base_object: Option<String>,
#[serde(default, with = "table_reference")]
pub table_reference: Option<String>,
pub columns: Vec<Arc<Column>>,
#[serde(default)]
pub primary_key: Option<String>,
#[serde(default, with = "bool_from_int")]
pub cached: bool,
#[serde(default)]
pub refresh_time: Option<String>,
}
};
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<String>,
#[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<String>,
#[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<String>,
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<Arc<Column>>,
pub measure: Vec<Arc<Column>>,
pub time_grain: Vec<TimeGrain>,
#[serde(default, with = "bool_from_int")]
pub cached: bool,
pub refresh_time: Option<String>,
}
};
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<TimeUnit>,
}
};
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)
}
+1
View File
@@ -0,0 +1 @@
pub mod mdl;
@@ -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::<Manifest>(&mdl_json).unwrap();
@@ -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<Arc<Model>>,
#[serde(default)]
pub relationships: Vec<Arc<Relationship>>,
#[serde(default)]
pub metrics: Vec<Arc<Metric>>,
#[serde(default)]
pub views: Vec<Arc<View>>,
pub data_source: Option<DataSource>,
#[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<String>,
#[serde(default)]
pub base_object: Option<String>,
#[serde(default, with = "table_reference")]
pub table_reference: Option<String>,
pub columns: Vec<Arc<Column>>,
#[serde(default)]
pub primary_key: Option<String>,
#[serde(default, with = "bool_from_int")]
pub cached: bool,
#[serde(default)]
pub refresh_time: Option<String>,
}
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<S>(
table_ref: &Option<String>,
serializer: S,
) -> Result<S::Ok, S::Error>
pub fn serialize<S>(table_ref: &Option<String>, serializer: S) -> Result<S::Ok, S::Error>
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<String>,
#[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<String>,
#[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<String>,
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<Arc<Column>>,
pub measure: Vec<Arc<Column>>,
pub time_grain: Vec<TimeGrain>,
#[serde(default, with = "bool_from_int")]
pub cached: bool,
pub refresh_time: Option<String>,
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<Arc<Column>> {
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<Item = Arc<Column>> + '_ {
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<Arc<Column>> {
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<TimeUnit>,
}
#[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);
});
}
+25
View File
@@ -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::*;
+61
View File
@@ -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<String> {
Ok(self.catalog.clone())
}
#[getter]
fn schema(&self) -> PyResult<String> {
Ok(self.schema.clone())
}
#[getter]
fn models(&self) -> PyResult<Vec<Model>> {
Ok(self
.models
.iter()
.map(|m| Arc::unwrap_or_clone(Arc::clone(m)))
.collect())
}
#[getter]
fn data_source(&self) -> PyResult<Option<DataSource>> {
Ok(self.data_source)
}
}
#[pymethods]
impl Model {
#[getter]
fn get_name(&self) -> PyResult<String> {
Ok(self.name.clone())
}
}
}
+45 -24
View File
@@ -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]]
+1
View File
@@ -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"
+10 -12
View File
@@ -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<String>,
) -> Result<PyManifest, CoreError> {
pub fn extract_by(&self, used_datasets: Vec<String>) -> Result<Manifest, CoreError> {
extract_manifest(&self.mdl, &used_datasets)
}
}
@@ -69,19 +67,19 @@ fn resolve_used_table_names(mdl: &WrenMDL, sql: &str) -> Result<Vec<String>, Cor
fn extract_manifest(
mdl: &WrenMDL,
used_datasets: &[String],
) -> Result<PyManifest, CoreError> {
) -> Result<Manifest, CoreError> {
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]
+1 -1
View File
@@ -14,7 +14,7 @@ fn wren_core_wrapper(m: &Bound<'_, PyModule>) -> PyResult<()> {
env_logger::init();
m.add_class::<context::PySessionContext>()?;
m.add_class::<PyRemoteFunction>()?;
m.add_class::<manifest::PyManifest>()?;
m.add_class::<manifest::Manifest>()?;
m.add_class::<extractor::PyManifestExtractor>()?;
m.add_function(wrap_pyfunction!(manifest::to_json_base64, m)?)?;
Ok(())
+6 -337
View File
@@ -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<String, CoreError> {
pub fn to_json_base64(mdl: Manifest) -> Result<String, CoreError> {
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<Manifest, CoreError> {
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<DataSource>,
pub models: Vec<Arc<Model>>,
pub relationships: Vec<Arc<Relationship>>,
pub metrics: Vec<Arc<Metric>>,
pub views: Vec<Arc<View>>,
}
#[pymethods]
impl PyManifest {
#[getter]
fn catalog(&self) -> PyResult<String> {
Ok(self.catalog.clone())
}
#[getter]
fn schema(&self) -> PyResult<String> {
Ok(self.schema.clone())
}
#[getter]
fn models(&self) -> PyResult<Vec<PyModel>> {
Ok(self
.models
.iter()
.map(|m| PyModel::from(m.as_ref()))
.collect())
}
#[getter]
fn relationships(&self) -> PyResult<Vec<PyRelationship>> {
Ok(self
.relationships
.iter()
.map(|r| PyRelationship::from(r.as_ref()))
.collect())
}
#[getter]
fn metrics(&self) -> PyResult<Vec<PyMetric>> {
Ok(self
.metrics
.iter()
.map(|m| PyMetric::from(m.as_ref()))
.collect())
}
#[getter]
fn views(&self) -> PyResult<Vec<PyView>> {
Ok(self
.views
.iter()
.map(|v| PyView::from(v.as_ref()))
.collect())
}
#[getter]
fn data_source(&self) -> PyResult<Option<PyDataSource>> {
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<String>,
#[pyo3(get)]
pub base_object: Option<String>,
#[pyo3(get)]
pub table_reference: Option<String>,
pub columns: Vec<Arc<Column>>,
#[pyo3(get)]
pub primary_key: Option<String>,
#[pyo3(get)]
pub cached: bool,
#[pyo3(get)]
pub refresh_time: Option<String>,
}
#[pymethods]
impl PyModel {
#[getter]
fn columns(&self) -> PyResult<Vec<PyColumn>> {
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<String>,
#[pyo3(get)]
pub is_calculated: bool,
#[pyo3(get)]
pub not_null: bool,
#[pyo3(get)]
pub expression: Option<String>,
#[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<String>,
pub join_type: JoinType,
#[pyo3(get)]
pub condition: String,
}
#[pymethods]
impl PyRelationship {
#[getter]
fn join_type(&self) -> PyResult<PyJoinType> {
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<Arc<Column>>,
pub measure: Vec<Arc<Column>>,
pub time_grain: Vec<TimeGrain>,
#[pyo3(get)]
pub cached: bool,
#[pyo3(get)]
pub refresh_time: Option<String>,
}
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<TimeUnit>,
}
#[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<DataSource> 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![
+1
View File
@@ -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" }
+1
View File
@@ -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 }
+4 -92
View File
@@ -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<Arc<Column>> {
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<Item = Arc<Column>> + '_ {
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<Arc<Column>> {
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<Field> {
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<Vec<Field>> {
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::<Result<_>>()
} else {
Ok(vec![self.to_field()?])
}
}
fn collect_columns(expr: Expr) -> Vec<Ident> {
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<Model>),
@@ -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::<Result<_>>()?;
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::<Result<Vec<Vec<Field>>>>()?
.iter()
.flat_map(|c| c.clone())
+12 -6
View File
@@ -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<RwLock<SessionState>>;
@@ -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<DataSource> {
&self.manifest.data_source
pub fn data_source(&self) -> Option<DataSource> {
self.manifest.data_source
}
pub fn get_model(&self, name: &str) -> Option<Arc<Model>> {
+45 -5
View File
@@ -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<Field> {
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<Vec<Field>> {
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::<Result<_>>()
} else {
Ok(vec![to_field(column)?])
}
}
fn collect_columns(expr: datafusion::logical_expr::sqlparser::ast::Expr) -> Vec<Ident> {
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;