use anda_core::{
AgentOutput, BoxError, CompletionRequest, ContentPart, FunctionDefinition, Json, Message,
ModelEffort,
};
use log::{Level::Debug, log_enabled};
use serde::{Serialize, de::DeserializeOwned};
use std::collections::{HashMap, VecDeque};
use crate::model::{
execute_completion_request_with_retry, read_completion_response_bytes, read_completion_stream,
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, model: &str) -> Result<(), BoxError>;
fn apply_sampling(
r: &mut Self::Request,
options: SamplingOptions,
model: &str,
) -> Result<(), BoxError>;
fn apply_tools(r: &mut Self::Request, tools: Vec<FunctionDefinition>, required: bool);
fn finalize_request(_r: &mut Self::Request) -> Result<(), BoxError> {
Ok(())
}
fn request_headers(_r: &Self::Request) -> Vec<(&'static str, String)> {
Vec::new()
}
fn is_stream(r: &Self::Request) -> bool;
fn endpoint(r: &Self::Request, model: &str) -> String;
fn aggregate_stream(
items: Vec<Self::StreamItem>,
done: bool,
) -> 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, &model)?;
}
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, &model)?;
}
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, &model)?;
}
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,
},
&model,
)?;
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 headers = W::request_headers(&r);
let (res, assistant_raw_message) = execute_completion_request_with_retry(
&model,
|| {
let mut request = post(&path).json(&r);
for (name, value) in &headers {
request = request.header(*name, value);
}
if stream {
request = streaming_completion_request(request);
}
request
},
|response| async {
if stream {
let (items, done) =
read_completion_stream::<W::StreamItem>(response, &model).await?;
Ok((W::aggregate_stream(items, done)?, 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)
}
pub(crate) fn assign_tool_call_ids(req: &mut CompletionRequest) {
let mut parts: Vec<_> = req
.chat_history
.iter_mut()
.flat_map(|msg| &mut msg.content)
.chain(&mut req.content)
.collect();
let mut pending = HashMap::<String, VecDeque<(usize, Option<String>)>>::new();
for index in 0..parts.len() {
match &*parts[index] {
ContentPart::ToolCall { name, call_id, .. } => {
pending
.entry(name.clone())
.or_default()
.push_back((index, call_id.clone().filter(|id| !id.is_empty())));
}
ContentPart::ToolOutput { name, call_id, .. } => {
let result_id = call_id.clone().filter(|id| !id.is_empty());
let queue = pending.entry(name.clone()).or_default();
let matched = if let Some(id) = &result_id {
queue
.iter()
.position(|(_, call_id)| call_id.as_ref() == Some(id))
.or_else(|| queue.iter().position(|(_, call_id)| call_id.is_none()))
.and_then(|position| queue.remove(position))
} else {
queue.pop_front()
};
if let Some((call_index, call_id)) = matched {
let id = result_id
.or(call_id)
.unwrap_or_else(|| format!("call_anda_{:032x}", rand::random::<u128>()));
if let ContentPart::ToolCall { call_id, .. } = &mut parts[call_index] {
*call_id = Some(id.clone());
}
if let ContentPart::ToolOutput { call_id, .. } = &mut parts[index] {
*call_id = Some(id);
}
}
}
_ => {}
}
}
for part in parts {
if let ContentPart::ToolCall { call_id, .. } = part
&& call_id.as_ref().is_none_or(String::is_empty)
{
*call_id = Some(format!("call_anda_{:032x}", rand::random::<u128>()));
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use serde_json::json;
#[test]
fn replay_reuses_ids_already_persisted_on_results() {
let mut req = CompletionRequest {
content: vec![
ContentPart::ToolCall {
name: "lookup".into(),
args: json!({}),
call_id: None,
},
ContentPart::ToolOutput {
name: "lookup".into(),
output: json!("ok"),
call_id: Some("persisted".into()),
is_error: None,
remote_id: None,
},
],
..Default::default()
};
assign_tool_call_ids(&mut req);
assert!(
matches!(&req.content[0], ContentPart::ToolCall {call_id:Some(id),..} if id == "persisted")
);
let first = req.content.clone();
assign_tool_call_ids(&mut req);
assert_eq!(req.content, first);
}
#[test]
fn results_answering_native_calls_and_raw_history_are_untouched() {
let raw = json!({"role":"assistant","tool_calls":[{"id":"","type":"function",
"function":{"name":"lookup","arguments":"{}"}}]});
let content = vec![
ContentPart::ToolOutput {
name: "lookup".into(),
output: json!("ok"),
call_id: Some(String::new()),
is_error: None,
remote_id: None,
},
ContentPart::ToolOutput {
name: "legacy".into(),
output: json!("ok"),
call_id: None,
is_error: None,
remote_id: None,
},
];
let mut req = CompletionRequest {
raw_history: vec![raw.clone()],
content: content.clone(),
..Default::default()
};
assign_tool_call_ids(&mut req);
assert_eq!(req.content, content);
assert_eq!(req.raw_history, vec![raw]);
}
}