use chrono::{DateTime, Utc};
use serde::{Deserialize, Serialize};
use serde_json::Map;
use starweaver_core::{ConversationId, RunId};
use starweaver_usage::Usage;
use super::{
ContentPart, FinishReason, Metadata, ModelRequestPart, ModelResponsePart, ProviderInfo,
ToolCallPart,
};
#[derive(Clone, Debug, Deserialize, Eq, PartialEq, Serialize)]
#[serde(tag = "kind", rename_all = "snake_case")]
pub enum ModelMessage {
Request(ModelRequest),
Response(ModelResponse),
}
#[derive(Clone, Debug, Deserialize, Eq, PartialEq, Serialize)]
pub struct ModelRequest {
pub parts: Vec<ModelRequestPart>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub timestamp: Option<DateTime<Utc>>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub instructions: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub run_id: Option<RunId>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub conversation_id: Option<ConversationId>,
#[serde(default, skip_serializing_if = "Map::is_empty")]
pub metadata: Metadata,
}
impl ModelRequest {
#[must_use]
pub fn user_text(text: impl Into<String>) -> Self {
Self {
parts: vec![ModelRequestPart::UserPrompt {
content: vec![ContentPart::Text { text: text.into() }],
name: None,
metadata: Metadata::default(),
}],
timestamp: None,
instructions: None,
run_id: None,
conversation_id: None,
metadata: Metadata::default(),
}
}
}
#[derive(Clone, Debug, Deserialize, Eq, PartialEq, Serialize)]
pub struct ModelResponse {
pub parts: Vec<ModelResponsePart>,
#[serde(default)]
pub usage: Usage,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub model_name: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub provider: Option<ProviderInfo>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub finish_reason: Option<FinishReason>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub timestamp: Option<DateTime<Utc>>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub run_id: Option<RunId>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub conversation_id: Option<ConversationId>,
#[serde(default, skip_serializing_if = "Map::is_empty")]
pub metadata: Metadata,
}
impl ModelResponse {
#[must_use]
pub fn text(text: impl Into<String>) -> Self {
Self {
parts: vec![ModelResponsePart::Text { text: text.into() }],
usage: Usage::default(),
model_name: None,
provider: None,
finish_reason: None,
timestamp: None,
run_id: None,
conversation_id: None,
metadata: Metadata::default(),
}
}
#[must_use]
pub fn text_output(&self) -> String {
self.parts
.iter()
.filter_map(ModelResponsePart::text)
.collect::<Vec<_>>()
.join("")
}
#[must_use]
pub fn tool_calls(&self) -> Vec<ToolCallPart> {
self.parts
.iter()
.filter_map(ModelResponsePart::tool_call)
.cloned()
.collect()
}
}