Skip to main content

tokio_prompt_orchestrator/
persona_manager.rs

1//! Persona management for AI assistants.
2//!
3//! Provides [`PersonaManager`] to register, activate, and apply distinct AI
4//! personas with configurable voice, style constraints, and system prompts.
5
6use std::collections::HashMap;
7
8// ---------------------------------------------------------------------------
9// PersonaVoice
10// ---------------------------------------------------------------------------
11
12/// The conversational voice / style register for a persona.
13#[derive(Debug, Clone, PartialEq, Eq, Hash)]
14pub enum PersonaVoice {
15    /// Formal, professional language with precise vocabulary.
16    Formal,
17    /// Relaxed, friendly language with contractions and colloquialisms.
18    Casual,
19    /// Jargon-heavy, detail-oriented language suited for developers.
20    Technical,
21    /// Imaginative, expressive language that bends conventions.
22    Creative,
23    /// Warm, supportive language that acknowledges feelings first.
24    Empathetic,
25    /// Confident, decisive language that commands attention.
26    Authoritative,
27}
28
29impl PersonaVoice {
30    /// Human-readable notes about what this voice implies for generation.
31    pub fn style_notes(&self) -> &str {
32        match self {
33            PersonaVoice::Formal => {
34                "Use complete sentences, avoid contractions, prefer Latinate vocabulary, \
35                 and maintain a respectful, impersonal tone."
36            }
37            PersonaVoice::Casual => {
38                "Use contractions freely, short sentences, everyday words, \
39                 and a friendly conversational tone."
40            }
41            PersonaVoice::Technical => {
42                "Prefer precise technical terms, include code snippets where relevant, \
43                 cite specifications, and assume an expert audience."
44            }
45            PersonaVoice::Creative => {
46                "Employ vivid metaphors, varied sentence rhythm, unexpected word choices, \
47                 and embrace ambiguity to spark imagination."
48            }
49            PersonaVoice::Empathetic => {
50                "Acknowledge emotions before facts, validate the user's perspective, \
51                 use 'I understand' and similar phrases, and keep a gentle pace."
52            }
53            PersonaVoice::Authoritative => {
54                "State conclusions first, back them with evidence, avoid hedging language, \
55                 and project confidence in recommendations."
56            }
57        }
58    }
59}
60
61// ---------------------------------------------------------------------------
62// PersonaConstraints
63// ---------------------------------------------------------------------------
64
65/// Behavioural guardrails for a [`Persona`].
66#[derive(Debug, Clone)]
67pub struct PersonaConstraints {
68    /// Hard cap on the character length of any response.
69    pub max_response_length: usize,
70    /// Topics the persona must refuse to engage with.
71    pub forbidden_topics: Vec<String>,
72    /// Strings appended verbatim to every response.
73    pub required_disclaimers: Vec<String>,
74    /// When `true`, the persona must structure output as bullet points.
75    pub always_use_bullet_points: bool,
76    /// Suggested sampling temperature hint `[0.0, 1.0]`.
77    pub temperature_hint: f64,
78}
79
80impl Default for PersonaConstraints {
81    fn default() -> Self {
82        Self {
83            max_response_length: 4096,
84            forbidden_topics: Vec::new(),
85            required_disclaimers: Vec::new(),
86            always_use_bullet_points: false,
87            temperature_hint: 0.7,
88        }
89    }
90}
91
92// ---------------------------------------------------------------------------
93// Persona
94// ---------------------------------------------------------------------------
95
96/// A fully-specified AI persona.
97#[derive(Debug, Clone)]
98pub struct Persona {
99    /// Unique identifier used to look up the persona.
100    pub id: String,
101    /// Human-readable display name.
102    pub name: String,
103    /// Short description of what this persona is intended for.
104    pub description: String,
105    /// System prompt injected before every user message.
106    pub system_prompt: String,
107    /// Voice / style register.
108    pub voice: PersonaVoice,
109    /// Behavioural constraints.
110    pub constraints: PersonaConstraints,
111    /// Whether this persona is currently selected as the active one.
112    pub active: bool,
113}
114
115// ---------------------------------------------------------------------------
116// PersonaManager
117// ---------------------------------------------------------------------------
118
119/// Registry and runtime controller for AI personas.
120///
121/// # Example
122/// ```
123/// use tokio_prompt_orchestrator::persona_manager::{PersonaManager, default_personas};
124///
125/// let mut mgr = PersonaManager::default();
126/// for p in default_personas() { mgr.register(p); }
127/// assert!(mgr.activate("assistant"));
128/// let persona = mgr.active_persona().unwrap();
129/// let prompt = mgr.apply_persona("What is Rust?", persona);
130/// assert!(prompt.contains("What is Rust?"));
131/// ```
132#[derive(Debug, Default)]
133pub struct PersonaManager {
134    personas: HashMap<String, Persona>,
135    active_id: Option<String>,
136}
137
138impl PersonaManager {
139    /// Register a new persona. If a persona with the same id already exists it
140    /// is replaced.
141    pub fn register(&mut self, persona: Persona) {
142        self.personas.insert(persona.id.clone(), persona);
143    }
144
145    /// Activate the persona with the given `id`.
146    ///
147    /// Returns `true` if the persona was found and activated, `false` otherwise.
148    pub fn activate(&mut self, id: &str) -> bool {
149        if !self.personas.contains_key(id) {
150            return false;
151        }
152        // Deactivate all others.
153        for p in self.personas.values_mut() {
154            p.active = false;
155        }
156        if let Some(p) = self.personas.get_mut(id) {
157            p.active = true;
158        }
159        self.active_id = Some(id.to_string());
160        true
161    }
162
163    /// Return a reference to the currently active persona, if any.
164    pub fn active_persona(&self) -> Option<&Persona> {
165        self.active_id
166            .as_deref()
167            .and_then(|id| self.personas.get(id))
168    }
169
170    /// Build the final prompt string by prepending the system context and
171    /// appending required disclaimers.
172    ///
173    /// Format:
174    /// ```text
175    /// [SYSTEM]: <system_prompt>
176    /// [STYLE]: <style_notes>
177    ///
178    /// <prompt>
179    ///
180    /// <disclaimer_1>
181    /// <disclaimer_2>
182    /// ```
183    pub fn apply_persona(&self, prompt: &str, persona: &Persona) -> String {
184        let mut out = String::new();
185        out.push_str("[SYSTEM]: ");
186        out.push_str(&persona.system_prompt);
187        out.push('\n');
188        out.push_str("[STYLE]: ");
189        out.push_str(persona.voice.style_notes());
190        out.push_str("\n\n");
191        out.push_str(prompt);
192        for disclaimer in &persona.constraints.required_disclaimers {
193            out.push('\n');
194            out.push_str(disclaimer);
195        }
196        out
197    }
198
199    /// Validate a model response against persona constraints.
200    ///
201    /// Returns a list of human-readable violation messages (empty = valid).
202    pub fn validate_response(&self, response: &str, persona: &Persona) -> Vec<String> {
203        let mut violations = Vec::new();
204
205        if response.len() > persona.constraints.max_response_length {
206            violations.push(format!(
207                "Response length {} exceeds max_response_length {}",
208                response.len(),
209                persona.constraints.max_response_length
210            ));
211        }
212
213        let lower = response.to_lowercase();
214        for topic in &persona.constraints.forbidden_topics {
215            if lower.contains(&topic.to_lowercase()) {
216                violations.push(format!("Response contains forbidden topic: '{topic}'"));
217            }
218        }
219
220        if persona.constraints.always_use_bullet_points {
221            let has_bullets = response.lines().any(|l| {
222                let t = l.trim_start();
223                t.starts_with("- ") || t.starts_with("* ") || t.starts_with("• ")
224            });
225            if !has_bullets {
226                violations.push(
227                    "Response must use bullet points (always_use_bullet_points = true)".to_string(),
228                );
229            }
230        }
231
232        violations
233    }
234
235    /// Create a blended persona by interpolating the numeric constraints of
236    /// persona `a_id` (weight `alpha`) and persona `b_id` (weight `1 - alpha`).
237    ///
238    /// The resulting persona's system prompt is the concatenation of both
239    /// system prompts separated by a newline. The voice of persona `a` is used.
240    /// Returns `None` if either id is not found or `alpha` is outside `[0, 1]`.
241    pub fn blend_personas(&self, a_id: &str, b_id: &str, alpha: f64) -> Option<Persona> {
242        if !(0.0..=1.0).contains(&alpha) {
243            return None;
244        }
245        let a = self.personas.get(a_id)?;
246        let b = self.personas.get(b_id)?;
247        let beta = 1.0 - alpha;
248
249        let blended_max_len = (a.constraints.max_response_length as f64 * alpha
250            + b.constraints.max_response_length as f64 * beta)
251            .round() as usize;
252        let blended_temp = a.constraints.temperature_hint * alpha
253            + b.constraints.temperature_hint * beta;
254
255        let mut forbidden = a.constraints.forbidden_topics.clone();
256        for t in &b.constraints.forbidden_topics {
257            if !forbidden.contains(t) {
258                forbidden.push(t.clone());
259            }
260        }
261        let mut disclaimers = a.constraints.required_disclaimers.clone();
262        for d in &b.constraints.required_disclaimers {
263            if !disclaimers.contains(d) {
264                disclaimers.push(d.clone());
265            }
266        }
267
268        Some(Persona {
269            id: format!("{a_id}_{b_id}_blend"),
270            name: format!("{} + {} Blend", a.name, b.name),
271            description: format!(
272                "Blended persona: {} ({:.0}%) and {} ({:.0}%)",
273                a.name,
274                alpha * 100.0,
275                b.name,
276                beta * 100.0
277            ),
278            system_prompt: format!("{}\n{}", a.system_prompt, b.system_prompt),
279            voice: a.voice.clone(),
280            constraints: PersonaConstraints {
281                max_response_length: blended_max_len,
282                forbidden_topics: forbidden,
283                required_disclaimers: disclaimers,
284                always_use_bullet_points: a.constraints.always_use_bullet_points
285                    || b.constraints.always_use_bullet_points,
286                temperature_hint: blended_temp,
287            },
288            active: false,
289        })
290    }
291
292    /// Return all personas that match the given voice.
293    pub fn list_by_voice(&self, voice: &PersonaVoice) -> Vec<&Persona> {
294        let mut result: Vec<&Persona> = self
295            .personas
296            .values()
297            .filter(|p| &p.voice == voice)
298            .collect();
299        result.sort_by(|a, b| a.id.cmp(&b.id));
300        result
301    }
302}
303
304// ---------------------------------------------------------------------------
305// Pre-built personas
306// ---------------------------------------------------------------------------
307
308/// Return the standard set of pre-built personas.
309///
310/// Personas included:
311/// - `assistant` — general-purpose formal assistant
312/// - `code_helper` — technical coding assistant
313/// - `creative_writer` — imaginative creative writing partner
314/// - `fact_checker` — authoritative fact-checking advisor
315pub fn default_personas() -> Vec<Persona> {
316    vec![
317        Persona {
318            id: "assistant".to_string(),
319            name: "Assistant".to_string(),
320            description: "General-purpose helpful assistant with a formal voice.".to_string(),
321            system_prompt: "You are a helpful, harmless, and honest AI assistant. \
322                            Answer questions accurately and concisely."
323                .to_string(),
324            voice: PersonaVoice::Formal,
325            constraints: PersonaConstraints {
326                max_response_length: 4096,
327                forbidden_topics: vec!["illegal activities".to_string()],
328                required_disclaimers: vec![],
329                always_use_bullet_points: false,
330                temperature_hint: 0.7,
331            },
332            active: false,
333        },
334        Persona {
335            id: "code_helper".to_string(),
336            name: "CodeHelper".to_string(),
337            description: "Expert coding assistant focused on correctness and best practices."
338                .to_string(),
339            system_prompt: "You are an expert software engineer. Provide correct, idiomatic code \
340                            with explanations. Prefer Rust, Python, and TypeScript unless asked otherwise."
341                .to_string(),
342            voice: PersonaVoice::Technical,
343            constraints: PersonaConstraints {
344                max_response_length: 8192,
345                forbidden_topics: vec![],
346                required_disclaimers: vec![
347                    "Note: always test generated code before deploying to production.".to_string(),
348                ],
349                always_use_bullet_points: false,
350                temperature_hint: 0.2,
351            },
352            active: false,
353        },
354        Persona {
355            id: "creative_writer".to_string(),
356            name: "CreativeWriter".to_string(),
357            description: "Imaginative creative writing partner.".to_string(),
358            system_prompt: "You are a creative writing collaborator. Embrace metaphor, subtext, \
359                            and narrative tension. Help the user craft compelling stories."
360                .to_string(),
361            voice: PersonaVoice::Creative,
362            constraints: PersonaConstraints {
363                max_response_length: 6000,
364                forbidden_topics: vec!["graphic violence".to_string(), "hate speech".to_string()],
365                required_disclaimers: vec![],
366                always_use_bullet_points: false,
367                temperature_hint: 0.9,
368            },
369            active: false,
370        },
371        Persona {
372            id: "fact_checker".to_string(),
373            name: "FactChecker".to_string(),
374            description: "Authoritative fact-checking advisor with citations.".to_string(),
375            system_prompt: "You are a rigorous fact-checker. Cite sources, distinguish between \
376                            established fact and opinion, and flag uncertainty explicitly."
377                .to_string(),
378            voice: PersonaVoice::Authoritative,
379            constraints: PersonaConstraints {
380                max_response_length: 3000,
381                forbidden_topics: vec![],
382                required_disclaimers: vec![
383                    "Disclaimer: verify all claims with primary sources.".to_string(),
384                ],
385                always_use_bullet_points: true,
386                temperature_hint: 0.1,
387            },
388            active: false,
389        },
390    ]
391}
392
393// ---------------------------------------------------------------------------
394// Tests
395// ---------------------------------------------------------------------------
396
397#[cfg(test)]
398mod tests {
399    use super::*;
400
401    fn make_manager() -> PersonaManager {
402        let mut mgr = PersonaManager::default();
403        for p in default_personas() {
404            mgr.register(p);
405        }
406        mgr
407    }
408
409    #[test]
410    fn test_register_and_activate() {
411        let mut mgr = make_manager();
412        assert!(mgr.active_persona().is_none());
413        assert!(mgr.activate("assistant"));
414        assert_eq!(mgr.active_persona().unwrap().id, "assistant");
415    }
416
417    #[test]
418    fn test_activate_unknown_returns_false() {
419        let mut mgr = make_manager();
420        assert!(!mgr.activate("nonexistent"));
421        assert!(mgr.active_persona().is_none());
422    }
423
424    #[test]
425    fn test_only_one_active_at_a_time() {
426        let mut mgr = make_manager();
427        mgr.activate("assistant");
428        mgr.activate("code_helper");
429        let active: Vec<&Persona> = mgr.personas.values().filter(|p| p.active).collect();
430        assert_eq!(active.len(), 1);
431        assert_eq!(active[0].id, "code_helper");
432    }
433
434    #[test]
435    fn test_apply_persona_contains_system_prompt() {
436        let mgr = make_manager();
437        let p = mgr.personas.get("fact_checker").unwrap();
438        let out = mgr.apply_persona("Is the Earth flat?", p);
439        assert!(out.contains("[SYSTEM]:"));
440        assert!(out.contains("Is the Earth flat?"));
441        assert!(out.contains("Disclaimer: verify all claims"));
442    }
443
444    #[test]
445    fn test_validate_response_length_violation() {
446        let mgr = make_manager();
447        let p = mgr.personas.get("fact_checker").unwrap();
448        let long = "x".repeat(4000);
449        let violations = mgr.validate_response(&long, p);
450        assert!(violations.iter().any(|v| v.contains("exceeds max_response_length")));
451    }
452
453    #[test]
454    fn test_validate_response_forbidden_topic() {
455        let mgr = make_manager();
456        let p = mgr.personas.get("assistant").unwrap();
457        let violations = mgr.validate_response("Here is how to do illegal activities cheaply", p);
458        assert!(violations.iter().any(|v| v.contains("forbidden topic")));
459    }
460
461    #[test]
462    fn test_validate_response_bullet_point_violation() {
463        let mgr = make_manager();
464        let p = mgr.personas.get("fact_checker").unwrap();
465        let violations = mgr.validate_response("The earth is round.", p);
466        assert!(violations.iter().any(|v| v.contains("bullet points")));
467    }
468
469    #[test]
470    fn test_validate_response_with_bullets_passes() {
471        let mgr = make_manager();
472        let p = mgr.personas.get("fact_checker").unwrap();
473        let violations = mgr.validate_response("- The earth is an oblate spheroid.", p);
474        assert!(!violations.iter().any(|v| v.contains("bullet points")));
475    }
476
477    #[test]
478    fn test_blend_personas() {
479        let mgr = make_manager();
480        let blended = mgr.blend_personas("assistant", "code_helper", 0.5).unwrap();
481        assert_eq!(blended.id, "assistant_code_helper_blend");
482        // Blended max_response_length should be midpoint of 4096 and 8192.
483        assert_eq!(blended.constraints.max_response_length, 6144);
484        // Both system prompts should be present.
485        assert!(blended.system_prompt.contains("helpful"));
486        assert!(blended.system_prompt.contains("software engineer"));
487    }
488
489    #[test]
490    fn test_blend_personas_invalid_alpha() {
491        let mgr = make_manager();
492        assert!(mgr.blend_personas("assistant", "code_helper", 1.5).is_none());
493        assert!(mgr.blend_personas("assistant", "code_helper", -0.1).is_none());
494    }
495
496    #[test]
497    fn test_blend_personas_unknown_id() {
498        let mgr = make_manager();
499        assert!(mgr.blend_personas("assistant", "unknown", 0.5).is_none());
500    }
501
502    #[test]
503    fn test_list_by_voice() {
504        let mgr = make_manager();
505        let technical = mgr.list_by_voice(&PersonaVoice::Technical);
506        assert!(technical.iter().any(|p| p.id == "code_helper"));
507        // No other default persona has Technical voice.
508        assert_eq!(technical.len(), 1);
509    }
510
511    #[test]
512    fn test_style_notes_non_empty() {
513        for voice in [
514            PersonaVoice::Formal,
515            PersonaVoice::Casual,
516            PersonaVoice::Technical,
517            PersonaVoice::Creative,
518            PersonaVoice::Empathetic,
519            PersonaVoice::Authoritative,
520        ] {
521            assert!(!voice.style_notes().is_empty());
522        }
523    }
524}