use super::message::{AssistantContent, DocumentMediaType};
use crate::http_client;
use crate::message::ToolChoice;
use crate::provider_response;
use crate::streaming::StreamingCompletionResponse;
use crate::wasm_compat::{WasmCompatSend, WasmCompatSync};
use crate::{
json_utils,
message::{Message, UserContent},
};
use serde::{Deserialize, Serialize};
use std::collections::HashMap;
use std::ops::{Add, AddAssign};
use thiserror::Error;
#[derive(Debug, Error)]
pub enum CompletionError {
#[error("HttpError: {0}")]
HttpError(#[from] http_client::Error),
#[error("JsonError: {0}")]
JsonError(#[from] serde_json::Error),
#[error("UrlError: {0}")]
UrlError(#[from] url::ParseError),
#[cfg(not(target_family = "wasm"))]
#[error("RequestError: {0}")]
RequestError(#[from] Box<dyn std::error::Error + Send + Sync + 'static>),
#[cfg(target_family = "wasm")]
#[error("RequestError: {0}")]
RequestError(#[from] Box<dyn std::error::Error + 'static>),
#[error("ResponseError: {0}")]
ResponseError(String),
#[error("ProviderError: {0}")]
ProviderError(String),
#[error("ProviderResponseError: {0}")]
ProviderResponse(provider_response::ProviderResponseError),
}
crate::provider_response::impl_provider_response_helpers!(CompletionError);
impl CompletionError {
pub(crate) fn from_stream_transport(error: http_client::Error) -> Self {
if error.non_success_status().is_some() {
Self::HttpError(error)
} else {
Self::ProviderError(error.to_string())
}
}
}
#[derive(Clone, Debug, PartialEq, Deserialize, Serialize)]
pub struct Document {
pub id: String,
pub text: String,
#[serde(flatten)]
pub additional_props: HashMap<String, String>,
}
impl std::fmt::Display for Document {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
write!(
f,
concat!("<file id: {}>\n", "{}\n", "</file>\n"),
self.id,
if self.additional_props.is_empty() {
self.text.clone()
} else {
let mut sorted_props = self.additional_props.iter().collect::<Vec<_>>();
sorted_props.sort_by(|a, b| a.0.cmp(b.0));
let metadata = sorted_props
.iter()
.map(|(k, v)| format!("{k}: {v:?}"))
.collect::<Vec<_>>()
.join(" ");
format!("<metadata {} />\n{}", metadata, self.text)
}
)
}
}
#[derive(Clone, Debug, Deserialize, Serialize, PartialEq)]
pub struct ToolDefinition {
pub name: String,
pub description: String,
pub parameters: serde_json::Value,
}
#[derive(Clone, Debug, Deserialize, Serialize, PartialEq)]
pub struct ProviderToolDefinition {
#[serde(rename = "type")]
pub kind: String,
#[serde(flatten, default, skip_serializing_if = "serde_json::Map::is_empty")]
pub config: serde_json::Map<String, serde_json::Value>,
}
impl ProviderToolDefinition {
pub fn new(kind: impl Into<String>) -> Self {
Self {
kind: kind.into(),
config: serde_json::Map::new(),
}
}
pub fn with_config(mut self, key: impl Into<String>, value: serde_json::Value) -> Self {
self.config.insert(key.into(), value);
self
}
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
pub enum FinishReason {
Stop,
Length,
ToolCalls,
ContentFilter,
Other(String),
}
impl FinishReason {
pub fn reconcile_with_output(self, has_tool_call: bool) -> Self {
if has_tool_call && matches!(self, Self::Stop) {
Self::ToolCalls
} else {
self
}
}
pub fn truncated_output(&self) -> bool {
matches!(self, Self::Length | Self::ContentFilter)
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(from = "CompletionResponseRepr")]
pub struct CompletionResponse {
pub choice: Vec<AssistantContent>,
pub usage: Usage,
#[serde(default)]
pub message_id: Option<String>,
#[serde(default)]
pub response_id: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub provider_request_id: Option<String>,
#[serde(default)]
finish_reason: Option<FinishReason>,
pub provider: String,
#[serde(default)]
pub model: Option<String>,
#[serde(default, skip_serializing_if = "serde_json::Value::is_null")]
pub raw: serde_json::Value,
}
#[derive(Clone, Debug, Default, PartialEq, Eq, Serialize, Deserialize)]
pub struct ResponseIdentity {
#[serde(default, skip_serializing_if = "Option::is_none")]
pub message_id: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub response_id: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub provider_request_id: Option<String>,
}
impl CompletionResponse {
pub fn new(choice: Vec<AssistantContent>, usage: Usage, provider: impl Into<String>) -> Self {
Self {
choice,
usage,
message_id: None,
response_id: None,
provider_request_id: None,
finish_reason: None,
provider: provider.into(),
model: None,
raw: serde_json::Value::Null,
}
}
pub fn finish_reason(&self) -> Option<FinishReason> {
self.finish_reason.clone()
}
pub fn identity(&self) -> ResponseIdentity {
ResponseIdentity {
message_id: self.message_id.clone(),
response_id: self.response_id.clone(),
provider_request_id: self.provider_request_id.clone(),
}
}
pub fn with_finish_reason(self, finish_reason: FinishReason) -> Self {
self.with_optional_finish_reason(Some(finish_reason))
}
pub fn with_optional_finish_reason(mut self, finish_reason: Option<FinishReason>) -> Self {
let has_tool_call = self
.choice
.iter()
.any(|content| matches!(content, AssistantContent::ToolCall(_)));
self.finish_reason =
finish_reason.map(|reason| reason.reconcile_with_output(has_tool_call));
self
}
}
crate::provider_response::response_metadata_setters!(CompletionResponse);
#[derive(Deserialize)]
struct CompletionResponseRepr {
choice: Vec<AssistantContent>,
usage: Usage,
#[serde(default)]
message_id: Option<String>,
#[serde(default)]
response_id: Option<String>,
#[serde(default)]
provider_request_id: Option<String>,
#[serde(default)]
finish_reason: Option<FinishReason>,
provider: String,
#[serde(default)]
model: Option<String>,
#[serde(default)]
raw: serde_json::Value,
}
impl From<CompletionResponseRepr> for CompletionResponse {
fn from(repr: CompletionResponseRepr) -> Self {
let CompletionResponseRepr {
choice,
usage,
message_id,
response_id,
provider_request_id,
finish_reason,
provider,
model,
raw,
} = repr;
Self::new(choice, usage, provider)
.with_optional_message_id(message_id)
.with_optional_response_id(response_id)
.with_optional_provider_request_id(provider_request_id)
.with_optional_finish_reason(finish_reason)
.with_optional_model(model)
.with_raw(raw)
}
}
pub trait NormalizeCompletionResponse {
fn normalize(self, provider: &str) -> Result<CompletionResponse, CompletionError>;
}
#[derive(Debug, PartialEq, Eq, Clone, Copy, Serialize, Deserialize)]
pub struct Usage {
pub input_tokens: u64,
pub output_tokens: u64,
pub total_tokens: u64,
pub cached_input_tokens: u64,
pub cache_creation_input_tokens: u64,
#[serde(default)]
pub tool_use_prompt_tokens: u64,
pub reasoning_tokens: u64,
}
impl Usage {
pub fn new() -> Self {
Self {
input_tokens: 0,
output_tokens: 0,
total_tokens: 0,
cached_input_tokens: 0,
cache_creation_input_tokens: 0,
tool_use_prompt_tokens: 0,
reasoning_tokens: 0,
}
}
pub fn has_values(&self) -> bool {
*self != Self::new()
}
}
impl Default for Usage {
fn default() -> Self {
Self::new()
}
}
impl Add for Usage {
type Output = Self;
fn add(mut self, other: Self) -> Self::Output {
self += other;
self
}
}
impl AddAssign for Usage {
fn add_assign(&mut self, other: Self) {
self.input_tokens += other.input_tokens;
self.output_tokens += other.output_tokens;
self.total_tokens += other.total_tokens;
self.cached_input_tokens += other.cached_input_tokens;
self.cache_creation_input_tokens += other.cache_creation_input_tokens;
self.tool_use_prompt_tokens += other.tool_use_prompt_tokens;
self.reasoning_tokens += other.reasoning_tokens;
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
pub struct ProviderCapabilities {
pub composes_native_output_with_tools: bool,
}
impl ProviderCapabilities {
pub const fn new() -> Self {
Self {
composes_native_output_with_tools: false,
}
}
pub const fn with_native_output_tool_composition(mut self, supported: bool) -> Self {
self.composes_native_output_with_tools = supported;
self
}
}
pub trait CompletionModel: WasmCompatSend + WasmCompatSync {
fn completion(
&self,
request: CompletionRequest,
) -> impl std::future::Future<Output = Result<CompletionResponse, CompletionError>> + WasmCompatSend;
fn stream(
&self,
request: CompletionRequest,
) -> impl std::future::Future<Output = Result<StreamingCompletionResponse, CompletionError>>
+ WasmCompatSend;
fn completion_request(&self, prompt: impl Into<Message>) -> CompletionRequestBuilder<Self>
where
Self: Sized + Clone,
{
CompletionRequestBuilder::new(self.clone(), prompt)
}
fn capabilities(&self) -> ProviderCapabilities {
ProviderCapabilities::default()
}
}
impl<M: CompletionModel + ?Sized> CompletionModel for std::sync::Arc<M> {
fn completion(
&self,
request: CompletionRequest,
) -> impl std::future::Future<Output = Result<CompletionResponse, CompletionError>> + WasmCompatSend
{
(**self).completion(request)
}
fn stream(
&self,
request: CompletionRequest,
) -> impl std::future::Future<Output = Result<StreamingCompletionResponse, CompletionError>>
+ WasmCompatSend {
(**self).stream(request)
}
fn capabilities(&self) -> ProviderCapabilities {
(**self).capabilities()
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct CompletionRequest {
pub model: Option<String>,
pub preamble: Option<String>,
pub chat_history: Vec<Message>,
pub documents: Vec<Document>,
pub tools: Vec<ToolDefinition>,
pub temperature: Option<f64>,
pub max_tokens: Option<u64>,
pub tool_choice: Option<ToolChoice>,
pub additional_params: Option<serde_json::Value>,
pub output_schema: Option<schemars::Schema>,
#[serde(skip)]
pub record_telemetry_content: bool,
}
impl CompletionRequest {
pub fn validate_message_content(&self) -> Result<(), CompletionError> {
if self.chat_history.is_empty() {
return Err(CompletionError::RequestError(
"request has an empty chat history; providers require at least one message"
.to_owned()
.into(),
));
}
let empty_message = |role: &str, index: usize| {
CompletionError::RequestError(
format!(
"{role} message at index {index} has no content; \
providers reject empty content blocks"
)
.into(),
)
};
for (index, message) in self.chat_history.iter().enumerate() {
match message {
Message::System { .. } => {}
Message::Assistant { content, .. } => {
if content.is_empty() {
return Err(empty_message("assistant", index));
}
}
Message::User { content } => {
if content.is_empty() {
return Err(empty_message("user", index));
}
for (position, item) in content.iter().enumerate() {
match item {
UserContent::ToolResult(result) if result.content.is_empty() => {
let name = &result.name;
return Err(CompletionError::RequestError(
format!(
"tool result for `{name}` at index {position} of the \
user message at index {index} has no content; \
providers reject empty content blocks"
)
.into(),
));
}
UserContent::ToolResult(_)
| UserContent::Text(_)
| UserContent::Image(_)
| UserContent::Audio(_)
| UserContent::Video(_)
| UserContent::Document(_) => {}
}
}
}
}
}
Ok(())
}
pub fn output_schema_name(&self) -> Option<String> {
self.output_schema.as_ref().map(|schema| {
schema
.as_object()
.and_then(|o| o.get("title"))
.and_then(|v| v.as_str())
.unwrap_or("response_schema")
.to_string()
})
}
pub fn normalized_documents(&self) -> Option<Message> {
Self::normalized_documents_from(&self.documents)
}
fn normalized_documents_from(documents: &[Document]) -> Option<Message> {
if documents.is_empty() {
return None;
}
let messages = documents
.iter()
.map(|doc| {
UserContent::document(
doc.to_string(),
Some(DocumentMediaType::TXT),
)
})
.collect::<Vec<_>>();
crate::message::non_empty(messages).map(|content| Message::User { content })
}
pub(crate) fn chat_history_with_documents(&self) -> Vec<Message> {
let mut chat_history = self.chat_history.clone();
if let Some(documents) = self.normalized_documents() {
insert_after_leading_system(&mut chat_history, documents);
}
chat_history
}
}
fn insert_after_leading_system(chat_history: &mut Vec<Message>, message: Message) {
let insert_at = chat_history
.iter()
.position(|message| !matches!(message, Message::System { .. }))
.unwrap_or(chat_history.len());
chat_history.insert(insert_at, message);
}
fn merge_provider_tools_into_additional_params(
additional_params: Option<serde_json::Value>,
provider_tools: Vec<ProviderToolDefinition>,
) -> Option<serde_json::Value> {
if provider_tools.is_empty() {
return additional_params;
}
let mut provider_tools_json = provider_tools
.into_iter()
.map(|ProviderToolDefinition { kind, mut config }| {
config.insert("type".to_string(), serde_json::Value::String(kind));
serde_json::Value::Object(config)
})
.collect::<Vec<_>>();
let mut params_map = match additional_params {
Some(serde_json::Value::Object(map)) => map,
Some(serde_json::Value::Bool(stream)) => {
let mut map = serde_json::Map::new();
map.insert("stream".to_string(), serde_json::Value::Bool(stream));
map
}
_ => serde_json::Map::new(),
};
let mut merged_tools = match params_map.remove("tools") {
Some(serde_json::Value::Array(existing)) => existing,
_ => Vec::new(),
};
merged_tools.append(&mut provider_tools_json);
params_map.insert("tools".to_string(), serde_json::Value::Array(merged_tools));
Some(serde_json::Value::Object(params_map))
}
pub struct CompletionRequestBuilder<M: CompletionModel> {
model: M,
prompt: Message,
request_model: Option<String>,
preamble: Option<String>,
chat_history: Vec<Message>,
documents: Vec<Document>,
tools: Vec<ToolDefinition>,
provider_tools: Vec<ProviderToolDefinition>,
temperature: Option<f64>,
max_tokens: Option<u64>,
tool_choice: Option<ToolChoice>,
additional_params: Option<serde_json::Value>,
output_schema: Option<schemars::Schema>,
record_telemetry_content: bool,
}
impl<M: CompletionModel> CompletionRequestBuilder<M> {
pub fn new(model: M, prompt: impl Into<Message>) -> Self {
Self {
model,
prompt: prompt.into(),
request_model: None,
preamble: None,
chat_history: Vec::new(),
documents: Vec::new(),
tools: Vec::new(),
provider_tools: Vec::new(),
temperature: None,
max_tokens: None,
tool_choice: None,
additional_params: None,
output_schema: None,
record_telemetry_content: false,
}
}
pub fn preamble(mut self, preamble: String) -> Self {
self.preamble = Some(preamble);
self
}
pub fn model(mut self, model: impl Into<String>) -> Self {
self.request_model = Some(model.into());
self
}
pub fn model_opt(mut self, model: Option<String>) -> Self {
self.request_model = model;
self
}
pub fn without_preamble(mut self) -> Self {
self.preamble = None;
self
}
pub fn message(mut self, message: Message) -> Self {
self.chat_history.push(message);
self
}
pub fn messages(mut self, messages: impl IntoIterator<Item = Message>) -> Self {
self.chat_history.extend(messages);
self
}
pub fn document(mut self, document: Document) -> Self {
self.documents.push(document);
self
}
pub fn documents(self, documents: impl IntoIterator<Item = Document>) -> Self {
documents
.into_iter()
.fold(self, |builder, doc| builder.document(doc))
}
pub fn tool(mut self, tool: ToolDefinition) -> Self {
self.tools.push(tool);
self
}
pub fn tools(self, tools: Vec<ToolDefinition>) -> Self {
tools
.into_iter()
.fold(self, |builder, tool| builder.tool(tool))
}
pub fn provider_tool(mut self, tool: ProviderToolDefinition) -> Self {
self.provider_tools.push(tool);
self
}
pub fn provider_tools(self, tools: Vec<ProviderToolDefinition>) -> Self {
tools
.into_iter()
.fold(self, |builder, tool| builder.provider_tool(tool))
}
pub fn additional_params(mut self, additional_params: serde_json::Value) -> Self {
match self.additional_params {
Some(params) => {
self.additional_params = Some(json_utils::merge(params, additional_params));
}
None => {
self.additional_params = Some(additional_params);
}
}
self
}
pub fn additional_params_opt(mut self, additional_params: Option<serde_json::Value>) -> Self {
self.additional_params = additional_params;
self
}
pub fn temperature(mut self, temperature: f64) -> Self {
self.temperature = Some(temperature);
self
}
pub fn temperature_opt(mut self, temperature: Option<f64>) -> Self {
self.temperature = temperature;
self
}
pub fn max_tokens(mut self, max_tokens: u64) -> Self {
self.max_tokens = Some(max_tokens);
self
}
pub fn max_tokens_opt(mut self, max_tokens: Option<u64>) -> Self {
self.max_tokens = max_tokens;
self
}
pub fn tool_choice(mut self, tool_choice: ToolChoice) -> Self {
self.tool_choice = Some(tool_choice);
self
}
pub fn output_schema(mut self, schema: schemars::Schema) -> Self {
self.output_schema = Some(schema);
self
}
pub fn output_schema_opt(mut self, schema: Option<schemars::Schema>) -> Self {
self.output_schema = schema;
self
}
pub fn record_content_telemetry(mut self, enabled: bool) -> Self {
self.record_telemetry_content = enabled;
self
}
pub fn messages_for_telemetry(&self) -> Vec<Message> {
let mut chat_history = self.chat_history.clone();
if let Some(preamble) = &self.preamble {
chat_history.insert(0, Message::system(preamble.clone()));
}
chat_history.push(self.prompt.clone());
if let Some(documents) = CompletionRequest::normalized_documents_from(&self.documents) {
insert_after_leading_system(&mut chat_history, documents);
}
chat_history
}
pub fn build(self) -> CompletionRequest {
self.into_model_and_request().1
}
fn into_model_and_request(self) -> (M, CompletionRequest) {
let model = self.model;
let mut chat_history = self.chat_history;
let prompt = self.prompt;
if let Some(preamble) = self.preamble {
chat_history.insert(0, Message::system(preamble));
}
chat_history.push(prompt);
let additional_params = merge_provider_tools_into_additional_params(
self.additional_params,
self.provider_tools,
);
let request = CompletionRequest {
model: self.request_model,
preamble: None,
chat_history,
documents: self.documents,
tools: self.tools,
temperature: self.temperature,
max_tokens: self.max_tokens,
tool_choice: self.tool_choice,
additional_params,
output_schema: self.output_schema,
record_telemetry_content: self.record_telemetry_content,
};
(model, request)
}
pub async fn send(self) -> Result<CompletionResponse, CompletionError> {
let (model, request) = self.into_model_and_request();
request.validate_message_content()?;
model.completion(request).await
}
pub async fn stream(self) -> Result<StreamingCompletionResponse, CompletionError> {
let (model, request) = self.into_model_and_request();
request.validate_message_content()?;
model.stream(request).await
}
}
#[cfg(test)]
mod tests {
use super::{CompletionResponse, FinishReason, ProviderCapabilities, Usage};
use crate::message::AssistantContent;
mod message_content_validation {
use super::super::CompletionRequest;
use crate::message::{AssistantContent, Message, UserContent};
fn request(chat_history: Vec<Message>) -> CompletionRequest {
CompletionRequest {
model: None,
preamble: None,
chat_history,
documents: Vec::new(),
tools: Vec::new(),
temperature: None,
max_tokens: None,
tool_choice: None,
additional_params: None,
output_schema: None,
record_telemetry_content: false,
}
}
#[test]
fn a_populated_history_passes() {
let request = request(vec![Message::user("hello")]);
assert!(request.validate_message_content().is_ok());
}
#[test]
fn an_empty_history_is_rejected() {
let error = request(Vec::new())
.validate_message_content()
.expect_err("an empty history must not reach a provider");
assert!(
error.to_string().contains("empty chat history"),
"unexpected error: {error}"
);
}
#[test]
fn an_empty_user_message_is_rejected_by_index_and_role() {
let error = request(vec![
Message::user("hello"),
Message::User {
content: Vec::new(),
},
])
.validate_message_content()
.expect_err("an empty user message must not reach a provider");
let message = error.to_string();
assert!(message.contains("user message at index 1"), "{message}");
}
#[test]
fn an_empty_assistant_message_is_rejected_by_index_and_role() {
let error = request(vec![Message::Assistant {
id: None,
content: Vec::new(),
}])
.validate_message_content()
.expect_err("an empty assistant message must not reach a provider");
let message = error.to_string();
assert!(
message.contains("assistant message at index 0"),
"{message}"
);
}
#[test]
fn an_empty_system_message_is_not_rejected() {
let request = request(vec![
Message::System {
content: String::new(),
},
Message::user("hello"),
]);
assert!(request.validate_message_content().is_ok());
}
#[test]
fn a_block_less_tool_result_is_rejected_naming_the_tool() {
use crate::message::{ToolCallId, ToolResult, ToolResultContent};
let error = request(vec![
Message::user("hello"),
Message::User {
content: vec![UserContent::ToolResult(ToolResult {
call: ToolCallId::new_or_mint("call_1"),
provider: None,
name: "lookup".to_owned(),
content: Vec::<ToolResultContent>::new(),
})],
},
])
.validate_message_content()
.expect_err("a block-less tool result must not reach a provider");
let message = error.to_string();
assert!(message.contains("`lookup`"), "{message}");
assert!(message.contains("index 0"), "{message}");
assert!(message.contains("user message at index 1"), "{message}");
}
#[test]
fn a_tool_result_with_one_empty_string_block_is_accepted() {
use crate::message::{ToolCallId, ToolResult, ToolResultContent};
let request = request(vec![Message::User {
content: vec![UserContent::ToolResult(ToolResult {
call: ToolCallId::new_or_mint("call_1"),
provider: None,
name: "lookup".to_owned(),
content: vec![ToolResultContent::text("")],
})],
}]);
assert!(request.validate_message_content().is_ok());
}
#[test]
fn the_legacy_fabricated_sentinel_still_passes() {
let request = request(vec![
Message::user("hello"),
Message::Assistant {
id: None,
content: vec![AssistantContent::text("")],
},
Message::User {
content: vec![UserContent::text("and again")],
},
]);
assert!(request.validate_message_content().is_ok());
}
}
fn tool_call_choice() -> Vec<AssistantContent> {
vec![AssistantContent::tool_call(
"call_1",
"lookup",
serde_json::json!({"query": "rig"}),
)]
}
#[test]
fn normalized_response_round_trips_through_serde() {
let response = CompletionResponse::new(
vec![AssistantContent::text("hello")],
Usage {
input_tokens: 3,
output_tokens: 2,
total_tokens: 5,
cached_input_tokens: 1,
cache_creation_input_tokens: 0,
tool_use_prompt_tokens: 0,
reasoning_tokens: 1,
},
"example",
)
.with_message_id("msg_123")
.with_finish_reason(FinishReason::Stop)
.with_model("provider-model-v2");
let encoded = serde_json::to_value(&response).expect("serialize response");
let decoded =
serde_json::from_value::<CompletionResponse>(encoded.clone()).expect("deserialize");
assert_eq!(
serde_json::to_value(decoded).expect("re-serialize"),
encoded
);
}
#[test]
fn deserializing_stop_with_a_tool_call_reconciles_to_tool_calls() {
let mut encoded = serde_json::to_value(CompletionResponse::new(
tool_call_choice(),
Usage::new(),
"example",
))
.expect("serialize response");
encoded["finish_reason"] = serde_json::json!("stop");
let decoded =
serde_json::from_value::<CompletionResponse>(encoded).expect("deserialize response");
assert_eq!(decoded.finish_reason(), Some(FinishReason::ToolCalls));
}
#[test]
fn deserializing_empty_identifiers_yields_none() {
let mut encoded = serde_json::to_value(CompletionResponse::new(
vec![AssistantContent::text("hello")],
Usage::new(),
"example",
))
.expect("serialize response");
encoded["message_id"] = serde_json::json!("");
encoded["response_id"] = serde_json::json!("");
encoded["model"] = serde_json::json!("");
let decoded =
serde_json::from_value::<CompletionResponse>(encoded).expect("deserialize response");
assert_eq!(decoded.message_id, None);
assert_eq!(decoded.response_id, None);
assert_eq!(decoded.model, None);
}
#[test]
fn unknown_finish_reason_survives_a_serde_round_trip_verbatim() {
let reason = FinishReason::Other("provider_specific_stop".to_owned());
let encoded = serde_json::to_string(&reason).expect("serialize");
let decoded = serde_json::from_str::<FinishReason>(&encoded).expect("deserialize");
assert_eq!(decoded, reason);
}
#[test]
fn stop_with_a_tool_call_reconciles_to_tool_calls() {
let response = CompletionResponse::new(tool_call_choice(), Usage::new(), "example")
.with_finish_reason(FinishReason::Stop);
assert_eq!(response.finish_reason, Some(FinishReason::ToolCalls));
}
#[test]
fn optional_setter_reconciles_exactly_like_the_plain_setter() {
let via_option = CompletionResponse::new(tool_call_choice(), Usage::new(), "example")
.with_optional_finish_reason(Some(FinishReason::Stop));
let via_plain = CompletionResponse::new(tool_call_choice(), Usage::new(), "example")
.with_finish_reason(FinishReason::Stop);
assert_eq!(via_option.finish_reason, Some(FinishReason::ToolCalls));
assert_eq!(via_option.finish_reason, via_plain.finish_reason);
}
#[test]
fn reconciliation_only_upgrades_a_natural_stop() {
for reason in [
FinishReason::Length,
FinishReason::ContentFilter,
FinishReason::Other("provider_specific".to_owned()),
] {
let response = CompletionResponse::new(tool_call_choice(), Usage::new(), "example")
.with_finish_reason(reason.clone());
assert_eq!(response.finish_reason, Some(reason));
}
}
#[test]
fn reconciliation_leaves_a_stop_without_tool_calls_alone() {
let response = CompletionResponse::new(
vec![AssistantContent::text("done")],
Usage::new(),
"example",
)
.with_finish_reason(FinishReason::Stop);
assert_eq!(response.finish_reason, Some(FinishReason::Stop));
}
#[test]
fn provider_capabilities_are_externally_configurable_from_default() {
let capabilities =
ProviderCapabilities::default().with_native_output_tool_composition(true);
assert!(capabilities.composes_native_output_with_tools);
assert!(!ProviderCapabilities::new().composes_native_output_with_tools);
assert_eq!(ProviderCapabilities::new(), ProviderCapabilities::default());
}
#[test]
fn usage_has_values_reflects_the_zero_sentinel() {
use super::Usage;
assert!(!Usage::new().has_values());
let mut usage = Usage::new();
usage.reasoning_tokens = 1;
assert!(usage.has_values());
}
use super::*;
use crate::test_utils::MockCompletionModel;
#[test]
fn completion_request_content_telemetry_is_opt_in_and_not_serialized() {
let default_request =
CompletionRequestBuilder::new(MockCompletionModel::default(), "completion prompt")
.build();
assert!(!default_request.record_telemetry_content);
let default_json = serde_json::to_value(&default_request).expect("serialize request");
assert!(
default_json.get("record_telemetry_content").is_none(),
"safe default should not serialize the telemetry opt-in field"
);
let default_roundtrip: CompletionRequest =
serde_json::from_value(default_json).expect("deserialize default request");
assert!(!default_roundtrip.record_telemetry_content);
let opt_in_request =
CompletionRequestBuilder::new(MockCompletionModel::default(), "completion prompt")
.record_content_telemetry(true)
.build();
assert!(opt_in_request.record_telemetry_content);
let opt_in_json = serde_json::to_value(&opt_in_request).expect("serialize opt-in request");
assert!(
opt_in_json.get("record_telemetry_content").is_none(),
"local telemetry policy must not be serialized into provider requests"
);
let legacy_roundtrip: CompletionRequest =
serde_json::from_value(opt_in_json).expect("deserialize legacy request");
assert!(
!legacy_roundtrip.record_telemetry_content,
"missing field should deserialize to the safe default"
);
}
#[test]
fn normalized_response_raw_round_trips_through_serde_mirror() {
let payload = serde_json::json!({
"id": "chatcmpl-1",
"system_fingerprint": "fp_abc",
"choices": [{"finish_reason": "stop"}]
});
let response = CompletionResponse::new(
vec![AssistantContent::text("hello")],
Usage::new(),
"example",
)
.with_response_id("chatcmpl-1")
.with_raw(payload.clone());
let encoded = serde_json::to_value(&response).expect("serialize response");
assert_eq!(encoded["raw"], payload);
let decoded: CompletionResponse =
serde_json::from_value(encoded.clone()).expect("deserialize response");
assert_eq!(decoded.raw, payload);
assert_eq!(decoded.response_id.as_deref(), Some("chatcmpl-1"));
assert_eq!(
serde_json::to_value(&decoded).expect("re-serialize"),
encoded
);
let legacy = serde_json::json!({
"choice": [{"type": "text", "text": "hello"}],
"usage": serde_json::to_value(Usage::new()).unwrap(),
"provider": "example"
});
let decoded: CompletionResponse = serde_json::from_value(legacy).expect("legacy loads");
assert!(decoded.raw.is_null());
let bare = serde_json::to_value(CompletionResponse::new(
vec![AssistantContent::text("hello")],
Usage::new(),
"example",
))
.unwrap();
assert!(bare.get("raw").is_none());
}
fn test_document(id: &str, text: &str) -> Document {
Document {
id: id.to_string(),
text: text.to_string(),
additional_props: HashMap::new(),
}
}
#[test]
fn message_telemetry_includes_normalized_documents() {
let builder = CompletionRequestBuilder::new(MockCompletionModel::default(), "prompt")
.preamble("system".to_string())
.message(Message::user("history"))
.document(test_document("doc1", "static context secret"));
let messages = builder.messages_for_telemetry();
assert_eq!(messages.len(), 4);
assert!(matches!(messages[0], Message::System { .. }));
assert!(is_document_message(&messages[1], "doc1"));
assert!(matches!(
&messages[2],
Message::User { content }
if matches!(content.first(), Some(UserContent::Text(text)) if text.text == "history")
));
assert!(matches!(
&messages[3],
Message::User { content }
if matches!(content.first(), Some(UserContent::Text(text)) if text.text == "prompt")
));
let request = builder.build();
assert_eq!(messages, request.chat_history_with_documents());
}
fn is_document_message(message: &Message, expected_id: &str) -> bool {
let Message::User { content } = message else {
return false;
};
content.iter().any(|content| {
matches!(
content,
UserContent::Document(document)
if document.data.to_string().contains(&format!("<file id: {expected_id}>"))
)
})
}
#[test]
fn test_document_display_without_metadata() {
let doc = Document {
id: "123".to_string(),
text: "This is a test document.".to_string(),
additional_props: HashMap::new(),
};
let expected = "<file id: 123>\nThis is a test document.\n</file>\n";
assert_eq!(format!("{doc}"), expected);
}
#[test]
fn test_document_display_with_metadata() {
let mut additional_props = HashMap::new();
additional_props.insert("author".to_string(), "John Doe".to_string());
additional_props.insert("length".to_string(), "42".to_string());
let doc = Document {
id: "123".to_string(),
text: "This is a test document.".to_string(),
additional_props,
};
let expected = concat!(
"<file id: 123>\n",
"<metadata author: \"John Doe\" length: \"42\" />\n",
"This is a test document.\n",
"</file>\n"
);
assert_eq!(format!("{doc}"), expected);
}
#[test]
fn test_normalize_documents_with_documents() {
let doc1 = Document {
id: "doc1".to_string(),
text: "Document 1 text.".to_string(),
additional_props: HashMap::new(),
};
let doc2 = Document {
id: "doc2".to_string(),
text: "Document 2 text.".to_string(),
additional_props: HashMap::new(),
};
let request = CompletionRequest {
model: None,
preamble: None,
chat_history: vec!["What is the capital of France?".into()],
documents: vec![doc1, doc2],
tools: Vec::new(),
temperature: None,
max_tokens: None,
tool_choice: None,
additional_params: None,
output_schema: None,
record_telemetry_content: false,
};
let expected = Message::User {
content: vec![
UserContent::document(
"<file id: doc1>\nDocument 1 text.\n</file>\n".to_string(),
Some(DocumentMediaType::TXT),
),
UserContent::document(
"<file id: doc2>\nDocument 2 text.\n</file>\n".to_string(),
Some(DocumentMediaType::TXT),
),
],
};
assert_eq!(request.normalized_documents(), Some(expected));
}
#[test]
fn test_normalize_documents_without_documents() {
let request = CompletionRequest {
model: None,
preamble: None,
chat_history: vec!["What is the capital of France?".into()],
documents: Vec::new(),
tools: Vec::new(),
temperature: None,
max_tokens: None,
tool_choice: None,
additional_params: None,
output_schema: None,
record_telemetry_content: false,
};
assert_eq!(request.normalized_documents(), None);
}
#[test]
fn preamble_builder_funnels_to_system_message() {
let request =
CompletionRequestBuilder::new(MockCompletionModel::default(), Message::user("Prompt"))
.preamble("System prompt".to_string())
.message(Message::user("History"))
.build();
assert_eq!(request.preamble, None);
let history = request.chat_history.into_iter().collect::<Vec<_>>();
assert_eq!(history.len(), 3);
assert!(matches!(
&history[0],
Message::System { content } if content == "System prompt"
));
assert!(matches!(&history[1], Message::User { .. }));
assert!(matches!(&history[2], Message::User { .. }));
}
#[test]
fn without_preamble_removes_legacy_preamble_injection() {
let request =
CompletionRequestBuilder::new(MockCompletionModel::default(), Message::user("Prompt"))
.preamble("System prompt".to_string())
.without_preamble()
.build();
assert_eq!(request.preamble, None);
let history = request.chat_history.into_iter().collect::<Vec<_>>();
assert_eq!(history.len(), 1);
assert!(matches!(&history[0], Message::User { .. }));
}
#[test]
fn build_places_documents_after_preamble_system_message() {
let request =
CompletionRequestBuilder::new(MockCompletionModel::default(), Message::user("Prompt"))
.preamble("System prompt".to_string())
.document(test_document("doc1", "Document text."))
.build();
assert_eq!(request.documents.len(), 1);
let history = request.chat_history_with_documents();
let history = history.iter().collect::<Vec<_>>();
assert_eq!(history.len(), 3);
assert!(matches!(
history[0],
Message::System { content } if content == "System prompt"
));
assert!(is_document_message(history[1], "doc1"));
assert!(matches!(history[2], Message::User { .. }));
}
#[test]
fn build_places_documents_after_leading_system_messages_before_prior_history() {
let request =
CompletionRequestBuilder::new(MockCompletionModel::default(), Message::user("Prompt"))
.message(Message::system("System one"))
.message(Message::system("System two"))
.message(Message::user("Earlier user turn"))
.message(Message::assistant("Earlier assistant turn"))
.document(test_document("doc1", "Document text."))
.build();
let history = request.chat_history_with_documents();
let history = history.iter().collect::<Vec<_>>();
assert_eq!(history.len(), 6);
assert!(matches!(
history[0],
Message::System { content } if content == "System one"
));
assert!(matches!(
history[1],
Message::System { content } if content == "System two"
));
assert!(is_document_message(history[2], "doc1"));
assert!(matches!(history[3], Message::User { .. }));
assert!(matches!(history[4], Message::Assistant { .. }));
assert!(matches!(history[5], Message::User { .. }));
}
#[test]
fn build_without_documents_keeps_message_order_unchanged() {
let request =
CompletionRequestBuilder::new(MockCompletionModel::default(), Message::user("Prompt"))
.message(Message::system("System prompt"))
.message(Message::user("Earlier user turn"))
.build();
let history = request.chat_history.iter().collect::<Vec<_>>();
assert_eq!(history.len(), 3);
assert!(matches!(
history[0],
Message::System { content } if content == "System prompt"
));
assert!(matches!(history[1], Message::User { .. }));
assert!(matches!(history[2], Message::User { .. }));
}
#[test]
fn chat_history_with_documents_places_documents_after_leading_system_messages() {
let request = CompletionRequest {
model: None,
preamble: None,
chat_history: vec![
Message::system("System prompt"),
Message::assistant("Earlier assistant turn"),
Message::user("Earlier user turn"),
Message::user("Prompt"),
],
documents: vec![test_document("doc1", "Document text.")],
tools: Vec::new(),
temperature: None,
max_tokens: None,
tool_choice: None,
additional_params: None,
output_schema: None,
record_telemetry_content: false,
};
assert_eq!(request.documents.len(), 1);
let history = request.chat_history_with_documents();
let history = history.iter().collect::<Vec<_>>();
assert_eq!(history.len(), 5);
assert!(matches!(history[0], Message::System { .. }));
assert!(is_document_message(history[1], "doc1"));
assert!(matches!(history[2], Message::Assistant { .. }));
assert!(matches!(history[3], Message::User { .. }));
assert!(matches!(history[4], Message::User { .. }));
}
#[test]
fn chat_history_with_documents_places_documents_before_mid_conversation_system_messages() {
let request = CompletionRequest {
model: None,
preamble: None,
chat_history: vec![
Message::system("Leading system prompt"),
Message::assistant("Earlier assistant turn"),
Message::system("Mid-conversation instruction"),
Message::user("Prompt"),
],
documents: vec![test_document("doc1", "Document text.")],
tools: Vec::new(),
temperature: None,
max_tokens: None,
tool_choice: None,
additional_params: None,
output_schema: None,
record_telemetry_content: false,
};
let history = request.chat_history_with_documents();
let history = history.iter().collect::<Vec<_>>();
assert_eq!(history.len(), 5);
assert!(matches!(
history[0],
Message::System { content } if content == "Leading system prompt"
));
assert!(is_document_message(history[1], "doc1"));
assert!(matches!(history[2], Message::Assistant { .. }));
assert!(matches!(
history[3],
Message::System { content } if content == "Mid-conversation instruction"
));
assert!(matches!(history[4], Message::User { .. }));
}
#[test]
fn chat_history_with_documents_does_not_duplicate_documents() {
let request = CompletionRequest {
model: None,
preamble: None,
chat_history: vec![
Message::system("System prompt"),
Message::user("Earlier user turn"),
Message::assistant("Earlier assistant turn"),
Message::user("Prompt"),
],
documents: vec![test_document("doc1", "Document text.")],
tools: Vec::new(),
temperature: None,
max_tokens: None,
tool_choice: None,
additional_params: None,
output_schema: None,
record_telemetry_content: false,
};
let history = request.chat_history_with_documents();
let document_messages = history
.iter()
.filter(|message| is_document_message(message, "doc1"))
.count();
assert_eq!(document_messages, 1);
}
#[test]
fn completion_error_provider_response_helpers_with_preserved_json_body() {
let body = r#"{"error":{"code":"rate_limit","message":"slow down"}}"#;
let error = CompletionError::ProviderResponse(
provider_response::ProviderResponseError::without_status(body.to_string()),
);
assert_eq!(error.provider_response_body(), Some(body));
assert_eq!(error.provider_response_status(), None);
assert_eq!(
error
.provider_response_json()
.expect("fixture body should parse as valid JSON"),
Some(serde_json::json!({
"error": {
"code": "rate_limit",
"message": "slow down"
}
}))
);
}
#[test]
fn completion_error_provider_response_helpers_with_preserved_status() {
let body = r#"{"error":{"message":"too many requests"}}"#;
let error =
CompletionError::ProviderResponse(provider_response::ProviderResponseError::new(
http::StatusCode::TOO_MANY_REQUESTS,
body.to_string(),
));
assert_eq!(error.provider_response_body(), Some(body));
assert_eq!(
error.provider_response_status(),
Some(http::StatusCode::TOO_MANY_REQUESTS)
);
}
#[test]
fn completion_error_provider_response_helpers_with_preserved_plain_text_body() {
let error = CompletionError::ProviderResponse(
provider_response::ProviderResponseError::without_status(
"provider exploded".to_string(),
),
);
assert_eq!(error.provider_response_body(), Some("provider exploded"));
assert_eq!(error.provider_response_status(), None);
assert!(error.provider_response_json().is_err());
}
#[test]
fn completion_error_provider_error_is_not_a_provider_response() {
let error = CompletionError::ProviderError("stream transport failed".to_string());
assert_eq!(error.provider_response_body(), None);
assert_eq!(error.provider_response_status(), None);
assert_eq!(
error
.provider_response_json()
.expect("no body is not an error"),
None
);
}
#[test]
fn completion_error_provider_response_helpers_with_http_non_success_body_and_status() {
let body = r#"{"error":{"type":"invalid_request","message":"bad request"}}"#;
let error = CompletionError::HttpError(http_client::Error::InvalidStatusCodeWithMessage(
http::StatusCode::BAD_REQUEST,
body.to_string(),
));
assert_eq!(error.provider_response_body(), Some(body));
assert_eq!(
error.provider_response_status(),
Some(http::StatusCode::BAD_REQUEST)
);
assert_eq!(
error.provider_response_json().expect("valid JSON body"),
Some(serde_json::json!({
"error": {
"type": "invalid_request",
"message": "bad request"
}
}))
);
}
#[test]
fn completion_error_provider_response_helpers_with_unrelated_variant() {
let error = CompletionError::ResponseError("failed to parse provider response".to_string());
assert_eq!(error.provider_response_body(), None);
assert_eq!(error.provider_response_status(), None);
assert_eq!(
error
.provider_response_json()
.expect("no body is not an error"),
None
);
}
#[test]
fn provider_response_json_returns_none_for_empty_preserved_body() {
let error = CompletionError::ProviderResponse(
provider_response::ProviderResponseError::without_status(String::new()),
);
assert_eq!(error.provider_response_body(), Some(""));
assert_eq!(
error
.provider_response_json()
.expect("empty body is not a JSON parse error"),
None
);
}
}
#[cfg(test)]
mod response_identity_tests {
use super::*;
#[test]
fn completion_response_without_request_id_still_deserializes() {
let response: CompletionResponse = serde_json::from_str(
r#"{"choice": [{"type": "text", "text": "hi"}],
"usage": {"input_tokens": 1, "output_tokens": 1, "total_tokens": 2,
"cached_input_tokens": 0, "cache_creation_input_tokens": 0,
"reasoning_tokens": 0},
"provider": "test"}"#,
)
.expect("pre-identity CompletionResponse JSON should load");
assert_eq!(response.provider_request_id, None);
assert_eq!(response.identity(), ResponseIdentity::default());
}
#[test]
fn identity_accessor_mirrors_flat_fields() {
let response = CompletionResponse::new(
vec![crate::completion::AssistantContent::text("hi")],
Usage::new(),
"test",
)
.with_message_id("msg_1")
.with_response_id("resp_1")
.with_provider_request_id("req_1");
assert_eq!(
response.identity(),
ResponseIdentity {
message_id: Some("msg_1".into()),
response_id: Some("resp_1".into()),
provider_request_id: Some("req_1".into()),
}
);
}
}