use anda_core::{
AgentOutput, BoxError, CompletionRequest, FunctionDefinition, Json, Message, ModelEffort,
};
use log::{Level::Debug, log_enabled};
use serde::{Serialize, de::DeserializeOwned};
use crate::model::{
execute_completion_request_with_retry, read_completion_response_bytes, read_sse_json_events,
streaming_completion_request,
};
use crate::{rfc3339_datetime, unix_ms};
pub(crate) struct SamplingOptions {
pub temperature: Option<f64>,
pub max_output_tokens: Option<usize>,
pub effort: Option<ModelEffort>,
pub output_schema: Option<Json>,
pub stop: Option<Vec<String>>,
}
pub(crate) trait WireFormat {
type Request: Serialize + Send + Sync;
type Response: Send;
type StreamItem: DeserializeOwned + Send;
fn set_instructions(r: &mut Self::Request, instructions: String);
fn append_raw_history(r: &mut Self::Request, raw_history: Vec<Json>) -> usize;
fn push_message(r: &mut Self::Request, msg: Message) -> Result<(), BoxError>;
fn apply_sampling(r: &mut Self::Request, options: SamplingOptions) -> Result<(), BoxError>;
fn apply_tools(r: &mut Self::Request, tools: Vec<FunctionDefinition>, required: bool);
fn finalize_request(_r: &mut Self::Request) {}
fn is_stream(r: &Self::Request) -> bool;
fn endpoint(r: &Self::Request, model: &str) -> String;
fn aggregate_stream(items: Vec<Self::StreamItem>) -> Result<Self::Response, BoxError>;
fn parse_response(model: &str, data: &[u8])
-> Result<(Self::Response, Option<Json>), BoxError>;
fn maybe_failed(res: &Self::Response) -> bool;
fn sent_messages(r: Self::Request, skip_raw: usize) -> Vec<Json>;
fn into_output(
res: Self::Response,
sent_messages: Vec<Json>,
chat_history: Vec<Message>,
assistant_raw_message: Option<Json>,
) -> Result<AgentOutput, BoxError>;
}
pub(crate) async fn drive_completion<W: WireFormat>(
model: String,
post: impl Fn(&str) -> reqwest::RequestBuilder,
mut r: W::Request,
req: CompletionRequest,
) -> Result<AgentOutput, BoxError> {
let timestamp = unix_ms();
let mut chat_history: Vec<Message> = Vec::new();
if !req.instructions.is_empty() {
W::set_instructions(&mut r, req.instructions);
}
let skip_raw = W::append_raw_history(&mut r, req.raw_history);
for msg in req.chat_history {
W::push_message(&mut r, msg)?;
}
if let Some(mut msg) = req
.documents
.to_message(&rfc3339_datetime(timestamp).unwrap())
{
msg.timestamp = Some(timestamp);
chat_history.push(msg.clone());
W::push_message(&mut r, msg)?;
}
let mut content = req.content;
if !req.prompt.is_empty() {
content.insert(0, req.prompt.into());
}
if !content.is_empty() {
let msg = Message {
role: req.role.unwrap_or_else(|| "user".to_string()),
content,
timestamp: Some(timestamp),
..Default::default()
};
chat_history.push(msg.clone());
W::push_message(&mut r, msg)?;
}
W::apply_sampling(
&mut r,
SamplingOptions {
temperature: req.temperature,
max_output_tokens: req.max_output_tokens,
effort: req.effort,
output_schema: req.output_schema,
stop: req.stop,
},
)?;
if !req.tools.is_empty() {
W::apply_tools(&mut r, req.tools, req.tool_choice_required);
}
W::finalize_request(&mut r);
let stream = W::is_stream(&r);
let path = W::endpoint(&r, &model);
let (res, assistant_raw_message) = execute_completion_request_with_retry(
&model,
|| {
let mut request = post(&path).json(&r);
if stream {
request = streaming_completion_request(request);
}
request
},
|response| async {
if stream {
let items = read_sse_json_events::<W::StreamItem>(response, &model).await?;
Ok((W::aggregate_stream(items)?, None))
} else {
let data = read_completion_response_bytes(response, &model).await?;
W::parse_response(&model, &data)
}
},
)
.await?;
let failed = W::maybe_failed(&res);
let sent_messages = W::sent_messages(r, skip_raw);
let output = W::into_output(res, sent_messages, chat_history, assistant_raw_message)?;
if failed {
log::warn!(model = model, usage:serde = output.usage; "Completion maybe failed");
} else if log_enabled!(Debug) {
log::debug!(model = model, usage:serde = output.usage; "Completion response");
}
Ok(output)
}