Skip to main content

rskit_ai/prompt/
registry.rs

1//! Prompt registries and identities.
2
3use std::collections::BTreeMap;
4
5use semver::Version;
6use serde::{Deserialize, Serialize};
7
8use super::render::placeholders;
9use super::template::{PromptError, PromptTemplate, VariableDecl, VariableType};
10
11/// Stable prompt identity returned by registry listing.
12#[derive(Debug, Clone, PartialEq, Eq, PartialOrd, Ord, Serialize, Deserialize)]
13pub struct PromptIdentity {
14    /// Prompt name.
15    pub name: String,
16    /// Prompt version.
17    pub version: Version,
18}
19
20/// Versioned prompt registry.
21#[derive(Debug, Clone, Default)]
22pub struct Registry {
23    prompts: BTreeMap<(String, Version), PromptTemplate>,
24}
25
26impl Registry {
27    /// Create an empty registry.
28    #[must_use]
29    pub fn new() -> Self {
30        Self::default()
31    }
32
33    /// Register a prompt template.
34    pub fn register(
35        &mut self,
36        name: impl Into<String>,
37        version: impl AsRef<str>,
38        template: impl Into<String>,
39        output_schema: Option<serde_json::Value>,
40    ) -> Result<PromptTemplate, PromptError> {
41        let name = name.into();
42        let version_str = version.as_ref().to_owned();
43        let version =
44            Version::parse(&version_str).map_err(|source| PromptError::InvalidVersion {
45                version: version_str,
46                source,
47            })?;
48        if self.prompts.contains_key(&(name.clone(), version.clone())) {
49            return Err(PromptError::AlreadyRegistered { name, version });
50        }
51        let body = template.into();
52        let variables = placeholders(&body)
53            .into_iter()
54            .map(|name| VariableDecl {
55                name,
56                kind: VariableType::Any,
57                required: true,
58                default: None,
59            })
60            .collect::<Vec<_>>();
61        let prompt = PromptTemplate {
62            name: name.clone(),
63            version: version.clone(),
64            template: body,
65            variables,
66            output_schema,
67            description: String::new(),
68        };
69        self.prompts.insert((name, version), prompt.clone());
70        Ok(prompt)
71    }
72
73    /// Register an already-built prompt template.
74    pub fn register_template(&mut self, prompt: PromptTemplate) -> Result<(), PromptError> {
75        let key = (prompt.name.clone(), prompt.version.clone());
76        if self.prompts.contains_key(&key) {
77            return Err(PromptError::AlreadyRegistered {
78                name: key.0,
79                version: key.1,
80            });
81        }
82        self.prompts.insert(key, prompt);
83        Ok(())
84    }
85
86    /// Look up a prompt by exact version.
87    pub fn lookup(&self, name: &str, version: &Version) -> Result<&PromptTemplate, PromptError> {
88        self.prompts
89            .get(&(name.to_owned(), version.clone()))
90            .ok_or_else(|| PromptError::NotFound {
91                name: name.to_owned(),
92                version: version.clone(),
93            })
94    }
95
96    /// Look up the highest semver for a prompt name.
97    pub fn lookup_latest(&self, name: &str) -> Result<&PromptTemplate, PromptError> {
98        self.prompts
99            .iter()
100            .rev()
101            .find_map(|((prompt_name, _), prompt)| (prompt_name == name).then_some(prompt))
102            .ok_or_else(|| PromptError::NameNotFound(name.to_owned()))
103    }
104
105    /// List prompt identities in stable order.
106    #[must_use]
107    pub fn list(&self) -> Vec<PromptIdentity> {
108        self.prompts
109            .keys()
110            .map(|(name, version)| PromptIdentity {
111                name: name.clone(),
112                version: version.clone(),
113            })
114            .collect()
115    }
116
117    /// Return versions for one prompt name in ascending semver order.
118    #[must_use]
119    pub fn versions(&self, name: &str) -> Vec<Version> {
120        self.prompts
121            .keys()
122            .filter(|(prompt_name, _)| prompt_name == name)
123            .map(|(_, version)| version.clone())
124            .collect()
125    }
126}