rskit_ai/prompt/
registry.rs1use 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#[derive(Debug, Clone, PartialEq, Eq, PartialOrd, Ord, Serialize, Deserialize)]
13pub struct PromptIdentity {
14 pub name: String,
16 pub version: Version,
18}
19
20#[derive(Debug, Clone, Default)]
22pub struct Registry {
23 prompts: BTreeMap<(String, Version), PromptTemplate>,
24}
25
26impl Registry {
27 #[must_use]
29 pub fn new() -> Self {
30 Self::default()
31 }
32
33 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 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 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 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 #[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 #[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}