Skip to main content

mdbook_plotly/preprocessor/handlers/code_handler/
until.rs

1// Tmp fix
2#![allow(unexpected_cfgs)]
3
4use anyhow::{Context, Result, anyhow};
5use fasteval::{Compiler, EvalNamespace, Evaler, Parser, Slab};
6use serde::{Deserialize, Deserializer, Serialize, Serializer, de::DeserializeOwned};
7use serde_json::{Map as JsonMap, Value, value::Index};
8use std::{collections::BTreeMap, fmt::Debug, fmt::Display};
9
10#[cfg(feature = "map-parser-extensions")]
11use chrono::{DateTime, NaiveDate, NaiveDateTime, TimeDelta, Utc};
12#[cfg(feature = "map-parser-extensions")]
13use rand::{Rng, RngExt, SeedableRng, rngs::StdRng};
14
15pub type Map = JsonMap<String, Value>;
16
17type Vars = BTreeMap<String, f64>;
18
19#[derive(Clone, Debug)]
20pub enum DataPack<T> {
21    Data(T),
22    Index(String),
23}
24
25#[inline]
26pub fn must_translate<T, N>(obj: &mut Value, map: &Map, name: N) -> Result<T>
27where
28    T: DeserializeOwned + Serialize + Debug + Clone,
29    N: Index + Display,
30{
31    take_optional(obj, map, &name)?.ok_or_else(|| anyhow!("missing `{}` field", name))
32}
33
34#[inline]
35fn take_optional<T, N>(obj: &mut Value, map: &Map, name: &N) -> Result<Option<T>>
36where
37    T: DeserializeOwned + Serialize + Debug + Clone,
38    N: Index + Display,
39{
40    let Some(value) = obj.get_mut(name) else {
41        return Ok(None);
42    };
43
44    serde_json::from_value::<DataPack<T>>(value.take())
45        .with_context(|| format!("failed to deserialize field '{}'", name))?
46        .unwrap(map)
47        .with_context(|| format!("failed to unwrap DataPack for field '{}'", name))
48        .map(Some)
49}
50
51#[inline]
52fn try_deser<T: DeserializeOwned>(value: Value, context: &'static str) -> Result<T> {
53    serde_json::from_value(value).context(context)
54}
55
56#[inline]
57fn json_number(value: f64) -> Result<Value> {
58    serde_json::Number::from_f64(value)
59        .map(Value::Number)
60        .ok_or_else(|| anyhow!("failed to create JSON number from {}", value))
61}
62
63#[inline]
64fn usize_count(count: u64, field: &str) -> Result<usize> {
65    usize::try_from(count).with_context(|| format!("{} is too large for this platform", field))
66}
67
68fn value_to_f64(value: &Value) -> Option<f64> {
69    match value {
70        Value::Number(n) => n.as_f64(),
71        Value::Bool(v) => Some(if *v { 1.0 } else { 0.0 }),
72        Value::String(s) => s.parse::<f64>().ok(),
73        _ => None,
74    }
75}
76
77fn lookup_path<'a>(map: &'a Map, name: &str) -> Option<&'a Value> {
78    let path = name.strip_prefix("map.").unwrap_or(name);
79    let mut parts = path.split('.');
80    let first = parts.next()?;
81    let mut value = map.get(first)?;
82
83    for part in parts {
84        match value {
85            Value::Object(obj) => value = obj.get(part)?,
86            Value::Array(arr) => {
87                let idx = part.parse::<usize>().ok()?;
88                value = arr.get(idx)?;
89            }
90            _ => return None,
91        }
92    }
93
94    Some(value)
95}
96
97fn map_value<'a>(map: &'a Map, index: &str) -> Result<&'a Value> {
98    lookup_path(map, index).ok_or_else(|| anyhow!("missing map value `{}`", index))
99}
100
101struct MapNamespace<'a> {
102    map: &'a Map,
103    vars: &'a Vars,
104}
105
106impl<'a> MapNamespace<'a> {
107    fn new(map: &'a Map, vars: &'a Vars) -> Self {
108        Self { map, vars }
109    }
110}
111
112impl EvalNamespace for MapNamespace<'_> {
113    fn lookup(&mut self, name: &str, _args: Vec<f64>, _keybuf: &mut String) -> Option<f64> {
114        self.vars
115            .get(name)
116            .copied()
117            .or_else(|| lookup_path(self.map, name).and_then(value_to_f64))
118    }
119}
120
121struct EvalContext {
122    parser: Parser,
123    slab: Slab,
124}
125
126impl Default for EvalContext {
127    fn default() -> Self {
128        Self {
129            parser: Parser::new(),
130            slab: Slab::new(),
131        }
132    }
133}
134
135impl EvalContext {
136    fn eval(&mut self, expr: &str, map: &Map, vars: &Vars) -> Result<f64> {
137        let mut namespace = MapNamespace::new(map, vars);
138        let expr_ref = self
139            .parser
140            .parse(expr, &mut self.slab.ps)
141            .with_context(|| format!("failed to parse expression `{}`", expr))?
142            .from(&self.slab.ps);
143
144        expr_ref
145            .eval(&self.slab, &mut namespace)
146            .with_context(|| format!("failed to evaluate expression `{}`", expr))
147    }
148}
149
150impl<T> DataPack<T>
151where
152    T: DeserializeOwned + Serialize + Debug + Clone,
153{
154    pub fn unwrap(self, map: &Map) -> Result<T> {
155        match self {
156            Self::Data(data) => Ok(data),
157            Self::Index(index) => {
158                let value = map_value(map, &index)?.clone();
159                Self::parse_value(map, value)
160                    .with_context(|| format!("failed to resolve map value `{}`", index))
161            }
162        }
163    }
164
165    fn parse_value(map: &Map, value: Value) -> Result<T> {
166        if value.is_object() && value.get("type").is_some() {
167            Self::parse_map(map, value)
168        } else {
169            serde_json::from_value(value).context("failed to deserialize value")
170        }
171    }
172
173    fn parse_map(map: &Map, mut value: Value) -> Result<T> {
174        let value_type = value
175            .get("type")
176            .and_then(Value::as_str)
177            .ok_or_else(|| anyhow!("`type` must be a string"))?
178            .to_owned();
179
180        let mut eval = EvalContext::default();
181        let vars = Vars::new();
182
183        match value_type.as_str() {
184            "raw" => must_translate(&mut value, map, "data"),
185            "g-number" => parse_g_number(map, &mut value, &mut eval, &vars),
186            "g-number-list" => parse_g_number_list(map, &mut value, &mut eval),
187            "g-range" => parse_g_range(map, &mut value),
188            "g-repeat" => parse_g_repeat(map, &mut value),
189            "g-linear" => parse_g_linear(map, &mut value),
190            "if" => parse_if(map, &mut value, &mut eval, &vars),
191            #[cfg(feature = "map-parser-extensions")]
192            "time" => parse_time(map, &mut value),
193            #[cfg(feature = "map-parser-extensions")]
194            "g-random" => parse_g_random(map, &mut value),
195            #[cfg(feature = "map-parser-extensions")]
196            "g-choose" => parse_g_choose(map, &mut value),
197            "g-env" => parse_g_env(map, &mut value),
198            "g-join" => parse_g_join(map, &mut value),
199            _ => Err(anyhow!("unknown type `{}`", value_type)),
200        }
201    }
202}
203
204fn parse_g_number<T>(map: &Map, value: &mut Value, eval: &mut EvalContext, vars: &Vars) -> Result<T>
205where
206    T: DeserializeOwned,
207{
208    let expr: String = must_translate(value, map, "expr")?;
209    try_deser(
210        json_number(eval.eval(&expr, map, vars)?)?,
211        "failed to deserialize generated number",
212    )
213}
214
215fn parse_g_number_list<T>(map: &Map, value: &mut Value, eval: &mut EvalContext) -> Result<T>
216where
217    T: DeserializeOwned,
218{
219    let index_begin: u64 = must_translate(value, map, "begin")?;
220    let index_end: u64 = must_translate(value, map, "end")?;
221    let expr: String = must_translate(value, map, "expr")?;
222
223    let len = index_end.saturating_sub(index_begin);
224    let mut result = Vec::with_capacity(usize_count(len, "g-number-list length")?);
225    let mut vars = Vars::new();
226
227    let compiled = eval
228        .parser
229        .parse(&expr, &mut eval.slab.ps)
230        .with_context(|| format!("failed to parse expression `{}`", expr))?
231        .from(&eval.slab.ps)
232        .compile(&eval.slab.ps, &mut eval.slab.cs);
233
234    for i in index_begin..index_end {
235        vars.insert("i".to_owned(), i as f64);
236        let mut namespace = MapNamespace::new(map, &vars);
237        result.push(json_number(fasteval::eval_compiled!(
238            compiled,
239            &eval.slab,
240            &mut namespace
241        ))?);
242    }
243
244    try_deser(
245        Value::Array(result),
246        "failed to deserialize generated number list",
247    )
248}
249
250fn parse_g_range<T>(map: &Map, value: &mut Value) -> Result<T>
251where
252    T: DeserializeOwned,
253{
254    let begin: f64 = must_translate(value, map, "begin")?;
255    let end: f64 = must_translate(value, map, "end")?;
256    let step: f64 = take_optional(value, map, &"step")?.unwrap_or(1.0);
257
258    if step <= 0.0 {
259        return Err(anyhow!("step must be positive"));
260    }
261
262    let capacity = if end > begin {
263        ((end - begin) / step).ceil() as usize
264    } else {
265        0
266    };
267    let mut result = Vec::with_capacity(capacity);
268    let mut current = begin;
269
270    while current < end {
271        result.push(json_number(current)?);
272        current += step;
273    }
274
275    try_deser(
276        Value::Array(result),
277        "failed to deserialize generated range",
278    )
279}
280
281fn parse_g_repeat<T>(map: &Map, value: &mut Value) -> Result<T>
282where
283    T: DeserializeOwned + Serialize + Debug + Clone,
284{
285    let val: Value = must_translate(value, map, "value")?;
286    let count: u64 = must_translate(value, map, "count")?;
287    let count = usize_count(count, "count")?;
288    let result = vec![val; count];
289    try_deser(
290        Value::Array(result),
291        "failed to deserialize repeated values",
292    )
293}
294
295fn parse_g_linear<T>(map: &Map, value: &mut Value) -> Result<T>
296where
297    T: DeserializeOwned,
298{
299    let begin: f64 = must_translate(value, map, "begin")?;
300    let end: f64 = must_translate(value, map, "end")?;
301    let count: u64 = must_translate(value, map, "count")?;
302
303    if count == 0 {
304        return Err(anyhow!("count must be positive"));
305    }
306
307    let count_usize = usize_count(count, "count")?;
308    let mut result = Vec::with_capacity(count_usize);
309
310    if count == 1 {
311        result.push(json_number(begin)?);
312    } else {
313        let step = (end - begin) / ((count - 1) as f64);
314        for i in 0..count {
315            result.push(json_number(begin + (i as f64) * step)?);
316        }
317    }
318
319    try_deser(
320        Value::Array(result),
321        "failed to deserialize linear spaced values",
322    )
323}
324
325fn parse_if<T>(map: &Map, value: &mut Value, eval: &mut EvalContext, vars: &Vars) -> Result<T>
326where
327    T: DeserializeOwned + Serialize + Debug + Clone,
328{
329    let condition: String = must_translate(value, map, "condition")?;
330    let true_val: Value = must_translate(value, map, "true")?;
331    let false_val: Value = must_translate(value, map, "false")?;
332    let selected = if eval.eval(&condition, map, vars)? != 0.0 {
333        true_val
334    } else {
335        false_val
336    };
337
338    DataPack::<T>::parse_value(map, selected)
339}
340
341#[cfg(feature = "map-parser-extensions")]
342fn parse_time<T>(map: &Map, value: &mut Value) -> Result<T>
343where
344    T: DeserializeOwned,
345{
346    let start: String = must_translate(value, map, "start")?;
347    let end: String = must_translate(value, map, "end")?;
348    let interval: String = must_translate(value, map, "interval")?;
349    let format: Option<String> = take_optional(value, map, &"format")?;
350
351    let start_dt = parse_time_str(&start)?;
352    let end_dt = parse_time_str(&end)?;
353    let step = parse_duration_str(&interval)?;
354
355    if step <= TimeDelta::zero() {
356        return Err(anyhow!("interval must be positive"));
357    }
358
359    let mut result = Vec::new();
360    let mut current = start_dt;
361
362    while current <= end_dt {
363        let ts = format.as_ref().map_or_else(
364            || current.to_rfc3339(),
365            |fmt| current.format(fmt).to_string(),
366        );
367        result.push(Value::String(ts));
368
369        let Some(next) = current.checked_add_signed(step) else {
370            break;
371        };
372        current = next;
373    }
374
375    try_deser(Value::Array(result), "failed to deserialize time values")
376}
377
378#[cfg(feature = "map-parser-extensions")]
379fn parse_g_random<T>(map: &Map, value: &mut Value) -> Result<T>
380where
381    T: DeserializeOwned,
382{
383    let min: f64 = must_translate(value, map, "min")?;
384    let max: f64 = must_translate(value, map, "max")?;
385
386    if min >= max {
387        return Err(anyhow!("min ({}) must be less than max ({})", min, max));
388    }
389
390    let integer = value
391        .get("integer")
392        .and_then(Value::as_bool)
393        .unwrap_or(false);
394    let seed: Option<u64> = take_optional(value, map, &"seed")?;
395    let count: Option<u64> = take_optional(value, map, &"count")?;
396
397    let gen_value = |rng: &mut dyn Rng| -> Result<Value> {
398        if integer {
399            if min.fract() != 0.0 || max.fract() != 0.0 {
400                return Err(anyhow!("integer random bounds must be whole numbers"));
401            }
402            Ok(Value::from(rng.random_range(min as i64..max as i64)))
403        } else {
404            json_number(rng.random_range(min..max))
405        }
406    };
407
408    match count {
409        Some(0) => Err(anyhow!("count must be positive")),
410        Some(count) => {
411            let count = usize_count(count, "count")?;
412            let values = with_rng(seed, |rng| {
413                (0..count)
414                    .map(|_| gen_value(rng))
415                    .collect::<Result<Vec<_>>>()
416            })?;
417            try_deser(Value::Array(values), "failed to deserialize g-random array")
418        }
419        None => {
420            let value = with_rng(seed, |rng| gen_value(rng))?;
421            try_deser(value, "failed to deserialize g-random single value")
422        }
423    }
424}
425
426#[cfg(feature = "map-parser-extensions")]
427fn parse_g_choose<T>(map: &Map, value: &mut Value) -> Result<T>
428where
429    T: DeserializeOwned + Serialize + Debug + Clone,
430{
431    let options: Vec<Value> = must_translate(value, map, "options")?;
432    if options.is_empty() {
433        return Err(anyhow!("options must not be empty for g-choose"));
434    }
435
436    let seed: Option<u64> = take_optional(value, map, &"seed")?;
437    let count: Option<u64> = take_optional(value, map, &"count")?;
438
439    match count {
440        Some(0) => Err(anyhow!("count must be positive")),
441        Some(count) => {
442            let count = usize_count(count, "count")?;
443            let selected = with_rng(seed, |rng| {
444                (0..count)
445                    .map(|_| options[rng.random_range(0..options.len())].clone())
446                    .collect::<Vec<_>>()
447            });
448            try_deser(
449                Value::Array(selected),
450                "failed to deserialize g-choose array",
451            )
452        }
453        None => {
454            let picked = with_rng(seed, |rng| {
455                options[rng.random_range(0..options.len())].clone()
456            });
457            try_deser(picked, "failed to deserialize g-choose single value")
458        }
459    }
460}
461
462fn parse_g_env<T>(map: &Map, value: &mut Value) -> Result<T>
463where
464    T: DeserializeOwned,
465{
466    let name: String = must_translate(value, map, "name")?;
467    let default: Option<String> = take_optional(value, map, &"default")?;
468    let env_val = std::env::var(&name).ok().or(default).ok_or_else(|| {
469        anyhow!(
470            "environment variable '{}' is not set and no default provided",
471            name
472        )
473    })?;
474
475    try_deser(Value::String(env_val), "failed to deserialize env value")
476}
477
478fn parse_g_join<T>(map: &Map, value: &mut Value) -> Result<T>
479where
480    T: DeserializeOwned,
481{
482    let values: Vec<String> = must_translate(value, map, "values")?;
483    let separator: String = take_optional(value, map, &"separator")?.unwrap_or_default();
484    try_deser(
485        Value::String(values.join(&separator)),
486        "failed to deserialize joined string",
487    )
488}
489
490#[cfg(feature = "map-parser-extensions")]
491fn parse_duration_str(s: &str) -> Result<TimeDelta> {
492    let s = s.trim();
493    if s.is_empty() {
494        return Err(anyhow!("duration string is empty"));
495    }
496
497    let mut total = TimeDelta::zero();
498    let mut num_str = String::new();
499
500    for ch in s.chars() {
501        if ch.is_ascii_digit() || ch == '.' {
502            num_str.push(ch);
503            continue;
504        }
505
506        if ch.is_whitespace() {
507            continue;
508        }
509
510        if !ch.is_alphabetic() {
511            return Err(anyhow!("unexpected character '{}' in duration string", ch));
512        }
513
514        if num_str.is_empty() {
515            return Err(anyhow!("missing number before duration unit '{}'", ch));
516        }
517
518        let num: f64 = num_str
519            .parse()
520            .with_context(|| format!("invalid number in duration: '{}'", num_str))?;
521        if num <= 0.0 {
522            return Err(anyhow!("duration components must be positive"));
523        }
524        num_str.clear();
525
526        let seconds = match ch.to_ascii_lowercase() {
527            's' => num,
528            'm' => num * 60.0,
529            'h' => num * 3_600.0,
530            'd' => num * 86_400.0,
531            'w' => num * 604_800.0,
532            other => return Err(anyhow!("unknown duration unit: '{}'", other)),
533        };
534
535        let delta = TimeDelta::try_seconds(seconds as i64)
536            .ok_or_else(|| anyhow!("duration overflow: {}{}", num, ch))?;
537        total = total
538            .checked_add(&delta)
539            .ok_or_else(|| anyhow!("duration overflow"))?;
540    }
541
542    if !num_str.is_empty() {
543        return Err(anyhow!("trailing number without unit: '{}'", num_str));
544    }
545    if total.is_zero() {
546        return Err(anyhow!("duration must be positive, got: '{}'", s));
547    }
548
549    Ok(total)
550}
551
552#[cfg(feature = "map-parser-extensions")]
553fn parse_time_str(s: &str) -> Result<DateTime<Utc>> {
554    let s = s.trim();
555
556    if let Some(rest) = s.strip_prefix("now") {
557        let base = Utc::now();
558        if rest.is_empty() {
559            return Ok(base);
560        }
561
562        let sign_char = rest
563            .chars()
564            .next()
565            .ok_or_else(|| anyhow!("expected '+' or '-' after 'now'"))?;
566        let duration_str = &rest[sign_char.len_utf8()..];
567        let delta = parse_duration_str(duration_str)?;
568
569        return match sign_char {
570            '+' => base
571                .checked_add_signed(delta)
572                .ok_or_else(|| anyhow!("time overflow for '{}'", s)),
573            '-' => base
574                .checked_sub_signed(delta)
575                .ok_or_else(|| anyhow!("time overflow for '{}'", s)),
576            _ => Err(anyhow!("expected '+' or '-' after 'now', got '{}'", rest)),
577        };
578    }
579
580    if let Ok(dt) = DateTime::parse_from_rfc3339(s) {
581        return Ok(dt.with_timezone(&Utc));
582    }
583
584    for fmt in [
585        "%Y-%m-%dT%H:%M:%S%.f%:z",
586        "%Y-%m-%dT%H:%M:%S%.f",
587        "%Y-%m-%dT%H:%M:%S%:z",
588        "%Y-%m-%dT%H:%M:%S",
589        "%Y-%m-%d %H:%M:%S",
590    ] {
591        if let Ok(dt) = DateTime::parse_from_str(s, fmt) {
592            return Ok(dt.with_timezone(&Utc));
593        }
594        if let Ok(naive) = NaiveDateTime::parse_from_str(s, fmt) {
595            return Ok(naive.and_utc());
596        }
597    }
598
599    if let Ok(date) = NaiveDate::parse_from_str(s, "%Y-%m-%d") {
600        return date
601            .and_hms_opt(0, 0, 0)
602            .map(|s| s.and_utc())
603            .ok_or_else(|| anyhow!("invalid date: '{}'", s));
604    }
605
606    Err(anyhow!(
607        "unable to parse time string: '{}'. Supported formats: RFC 3339, \
608         'YYYY-MM-DDTHH:MM:SS', 'YYYY-MM-DD HH:MM:SS', 'YYYY-MM-DD', \
609         'now', 'now+duration', 'now-duration'",
610        s
611    ))
612}
613
614#[cfg(feature = "map-parser-extensions")]
615fn with_rng<F, R>(seed: Option<u64>, f: F) -> R
616where
617    F: FnOnce(&mut dyn Rng) -> R,
618{
619    if let Some(seed) = seed {
620        let mut rng = StdRng::seed_from_u64(seed);
621        f(&mut rng)
622    } else {
623        let mut rng = rand::rng();
624        f(&mut rng)
625    }
626}
627
628impl<'de, T> Deserialize<'de> for DataPack<T>
629where
630    T: DeserializeOwned + Serialize + Debug + Clone,
631{
632    fn deserialize<D>(deserializer: D) -> std::result::Result<Self, D::Error>
633    where
634        D: Deserializer<'de>,
635    {
636        let value = Value::deserialize(deserializer)?;
637
638        if let Some(index) = value.as_str().and_then(|s| s.strip_prefix("map.")) {
639            return Ok(Self::Index(index.to_owned()));
640        }
641
642        serde_json::from_value::<T>(value)
643            .map(Self::Data)
644            .map_err(serde::de::Error::custom)
645    }
646}
647
648impl<T> Serialize for DataPack<T>
649where
650    T: DeserializeOwned + Serialize + Debug + Clone,
651{
652    fn serialize<S>(&self, serializer: S) -> std::result::Result<S::Ok, S::Error>
653    where
654        S: Serializer,
655    {
656        match self {
657            Self::Data(data) => data.serialize(serializer),
658            Self::Index(index) => serializer.serialize_str(&format!("map.{index}")),
659        }
660    }
661}
662
663use plotly::color;
664
665// This is to make Json look clearer when it is written.
666#[allow(clippy::enum_variant_names)]
667#[derive(Clone, Debug, Serialize)]
668#[serde(rename_all = "snake_case")]
669pub enum Color {
670    NamedColor(color::NamedColor),
671    RgbColor(color::Rgb),
672    RgbaColor(color::Rgba),
673}
674
675impl color::Color for Color {}
676
677impl<'de> Deserialize<'de> for Color {
678    fn deserialize<D>(deserializer: D) -> std::result::Result<Self, D::Error>
679    where
680        D: Deserializer<'de>,
681    {
682        let value = Value::deserialize(deserializer)?;
683
684        if let Some(s) = value.as_str() 
685            && let Ok(named) = serde_json::from_str::<color::NamedColor>(&format!("\"{s}\"")) {
686                return Ok(Self::NamedColor(named));
687        }
688
689        if let Ok(rgb) = serde_json::from_value::<color::Rgb>(value.clone()) {
690            return Ok(Self::RgbColor(rgb));
691        }
692
693        if let Ok(rgba) = serde_json::from_value::<color::Rgba>(value) {
694            return Ok(Self::RgbaColor(rgba));
695        }
696
697        Err(serde::de::Error::custom("invalid color format"))
698    }
699}