use serde::{Deserialize, Serialize};
use crate::capabilities::{CapabilityRequirements, StructuredOutputCapability};
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, PartialOrd, Ord, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
#[non_exhaustive]
pub enum ModelPurpose {
OfflineEvaluate,
Segment,
Coverage,
Route,
Locate,
Extract,
Verify,
QuestionFrame,
CrossCheck,
Respects,
Investigate,
Acknowledge,
Answer,
Review,
Progress,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Default, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
pub enum SafetyMode {
#[default]
Default,
UnsafeExperimental,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
pub enum LogDetail {
Metadata,
Redacted,
Full,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Serialize, Deserialize)]
pub struct LoggingPolicy {
pub prompt: LogDetail,
pub output: LogDetail,
pub retain_raw_bodies: bool,
}
impl LoggingPolicy {
pub const METADATA_ONLY: Self = Self {
prompt: LogDetail::Metadata,
output: LogDetail::Metadata,
retain_raw_bodies: false,
};
pub const REDACTED_OUTPUT: Self = Self {
prompt: LogDetail::Metadata,
output: LogDetail::Redacted,
retain_raw_bodies: false,
};
pub const FULL: Self = Self {
prompt: LogDetail::Full,
output: LogDetail::Full,
retain_raw_bodies: false,
};
#[must_use]
pub const fn logs_content(&self) -> bool {
!matches!(self.prompt, LogDetail::Metadata) || !matches!(self.output, LogDetail::Metadata)
}
}
pub const MUTATION_SAFE_STRUCTURED_OUTPUT: [StructuredOutputCapability; 3] = [
StructuredOutputCapability::NativeJsonSchema,
StructuredOutputCapability::NativeFunctionSchema,
StructuredOutputCapability::GrammarConstrained,
];
pub const READ_ONLY_STRUCTURED_OUTPUT: [StructuredOutputCapability; 4] = [
StructuredOutputCapability::NativeJsonSchema,
StructuredOutputCapability::NativeFunctionSchema,
StructuredOutputCapability::GrammarConstrained,
StructuredOutputCapability::JsonObject,
];
impl ModelPurpose {
pub const ALL: [Self; 15] = [
Self::OfflineEvaluate,
Self::Segment,
Self::Coverage,
Self::Route,
Self::Locate,
Self::Extract,
Self::Verify,
Self::QuestionFrame,
Self::CrossCheck,
Self::Respects,
Self::Investigate,
Self::Acknowledge,
Self::Answer,
Self::Review,
Self::Progress,
];
#[must_use]
pub const fn is_understanding(self) -> bool {
matches!(
self,
Self::Segment
| Self::Coverage
| Self::Route
| Self::Locate
| Self::Extract
| Self::Verify
| Self::QuestionFrame
| Self::CrossCheck
| Self::Respects
| Self::Investigate
)
}
#[must_use]
pub const fn as_str(self) -> &'static str {
match self {
Self::OfflineEvaluate => "offline_evaluate",
Self::Segment => "segment",
Self::Coverage => "coverage",
Self::Route => "route",
Self::Locate => "locate",
Self::Extract => "extract",
Self::Verify => "verify",
Self::QuestionFrame => "question_frame",
Self::CrossCheck => "cross_check",
Self::Respects => "respects",
Self::Investigate => "investigate",
Self::Acknowledge => "acknowledge",
Self::Answer => "answer",
Self::Review => "review",
Self::Progress => "progress",
}
}
#[must_use]
pub const fn is_critical(self) -> bool {
self.is_understanding()
}
#[must_use]
pub fn requirements(self) -> CapabilityRequirements {
self.requirements_in(SafetyMode::Default)
}
#[must_use]
pub fn requirements_in(self, mode: SafetyMode) -> CapabilityRequirements {
let structured_output: Vec<StructuredOutputCapability> = match (self, mode) {
(_, SafetyMode::UnsafeExperimental) | (Self::OfflineEvaluate, _) => Vec::new(),
(
Self::Investigate
| Self::Acknowledge
| Self::Answer
| Self::Review
| Self::Progress,
SafetyMode::Default,
) => READ_ONLY_STRUCTURED_OUTPUT.to_vec(),
(_, SafetyMode::Default) => MUTATION_SAFE_STRUCTURED_OUTPUT.to_vec(),
};
CapabilityRequirements {
structured_output,
needs_tools: false,
needs_streaming: false,
min_context_tokens: None,
needs_vision: false,
needs_documents: false,
}
}
#[must_use]
pub const fn logging_policy(self) -> LoggingPolicy {
match self {
Self::Segment
| Self::Coverage
| Self::Route
| Self::Locate
| Self::Extract
| Self::Verify
| Self::QuestionFrame
| Self::CrossCheck
| Self::Respects
| Self::Investigate => LoggingPolicy::REDACTED_OUTPUT,
Self::Acknowledge | Self::Answer | Self::Review | Self::Progress => {
LoggingPolicy::METADATA_ONLY
}
Self::OfflineEvaluate => LoggingPolicy::FULL,
}
}
}
impl std::fmt::Display for ModelPurpose {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.write_str(self.as_str())
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::capabilities::ProviderCapabilities;
fn caps(structured: StructuredOutputCapability) -> ProviderCapabilities {
ProviderCapabilities::minimal().with_structured_output(structured)
}
#[test]
fn understanding_rejects_prompt_only_json_object_and_none() {
let requirements = ModelPurpose::Extract.requirements();
for unsafe_transport in [
StructuredOutputCapability::PromptOnly,
StructuredOutputCapability::JsonObject,
StructuredOutputCapability::None,
] {
assert!(requirements.satisfied_by(&caps(unsafe_transport)).is_err());
}
for safe in MUTATION_SAFE_STRUCTURED_OUTPUT {
assert!(requirements.satisfied_by(&caps(safe)).is_ok());
}
}
#[test]
fn unsafe_experimental_is_an_explicit_opt_in() {
let requirements = ModelPurpose::Extract.requirements_in(SafetyMode::UnsafeExperimental);
assert!(requirements.structured_output.is_empty());
assert!(
requirements
.satisfied_by(&caps(StructuredOutputCapability::PromptOnly))
.is_ok()
);
}
#[test]
fn a_read_only_task_accepts_json_object_but_not_prompt_only() {
for purpose in [ModelPurpose::Investigate, ModelPurpose::Acknowledge] {
let requirements = purpose.requirements();
assert_eq!(
requirements.structured_output,
READ_ONLY_STRUCTURED_OUTPUT.to_vec()
);
assert!(
requirements
.satisfied_by(&caps(StructuredOutputCapability::PromptOnly))
.is_err()
);
}
}
#[test]
fn narration_logs_metadata_only_and_is_not_critical() {
assert!(!ModelPurpose::Acknowledge.logging_policy().logs_content());
assert!(!ModelPurpose::Acknowledge.is_critical());
assert!(ModelPurpose::Extract.is_critical());
assert!(
ModelPurpose::OfflineEvaluate
.requirements()
.structured_output
.is_empty()
);
}
#[test]
fn labels_are_unique_and_never_retain_raw_bodies() {
let mut labels: Vec<&str> = ModelPurpose::ALL.iter().map(|p| p.as_str()).collect();
labels.sort_unstable();
labels.dedup();
assert_eq!(labels.len(), ModelPurpose::ALL.len());
for purpose in ModelPurpose::ALL {
assert!(!purpose.logging_policy().retain_raw_bodies);
let json = serde_json::to_string(&purpose).unwrap();
assert_eq!(json, format!("\"{}\"", purpose.as_str()));
}
}
}