mirror of
https://github.com/Canner/WrenAI.git
synced 2026-09-24 23:29:49 +08:00
refactor(core): introduce wren-core-base module to collect the common structs and utilities (#1003)
This commit is contained in:
@@ -0,0 +1,4 @@
|
||||
Cargo.lock
|
||||
target/
|
||||
manifest-macro/Cargo.lock
|
||||
manifest-macro/target/
|
||||
@@ -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"
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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();
|
||||
|
||||
+118
-153
@@ -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);
|
||||
});
|
||||
}
|
||||
@@ -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::*;
|
||||
@@ -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())
|
||||
}
|
||||
}
|
||||
}
|
||||
Generated
+45
-24
@@ -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]]
|
||||
|
||||
@@ -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"
|
||||
|
||||
@@ -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]
|
||||
|
||||
@@ -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(())
|
||||
|
||||
@@ -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![
|
||||
|
||||
@@ -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" }
|
||||
|
||||
@@ -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 }
|
||||
|
||||
@@ -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())
|
||||
|
||||
@@ -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>> {
|
||||
|
||||
@@ -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;
|
||||
|
||||
Reference in New Issue
Block a user