use serde::{
de::Error as _,
ser::{SerializeMap, SerializeSeq},
Deserialize, Deserializer, Serialize, Serializer,
};
use serde_json::value::Value;
#[derive(Clone, Debug, PartialEq, Eq)]
pub enum Message {
System(SystemMessage),
User(UserMessage),
Assistant(AssistantMessage),
Tool(ToolMessage),
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct SystemMessage {
pub content: String,
pub name: Option<String>,
}
impl SystemMessage {
pub fn new(content: String) -> Self {
Self {
content,
name: None,
}
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct UserMessage {
pub content: Content,
pub name: Option<String>,
}
impl UserMessage {
pub fn new(content: Content) -> Self {
Self {
content,
name: None,
}
}
#[cfg(test)]
pub fn new_from_str(content: &str) -> Self {
Self {
content: Content::Text(content.to_string()),
name: None,
}
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct AssistantMessage {
pub content: Content,
pub name: Option<String>,
pub refusal: Option<String>,
pub tool_calls: Option<Value>,
}
impl AssistantMessage {
pub fn new(content: Content) -> Self {
Self {
content,
name: None,
refusal: None,
tool_calls: None,
}
}
#[cfg(test)]
pub fn new_from_str(content: &str) -> Self {
Self {
content: Content::Text(content.to_string()),
name: None,
refusal: None,
tool_calls: None,
}
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct ToolMessage {
pub content: String,
pub tool_call_id: String,
}
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)
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
#[serde(rename_all = "lowercase")]
pub enum Role {
System,
User,
Assistant,
Tool,
}
#[derive(Debug, Clone, Deserialize, PartialEq, Eq)]
pub struct ResponseGenericMessage {
role: Role,
content: Option<String>,
name: Option<String>,
refusal: Option<String>,
tool_calls: Option<Value>,
tool_call_id: Option<String>,
reasoning: Option<String>,
images: Option<Vec<ImagePart>>,
}
#[derive(Debug, Clone, Serialize, PartialEq, Eq)]
pub struct RequestGenericMessage {
role: Role,
content: Content,
#[serde(skip_serializing_if = "Option::is_none")]
name: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
tool_call_id: Option<String>,
}
impl From<Message> for RequestGenericMessage {
fn from(message: Message) -> Self {
match message {
Message::System(m) => m.into(),
Message::User(m) => m.into(),
Message::Assistant(m) => m.into(),
Message::Tool(m) => m.into(),
}
}
}
impl From<SystemMessage> for RequestGenericMessage {
fn from(SystemMessage { content, name }: SystemMessage) -> Self {
Self {
role: Role::System,
content: Content::Text(content),
name,
tool_call_id: None,
}
}
}
impl From<UserMessage> for RequestGenericMessage {
fn from(UserMessage { content, name }: UserMessage) -> Self {
Self {
role: Role::User,
content,
name,
tool_call_id: None,
}
}
}
impl From<AssistantMessage> for RequestGenericMessage {
fn from(
AssistantMessage {
content,
name,
refusal: _,
tool_calls: _,
}: AssistantMessage,
) -> Self {
Self {
role: Role::Assistant,
content,
name,
tool_call_id: None,
}
}
}
impl From<ToolMessage> for RequestGenericMessage {
fn from(
ToolMessage {
content,
tool_call_id,
}: ToolMessage,
) -> Self {
Self {
role: Role::Tool,
content: Content::Text(content),
name: None,
tool_call_id: Some(tool_call_id),
}
}
}
#[derive(Debug, Clone, Eq, PartialEq, Serialize)]
pub struct ImagePart {
pub url: String,
pub detail: Option<String>,
}
#[derive(Debug, Clone, Eq, PartialEq, Serialize)]
pub struct FilePart {
pub file_data: String,
pub filename: Option<String>,
}
#[derive(Debug, Clone, Eq, PartialEq)]
pub enum ContentPart {
Text(String),
Image(ImagePart),
File(FilePart),
}
impl Serialize for ContentPart {
fn serialize<S>(&self, serializer: S) -> Result<S::Ok, S::Error>
where
S: Serializer,
{
match self {
ContentPart::Text(text) => {
let mut map = serializer.serialize_map(Some(2))?;
map.serialize_entry("type", "text")?;
map.serialize_entry("text", text)?;
map.end()
}
ContentPart::Image(image) => {
let mut map = serializer.serialize_map(Some(2))?;
map.serialize_entry("type", "image_url")?;
map.serialize_entry("image_url", image)?;
map.end()
}
ContentPart::File(file) => {
let mut map = serializer.serialize_map(Some(2))?;
map.serialize_entry("type", "file")?;
map.serialize_entry("file", file)?;
map.end()
}
}
}
}
#[derive(Debug, Clone, Eq, PartialEq)]
pub enum Content {
Text(String),
ContentParts(Vec<ContentPart>),
}
impl Serialize for Content {
fn serialize<S>(&self, serializer: S) -> Result<S::Ok, S::Error>
where
S: Serializer,
{
match self {
Content::Text(s) => serializer.serialize_str(s),
Content::ContentParts(parts) => {
let mut seq = serializer.serialize_seq(Some(parts.len()))?;
for part in parts {
seq.serialize_element(part)?;
}
seq.end()
}
}
}
}
#[derive(Debug, Clone, Eq, PartialEq, serde_query::Deserialize)]
pub struct IntermediateImagePart {
#[query(".type")]
ty: String,
#[query(".image_url.url")]
url: String,
}
impl<'de> Deserialize<'de> for ImagePart {
fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
where
D: Deserializer<'de>,
{
let part = IntermediateImagePart::deserialize(deserializer)?;
match part.ty.as_ref() {
"image_url" => Ok(ImagePart {
url: part.url,
detail: None,
}),
ty => Err(D::Error::custom(format!("unsupported type `{}`", ty))),
}
}
}
#[derive(Debug, thiserror::Error)]
pub enum Error {
#[error("Missing mandatory field `{0}`")]
MissingField(&'static str),
#[error("Expected role {0:?}, got {1:?}")]
RoleMismatch(Role, Role),
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct ResponseAssistantMessage {
pub content: Option<String>,
pub name: Option<String>,
pub refusal: Option<String>,
pub tool_calls: Option<Value>,
pub reasoning: Option<String>,
pub images: Option<Vec<ImagePart>>,
}
impl TryFrom<ResponseGenericMessage> for ResponseAssistantMessage {
type Error = Error;
fn try_from(
ResponseGenericMessage {
role,
content,
name,
refusal,
tool_calls,
tool_call_id: _,
reasoning,
images,
}: ResponseGenericMessage,
) -> Result<Self, Error> {
if role == Role::Assistant {
Ok(Self {
content,
name,
refusal,
tool_calls,
reasoning,
images,
})
} else {
Err(Error::RoleMismatch(Role::Assistant, role))
}
}
}