Skip to main content

ai_agents_context/
render.rs

1use std::collections::HashMap;
2
3use minijinja::{Environment, Value as MJValue};
4use serde_json::Value;
5
6use ai_agents_core::{AgentError, Result};
7
8pub struct TemplateRenderer {
9    env: Environment<'static>,
10}
11
12impl Default for TemplateRenderer {
13    fn default() -> Self {
14        Self::new()
15    }
16}
17
18impl TemplateRenderer {
19    pub fn new() -> Self {
20        let mut env = Environment::new();
21        env.set_auto_escape_callback(|_| minijinja::AutoEscape::None);
22        Self { env }
23    }
24
25    pub fn render(&self, template: &str, context: &HashMap<String, Value>) -> Result<String> {
26        let mut ctx = HashMap::new();
27
28        // Build context map - all values available as {{ context.<key> }}
29        let mut context_map = serde_json::Map::new();
30        for (key, value) in context {
31            context_map.insert(key.clone(), value.clone());
32        }
33        ctx.insert("context", json_to_minijinja(&Value::Object(context_map)));
34
35        if let Some(env_vars) = context.get("env") {
36            ctx.insert("env", json_to_minijinja(env_vars));
37        }
38
39        if let Some(state) = context.get("state") {
40            ctx.insert("state", json_to_minijinja(state));
41        }
42
43        // Hoist a fixed set of well-known top-level variables so system prompt
44        // templates can use {{ actor_facts }} directly instead of
45        // {{ context.actor_facts }}. Any key listed here is also accessible via
46        // the context. prefix, so both forms work.
47        for key in &["actor_facts", "relationship_memory"] {
48            if let Some(value) = context.get(*key) {
49                ctx.insert(key, json_to_minijinja(value));
50            }
51        }
52
53        let tmpl = self
54            .env
55            .template_from_str(template)
56            .map_err(|e| AgentError::TemplateError(e.to_string()))?;
57
58        tmpl.render(&ctx)
59            .map_err(|e| AgentError::TemplateError(e.to_string()))
60    }
61
62    pub fn render_path(
63        &self,
64        path_template: &str,
65        context: &HashMap<String, Value>,
66    ) -> Result<String> {
67        self.render(path_template, context)
68    }
69
70    pub fn render_with_state(
71        &self,
72        template: &str,
73        context: &HashMap<String, Value>,
74        state_name: &str,
75        previous_state: Option<&str>,
76        turn_count: u32,
77        max_turns: Option<u32>,
78    ) -> Result<String> {
79        let mut full_context = context.clone();
80
81        let mut state_ctx = serde_json::Map::new();
82        state_ctx.insert("name".into(), Value::String(state_name.to_string()));
83        state_ctx.insert(
84            "previous".into(),
85            Value::String(previous_state.unwrap_or("none").to_string()),
86        );
87        state_ctx.insert("turn_count".into(), Value::Number(turn_count.into()));
88        if let Some(max) = max_turns {
89            state_ctx.insert("max_turns".into(), Value::Number(max.into()));
90        }
91        full_context.insert("state".into(), Value::Object(state_ctx));
92
93        self.render(template, &full_context)
94    }
95}
96
97fn json_to_minijinja(value: &Value) -> MJValue {
98    match value {
99        Value::Null => MJValue::from(()),
100        Value::Bool(b) => MJValue::from(*b),
101        Value::Number(n) => {
102            if let Some(i) = n.as_i64() {
103                MJValue::from(i)
104            } else if let Some(u) = n.as_u64() {
105                MJValue::from(u)
106            } else if let Some(f) = n.as_f64() {
107                MJValue::from(f)
108            } else {
109                MJValue::from(())
110            }
111        }
112        Value::String(s) => MJValue::from(s.as_str()),
113        Value::Array(arr) => {
114            let items: Vec<MJValue> = arr.iter().map(json_to_minijinja).collect();
115            MJValue::from(items)
116        }
117        Value::Object(obj) => {
118            let map: std::collections::BTreeMap<String, MJValue> = obj
119                .iter()
120                .map(|(k, v)| (k.clone(), json_to_minijinja(v)))
121                .collect();
122            MJValue::from_iter(map)
123        }
124    }
125}
126
127#[cfg(test)]
128mod tests {
129    use super::*;
130    use serde_json::json;
131
132    #[test]
133    fn test_simple_variable() {
134        let renderer = TemplateRenderer::new();
135        let mut context = HashMap::new();
136        context.insert("user".into(), json!({"name": "Alice", "tier": "premium"}));
137
138        let template = "Hello, {{ context.user.name }}!";
139        let result = renderer.render(template, &context).unwrap();
140        assert_eq!(result, "Hello, Alice!");
141    }
142
143    #[test]
144    fn test_nested_variable() {
145        let renderer = TemplateRenderer::new();
146        let mut context = HashMap::new();
147        context.insert(
148            "user".into(),
149            json!({"preferences": {"theme": "dark", "language": "ko"}}),
150        );
151
152        let template = "Theme: {{ context.user.preferences.theme }}";
153        let result = renderer.render(template, &context).unwrap();
154        assert_eq!(result, "Theme: dark");
155    }
156
157    #[test]
158    fn test_conditional() {
159        let renderer = TemplateRenderer::new();
160        let mut context = HashMap::new();
161        context.insert("user".into(), json!({"tier": "premium"}));
162
163        let template = r#"{% if context.user.tier == "premium" %}Premium user{% else %}Regular user{% endif %}"#;
164        let result = renderer.render(template, &context).unwrap();
165        assert_eq!(result, "Premium user");
166    }
167
168    #[test]
169    fn test_loop() {
170        let renderer = TemplateRenderer::new();
171        let mut context = HashMap::new();
172        context.insert("items".into(), json!([{"name": "A"}, {"name": "B"}]));
173
174        let template = "{% for item in context.items %}{{ item.name }}{% endfor %}";
175        let result = renderer.render(template, &context).unwrap();
176        assert_eq!(result, "AB");
177    }
178
179    #[test]
180    fn test_state_variables() {
181        let renderer = TemplateRenderer::new();
182        let context = HashMap::new();
183
184        let template = "State: {{ state.name }}, Turn: {{ state.turn_count }}";
185        let result = renderer
186            .render_with_state(template, &context, "support", Some("greeting"), 2, Some(5))
187            .unwrap();
188        assert_eq!(result, "State: support, Turn: 2");
189    }
190
191    #[test]
192    fn test_korean_content() {
193        let renderer = TemplateRenderer::new();
194        let mut context = HashMap::new();
195        context.insert("user".into(), json!({"name": "김철수", "language": "ko"}));
196
197        let template = "안녕하세요, {{ context.user.name }}님! 언어: {{ context.user.language }}";
198        let result = renderer.render(template, &context).unwrap();
199        assert_eq!(result, "안녕하세요, 김철수님! 언어: ko");
200    }
201
202    #[test]
203    fn test_path_rendering() {
204        let renderer = TemplateRenderer::new();
205        let mut context = HashMap::new();
206        context.insert("user".into(), json!({"language": "ja"}));
207
208        let path = "./rules/{{ context.user.language }}/support.txt";
209        let result = renderer.render_path(path, &context).unwrap();
210        assert_eq!(result, "./rules/ja/support.txt");
211    }
212
213    #[test]
214    fn test_default_filter() {
215        let renderer = TemplateRenderer::new();
216        let context = HashMap::new();
217
218        let template = "{{ context.missing | default('N/A') }}";
219        let result = renderer.render(template, &context).unwrap();
220        assert_eq!(result, "N/A");
221    }
222
223    // actor_facts must be accessible as a top-level variable {{ actor_facts }}
224    // because that is the form used in system prompt templates.
225    // It is also accessible as {{ context.actor_facts }}.
226    #[test]
227    fn test_actor_facts_top_level_variable() {
228        let renderer = TemplateRenderer::new();
229        let mut context = HashMap::new();
230        context.insert(
231            "actor_facts".into(),
232            json!("- User name is Jay.\n- User works as an AI engineer.\n"),
233        );
234
235        // Top-level form used in system prompts.
236        let template = "{% if actor_facts %}Known facts:\n{{ actor_facts }}{% endif %}";
237        let result = renderer.render(template, &context).unwrap();
238        assert!(
239            result.contains("User name is Jay"),
240            "actor_facts must render at top level without context. prefix"
241        );
242        assert!(
243            result.contains("Known facts"),
244            "{{% if actor_facts %}} must evaluate to true for non-empty string"
245        );
246    }
247
248    #[test]
249    fn test_actor_facts_also_accessible_via_context_prefix() {
250        let renderer = TemplateRenderer::new();
251        let mut context = HashMap::new();
252        context.insert("actor_facts".into(), json!("- User name is Jay.\n"));
253
254        // Both forms must work.
255        let template_top = "{{ actor_facts }}";
256        let template_ctx = "{{ context.actor_facts }}";
257        let result_top = renderer.render(template_top, &context).unwrap();
258        let result_ctx = renderer.render(template_ctx, &context).unwrap();
259        assert_eq!(result_top, result_ctx);
260    }
261
262    #[test]
263    fn test_actor_facts_if_block_false_when_absent() {
264        let renderer = TemplateRenderer::new();
265        let context = HashMap::new(); // no actor_facts key
266
267        let template = "base{% if actor_facts %} facts: {{ actor_facts }}{% endif %} end";
268        let result = renderer.render(template, &context).unwrap();
269        assert_eq!(
270            result, "base end",
271            "{{% if actor_facts %}} must be false when actor_facts is not in context"
272        );
273    }
274}