use serde::{Deserialize, Serialize};
use serde_json::Value;
use crate::message::ModelResponse;
#[derive(Clone, Debug, Deserialize, Eq, PartialEq, Serialize)]
#[serde(tag = "kind", rename_all = "snake_case")]
pub enum ModelResponseStreamEvent {
PartStart(PartStart),
PartDelta(PartDelta),
PartEnd(PartEnd),
FinalResult(Box<ModelResponse>),
}
#[derive(Clone, Debug, Deserialize, Eq, PartialEq, Serialize)]
pub struct PartStart {
pub index: usize,
pub part_kind: String,
}
#[derive(Clone, Debug, Deserialize, Eq, PartialEq, Serialize)]
pub struct PartDelta {
pub index: usize,
#[serde(flatten)]
pub delta: StreamDelta,
}
impl PartDelta {
#[must_use]
pub fn text(index: usize, text: impl Into<String>) -> Self {
Self {
index,
delta: StreamDelta::Text { text: text.into() },
}
}
#[must_use]
pub fn thinking(index: usize, text: impl Into<String>) -> Self {
Self {
index,
delta: StreamDelta::Thinking { text: text.into() },
}
}
#[must_use]
pub fn as_text(&self) -> String {
match &self.delta {
StreamDelta::Text { text }
| StreamDelta::Thinking { text }
| StreamDelta::ToolCallName { name: text }
| StreamDelta::ToolCallArguments {
arguments_delta: text,
} => text.clone(),
StreamDelta::NativePayload { payload } | StreamDelta::FileMetadata { payload } => {
payload.to_string()
}
}
}
}
#[derive(Clone, Debug, Deserialize, Eq, PartialEq, Serialize)]
#[serde(tag = "delta_kind", rename_all = "snake_case")]
pub enum StreamDelta {
Text {
text: String,
},
Thinking {
text: String,
},
ToolCallName {
name: String,
},
ToolCallArguments {
arguments_delta: String,
},
NativePayload {
payload: Value,
},
FileMetadata {
payload: Value,
},
}
#[derive(Clone, Debug, Deserialize, Eq, PartialEq, Serialize)]
pub struct PartEnd {
pub index: usize,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub part_kind: Option<String>,
}
impl PartEnd {
#[must_use]
pub const fn new(index: usize) -> Self {
Self {
index,
part_kind: None,
}
}
#[must_use]
pub fn with_kind(index: usize, part_kind: impl Into<String>) -> Self {
Self {
index,
part_kind: Some(part_kind.into()),
}
}
}
#[derive(Clone, Copy, Debug, Default, Deserialize, Eq, PartialEq, Serialize)]
#[serde(rename_all = "snake_case")]
pub enum StreamLifecycle {
#[default]
Incomplete,
Complete,
Interrupted,
}
#[derive(Clone, Debug, Default, Deserialize, Eq, PartialEq, Serialize)]
pub struct ModelStreamState {
pub lifecycle: StreamLifecycle,
pub started_parts: usize,
pub ended_parts: usize,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub final_response: Option<Box<ModelResponse>>,
}
impl ModelStreamState {
pub fn apply(&mut self, event: &ModelResponseStreamEvent) {
match event {
ModelResponseStreamEvent::PartStart(_) => {
self.started_parts += 1;
}
ModelResponseStreamEvent::PartDelta(_) => {}
ModelResponseStreamEvent::PartEnd(_) => {
self.ended_parts += 1;
}
ModelResponseStreamEvent::FinalResult(response) => {
self.lifecycle = StreamLifecycle::Complete;
self.final_response = Some(response.clone());
}
}
}
pub const fn interrupt(&mut self) {
self.lifecycle = StreamLifecycle::Interrupted;
}
}