use crate::models::ModelInfo;
use crate::tool_choice::ToolChoice;
use mcp_protocol::tool::{Tool, ToolContent};
use serde::{Deserialize, Serialize};
#[derive(Serialize, Deserialize, Debug, Clone)]
#[serde(tag = "type", rename_all = "snake_case")]
pub enum MessageContent {
Text {
text: String,
},
ToolUse {
id: String,
name: String,
input: serde_json::Value,
},
ToolResult {
tool_use_id: String,
content: Vec<ToolContent>,
is_error: Option<bool>,
},
}
#[derive(Serialize, Deserialize, Debug, Clone)]
pub struct Message {
pub role: Role,
pub content: Vec<MessageContent>,
}
#[derive(Serialize, Deserialize, Debug, Clone)]
pub enum Role {
#[serde(rename = "user")]
User,
#[serde(rename = "assistant")]
Assistant,
#[serde(rename = "system")]
System,
}
impl Message {
pub fn new_structured(role: impl Into<String>, content: Vec<MessageContent>) -> Self {
Self {
role: match role.into().as_str() {
"user" => Role::User,
"assistant" => Role::Assistant,
"system" => Role::System,
_ => panic!("Invalid role"),
},
content,
}
}
}
#[derive(Serialize, Deserialize, Debug, Clone)]
pub struct CompletionRequest {
pub model: String,
pub messages: Vec<Message>,
pub max_tokens: u32,
pub temperature: Option<f32>,
pub system: Option<String>,
pub tools: Option<Vec<Tool>>,
pub tool_choice: Option<ToolChoice>,
pub disable_parallel_tool_use: Option<bool>,
}
#[derive(Serialize, Deserialize, Debug, Clone)]
pub struct Usage {
pub input_tokens: u32,
pub output_tokens: u32,
}
#[derive(Serialize, Deserialize, Debug, Clone)]
pub struct CompletionResponse {
pub content: Vec<MessageContent>,
pub id: String,
pub model: String,
pub role: Role,
pub stop_reason: StopReason,
pub stop_sequence: Option<String>,
#[serde(rename = "type")]
pub message_type: String,
pub usage: Usage,
}
impl From<CompletionResponse> for Message {
fn from(request: CompletionResponse) -> Self {
Self {
role: request.role,
content: request.content,
}
}
}
#[derive(Serialize, Deserialize, Debug, Clone, PartialEq, Eq)]
pub enum StopReason {
#[serde(rename = "end_turn")]
EndTurn,
#[serde(rename = "max_tokens")]
MaxTokens,
#[serde(rename = "stop_sequence")]
StopSequence,
#[serde(rename = "tool_use")]
ToolUse,
#[serde(rename = "other")]
Other(String),
}
#[derive(Serialize, Deserialize, Debug, Clone)]
pub enum ProxyRequest {
ListModels,
GenerateCompletion { request: CompletionRequest },
}
#[derive(Serialize, Deserialize, Debug, Clone)]
pub enum ProxyResponse {
ListModels { models: Vec<ModelInfo> },
Completion { completion: CompletionResponse },
Error { error: String },
}