use std::fmt;
use schemars::JsonSchema;
use serde::{Deserialize, Serialize};
use crate::hash::Digest;
use crate::ids::string_id;
string_id! {
PromptName
}
string_id! {
PromptVersion
}
#[derive(
Debug, Clone, PartialEq, Eq, Hash, PartialOrd, Ord, Serialize, Deserialize, JsonSchema,
)]
pub struct PromptRef {
pub name: PromptName,
pub version: PromptVersion,
pub hash: Digest,
}
impl PromptRef {
#[must_use]
pub fn of_text(
name: impl Into<PromptName>,
version: impl Into<PromptVersion>,
text: &str,
) -> Self {
Self {
name: name.into(),
version: version.into(),
hash: Digest::of_bytes(text.as_bytes()),
}
}
#[must_use]
pub fn new(
name: impl Into<PromptName>,
version: impl Into<PromptVersion>,
hash: Digest,
) -> Self {
Self {
name: name.into(),
version: version.into(),
hash,
}
}
#[must_use]
pub fn matches(&self, text: &str) -> bool {
self.hash == Digest::of_bytes(text.as_bytes())
}
#[must_use]
pub fn label(&self) -> String {
format!("{}@{}", self.name, self.version)
}
}
impl fmt::Display for PromptRef {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
let short: String = self.hash.as_str().chars().take(8).collect();
write!(f, "{}@{}#{short}", self.name, self.version)
}
}
#[derive(Debug, Clone, PartialEq, Eq, Hash, PartialOrd, Ord, Default)]
pub enum PromptSelector {
#[default]
Latest,
Label(String),
Version(PromptVersion),
}
impl PromptSelector {
#[must_use]
pub fn label(label: impl Into<String>) -> Self {
Self::Label(label.into())
}
#[must_use]
pub fn version(version: impl Into<PromptVersion>) -> Self {
Self::Version(version.into())
}
#[must_use]
pub fn as_key(&self) -> String {
match self {
Self::Latest => "latest".to_owned(),
Self::Label(label) => format!("label:{label}"),
Self::Version(version) => format!("version:{version}"),
}
}
#[must_use]
pub const fn is_pinned(&self) -> bool {
matches!(self, Self::Version(_))
}
}
impl fmt::Display for PromptSelector {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.write_str(&self.as_key())
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct LoadedPrompt {
reference: PromptRef,
text: String,
}
impl LoadedPrompt {
#[must_use]
pub fn new(
name: impl Into<PromptName>,
version: impl Into<PromptVersion>,
text: impl Into<String>,
) -> Self {
let text = text.into();
let reference = PromptRef::of_text(name, version, &text);
Self { reference, text }
}
pub fn from_parts(reference: PromptRef, text: impl Into<String>) -> Result<Self, PromptError> {
let text = text.into();
if !reference.matches(&text) {
return Err(PromptError::InconsistentReference {
name: reference.name,
});
}
Ok(Self { reference, text })
}
#[must_use]
pub fn text(&self) -> &str {
&self.text
}
#[must_use]
pub fn reference(&self) -> &PromptRef {
&self.reference
}
#[must_use]
pub fn name(&self) -> &PromptName {
&self.reference.name
}
#[must_use]
pub fn version(&self) -> &PromptVersion {
&self.reference.version
}
#[must_use]
pub fn into_text(self) -> String {
self.text
}
#[must_use]
pub fn into_parts(self) -> (PromptRef, String) {
(self.reference, self.text)
}
}
#[derive(Debug, Clone, PartialEq, Eq, thiserror::Error)]
#[non_exhaustive]
pub enum PromptError {
#[error("no prompt named {name}")]
NotFound {
name: PromptName,
},
#[error("prompt {name} has no version {version}")]
VersionNotFound {
name: PromptName,
version: PromptVersion,
},
#[error("prompt {name} has no version labelled {label}")]
LabelNotFound {
name: PromptName,
label: String,
},
#[error("the prompt source rejected the credentials")]
Unauthorized,
#[error("the credentials are not entitled to prompt {name}")]
Forbidden {
name: PromptName,
},
#[error("the prompt source rate-limited the request")]
RateLimited,
#[error("the prompt source could not be reached: {code}")]
Transport {
code: &'static str,
},
#[error("the prompt source returned an answer this adapter cannot read: {code}")]
Malformed {
code: &'static str,
},
#[error("the prompt source cannot serve this prompt here: {reason}")]
Unsupported {
reason: &'static str,
},
#[error("prompt {name} came back with a reference that does not match its text")]
InconsistentReference {
name: PromptName,
},
}
impl PromptError {
#[must_use]
pub const fn code(&self) -> &'static str {
match self {
Self::NotFound { .. } => "not_found",
Self::VersionNotFound { .. } => "version_not_found",
Self::LabelNotFound { .. } => "label_not_found",
Self::Unauthorized => "unauthorized",
Self::Forbidden { .. } => "forbidden",
Self::RateLimited => "rate_limited",
Self::Transport { .. } => "transport",
Self::Malformed { .. } => "malformed",
Self::Unsupported { .. } => "unsupported",
Self::InconsistentReference { .. } => "inconsistent_reference",
}
}
#[must_use]
pub const fn is_transient(&self) -> bool {
matches!(self, Self::RateLimited | Self::Transport { .. })
}
}
#[async_trait::async_trait]
pub trait PromptSource: Send + Sync + fmt::Debug {
async fn load(
&self,
name: &PromptName,
selector: &PromptSelector,
) -> Result<LoadedPrompt, PromptError>;
fn describe(&self) -> &'static str;
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn a_reference_built_from_text_matches_that_text_and_no_other() {
let reference = PromptRef::of_text("interpret.system", "v1", "one");
assert!(reference.matches("one"));
assert!(!reference.matches("two"));
assert!(!reference.matches("one "));
}
#[test]
fn the_same_text_under_two_names_hashes_the_same_but_is_a_different_reference() {
let a = PromptRef::of_text("a", "v1", "shared");
let b = PromptRef::of_text("b", "v1", "shared");
assert_eq!(a.hash, b.hash);
assert_ne!(a, b);
}
#[test]
fn rendering_carries_identifiers_and_never_the_text() {
let reference = PromptRef::of_text("interpret.system", "v7", "SECRET INSTRUCTIONS");
let rendered = reference.to_string();
assert!(rendered.starts_with("interpret.system@v7#"));
assert!(!rendered.contains("SECRET"));
assert_eq!(reference.label(), "interpret.system@v7");
assert!(!format!("{reference:?}").contains("SECRET"));
}
#[test]
fn a_reference_round_trips() {
let reference = PromptRef::of_text("n", "v", "text");
let json = serde_json::to_string(&reference).unwrap();
assert_eq!(serde_json::from_str::<PromptRef>(&json).unwrap(), reference);
}
#[test]
fn a_loaded_prompt_carries_the_hash_of_its_own_text() {
let loaded = LoadedPrompt::new("interpret.system", "v1", "Answer with the plan only.");
assert!(loaded.reference().matches(loaded.text()));
assert_eq!(loaded.name().as_str(), "interpret.system");
assert_eq!(loaded.version().as_str(), "v1");
assert_eq!(loaded.clone().into_text(), "Answer with the plan only.");
let (reference, text) = LoadedPrompt::new("n", "v", "t").into_parts();
assert!(reference.matches(&text));
}
#[test]
fn a_reference_that_names_another_text_is_refused() {
let honest = PromptRef::of_text("n", "v", "the real text");
assert!(LoadedPrompt::from_parts(honest.clone(), "the real text").is_ok());
let error = LoadedPrompt::from_parts(honest, "a different text").unwrap_err();
assert_eq!(error.code(), "inconsistent_reference");
}
#[test]
fn selector_keys_are_distinct_stable_and_say_whether_they_pin() {
assert_eq!(PromptSelector::Latest.as_key(), "latest");
assert_eq!(
PromptSelector::label("production").as_key(),
"label:production"
);
assert_eq!(PromptSelector::version("7").as_key(), "version:7");
assert_ne!(
PromptSelector::label("7").as_key(),
PromptSelector::version("7").as_key()
);
assert_eq!(PromptSelector::default(), PromptSelector::Latest);
assert_eq!(PromptSelector::Latest.to_string(), "latest");
assert!(PromptSelector::version("7").is_pinned());
assert!(!PromptSelector::label("production").is_pinned());
assert!(!PromptSelector::Latest.is_pinned());
}
#[test]
fn every_error_variant_renders_identifiers_and_codes_only() {
let planted = "sk-live-0123456789abcdef";
let errors = [
PromptError::NotFound {
name: PromptName::from("interpret.system"),
},
PromptError::VersionNotFound {
name: PromptName::from("interpret.system"),
version: PromptVersion::from("7"),
},
PromptError::LabelNotFound {
name: PromptName::from("interpret.system"),
label: "production".to_owned(),
},
PromptError::Unauthorized,
PromptError::Forbidden {
name: PromptName::from("interpret.system"),
},
PromptError::RateLimited,
PromptError::Transport { code: "timeout" },
PromptError::Malformed { code: "not_json" },
PromptError::Unsupported {
reason: "chat_prompt",
},
PromptError::InconsistentReference {
name: PromptName::from("interpret.system"),
},
];
for error in errors {
let rendered = format!("{error} {error:?}");
assert!(!rendered.contains(planted), "{rendered}");
assert!(!error.code().is_empty());
}
}
#[test]
fn only_a_transport_fault_and_a_rate_limit_are_worth_repeating() {
assert!(PromptError::RateLimited.is_transient());
assert!(PromptError::Transport { code: "connect" }.is_transient());
assert!(!PromptError::Unauthorized.is_transient());
assert!(
!PromptError::NotFound {
name: PromptName::from("x"),
}
.is_transient()
);
}
}