mod stderr;
use super::staging::{StagingWriter, staging_path_for};
use super::stop_signal;
use crate::config::RetryConfig;
use crate::prompt::Error;
use crate::prompt::adapter::AdapterRunner;
use brazen::{
CanonicalError, CanonicalRequest, Content, EVENT_SCHEMA_VERSION, Event, Message, Tool,
};
use std::ffi::OsString;
use std::fs::File;
use std::io::Write;
use std::path::Path;
use std::sync::atomic::AtomicBool;
use std::time::Duration;
enum SegmentOutcome {
Complete { handshake_v: Option<u8> },
Failed(CanonicalError),
HalfStream { stderr_tail: String },
}
pub trait Sleeper {
fn sleep(&self, dur: Duration);
}
#[derive(Debug, Clone, Copy)]
pub struct RealSleeper;
impl Sleeper for RealSleeper {
fn sleep(&self, dur: Duration) {
std::thread::sleep(dur);
}
}
pub(super) struct ModelCall<'a> {
pub(super) adapter: &'a dyn AdapterRunner,
pub(super) sleeper: &'a dyn Sleeper,
pub(super) binary: &'a OsString,
pub(super) provider_row: &'a str,
pub(super) retry: RetryConfig,
pub(super) stop: &'a AtomicBool,
pub(super) expect_handshake: bool,
}
pub(super) fn build_request(
model_id: &str,
system: &str,
messages: Vec<Message>,
tools: Vec<Tool>,
max_tokens: u32,
) -> CanonicalRequest {
CanonicalRequest {
model: model_id.to_string(),
system: Some(vec![Content::Text(system.to_string())]),
messages,
tools,
max_tokens: Some(max_tokens),
..CanonicalRequest::default()
}
}
pub(super) fn run(
call: &ModelCall<'_>,
request_bytes: &[u8],
response_path: &Path,
) -> Result<(), Error> {
if let Some(parent) = response_path.parent() {
std::fs::create_dir_all(parent)?;
}
let mut response_file = File::create(response_path)?;
let stderr_path = response_path.with_file_name(crate::prompt::step::STDERR_FILE);
let mut stderr_file = File::create(&stderr_path)?;
let mut staging = StagingWriter::create(&staging_path_for(response_path))?;
let args = ["--json", "--provider", call.provider_row];
let max = call.retry.max_attempts.max(1);
let mut attempt = 1;
loop {
staging.begin_segment();
let outcome = run_attempt(
call,
&args,
request_bytes,
&mut response_file,
&mut stderr_file,
&mut staging,
)?;
match outcome {
SegmentOutcome::Complete { handshake_v } => {
check_handshake(call.expect_handshake, handshake_v)?;
staging.seal()?;
drop(response_file);
return Ok(());
}
SegmentOutcome::Failed(err) => {
staging.truncate_segment()?;
if err.retryable() && attempt < max && !stop_signal::stopped(call.stop) {
let d = call.retry.backoff.delay(attempt, err.retry_after_seconds);
call.sleeper.sleep(d);
if !stop_signal::stopped(call.stop) {
attempt += 1;
continue;
}
}
drop(response_file);
return Err(Error::AdapterError {
kind: format!("{:?}", err.kind),
message: err.message,
});
}
SegmentOutcome::HalfStream { stderr_tail } => {
drop(response_file);
return Err(Error::AdapterHalfStream {
stderr_log: stderr_path,
tail: stderr_tail,
});
}
}
}
}
fn run_attempt(
call: &ModelCall<'_>,
args: &[&str],
request_bytes: &[u8],
response_file: &mut File,
stderr_file: &mut File,
staging: &mut StagingWriter,
) -> Result<SegmentOutcome, Error> {
let mut feed_err: Option<serde_json::Error> = None;
let mut staging_err: Option<Error> = None;
let mut error: Option<CanonicalError> = None;
let mut ended = false;
let mut handshake_v: Option<u8> = None;
let stderr = call
.adapter
.run(call.binary, args, request_bytes, &mut |line| {
response_file.write_all(line)?;
response_file.write_all(b"\n")?;
if feed_err.is_none() && staging_err.is_none() && !ended {
match serde_json::from_slice::<Event>(line) {
Ok(event) => {
match &event {
Event::MessageStart { v, .. } => handshake_v = Some(*v),
Event::Error(e) => error = Some(e.clone()),
Event::End => ended = true,
_ => {}
}
if let Err(e) = staging.feed(&event) {
staging_err = Some(e);
}
}
Err(e) => feed_err = Some(e),
}
}
Ok(())
})
.map_err(|e| crate::prompt::adapter::spawn_error(call.binary, e))?;
stderr_file.write_all(&stderr)?;
if let Some(e) = feed_err {
return Err(Error::AdapterJson(e));
}
if let Some(e) = staging_err {
return Err(e);
}
if let Some(err) = error {
return Ok(SegmentOutcome::Failed(err));
}
if !ended {
return Ok(SegmentOutcome::HalfStream {
stderr_tail: stderr::tail(&stderr),
});
}
Ok(SegmentOutcome::Complete { handshake_v })
}
fn check_handshake(expect: bool, handshake_v: Option<u8>) -> Result<(), Error> {
if expect && handshake_v != Some(EVENT_SCHEMA_VERSION) {
return Err(Error::HandshakeMismatch {
found: handshake_v,
expected: EVENT_SCHEMA_VERSION,
});
}
Ok(())
}
#[cfg(test)]
mod tests;