use std::path::{Path, PathBuf};
use bytes::Bytes;
use serde::{Deserialize, Serialize};
use crate::{Completion, Reasoning, ToolCall};
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
#[non_exhaustive]
pub enum Role {
User,
Assistant,
Tool,
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
#[non_exhaustive]
pub struct Message {
pub role: Role,
pub parts: Vec<Part>,
}
impl Message {
pub fn new(role: Role) -> Self {
Self {
role,
parts: Vec::new(),
}
}
pub fn user(text: impl Into<String>) -> Self {
Self::new(Role::User).with_text(text.into())
}
pub fn assistant(text: impl Into<String>) -> Self {
Self::new(Role::Assistant).with_text(text.into())
}
pub fn tool_result(result: ToolResult) -> Self {
Self::new(Role::Tool).with(result)
}
pub fn with(mut self, part: impl Into<Part>) -> Self {
self.parts.push(part.into());
self
}
fn with_text(self, text: String) -> Self {
if text.is_empty() {
self
} else {
self.with(text)
}
}
}
impl From<Completion> for Message {
fn from(done: Completion) -> Self {
let mut parts = Vec::with_capacity(done.reasoning.len() + 1 + done.calls.len());
parts.extend(done.reasoning.into_iter().map(Part::Reasoning));
if !done.text.is_empty() {
parts.push(Part::Text(done.text));
}
parts.extend(done.calls.into_iter().map(Part::ToolCall));
Self {
role: Role::Assistant,
parts,
}
}
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
#[non_exhaustive]
pub enum Part {
Text(String),
Image(Image),
File(TextFile),
Reasoning(Reasoning),
ToolCall(ToolCall),
ToolResult(ToolResult),
}
impl From<String> for Part {
fn from(text: String) -> Self {
Self::Text(text)
}
}
impl From<&str> for Part {
fn from(text: &str) -> Self {
Self::Text(text.to_owned())
}
}
impl From<Image> for Part {
fn from(image: Image) -> Self {
Self::Image(image)
}
}
impl From<TextFile> for Part {
fn from(file: TextFile) -> Self {
Self::File(file)
}
}
impl From<Reasoning> for Part {
fn from(reasoning: Reasoning) -> Self {
Self::Reasoning(reasoning)
}
}
impl From<ToolCall> for Part {
fn from(call: ToolCall) -> Self {
Self::ToolCall(call)
}
}
impl From<ToolResult> for Part {
fn from(result: ToolResult) -> Self {
Self::ToolResult(result)
}
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
#[non_exhaustive]
pub enum Source {
Path(PathBuf),
Bytes(#[serde(with = "base64_bytes")] Bytes),
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
#[non_exhaustive]
pub struct Image {
#[serde(flatten)]
pub source: Source,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub media_type: Option<String>,
}
impl Image {
pub fn path(path: impl Into<PathBuf>) -> Self {
let path = path.into();
let media_type = media_type_of(&path);
Self {
source: Source::Path(path),
media_type,
}
}
pub fn bytes(data: impl Into<Bytes>, media_type: impl Into<String>) -> Self {
Self {
source: Source::Bytes(data.into()),
media_type: Some(media_type.into()),
}
}
pub fn media_type(mut self, media_type: impl Into<String>) -> Self {
self.media_type = Some(media_type.into());
self
}
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
#[non_exhaustive]
pub struct TextFile {
pub name: String,
#[serde(flatten)]
pub source: Source,
}
impl TextFile {
pub fn path(path: impl Into<PathBuf>) -> Self {
let path = path.into();
let name = path
.file_name()
.map(|name| name.to_string_lossy().into_owned())
.unwrap_or_else(|| path.display().to_string());
Self {
name,
source: Source::Path(path),
}
}
pub fn text(name: impl Into<String>, text: impl Into<String>) -> Self {
Self {
name: name.into(),
source: Source::Bytes(Bytes::from(text.into())),
}
}
pub fn name(mut self, name: impl Into<String>) -> Self {
self.name = name.into();
self
}
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
#[non_exhaustive]
pub struct ToolResult {
pub call_id: String,
pub content: String,
}
impl ToolResult {
pub fn new(call_id: impl Into<String>, content: impl Into<String>) -> Self {
Self {
call_id: call_id.into(),
content: content.into(),
}
}
}
fn media_type_of(path: &Path) -> Option<String> {
let extension = path.extension()?.to_str()?.to_ascii_lowercase();
let media_type = match extension.as_str() {
"png" => "image/png",
"jpg" | "jpeg" => "image/jpeg",
"gif" => "image/gif",
"webp" => "image/webp",
_ => return None,
};
Some(media_type.to_owned())
}
mod base64_bytes {
use base64::{Engine, engine::general_purpose::STANDARD};
use bytes::Bytes;
use serde::{Deserialize, Deserializer, Serializer, de::Error};
pub(super) fn serialize<S: Serializer>(
value: &Bytes,
serializer: S,
) -> Result<S::Ok, S::Error> {
serializer.serialize_str(&STANDARD.encode(value))
}
pub(super) fn deserialize<'de, D: Deserializer<'de>>(
deserializer: D,
) -> Result<Bytes, D::Error> {
let text = String::deserialize(deserializer)?;
STANDARD
.decode(text)
.map(Bytes::from)
.map_err(D::Error::custom)
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::{FinishReason, ReasoningSource};
#[test]
fn parts_keep_the_order_given() {
let msg = Message::user("What changed?")
.with(Image::path("before.png"))
.with(TextFile::path("dir/diff.patch"))
.with(Image::path("after.JPG"));
let kinds: Vec<_> = msg
.parts
.iter()
.map(|p| match p {
Part::Text(_) => "text",
Part::Image(_) => "image",
Part::File(_) => "file",
_ => "other",
})
.collect();
assert_eq!(kinds, ["text", "image", "file", "image"]);
}
#[test]
fn empty_text_adds_no_part() {
let msg = Message::user("").with(Image::path("a.png"));
assert_eq!(msg.parts.len(), 1);
assert!(Message::assistant("").parts.is_empty());
}
#[test]
fn attachments_know_their_media_type_and_name_without_io() {
assert_eq!(
Image::path("a.png").media_type.as_deref(),
Some("image/png")
);
assert_eq!(
Image::path("a.JPEG").media_type.as_deref(),
Some("image/jpeg")
);
assert_eq!(Image::path("a.heic").media_type, None);
assert_eq!(
Image::path("a.heic")
.media_type("image/heic")
.media_type
.as_deref(),
Some("image/heic")
);
assert_eq!(TextFile::path("src/diff.patch").name, "diff.patch");
assert_eq!(
TextFile::path("src/diff.patch").name("changes").name,
"changes"
);
}
#[test]
fn a_completion_becomes_a_message_with_nothing_lost() {
let mut done = Completion::new(FinishReason::ToolCalls);
done.text = "Let me look.".into();
done.reasoning = vec![Reasoning::new(ReasoningSource::ReasoningContent, "hmm")];
done.calls = vec![ToolCall::new("call-a", "lookup", r#"{"value":1}"#)];
let msg = Message::from(done);
assert_eq!(msg.role, Role::Assistant);
assert_eq!(
msg.parts,
[
Part::Reasoning(Reasoning::new(ReasoningSource::ReasoningContent, "hmm")),
Part::Text("Let me look.".into()),
Part::ToolCall(ToolCall::new("call-a", "lookup", r#"{"value":1}"#)),
]
);
}
}