use crate::tool::{ToolDefinition, ToolInvocation};
use anyhow::Result;
use async_trait::async_trait;
use futures_util::StreamExt;
use futures_util::stream::BoxStream;
use serde::{Deserialize, Serialize};
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
#[serde(tag = "type", rename_all = "snake_case")]
pub enum ContentPart {
Text {
text: String,
},
Image {
media_type: String,
data: String,
},
}
impl ContentPart {
pub fn is_text(&self) -> bool {
matches!(self, Self::Text { .. })
}
pub fn is_image(&self) -> bool {
matches!(self, Self::Image { .. })
}
pub fn as_text(&self) -> Option<&str> {
match self {
Self::Text { text } => Some(text),
_ => None,
}
}
pub fn text_from_parts(parts: &[ContentPart]) -> String {
parts
.iter()
.filter_map(|p| p.as_text())
.collect::<Vec<_>>()
.join("")
}
}
#[async_trait]
pub trait ModelProvider: Send + Sync {
fn stream(&self, request: ModelRequest) -> ModelStream;
async fn complete(&self, request: ModelRequest) -> Result<ModelResponse> {
let mut stream = self.stream(request);
let mut text = String::new();
while let Some(event) = stream.next().await {
match event? {
ModelStreamEvent::TextDelta { text: delta } => text.push_str(&delta),
ModelStreamEvent::Done => break,
ModelStreamEvent::Status { .. }
| ModelStreamEvent::Usage { .. }
| ModelStreamEvent::ThinkingDelta { .. }
| ModelStreamEvent::ToolCall(_) => {}
}
}
Ok(ModelResponse { text })
}
async fn list_models(&self) -> Result<Vec<String>> {
anyhow::bail!("listing models is not supported by this provider")
}
}
pub type ModelStream = BoxStream<'static, Result<ModelStreamEvent>>;
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ModelRequest {
pub model: String,
pub messages: Vec<ModelMessage>,
pub thinking: ThinkingConfig,
#[serde(default)]
pub tools: Vec<ToolDefinition>,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ModelMessage {
pub role: ModelRole,
pub content: String,
#[serde(default, skip_serializing_if = "Vec::is_empty")]
pub content_parts: Vec<ContentPart>,
#[serde(default)]
pub tool_call_id: Option<String>,
#[serde(default)]
pub tool_name: Option<String>,
#[serde(default)]
pub tool_calls: Vec<ToolInvocation>,
#[serde(default, skip_serializing, skip_deserializing)]
pub created_at: Option<u64>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub thinking_content: Option<String>,
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
#[serde(rename_all = "lowercase")]
pub enum ModelRole {
System,
User,
Assistant,
Tool,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ModelResponse {
pub text: String,
}
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
pub enum ModelStreamEvent {
TextDelta {
text: String,
},
ThinkingDelta {
text: String,
},
Status {
label: String,
},
Usage {
input_tokens: Option<u64>,
output_tokens: Option<u64>,
cache_creation_tokens: Option<u64>,
cache_read_tokens: Option<u64>,
},
ToolCall(ToolInvocation),
Done,
}
impl ModelMessage {
pub fn system(content: impl Into<String>) -> Self {
Self::new(ModelRole::System, content)
}
pub fn user(content: impl Into<String>) -> Self {
Self::new(ModelRole::User, content)
}
pub fn user_multimodal(content: impl Into<String>, parts: Vec<ContentPart>) -> Self {
Self {
role: ModelRole::User,
content: content.into(),
content_parts: parts,
tool_call_id: None,
tool_name: None,
tool_calls: Vec::new(),
created_at: Some(current_unix_millis()),
thinking_content: None,
}
}
pub fn assistant(content: impl Into<String>) -> Self {
Self {
thinking_content: None,
..Self::new(ModelRole::Assistant, content)
}
}
pub fn assistant_with_thinking(content: impl Into<String>, thinking: Option<String>) -> Self {
Self {
thinking_content: thinking,
..Self::new(ModelRole::Assistant, content)
}
}
pub fn tool_result(
tool_call_id: impl Into<String>,
tool_name: impl Into<String>,
content: impl Into<String>,
) -> Self {
Self {
role: ModelRole::Tool,
content: content.into(),
content_parts: Vec::new(),
tool_call_id: Some(tool_call_id.into()),
tool_name: Some(tool_name.into()),
tool_calls: Vec::new(),
created_at: Some(current_unix_millis()),
thinking_content: None,
}
}
pub fn assistant_tool_call(invocation: ToolInvocation) -> Self {
Self::assistant_tool_call_with_context(invocation, String::new(), None)
}
pub fn assistant_tool_call_with_context(
invocation: ToolInvocation,
content: impl Into<String>,
thinking: Option<String>,
) -> Self {
Self::assistant_tool_calls_with_context(vec![invocation], content, thinking)
}
pub fn assistant_tool_calls_with_context(
invocations: Vec<ToolInvocation>,
content: impl Into<String>,
thinking: Option<String>,
) -> Self {
Self {
role: ModelRole::Assistant,
content: content.into(),
content_parts: Vec::new(),
tool_call_id: None,
tool_name: None,
tool_calls: invocations,
created_at: Some(current_unix_millis()),
thinking_content: thinking,
}
}
fn new(role: ModelRole, content: impl Into<String>) -> Self {
Self {
role,
content: content.into(),
content_parts: Vec::new(),
tool_call_id: None,
tool_name: None,
tool_calls: Vec::new(),
created_at: Some(current_unix_millis()),
thinking_content: None,
}
}
}
fn current_unix_millis() -> u64 {
std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.unwrap_or_default()
.as_millis() as u64
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
#[serde(rename_all = "lowercase")]
pub enum ThinkingConfig {
Max,
High,
Medium,
Low,
Off,
Adaptive,
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct ThinkingRequest {
pub enabled: bool,
pub effort: Option<&'static str>,
pub budget_tokens: Option<u32>,
}
impl ThinkingConfig {
pub fn to_thinking_request(self) -> ThinkingRequest {
match self {
Self::Max => ThinkingRequest {
enabled: true,
effort: Some("high"),
budget_tokens: Some(32000),
},
Self::High => ThinkingRequest {
enabled: true,
effort: Some("high"),
budget_tokens: Some(10000),
},
Self::Medium => ThinkingRequest {
enabled: true,
effort: Some("medium"),
budget_tokens: Some(4096),
},
Self::Low => ThinkingRequest {
enabled: true,
effort: Some("low"),
budget_tokens: Some(1024),
},
Self::Off => ThinkingRequest {
enabled: false,
effort: None,
budget_tokens: None,
},
Self::Adaptive => Self::Medium.to_thinking_request(),
}
}
pub fn resolve_adaptive(
messages: &[ModelMessage],
tool_names: &[String],
_iteration: usize,
) -> Self {
let complexity = TaskComplexity::classify(messages, tool_names);
match complexity {
TaskComplexity::Simple => Self::Low,
TaskComplexity::Medium => Self::Medium,
TaskComplexity::Complex => Self::High,
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
enum TaskComplexity {
Simple,
Medium,
Complex,
}
impl TaskComplexity {
fn classify(messages: &[ModelMessage], tool_names: &[String]) -> Self {
let mut score: u32 = 0;
let msg_count = messages.len() as u32;
if msg_count > 20 {
score += 2;
} else if msg_count > 8 {
score += 1;
}
let has_write_tools = tool_names.iter().any(|t| {
matches!(
t.as_str(),
"write_file" | "apply_patch" | "write" | "code_edit"
)
});
let has_complex_tools = tool_names.iter().any(|t| matches!(t.as_str(), "bash"));
let has_read_only = tool_names.iter().any(|t| {
matches!(
t.as_str(),
"read_file" | "grep" | "fs_browser" | "read" | "search" | "code"
)
});
if has_write_tools {
score += 2;
}
if has_complex_tools {
score += 1;
}
if has_read_only && !has_write_tools && !has_complex_tools {
score = score.saturating_sub(1);
}
let recent_errors = messages
.iter()
.rev()
.take(6)
.filter(|m| {
m.role == ModelRole::Tool
&& m.content.contains("\"error\"")
&& !m.content.contains("[Old tool result content cleared]")
})
.count();
if recent_errors > 1 {
score += 1;
}
let last_user = messages.iter().rev().find(|m| m.role == ModelRole::User);
if let Some(user_msg) = last_user
&& user_msg.content.len() > 500
{
score += 1;
}
match score {
0..=1 => Self::Simple,
2..=3 => Self::Medium,
_ => Self::Complex,
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn regression_thinking_request_high_produces_effort_and_budget() {
let request = ThinkingConfig::High.to_thinking_request();
assert!(request.enabled);
assert_eq!(request.effort, Some("high"));
assert_eq!(request.budget_tokens, Some(10000));
}
#[test]
fn regression_thinking_request_max_produces_effort_and_budget() {
let request = ThinkingConfig::Max.to_thinking_request();
assert!(request.enabled);
assert_eq!(request.effort, Some("high"));
assert_eq!(request.budget_tokens, Some(32000));
}
#[test]
fn regression_thinking_request_off_produces_disabled() {
let request = ThinkingConfig::Off.to_thinking_request();
assert!(!request.enabled);
assert!(request.effort.is_none());
assert!(request.budget_tokens.is_none());
}
#[test]
fn regression_thinking_request_medium_produces_medium_effort() {
let request = ThinkingConfig::Medium.to_thinking_request();
assert!(request.enabled);
assert_eq!(request.effort, Some("medium"));
assert_eq!(request.budget_tokens, Some(4096));
}
#[test]
fn regression_thinking_request_low_produces_low_effort() {
let request = ThinkingConfig::Low.to_thinking_request();
assert!(request.enabled);
assert_eq!(request.effort, Some("low"));
assert_eq!(request.budget_tokens, Some(1024));
}
#[test]
fn regression_system_message_has_correct_role() {
let msg = ModelMessage::system("test".to_string());
assert_eq!(msg.role, ModelRole::System);
assert_eq!(msg.content, "test");
}
#[test]
fn regression_user_message_has_correct_role() {
let msg = ModelMessage::user("hello".to_string());
assert_eq!(msg.role, ModelRole::User);
assert_eq!(msg.content, "hello");
}
#[test]
fn regression_assistant_message_has_correct_role() {
let msg = ModelMessage::assistant("response".to_string());
assert_eq!(msg.role, ModelRole::Assistant);
assert_eq!(msg.content, "response");
}
#[test]
fn regression_tool_result_sets_call_id_and_name() {
let msg = ModelMessage::tool_result("call-1", "read_file", "content");
assert_eq!(msg.role, ModelRole::Tool);
assert_eq!(msg.tool_call_id.as_deref(), Some("call-1"));
assert_eq!(msg.tool_name.as_deref(), Some("read_file"));
assert_eq!(msg.content, "content");
}
#[test]
fn regression_assistant_tool_call_with_context_sets_fields() {
let inv = ToolInvocation {
id: "call-1".to_string(),
tool_name: "read_file".to_string(),
input: serde_json::json!({"path": "test.rs"}),
};
let msg = ModelMessage::assistant_tool_call_with_context(
inv,
"thinking text",
Some("reasoning".to_string()),
);
assert_eq!(msg.role, ModelRole::Assistant);
assert_eq!(msg.content, "thinking text");
assert_eq!(msg.thinking_content.as_deref(), Some("reasoning"));
assert_eq!(msg.tool_calls.len(), 1);
assert_eq!(msg.tool_calls[0].id, "call-1");
}
#[test]
fn regression_model_message_serialization_roundtrip() {
let msg = ModelMessage {
role: ModelRole::Assistant,
content: "hello".to_string(),
content_parts: Vec::new(),
tool_call_id: None,
tool_name: None,
tool_calls: vec![],
thinking_content: Some("thinking".to_string()),
created_at: Some(12345),
};
let json = serde_json::to_string(&msg).unwrap();
let deserialized: ModelMessage = serde_json::from_str(&json).unwrap();
assert_eq!(deserialized.role, msg.role);
assert_eq!(deserialized.content, msg.content);
assert_eq!(deserialized.thinking_content, msg.thinking_content);
assert!(deserialized.created_at.is_none());
}
#[test]
fn adaptive_thinking_simple_read_only_gives_low() {
let messages = vec![
ModelMessage::user("read this file"),
ModelMessage::assistant("ok"),
];
let tools = vec!["read_file".to_string()];
let result = ThinkingConfig::resolve_adaptive(&messages, &tools, 0);
assert_eq!(result, ThinkingConfig::Low);
}
#[test]
fn adaptive_thinking_write_tools_gives_medium_or_high() {
let messages = vec![
ModelMessage::user("refactor this module"),
ModelMessage::assistant("ok"),
];
let tools = vec!["write_file".to_string(), "read_file".to_string()];
let result = ThinkingConfig::resolve_adaptive(&messages, &tools, 0);
assert!(matches!(
result,
ThinkingConfig::Medium | ThinkingConfig::High
));
}
#[test]
fn adaptive_thinking_long_conversation_gives_higher() {
let mut messages = Vec::new();
for i in 0..25 {
messages.push(ModelMessage::user(format!("message {i}")));
messages.push(ModelMessage::assistant(format!("response {i}")));
}
let tools = vec!["bash".to_string(), "write_file".to_string()];
let result = ThinkingConfig::resolve_adaptive(&messages, &tools, 0);
assert!(matches!(
result,
ThinkingConfig::Medium | ThinkingConfig::High
));
}
#[test]
fn adaptive_thinking_errors_increase_complexity() {
let messages = vec![
ModelMessage::user("fix this"),
ModelMessage::tool_result("c1", "bash", "{\"error\": \"failed\"}"),
ModelMessage::tool_result("c2", "bash", "{\"error\": \"failed again\"}"),
];
let tools = vec!["bash".to_string()];
let result = ThinkingConfig::resolve_adaptive(&messages, &tools, 0);
assert!(matches!(
result,
ThinkingConfig::Medium | ThinkingConfig::High
));
}
#[test]
fn adaptive_falls_back_to_medium_in_to_thinking_request() {
let request = ThinkingConfig::Adaptive.to_thinking_request();
assert!(request.enabled);
assert_eq!(request.effort, Some("medium"));
}
#[test]
fn adaptive_serde_roundtrip() {
let config = ThinkingConfig::Adaptive;
let json = serde_json::to_string(&config).unwrap();
assert_eq!(json, "\"adaptive\"");
let deserialized: ThinkingConfig = serde_json::from_str(&json).unwrap();
assert_eq!(deserialized, ThinkingConfig::Adaptive);
}
}