use ferrin_spec::ProviderOptions;
use serde::Deserialize;
use serde::Serialize;
use crate::part::AssistantPart;
use crate::part::ToolApprovalResponse;
use crate::part::ToolPart;
use crate::part::UserPart;
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
#[serde(tag = "role", rename_all = "lowercase")]
#[non_exhaustive]
pub enum Message {
System(SystemMessage),
User(UserMessage),
Assistant(AssistantMessage),
Tool(ToolMessage),
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Serialize, Deserialize)]
#[serde(rename_all = "lowercase")]
#[non_exhaustive]
pub enum Role {
System,
User,
Assistant,
Tool,
}
impl Role {
#[must_use]
pub fn as_str(self) -> &'static str {
match self {
Self::System => "system",
Self::User => "user",
Self::Assistant => "assistant",
Self::Tool => "tool",
}
}
}
impl std::fmt::Display for Role {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.write_str(self.as_str())
}
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub struct SystemMessage {
pub content: String,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub provider_options: Option<ProviderOptions>,
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub struct UserMessage {
pub content: UserContent,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub provider_options: Option<ProviderOptions>,
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub struct AssistantMessage {
pub content: AssistantContent,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub provider_options: Option<ProviderOptions>,
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub struct ToolMessage {
pub content: Vec<ToolPart>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub provider_options: Option<ProviderOptions>,
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
#[serde(untagged)]
pub enum UserContent {
Text(String),
Parts(Vec<UserPart>),
}
impl UserContent {
#[must_use]
pub fn is_empty(&self) -> bool {
match self {
Self::Text(text) => text.is_empty(),
Self::Parts(parts) => parts.is_empty(),
}
}
#[must_use]
pub fn as_text(&self) -> Option<&str> {
match self {
Self::Text(text) => Some(text),
Self::Parts(_) => None,
}
}
#[must_use]
pub fn into_parts(self) -> Vec<UserPart> {
match self {
Self::Text(text) => vec![UserPart::text(text)],
Self::Parts(parts) => parts,
}
}
}
impl From<&str> for UserContent {
fn from(text: &str) -> Self {
Self::Text(text.to_owned())
}
}
impl From<String> for UserContent {
fn from(text: String) -> Self {
Self::Text(text)
}
}
impl From<Vec<UserPart>> for UserContent {
fn from(parts: Vec<UserPart>) -> Self {
Self::Parts(parts)
}
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
#[serde(untagged)]
pub enum AssistantContent {
Text(String),
Parts(Vec<AssistantPart>),
}
impl AssistantContent {
#[must_use]
pub fn is_empty(&self) -> bool {
match self {
Self::Text(text) => text.is_empty(),
Self::Parts(parts) => parts.is_empty(),
}
}
#[must_use]
pub fn as_text(&self) -> Option<&str> {
match self {
Self::Text(text) => Some(text),
Self::Parts(_) => None,
}
}
#[must_use]
pub fn into_parts(self) -> Vec<AssistantPart> {
match self {
Self::Text(text) => vec![AssistantPart::text(text)],
Self::Parts(parts) => parts,
}
}
#[must_use]
pub fn as_parts(&self) -> Option<&[AssistantPart]> {
match self {
Self::Text(_) => None,
Self::Parts(parts) => Some(parts),
}
}
}
impl From<&str> for AssistantContent {
fn from(text: &str) -> Self {
Self::Text(text.to_owned())
}
}
impl From<String> for AssistantContent {
fn from(text: String) -> Self {
Self::Text(text)
}
}
impl From<Vec<AssistantPart>> for AssistantContent {
fn from(parts: Vec<AssistantPart>) -> Self {
Self::Parts(parts)
}
}
impl Message {
#[must_use]
pub fn system(content: impl Into<String>) -> Self {
Self::System(SystemMessage {
content: content.into(),
provider_options: None,
})
}
#[must_use]
pub fn user(text: impl Into<String>) -> Self {
Self::User(UserMessage {
content: UserContent::Text(text.into()),
provider_options: None,
})
}
#[must_use]
pub fn user_parts(parts: impl IntoIterator<Item = impl Into<UserPart>>) -> Self {
Self::User(UserMessage {
content: UserContent::Parts(parts.into_iter().map(Into::into).collect()),
provider_options: None,
})
}
#[must_use]
pub fn assistant(text: impl Into<String>) -> Self {
Self::Assistant(AssistantMessage {
content: AssistantContent::Text(text.into()),
provider_options: None,
})
}
#[must_use]
pub fn assistant_parts(parts: impl IntoIterator<Item = impl Into<AssistantPart>>) -> Self {
Self::Assistant(AssistantMessage {
content: AssistantContent::Parts(parts.into_iter().map(Into::into).collect()),
provider_options: None,
})
}
#[must_use]
pub fn tool(parts: impl IntoIterator<Item = impl Into<ToolPart>>) -> Self {
Self::Tool(ToolMessage {
content: parts.into_iter().map(Into::into).collect(),
provider_options: None,
})
}
#[must_use]
pub fn role(&self) -> Role {
match self {
Self::System(_) => Role::System,
Self::User(_) => Role::User,
Self::Assistant(_) => Role::Assistant,
Self::Tool(_) => Role::Tool,
}
}
#[must_use]
pub fn is_empty(&self) -> bool {
match self {
Self::System(message) => message.content.is_empty(),
Self::User(message) => message.content.is_empty(),
Self::Assistant(message) => message.content.is_empty(),
Self::Tool(message) => message.content.is_empty(),
}
}
#[must_use]
pub fn provider_options(&self) -> Option<&ProviderOptions> {
match self {
Self::System(message) => message.provider_options.as_ref(),
Self::User(message) => message.provider_options.as_ref(),
Self::Assistant(message) => message.provider_options.as_ref(),
Self::Tool(message) => message.provider_options.as_ref(),
}
}
#[must_use]
pub fn with_provider_options(mut self, options: ProviderOptions) -> Self {
let slot = match &mut self {
Self::System(message) => &mut message.provider_options,
Self::User(message) => &mut message.provider_options,
Self::Assistant(message) => &mut message.provider_options,
Self::Tool(message) => &mut message.provider_options,
};
*slot = Some(options);
self
}
#[must_use]
pub fn as_system(&self) -> Option<&SystemMessage> {
match self {
Self::System(message) => Some(message),
_ => None,
}
}
#[must_use]
pub fn as_user(&self) -> Option<&UserMessage> {
match self {
Self::User(message) => Some(message),
_ => None,
}
}
#[must_use]
pub fn as_assistant(&self) -> Option<&AssistantMessage> {
match self {
Self::Assistant(message) => Some(message),
_ => None,
}
}
#[must_use]
pub fn as_tool(&self) -> Option<&ToolMessage> {
match self {
Self::Tool(message) => Some(message),
_ => None,
}
}
}
impl From<SystemMessage> for Message {
fn from(message: SystemMessage) -> Self {
Self::System(message)
}
}
impl From<UserMessage> for Message {
fn from(message: UserMessage) -> Self {
Self::User(message)
}
}
impl From<AssistantMessage> for Message {
fn from(message: AssistantMessage) -> Self {
Self::Assistant(message)
}
}
impl From<ToolMessage> for Message {
fn from(message: ToolMessage) -> Self {
Self::Tool(message)
}
}
pub trait MessagesExt {
fn push_approval_response(&mut self, response: ToolApprovalResponse);
fn pending_approval_requests(&self) -> Vec<&crate::part::ToolApprovalRequest>;
}
impl MessagesExt for Vec<Message> {
fn push_approval_response(&mut self, response: ToolApprovalResponse) {
if let Some(Message::Tool(tool)) = self.last_mut() {
tool.content.push(ToolPart::ToolApprovalResponse(response));
} else {
self.push(Message::tool([ToolPart::ToolApprovalResponse(response)]));
}
}
fn pending_approval_requests(&self) -> Vec<&crate::part::ToolApprovalRequest> {
let answered: std::collections::HashSet<&ferrin_spec::ApprovalId> = self
.iter()
.filter_map(Message::as_tool)
.flat_map(|tool| tool.content.iter())
.filter_map(ToolPart::as_tool_approval_response)
.map(|response| &response.approval_id)
.collect();
self.iter()
.filter_map(Message::as_assistant)
.filter_map(|assistant| assistant.content.as_parts())
.flatten()
.filter_map(AssistantPart::as_tool_approval_request)
.filter(|request| !answered.contains(&request.approval_id))
.collect()
}
}