use merman::{MermaidConfig, render::HeadlessRenderer};
use serde_json::{Map, Value, json};
use crate::{PreviewError, Result};
#[derive(Debug, Default)]
pub struct MermaidRenderOptions {
pub theme: Option<String>,
pub font_family: Option<String>,
pub font_size: Option<f64>,
}
pub fn render_mermaid_svg(source: &str, options: MermaidRenderOptions) -> Result<String> {
let config = render_config(options);
HeadlessRenderer::new()
.with_site_config(config)
.render_svg_sync(source)
.map_err(|error| PreviewError::Mermaid(error.to_string()))?
.ok_or_else(|| PreviewError::Mermaid("unsupported Mermaid diagram".to_string()))
}
fn render_config(options: MermaidRenderOptions) -> MermaidConfig {
let mut config = match options.theme.as_deref() {
Some("default") | _ => json!({ "theme": "default" }),
Some("dark") => json!({ "theme": "dark", "darkMode": true }),
};
let config = config.as_object_mut().expect("preview config is an object");
if let Some(font_family) = options.font_family.as_ref() {
config.insert("fontFamily".to_string(), Value::String(font_family.clone()));
}
let variables = config
.entry("themeVariables")
.or_insert_with(|| Value::Object(Map::new()))
.as_object_mut()
.expect("theme variables are an object");
if let Some(font_family) = options.font_family {
variables.insert("fontFamily".to_string(), Value::String(font_family));
}
if let Some(font_size) = options.font_size {
variables.insert("fontSize".to_string(), Value::String(format!("{font_size}px")));
}
MermaidConfig::from_value(Value::Object(config.clone()))
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn renders_supported_diagram_families() {
let diagrams = [
("architecture", "architecture-beta\nservice api[API]"),
("block", "block\n columns 2\n A --> B"),
("c4", "C4Context\nPerson(user, \"User\")"),
("class", "classDiagram\nAnimal <|-- Dog"),
("er", "erDiagram\nUSER ||--o{ ORDER : places"),
("flowchart", "flowchart TD\nA --> B"),
(
"gantt",
"gantt\ntitle Plan\ndateFormat YYYY-MM-DD\nsection Work\nTask :2026-01-01, 1d",
),
("gitgraph", "gitGraph\ncommit id: \"start\""),
("info", "info"),
("journey", "journey\ntitle Day\nsection Work\nTask: 5: Me"),
("kanban", "kanban\n Todo\n item1"),
("mindmap", "mindmap\n root((AFFiNE))\n Preview"),
("packet", "packet-beta\n0-7: \"Header\""),
("pie", "pie\n\"A\" : 1\n\"B\" : 2"),
(
"quadrantchart",
"quadrantChart\n x-axis Low --> High\n y-axis Low --> High\n A: [0.3, 0.6]",
),
(
"radar",
"radar-beta\n axis quality[\"Quality\"], speed[\"Speed\"]\n curve product[\"Product\"]{3, 4}\n max 5",
),
(
"requirement",
"requirementDiagram\nrequirement test_req {\nid: 1\ntext: Test\nrisk: low\nverifymethod: test\n}",
),
("sankey", "sankey-beta\nA,B,10\nA,C,5"),
("sequence", "sequenceDiagram\nAlice->>Bob: Hello"),
("state", "stateDiagram-v2\n[*] --> Ready"),
("timeline", "timeline\ntitle History\n2026 : Started"),
("treemap", "treemap\n \"Root\"\n \"Item\": 10"),
("venn", "venn-beta\n set A[\"Alpha\"]:20"),
("xychart", "xychart-beta\nx-axis [1, 2]\ny-axis 0 --> 2\nline [1, 2]"),
("zenuml", "zenuml\n Alice->Bob: Hello"),
];
let supported = merman::supported_diagrams();
assert_eq!(diagrams.len(), supported.len());
for (family, diagram) in diagrams {
assert!(supported.contains(&family), "missing corpus for {family}");
let svg = render_mermaid_svg(diagram, MermaidRenderOptions::default())
.unwrap_or_else(|error| panic!("failed to render {diagram:?}: {error}"));
assert!(svg.starts_with("<svg"), "{diagram}");
assert!(svg.trim_end().ends_with("</svg>"), "{diagram}");
}
}
#[test]
fn applies_theme_and_font_options() {
let svg = render_mermaid_svg(
"flowchart TD\nA[你好] --> B[Done]",
MermaidRenderOptions {
theme: Some("dark".to_string()),
font_family: Some("IBM Plex Mono".to_string()),
font_size: Some(18.0),
},
)
.unwrap();
assert!(svg.contains("IBM Plex Mono"));
assert!(svg.contains("你好"));
}
}