use anyhow::Result;
use async_trait::async_trait;
use serde::{Deserialize, Serialize};
use serde_json::{Value, json};
use crate::tools::ToolDefinition;
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
#[serde(rename_all = "lowercase")]
pub enum Role {
System,
User,
Assistant,
Tool,
}
impl std::fmt::Display for Role {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
Role::System => write!(f, "system"),
Role::User => write!(f, "user"),
Role::Assistant => write!(f, "assistant"),
Role::Tool => write!(f, "tool"),
}
}
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub struct Message {
pub role: Role,
pub content: String,
#[serde(default, skip_serializing_if = "Vec::is_empty")]
pub tool_calls: Vec<ToolCall>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub tool_call_id: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub name: Option<String>,
#[serde(default, skip_serializing_if = "std::ops::Not::not")]
pub is_error: bool,
#[serde(default, skip_serializing_if = "Vec::is_empty")]
pub thinking_blocks: Vec<Value>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub timestamp: Option<chrono::DateTime<chrono::FixedOffset>>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub log_line: Option<u64>,
#[serde(default, skip_serializing_if = "String::is_empty")]
pub thinking: String,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub usage: Option<TokenUsage>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub duration_ms: Option<u64>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub temperature: Option<String>,
}
impl Message {
fn new(role: Role, content: &str) -> Self {
Self {
role,
content: content.to_string(),
tool_calls: vec![],
is_error: false,
thinking_blocks: vec![],
tool_call_id: None,
name: None,
timestamp: None,
log_line: None,
thinking: String::new(),
usage: None,
duration_ms: None,
temperature: None,
}
}
pub fn system(content: &str) -> Self {
Self::new(Role::System, content)
}
pub fn user(content: &str) -> Self {
Self::new(Role::User, content)
}
pub fn assistant(content: &str) -> Self {
Self::new(Role::Assistant, content)
}
pub fn assistant_with_tools(content: &str, tool_calls: Vec<ToolCall>) -> Self {
Self { tool_calls, ..Self::new(Role::Assistant, content) }
}
pub fn tool_result(tool_call_id: &str, name: &str, content: &str) -> Self {
Self {
tool_call_id: Some(tool_call_id.to_string()),
name: Some(name.to_string()),
..Self::new(Role::Tool, content)
}
}
pub fn tool_error(tool_call_id: &str, name: &str, content: &str) -> Self {
Self { is_error: true, ..Self::tool_result(tool_call_id, name, content) }
}
}
#[derive(Debug, Clone, Default)]
pub struct LLMResponse {
pub content: String,
pub tool_calls: Vec<ToolCall>,
pub usage: Option<TokenUsage>,
pub stop_reason: Option<String>,
pub thinking: String,
pub thinking_blocks: Vec<Value>,
}
#[derive(Debug, Clone, Default, PartialEq, Serialize, Deserialize)]
pub struct TokenUsage {
pub prompt_tokens: i64,
pub completion_tokens: i64,
pub total_tokens: i64,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub aic: Option<f64>,
}
pub fn copilot_aic(value: &Value) -> Option<f64> {
let nano = value.pointer("/copilot_usage/total_nano_aiu").and_then(Value::as_f64)?;
Some(nano / 1e9)
}
#[derive(Debug, Clone, Default, PartialEq, Serialize, Deserialize)]
pub struct ToolCall {
pub id: String,
pub name: String,
pub arguments: Value,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub item_id: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub malformed_arguments: Option<String>,
}
impl ToolCall {
pub fn from_raw_arguments(id: String, name: String, raw: &str, item_id: Option<String>) -> Self {
let (arguments, malformed_arguments) = Self::split_arguments(raw);
Self { id, name, arguments, item_id, malformed_arguments }
}
fn split_arguments(raw: &str) -> (Value, Option<String>) {
if raw.trim().is_empty() {
return (json!({}), None);
}
match serde_json::from_str(raw) {
Ok(value) => (value, None),
Err(_) => (json!({}), Some(raw.to_string())),
}
}
pub fn encoded_arguments(&self) -> String {
match &self.arguments {
Value::String(raw) => raw.clone(),
other => other.to_string(),
}
}
pub fn invalid_arguments(&self) -> Option<&str> {
self.malformed_arguments.as_deref()
}
pub fn raw_arguments_error(&self, stop_reason: Option<&str>) -> Option<String> {
self.invalid_arguments().map(|raw| {
let shown: String = raw.chars().take(300).collect();
let truncated = if raw.chars().count() > 300 { "…" } else { "" };
if stop_reason_is_length(stop_reason) {
format!(
"the arguments for `{}` were cut off because the response hit the output-token \
limit (stop reason `{}`), so the call was not run: `{shown}{truncated}`. \
Retrying the same call will hit the same limit — make a smaller call instead, \
for example by splitting the work across multiple `{}` calls or reducing the \
argument size so the full JSON fits within the limit.",
self.name,
stop_reason.unwrap_or("length"),
self.name
)
} else {
format!(
"the arguments for `{}` arrived as malformed JSON (usually a truncated \
response), so the call was not run: `{shown}{truncated}`. \
If this was a transient transport error, simply retry the same `{}` call with \
the same arguments; if it keeps happening, the response was probably truncated \
by the output-token limit, so make a smaller call instead (for example by \
splitting the content).",
self.name, self.name
)
}
})
}
}
pub(crate) fn stop_reason_is_length(stop_reason: Option<&str>) -> bool {
stop_reason.is_some_and(|reason| {
matches!(reason.to_ascii_lowercase().as_str(), "length" | "max_tokens" | "max_output_tokens")
})
}
#[derive(Debug, Clone)]
pub struct ChatRequest<'a> {
pub messages: &'a [Message],
pub tools: &'a [ToolDefinition],
pub temperature: Option<f64>,
pub max_tokens: Option<i64>,
}
#[derive(Debug, Clone, Copy, PartialEq)]
pub enum StreamEvent<'a> {
Text(&'a str),
Thinking(&'a str),
}
impl StreamEvent<'_> {
pub fn has_content(&self) -> bool {
match self {
StreamEvent::Text(text) | StreamEvent::Thinking(text) => !text.is_empty(),
}
}
}
pub type StreamSink<'a> = &'a (dyn Fn(StreamEvent<'_>) + Send + Sync);
pub fn report_whole(sink: StreamSink<'_>, response: &LLMResponse) {
if !response.thinking.is_empty() {
sink(StreamEvent::Thinking(&response.thinking));
}
if !response.content.is_empty() {
sink(StreamEvent::Text(&response.content));
}
}
#[async_trait]
pub trait LLMClient: Send + Sync {
async fn chat(&self, request: &ChatRequest<'_>) -> Result<LLMResponse>;
async fn chat_stream(&self, request: &ChatRequest<'_>, sink: StreamSink<'_>) -> Result<LLMResponse> {
let response = self.chat(request).await?;
report_whole(sink, &response);
Ok(response)
}
async fn list_models(&self) -> Result<Vec<String>> {
anyhow::bail!("provider {:?} cannot list models", self.provider_name())
}
async fn detect_context_window(&self) -> Option<DetectedWindow> {
None
}
fn model_name(&self) -> &str;
fn provider_name(&self) -> &str;
fn kind(&self) -> Option<crate::providers::ProviderKind> {
None
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct DetectedWindow {
pub tokens: usize,
pub source: String,
pub cap: ContextCap,
pub total_tokens: Option<usize>,
}
impl DetectedWindow {
pub fn total(tokens: usize, source: impl Into<String>) -> Self {
DetectedWindow { tokens, source: source.into(), cap: ContextCap::Total, total_tokens: None }
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
pub enum ContextCap {
#[default]
Total,
Prompt,
}
#[derive(Debug, Default)]
pub struct ThinkSplitter {
in_think: bool,
pending: String,
}
impl ThinkSplitter {
pub fn push(&mut self, text: &str, emit: &mut dyn FnMut(bool, &str)) {
self.pending.push_str(text);
loop {
let tag = if self.in_think { "</think>" } else { "<think>" };
if let Some(at) = self.pending.find(tag) {
if at > 0 {
emit(self.in_think, &self.pending[..at]);
}
self.pending.drain(..at + tag.len());
self.in_think = !self.in_think;
continue;
}
let keep = (1..tag.len())
.rev()
.find(|&n| {
self.pending.len() >= n
&& self.pending.is_char_boundary(self.pending.len() - n)
&& tag.starts_with(&self.pending[self.pending.len() - n..])
})
.unwrap_or(0);
let split = self.pending.len() - keep;
if split > 0 {
emit(self.in_think, &self.pending[..split]);
self.pending.drain(..split);
}
return;
}
}
pub fn finish(&mut self, emit: &mut dyn FnMut(bool, &str)) {
if !self.pending.is_empty() {
let rest = std::mem::take(&mut self.pending);
emit(self.in_think, &rest);
}
}
pub fn split_all(text: &str) -> (String, String) {
if !text.contains("<think>") {
return (text.to_string(), String::new());
}
let (mut content, mut thinking) = (String::new(), String::new());
let mut splitter = Self::default();
let mut emit =
|think: bool, piece: &str| if think { thinking.push_str(piece) } else { content.push_str(piece) };
splitter.push(text, &mut emit);
splitter.finish(&mut emit);
(content.trim_start().to_string(), thinking.trim().to_string())
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn trajectory_fields_round_trip_and_stay_out_of_plain_messages() {
let message = Message {
thinking: "because".into(),
usage: Some(TokenUsage { prompt_tokens: 10, completion_tokens: 5, total_tokens: 15, aic: None }),
duration_ms: Some(42),
..Message::assistant("hi")
};
let json = serde_json::to_value(&message).unwrap();
assert_eq!(json["thinking"], "because");
assert_eq!(json["usage"], serde_json::json!({"prompt_tokens": 10, "completion_tokens": 5, "total_tokens": 15}));
assert_eq!(json["duration_ms"], 42);
assert_eq!(serde_json::from_value::<Message>(json).unwrap(), message);
let plain = serde_json::to_value(Message::assistant("hi")).unwrap();
assert_eq!(plain, serde_json::json!({"role": "assistant", "content": "hi"}));
assert_eq!(serde_json::from_value::<Message>(plain).unwrap(), Message::assistant("hi"));
}
#[test]
fn splits_think_tags_across_chunks() {
let mut splitter = ThinkSplitter::default();
let (mut content, mut thinking) = (String::new(), String::new());
let mut emit =
|think: bool, piece: &str| if think { thinking.push_str(piece) } else { content.push_str(piece) };
for chunk in ["<th", "ink>plan ", "it</thi", "nk>Answer <b>", "</b> done"] {
splitter.push(chunk, &mut emit);
}
splitter.finish(&mut emit);
assert_eq!(thinking, "plan it");
assert_eq!(content, "Answer <b></b> done");
assert_eq!(ThinkSplitter::split_all("<think>x</think>\n\nhi"), ("hi".into(), "x".into()));
assert_eq!(ThinkSplitter::split_all("plain"), ("plain".into(), String::new()));
}
#[test]
fn reads_copilot_aic_from_nano() {
let value = serde_json::json!({
"usage": {"prompt_tokens": 8, "completion_tokens": 10, "total_tokens": 18},
"copilot_usage": {"total_nano_aiu": 11_600_000}
});
assert_eq!(copilot_aic(&value), Some(0.0116));
assert_eq!(copilot_aic(&serde_json::json!({"usage": {}})), None);
assert_eq!(copilot_aic(&serde_json::json!({"copilot_usage": {"total_nano_aiu": 0}})), Some(0.0));
}
#[test]
fn decodes_valid_arguments_and_marks_invalid_ones() {
let valid = ToolCall::from_raw_arguments("c".into(), "write_file".into(), r#"{"path":"/tmp/x"}"#, None);
assert_eq!(valid.arguments, json!({"path": "/tmp/x"}));
assert_eq!(valid.invalid_arguments(), None);
assert!(ToolCall::from_raw_arguments("c".into(), "t".into(), "", None).arguments.is_object());
assert!(ToolCall::from_raw_arguments("c".into(), "t".into(), " ", None).invalid_arguments().is_none());
let raw = r#"{"path":"/tmp/x","content":"abc"#;
let call = ToolCall::from_raw_arguments("c1".into(), "write_file".into(), raw, None);
assert_eq!(call.arguments, json!({}), "malformed args leave the object empty");
assert_eq!(call.invalid_arguments(), Some(raw));
let error = call.raw_arguments_error(None).expect("malformed args report an error");
assert!(error.contains("write_file"), "names the tool: {error}");
assert!(error.contains("malformed JSON"), "explains the failure: {error}");
assert!(error.contains("retry the same"), "tells the model to retry: {error}");
assert!(error.contains(r#"{"path":"/tmp/x"#), "shows the raw text: {error}");
let ok = ToolCall {
id: "c2".into(),
name: "write_file".into(),
arguments: json!({"path": "/tmp/x"}),
item_id: None,
malformed_arguments: None,
};
assert!(ok.invalid_arguments().is_none());
assert!(ok.raw_arguments_error(None).is_none());
}
#[test]
fn raw_arguments_error_tailors_advice_to_stop_reason() {
let raw = r#"{"path":"/tmp/x","content":"abc"#;
let call = ToolCall::from_raw_arguments("c1".into(), "write_file".into(), raw, None);
for reason in ["length", "max_tokens", "max_output_tokens"] {
let error = call.raw_arguments_error(Some(reason)).expect("malformed args report an error");
assert!(error.contains("output-token limit"), "names the cause for {reason}: {error}");
assert!(error.contains("smaller call"), "advises shrinking for {reason}: {error}");
assert!(!error.contains("retry the same"), "does not tell it to repeat for {reason}: {error}");
}
for reason in ["tool_use", "incomplete"] {
let error = call.raw_arguments_error(Some(reason)).expect("malformed args report an error");
assert!(error.contains("retry the same"), "offers a retry for {reason}: {error}");
assert!(
error.contains("truncated by the output-token limit"),
"hedges on recurrence for {reason}: {error}"
);
}
}
#[test]
fn raw_arguments_error_truncates_long_payloads() {
let raw = format!(r#"{{"content":"{}"#, "x".repeat(1000));
let call = ToolCall::from_raw_arguments("c".into(), "write_file".into(), &raw, None);
let error = call.raw_arguments_error(None).unwrap();
assert!(error.contains('…'), "truncated payload is marked: {error}");
assert!(error.len() < raw.len() + 400, "error stays bounded");
}
}