use super::template::{PromptComposition, PromptError, PromptResult, PromptTemplate};
use serde::{Deserialize, Serialize};
use std::collections::HashMap;
use std::path::Path;
use std::sync::{Arc, RwLock};
#[derive(Default)]
pub struct PromptRegistry {
templates: HashMap<String, PromptTemplate>,
compositions: HashMap<String, PromptComposition>,
tag_index: HashMap<String, Vec<String>>,
}
impl PromptRegistry {
pub fn new() -> Self {
Self::default()
}
pub fn register(&mut self, template: PromptTemplate) {
let id = template.id.clone();
for tag in &template.tags {
self.tag_index
.entry(tag.clone())
.or_default()
.push(id.clone());
}
self.templates.insert(id, template);
}
pub fn register_composition(&mut self, composition: PromptComposition) {
self.compositions
.insert(composition.id.clone(), composition);
}
pub fn get(&self, id: &str) -> PromptResult<&PromptTemplate> {
self.templates
.get(id)
.ok_or_else(|| PromptError::TemplateNotFound(id.to_string()))
}
pub fn get_mut(&mut self, id: &str) -> PromptResult<&mut PromptTemplate> {
self.templates
.get_mut(id)
.ok_or_else(|| PromptError::TemplateNotFound(id.to_string()))
}
pub fn get_composition(&self, id: &str) -> PromptResult<&PromptComposition> {
self.compositions
.get(id)
.ok_or_else(|| PromptError::TemplateNotFound(format!("composition:{}", id)))
}
pub fn contains(&self, id: &str) -> bool {
self.templates.contains_key(id)
}
pub fn remove(&mut self, id: &str) -> Option<PromptTemplate> {
if let Some(template) = self.templates.remove(id) {
for tag in &template.tags {
if let Some(ids) = self.tag_index.get_mut(tag) {
ids.retain(|i| i != id);
}
}
Some(template)
} else {
None
}
}
pub fn list_ids(&self) -> Vec<&str> {
self.templates.keys().map(|s| s.as_str()).collect()
}
pub fn find_by_tag(&self, tag: &str) -> Vec<&PromptTemplate> {
self.tag_index
.get(tag)
.map(|ids| ids.iter().filter_map(|id| self.templates.get(id)).collect())
.unwrap_or_default()
}
pub fn search(&self, query: &str) -> Vec<&PromptTemplate> {
let query_lower = query.to_lowercase();
self.templates
.values()
.filter(|t| {
t.id.to_lowercase().contains(&query_lower)
|| t.name
.as_ref()
.is_some_and(|n| n.to_lowercase().contains(&query_lower))
|| t.description
.as_ref()
.is_some_and(|d| d.to_lowercase().contains(&query_lower))
})
.collect()
}
pub fn list_tags(&self) -> Vec<&str> {
self.tag_index.keys().map(|s| s.as_str()).collect()
}
pub fn render(&self, id: &str, vars: &[(&str, &str)]) -> PromptResult<String> {
self.get(id)?.render(vars)
}
pub fn render_composition(
&self,
composition_id: &str,
vars: &[(&str, &str)],
) -> PromptResult<String> {
let composition = self.get_composition(composition_id)?;
let mut results = Vec::new();
for template_id in &composition.template_ids {
let rendered = self.render(template_id, vars)?;
results.push(rendered);
}
Ok(results.join(&composition.separator))
}
pub fn load_from_file(&mut self, path: impl AsRef<Path>) -> PromptResult<()> {
let content = std::fs::read_to_string(path)?;
self.load_from_yaml(&content)
}
pub fn load_from_yaml(&mut self, yaml: &str) -> PromptResult<()> {
let config: PromptYamlConfig =
serde_yaml::from_str(yaml).map_err(|e| PromptError::YamlError(e.to_string()))?;
if let Some(templates) = config.templates {
for template in templates {
self.register(template);
}
}
if let Some(compositions) = config.compositions {
for composition in compositions {
self.register_composition(composition);
}
}
Ok(())
}
pub fn export_to_yaml(&self) -> PromptResult<String> {
let config = PromptYamlConfig {
templates: Some(self.templates.values().cloned().collect()),
compositions: Some(self.compositions.values().cloned().collect()),
};
serde_yaml::to_string(&config).map_err(|e| PromptError::YamlError(e.to_string()))
}
pub fn merge(&mut self, other: PromptRegistry) {
for (id, template) in other.templates {
self.templates.insert(id, template);
}
for (id, composition) in other.compositions {
self.compositions.insert(id, composition);
}
self.rebuild_tag_index();
}
fn rebuild_tag_index(&mut self) {
self.tag_index.clear();
for (id, template) in &self.templates {
for tag in &template.tags {
self.tag_index
.entry(tag.clone())
.or_default()
.push(id.clone());
}
}
}
pub fn len(&self) -> usize {
self.templates.len()
}
pub fn is_empty(&self) -> bool {
self.templates.is_empty()
}
pub fn clear(&mut self) {
self.templates.clear();
self.compositions.clear();
self.tag_index.clear();
}
}
#[derive(Debug, Serialize, Deserialize)]
struct PromptYamlConfig {
#[serde(default)]
templates: Option<Vec<PromptTemplate>>,
#[serde(default)]
compositions: Option<Vec<PromptComposition>>,
}
#[derive(Clone, Default)]
pub struct GlobalPromptRegistry {
inner: Arc<RwLock<PromptRegistry>>,
}
impl GlobalPromptRegistry {
pub fn new() -> Self {
Self::default()
}
pub fn register(&self, template: PromptTemplate) {
self.inner.write().unwrap().register(template);
}
pub fn get(&self, id: &str) -> PromptResult<PromptTemplate> {
self.inner.read().unwrap().get(id).cloned()
}
pub fn render(&self, id: &str, vars: &[(&str, &str)]) -> PromptResult<String> {
self.inner.read().unwrap().render(id, vars)
}
pub fn contains(&self, id: &str) -> bool {
self.inner.read().unwrap().contains(id)
}
pub fn remove(&self, id: &str) -> Option<PromptTemplate> {
self.inner.write().unwrap().remove(id)
}
pub fn load_from_file(&self, path: impl AsRef<Path>) -> PromptResult<()> {
self.inner.write().unwrap().load_from_file(path)
}
pub fn load_from_yaml(&self, yaml: &str) -> PromptResult<()> {
self.inner.write().unwrap().load_from_yaml(yaml)
}
pub fn list_ids(&self) -> Vec<String> {
self.inner
.read()
.unwrap()
.list_ids()
.iter()
.map(|s| s.to_string())
.collect()
}
pub fn find_by_tag(&self, tag: &str) -> Vec<PromptTemplate> {
self.inner
.read()
.unwrap()
.find_by_tag(tag)
.iter()
.map(|t| (*t).clone())
.collect()
}
pub fn search(&self, query: &str) -> Vec<PromptTemplate> {
self.inner
.read()
.unwrap()
.search(query)
.iter()
.map(|t| (*t).clone())
.collect()
}
pub fn len(&self) -> usize {
self.inner.read().unwrap().len()
}
pub fn is_empty(&self) -> bool {
self.inner.read().unwrap().is_empty()
}
pub fn clear(&self) {
self.inner.write().unwrap().clear();
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_registry_basic() {
let mut registry = PromptRegistry::new();
let template = PromptTemplate::new("greeting")
.with_content("Hello, {name}!")
.with_tag("basic");
registry.register(template);
assert!(registry.contains("greeting"));
assert_eq!(registry.len(), 1);
let result = registry.render("greeting", &[("name", "World")]).unwrap();
assert_eq!(result, "Hello, World!");
}
#[test]
fn test_registry_tags() {
let mut registry = PromptRegistry::new();
registry.register(
PromptTemplate::new("t1")
.with_content("Template 1")
.with_tag("tag-a")
.with_tag("tag-b"),
);
registry.register(
PromptTemplate::new("t2")
.with_content("Template 2")
.with_tag("tag-a"),
);
registry.register(
PromptTemplate::new("t3")
.with_content("Template 3")
.with_tag("tag-c"),
);
let tag_a_templates = registry.find_by_tag("tag-a");
assert_eq!(tag_a_templates.len(), 2);
let tag_b_templates = registry.find_by_tag("tag-b");
assert_eq!(tag_b_templates.len(), 1);
let tag_c_templates = registry.find_by_tag("tag-c");
assert_eq!(tag_c_templates.len(), 1);
}
#[test]
fn test_registry_search() {
let mut registry = PromptRegistry::new();
registry.register(
PromptTemplate::new("code-review")
.with_name("Code Review")
.with_description("Review code for issues"),
);
registry.register(
PromptTemplate::new("code-explain")
.with_name("Code Explanation")
.with_description("Explain code in detail"),
);
registry.register(
PromptTemplate::new("chat")
.with_name("Chat Assistant")
.with_description("General chat"),
);
let code_templates = registry.search("code");
assert_eq!(code_templates.len(), 2);
let review_templates = registry.search("review");
assert_eq!(review_templates.len(), 1);
}
#[test]
fn test_registry_yaml() {
let yaml = r#"
templates:
- id: greeting
name: Greeting
content: "Hello, {name}!"
tags:
- basic
variables:
- name: name
required: true
- id: farewell
content: "Goodbye, {name}!"
variables:
- name: name
default: friend
compositions:
- id: full-conversation
template_ids:
- greeting
- farewell
separator: "\n"
"#;
let mut registry = PromptRegistry::new();
registry.load_from_yaml(yaml).unwrap();
assert_eq!(registry.len(), 2);
assert!(registry.contains("greeting"));
assert!(registry.contains("farewell"));
let greeting = registry.render("greeting", &[("name", "Alice")]).unwrap();
assert_eq!(greeting, "Hello, Alice!");
let farewell = registry.render("farewell", &[]).unwrap();
assert_eq!(farewell, "Goodbye, friend!");
let composition = registry
.render_composition("full-conversation", &[("name", "Bob")])
.unwrap();
assert_eq!(composition, "Hello, Bob!\nGoodbye, Bob!");
}
#[test]
fn test_registry_remove() {
let mut registry = PromptRegistry::new();
registry.register(
PromptTemplate::new("test")
.with_content("Test")
.with_tag("removable"),
);
assert!(registry.contains("test"));
assert_eq!(registry.find_by_tag("removable").len(), 1);
let removed = registry.remove("test");
assert!(removed.is_some());
assert!(!registry.contains("test"));
assert_eq!(registry.find_by_tag("removable").len(), 0);
}
#[test]
fn test_global_registry() {
let registry = GlobalPromptRegistry::new();
registry.register(PromptTemplate::new("test").with_content("Hello, {name}!"));
assert!(registry.contains("test"));
let result = registry.render("test", &[("name", "World")]).unwrap();
assert_eq!(result, "Hello, World!");
}
}