1use warble::{Additivity, ContextLoader, DimensionInfo, LineageGraph, MetricInfo, ModelInfo};
11use wren_core_base::mdl::manifest::Manifest;
12
13use crate::lineage;
14use crate::project::{assemble, LoadError, ProjectSources};
15
16pub struct MdlContext {
19 parseable: bool,
20 parse_error: Option<String>,
21 metrics: Vec<MetricInfo>,
22 dimensions: Vec<DimensionInfo>,
23 time_dimensions: Vec<DimensionInfo>,
24 models: Vec<ModelInfo>,
25 lineage: LineageGraph,
26 lineage_diagnostics: Vec<String>,
27}
28
29impl MdlContext {
30 pub fn from_sources(sources: &ProjectSources) -> Self {
36 match assemble(sources) {
37 Ok(loaded) => Self::from_manifest_and_consumers(&loaded.manifest, sources),
38 Err(_) => Self::unparseable(),
39 }
40 }
41
42 pub fn try_from_sources(sources: &ProjectSources) -> Result<Self, LoadError> {
44 assemble(sources).map(|loaded| Self::from_manifest_and_consumers(&loaded.manifest, sources))
45 }
46
47 fn from_manifest_and_consumers(manifest: &Manifest, sources: &ProjectSources) -> Self {
51 let mut ctx = Self::from_manifest(manifest);
52 lineage::extend_with_consumers(
53 &mut ctx.lineage,
54 manifest,
55 sources,
56 &mut ctx.lineage_diagnostics,
57 );
58 ctx
59 }
60
61 pub fn unparseable() -> Self {
65 Self::unparseable_with_error(None)
66 }
67
68 pub fn unparseable_with_error(parse_error: Option<String>) -> Self {
72 MdlContext {
73 parseable: false,
74 parse_error,
75 metrics: Vec::new(),
76 dimensions: Vec::new(),
77 time_dimensions: Vec::new(),
78 models: Vec::new(),
79 lineage: LineageGraph::default(),
80 lineage_diagnostics: Vec::new(),
81 }
82 }
83
84 pub fn from_manifest(manifest: &Manifest) -> Self {
86 let mut metrics = Vec::new();
87 let mut dimensions = Vec::new();
88 let mut time_dimensions = Vec::new();
89 let mut models = Vec::new();
90
91 for cube in &manifest.cubes {
94 for measure in &cube.measures {
95 metrics.push(MetricInfo {
96 name: measure.name.clone(),
97 owner: cube.name.clone(),
98 declared: true,
99 additivity: Some(infer_additivity(&measure.expression)),
100 });
101 }
102 for dim in &cube.dimensions {
103 dimensions.push(DimensionInfo {
104 name: dim.name.clone(),
105 owner: cube.name.clone(),
106 is_temporal: false,
107 });
108 }
109 for tdim in &cube.time_dimensions {
110 let d = DimensionInfo {
111 name: tdim.name.clone(),
112 owner: cube.name.clone(),
113 is_temporal: true,
114 };
115 time_dimensions.push(d.clone());
116 dimensions.push(d);
117 }
118 }
119
120 for model in &manifest.models {
123 let mut has_timestamp = false;
124 let mut column_names = Vec::new();
125 for col in model.columns.iter().filter(|c| !c.is_hidden) {
126 column_names.push(col.name.clone());
127 if col.relationship.is_some() {
129 continue;
130 }
131 if is_temporal_type(&col.r#type) {
132 has_timestamp = true;
133 let d = DimensionInfo {
134 name: col.name.clone(),
135 owner: model.name.clone(),
136 is_temporal: true,
137 };
138 time_dimensions.push(d.clone());
139 dimensions.push(d);
140 } else if is_numeric_type(&col.r#type) {
141 metrics.push(MetricInfo {
142 name: col.name.clone(),
143 owner: model.name.clone(),
144 declared: false,
145 additivity: None,
146 });
147 } else {
148 dimensions.push(DimensionInfo {
150 name: col.name.clone(),
151 owner: model.name.clone(),
152 is_temporal: false,
153 });
154 }
155 }
156 models.push(ModelInfo {
157 name: model.name.clone(),
158 has_timestamp,
159 columns: column_names,
160 });
161 }
162
163 let (lineage, lineage_diagnostics) = lineage::build(manifest);
164 MdlContext {
165 parseable: true,
166 parse_error: None,
167 metrics,
168 dimensions,
169 time_dimensions,
170 models,
171 lineage,
172 lineage_diagnostics,
173 }
174 }
175}
176
177impl ContextLoader for MdlContext {
178 fn is_parseable(&self) -> bool {
179 self.parseable
180 }
181 fn parse_error(&self) -> Option<&str> {
182 self.parse_error.as_deref()
183 }
184 fn metrics(&self) -> &[MetricInfo] {
185 &self.metrics
186 }
187 fn dimensions(&self) -> &[DimensionInfo] {
188 &self.dimensions
189 }
190 fn time_dimensions(&self) -> &[DimensionInfo] {
191 &self.time_dimensions
192 }
193 fn models(&self) -> &[ModelInfo] {
194 &self.models
195 }
196 fn lineage(&self) -> &LineageGraph {
197 &self.lineage
198 }
199 fn lineage_diagnostics(&self) -> &[String] {
200 &self.lineage_diagnostics
201 }
202}
203
204pub fn infer_additivity(expression: &str) -> Additivity {
213 let upper = expression.to_uppercase();
214 if upper.contains("DISTINCT") || upper.contains('/') {
215 return Additivity::NonAdditive;
216 }
217 match leading_function(&upper).as_deref() {
218 Some("SUM") | Some("COUNT") => Additivity::Additive,
219 _ => Additivity::NonAdditive,
220 }
221}
222
223fn leading_function(upper_expr: &str) -> Option<String> {
226 let open = upper_expr.find('(')?;
227 let name = upper_expr[..open].trim();
228 if name.is_empty() || !name.chars().all(|c| c.is_ascii_alphanumeric() || c == '_') {
229 return None;
230 }
231 Some(name.to_string())
232}
233
234fn normalize_type(t: &str) -> String {
235 let base = t.split('(').next().unwrap_or(t).trim();
237 base.to_uppercase()
238}
239
240fn is_temporal_type(t: &str) -> bool {
241 let t = normalize_type(t);
242 t == "DATE" || t == "DATETIME" || t.starts_with("TIMESTAMP")
243}
244
245fn is_numeric_type(t: &str) -> bool {
246 matches!(
247 normalize_type(t).as_str(),
248 "INT"
249 | "INTEGER"
250 | "BIGINT"
251 | "SMALLINT"
252 | "TINYINT"
253 | "DECIMAL"
254 | "NUMERIC"
255 | "DOUBLE"
256 | "FLOAT"
257 | "REAL"
258 )
259}
260
261#[cfg(test)]
262mod tests {
263 use super::*;
264
265 #[test]
266 fn additivity_heuristic() {
267 assert_eq!(infer_additivity("SUM(amount)"), Additivity::Additive);
268 assert_eq!(infer_additivity("sum(amount)"), Additivity::Additive);
269 assert_eq!(infer_additivity("COUNT(*)"), Additivity::Additive);
270 assert_eq!(
271 infer_additivity("COUNT(DISTINCT customer_id)"),
272 Additivity::NonAdditive
273 );
274 assert_eq!(infer_additivity("AVG(amount)"), Additivity::NonAdditive);
275 assert_eq!(infer_additivity("MIN(amount)"), Additivity::NonAdditive);
276 assert_eq!(infer_additivity("MAX(amount)"), Additivity::NonAdditive);
277 assert_eq!(infer_additivity("SUM(a) / SUM(b)"), Additivity::NonAdditive);
279 assert_eq!(infer_additivity("weird(x)"), Additivity::NonAdditive);
281 assert_eq!(infer_additivity("bare_column"), Additivity::NonAdditive);
282 }
283
284 #[test]
285 fn cubeless_manifest_cannot_answer_metric_additive() {
286 use warble::ContextLoader;
287 use wren_core_base::mdl::manifest::Manifest;
288
289 let json = r#"{
292 "catalog":"wren","schema":"public",
293 "models":[{"name":"orders","tableReference":{"schema":"main","table":"orders"},
294 "columns":[{"name":"amount","type":"DOUBLE"},{"name":"status","type":"TEXT"}]}],
295 "relationships":[],"cubes":[],"views":[]
296 }"#;
297 let manifest: Manifest = serde_json::from_str(json).unwrap();
298 let ctx = MdlContext::from_manifest(&manifest);
299 assert!(!ctx.metrics().is_empty(), "amount is an implicit metric");
300 assert!(
301 !ctx.can_answer("metric_additive"),
302 "no declared measure ⇒ additivity unanswerable"
303 );
304 assert!(ctx
305 .metrics()
306 .iter()
307 .all(|m| !m.declared && m.additivity.is_none()));
308 }
309
310 #[test]
311 fn type_classification() {
312 assert!(is_temporal_type("DATE"));
313 assert!(is_temporal_type("timestamp"));
314 assert!(is_temporal_type("TIMESTAMP WITH TIME ZONE"));
315 assert!(!is_temporal_type("TEXT"));
316 assert!(is_numeric_type("INT"));
317 assert!(is_numeric_type("BIGINT"));
318 assert!(is_numeric_type("DECIMAL(10,2)"));
319 assert!(!is_numeric_type("TEXT"));
320 assert!(!is_numeric_type("DATE"));
321 }
322 #[test]
323 fn mdl_context_cannot_answer_raw_shape_predicates() {
324 use crate::MdlContext;
327 use wren_core_base::mdl::manifest::Manifest;
328
329 let json = r#"{
330 "catalog":"wren","schema":"public",
331 "models":[],"relationships":[],"cubes":[],"views":[]
332 }"#;
333 let manifest: Manifest = serde_json::from_str(json).unwrap();
334 let ctx = MdlContext::from_manifest(&manifest);
335 assert_eq!(ctx.source_introspectable(), None);
336 assert_eq!(ctx.raw_docs_readable(), None);
337 assert!(!ctx.can_answer("source_introspectable"));
338 assert!(!ctx.can_answer("raw_docs_readable"));
339 }
340}