1use std::collections::HashMap;
7
8#[derive(Debug, Clone, PartialEq, Eq, Hash)]
14pub enum PersonaVoice {
15 Formal,
17 Casual,
19 Technical,
21 Creative,
23 Empathetic,
25 Authoritative,
27}
28
29impl PersonaVoice {
30 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#[derive(Debug, Clone)]
67pub struct PersonaConstraints {
68 pub max_response_length: usize,
70 pub forbidden_topics: Vec<String>,
72 pub required_disclaimers: Vec<String>,
74 pub always_use_bullet_points: bool,
76 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#[derive(Debug, Clone)]
98pub struct Persona {
99 pub id: String,
101 pub name: String,
103 pub description: String,
105 pub system_prompt: String,
107 pub voice: PersonaVoice,
109 pub constraints: PersonaConstraints,
111 pub active: bool,
113}
114
115#[derive(Debug, Default)]
133pub struct PersonaManager {
134 personas: HashMap<String, Persona>,
135 active_id: Option<String>,
136}
137
138impl PersonaManager {
139 pub fn register(&mut self, persona: Persona) {
142 self.personas.insert(persona.id.clone(), persona);
143 }
144
145 pub fn activate(&mut self, id: &str) -> bool {
149 if !self.personas.contains_key(id) {
150 return false;
151 }
152 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 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 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 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 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 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
304pub 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#[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 assert_eq!(blended.constraints.max_response_length, 6144);
484 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 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}