use crate::prompt::Error;
use brazen::{Content, ContentKind, Delta, Event};
use std::fs::File;
use std::io::{Seek, SeekFrom, Write};
use std::path::{Path, PathBuf};
enum Block {
Text(String),
Thinking(String),
ToolUse {
id: String,
name: String,
json: String,
},
Skip,
}
pub(super) struct StagingWriter {
file: File,
len: u64,
sep: &'static str,
seg_len: u64,
seg_sep: &'static str,
cur: Option<Block>,
}
impl StagingWriter {
pub(super) fn create(path: &Path) -> Result<Self, Error> {
let mut file = File::create(path)?;
file.write_all(b"[")?;
Ok(Self {
file,
len: 1,
sep: "",
seg_len: 1,
seg_sep: "",
cur: None,
})
}
pub(super) fn begin_segment(&mut self) {
self.seg_len = self.len;
self.seg_sep = self.sep;
self.cur = None;
}
pub(super) fn feed(&mut self, event: &Event) -> Result<(), Error> {
match event {
Event::ContentStart { kind, .. } => self.cur = Some(open_block(kind)),
Event::ContentDelta { delta, .. } => self.on_delta(delta),
Event::ContentStop { .. } => self.finalize()?,
_ => {}
}
Ok(())
}
fn on_delta(&mut self, delta: &Delta) {
match (&mut self.cur, delta) {
(Some(Block::Text(s)), Delta::TextDelta(t)) => s.push_str(t),
(Some(Block::Thinking(s)), Delta::ThinkingDelta(t)) => s.push_str(t),
(Some(Block::ToolUse { json, .. }), Delta::JsonDelta(t)) => json.push_str(t),
_ => {}
}
}
fn finalize(&mut self) -> Result<(), Error> {
let block = match self.cur.take() {
Some(b) => b,
None => return Ok(()),
};
let content = match block {
Block::Text(text) => Content::Text(text),
Block::Thinking(text) => Content::Thinking {
text,
signature: None,
id: None,
encrypted_content: None,
},
Block::ToolUse { id, name, json } => {
let input = if json.is_empty() {
serde_json::Value::Object(serde_json::Map::new())
} else {
serde_json::from_str(&json).map_err(Error::AdapterJson)?
};
Content::ToolUse {
id,
name,
input,
signature: None,
}
}
Block::Skip => return Ok(()),
};
let bytes = serde_json::to_vec(&content).expect("Content serializes");
self.file.write_all(self.sep.as_bytes())?;
self.file.write_all(&bytes)?;
self.len += self.sep.len() as u64 + bytes.len() as u64;
self.sep = ",";
Ok(())
}
pub(super) fn truncate_segment(&mut self) -> Result<(), Error> {
self.file.set_len(self.seg_len)?;
self.file.seek(SeekFrom::Start(self.seg_len))?;
self.len = self.seg_len;
self.sep = self.seg_sep;
self.cur = None;
Ok(())
}
pub(super) fn seal(mut self) -> Result<(), Error> {
self.finalize()?;
self.file.write_all(b"]")?;
self.file.flush()?;
Ok(())
}
}
fn open_block(kind: &ContentKind) -> Block {
match kind {
ContentKind::Text {} => Block::Text(String::new()),
ContentKind::Thinking { .. } => Block::Thinking(String::new()),
ContentKind::ToolUse { id, name } => Block::ToolUse {
id: id.clone(),
name: name.clone(),
json: String::new(),
},
_ => Block::Skip,
}
}
pub(super) fn staging_path_for(response_path: &Path) -> PathBuf {
response_path.with_file_name(crate::prompt::step::STAGING_FILE)
}
#[cfg(test)]
mod tests;