use std::collections::BTreeSet;
use serde::{Deserialize, Serialize};
use thiserror::Error;
use crate::{PromptAuthority, PromptModuleId, PromptSegment};
pub const MAX_MODULE_SEGMENTS: usize = 256;
#[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord, Hash, Serialize, Deserialize)]
#[serde(transparent)]
pub struct PromptPriority(i16);
impl PromptPriority {
#[must_use]
pub const fn new(value: i16) -> Self {
Self(value)
}
#[must_use]
pub const fn get(self) -> i16 {
self.0
}
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize)]
#[serde(rename_all = "camelCase")]
pub struct PromptModule {
id: PromptModuleId,
authority: PromptAuthority,
priority: PromptPriority,
segments: Vec<PromptSegment>,
}
impl PromptModule {
pub fn new(
id: PromptModuleId,
authority: PromptAuthority,
priority: PromptPriority,
segments: Vec<PromptSegment>,
) -> Result<Self, ModuleError> {
if segments.is_empty() || segments.len() > MAX_MODULE_SEGMENTS {
return Err(ModuleError::InvalidSegmentCount);
}
let mut ids = BTreeSet::new();
if segments.iter().any(|segment| !ids.insert(segment.id())) {
return Err(ModuleError::DuplicateSegmentId);
}
Ok(Self {
id,
authority,
priority,
segments,
})
}
#[must_use]
pub const fn id(&self) -> &PromptModuleId {
&self.id
}
#[must_use]
pub const fn authority(&self) -> PromptAuthority {
self.authority
}
#[must_use]
pub const fn priority(&self) -> PromptPriority {
self.priority
}
#[must_use]
pub fn segments(&self) -> &[PromptSegment] {
&self.segments
}
}
#[derive(Deserialize)]
#[serde(rename_all = "camelCase")]
struct RawPromptModule {
id: PromptModuleId,
authority: PromptAuthority,
priority: PromptPriority,
segments: Vec<PromptSegment>,
}
impl<'de> Deserialize<'de> for PromptModule {
fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
where
D: serde::Deserializer<'de>,
{
let raw = RawPromptModule::deserialize(deserializer)?;
Self::new(raw.id, raw.authority, raw.priority, raw.segments)
.map_err(serde::de::Error::custom)
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Error)]
pub enum ModuleError {
#[error("prompt module segment count is invalid")]
InvalidSegmentCount,
#[error("prompt module contains a duplicate segment ID")]
DuplicateSegmentId,
}