use alloc::{borrow::Cow, boxed::Box, format, string::ToString, vec::Vec};
use core::{future::Future, pin::Pin};
use super::types::{
PromptArgumentDefinition, PromptDefinition, PromptOutput, RegistryItem, ToolError,
};
use crate::compat::{HashMap, HashMapIter};
pub trait RustPrompt: Send + Sync {
type Params: serde::de::DeserializeOwned + schemars::JsonSchema + Send;
const NAME: &'static str;
const DESCRIPTION: &'static str;
fn description(&self) -> Cow<'static, str> {
Cow::Borrowed(Self::DESCRIPTION)
}
fn render(
&self,
params: Self::Params,
) -> impl Future<Output = Result<PromptOutput, ToolError>> + Send;
}
pub fn definition_of_prompt<T: RustPrompt>(prompt: &T) -> Result<PromptDefinition, ToolError> {
let schema = schemars::schema_for!(T::Params);
let val = serde_json::to_value(&schema).map_err(|e| {
ToolError::new(format!(
"Failed to serialize schema for prompt '{}': {e}",
T::NAME
))
})?;
let mut arguments = Vec::new();
if let Some(obj) = val.as_object() {
let required_fields: Vec<&str> = match obj.get("required").and_then(|r| r.as_array()) {
Some(arr) => arr.iter().filter_map(|v| v.as_str()).collect(),
None => Vec::new(),
};
if let Some(props) = obj.get("properties").and_then(|p| p.as_object()) {
for (name, prop) in props {
let description = prop
.get("description")
.and_then(|d| d.as_str())
.unwrap_or("")
.to_string();
let required = required_fields.contains(&name.as_str());
arguments.push(PromptArgumentDefinition {
name: name.clone(),
description,
required,
});
}
}
}
Ok(PromptDefinition {
name: T::NAME.to_string(),
description: prompt.description().into_owned(),
arguments,
})
}
pub(crate) type BoxPromptFuture<'a> =
Pin<Box<dyn Future<Output = Result<PromptOutput, ToolError>> + Send + 'a>>;
pub(crate) trait ErasedPrompt: Send + Sync {
fn render_erased(&self, args: serde_json::Value) -> BoxPromptFuture<'_>;
}
impl<T: RustPrompt> ErasedPrompt for T {
fn render_erased(&self, args: serde_json::Value) -> BoxPromptFuture<'_> {
Box::pin(async move {
let params: T::Params = serde_json::from_value(args).map_err(|e| {
ToolError::new(format!("Failed to deserialize prompt parameters: {e}"))
})?;
self.render(params).await
})
}
}
struct RegisteredPrompt {
definition: PromptDefinition,
erased: Box<dyn ErasedPrompt>,
}
#[derive(Default)]
pub struct PromptRegistry {
prompts: HashMap<&'static str, RegisteredPrompt>,
}
impl core::fmt::Debug for PromptRegistry {
fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
let names: Vec<&str> = self.prompts.keys().copied().collect();
f.debug_struct("PromptRegistry")
.field("prompt_count", &self.prompts.len())
.field("prompt_names", &names)
.finish()
}
}
impl PromptRegistry {
#[must_use]
pub fn new() -> Self {
Self {
prompts: HashMap::new(),
}
}
pub fn register<P: RustPrompt + 'static>(&mut self, prompt: P) -> &mut Self {
if let Err(e) = self.try_register(prompt) {
panic!("Failed to build definition for prompt '{}': {e}", P::NAME);
}
self
}
pub fn try_register<P: RustPrompt + 'static>(
&mut self,
prompt: P,
) -> Result<&mut Self, ToolError> {
let definition = definition_of_prompt(&prompt)?;
self.prompts.insert(
P::NAME,
RegisteredPrompt {
definition,
erased: Box::new(prompt),
},
);
Ok(self)
}
#[must_use]
pub fn with_prompt<P: RustPrompt + 'static>(mut self, prompt: P) -> Self {
self.register(prompt);
self
}
#[must_use]
pub fn definitions(&self) -> Vec<PromptDefinition> {
self.prompts
.values()
.map(|entry| entry.definition.clone())
.collect()
}
pub fn remove(&mut self, name: &str) -> bool {
self.prompts.remove(name).is_some()
}
pub fn clear(&mut self) {
self.prompts.clear();
}
#[must_use]
pub fn len(&self) -> usize {
self.prompts.len()
}
#[must_use]
pub fn is_empty(&self) -> bool {
self.prompts.is_empty()
}
#[must_use]
pub fn contains(&self, name: &str) -> bool {
self.prompts.contains_key(name)
}
#[must_use]
pub fn definition(&self, name: &str) -> Option<&PromptDefinition> {
self.prompts.get(name).map(|entry| &entry.definition)
}
#[must_use]
pub fn iter(&self) -> PromptDefinitions<'_> {
PromptDefinitions {
inner: self.prompts.iter(),
}
}
pub async fn render(
&self,
name: &str,
args: serde_json::Value,
) -> Result<PromptOutput, ToolError> {
let Some(entry) = self.prompts.get(name) else {
return Err(ToolError::not_found(RegistryItem::Prompt, name));
};
entry.erased.render_erased(args).await
}
}
pub struct PromptDefinitions<'a> {
inner: HashMapIter<'a, &'static str, RegisteredPrompt>,
}
impl Iterator for PromptDefinitions<'_> {
type Item = (&'static str, PromptDefinition);
fn next(&mut self) -> Option<Self::Item> {
self.inner
.next()
.map(|(name, entry)| (*name, entry.definition.clone()))
}
fn size_hint(&self) -> (usize, Option<usize>) {
self.inner.size_hint()
}
}
impl ExactSizeIterator for PromptDefinitions<'_> {
fn len(&self) -> usize {
self.inner.len()
}
}
impl<'a> IntoIterator for &'a PromptRegistry {
type Item = (&'static str, PromptDefinition);
type IntoIter = PromptDefinitions<'a>;
fn into_iter(self) -> Self::IntoIter {
self.iter()
}
}