use async_trait::async_trait;
use serde::{Deserialize, Serialize};
use thiserror::Error;
use openclaw_core::types::TokenUsage;
use std::pin::Pin;
#[derive(Error, Debug)]
pub enum ProviderError {
#[error("API error: {status} - {message}")]
Api {
status: u16,
message: String,
},
#[error("Network error: {0}")]
Network(#[from] reqwest::Error),
#[error("Serialization error: {0}")]
Serialization(#[from] serde_json::Error),
#[error("Rate limited, retry after {retry_after_secs} seconds")]
RateLimited {
retry_after_secs: u64,
},
#[error("Invalid configuration: {0}")]
Config(String),
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct CompletionRequest {
pub model: String,
pub messages: Vec<Message>,
pub system: Option<String>,
pub max_tokens: u32,
pub temperature: f32,
pub stop: Option<Vec<String>>,
pub tools: Option<Vec<Tool>>,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct Message {
pub role: Role,
pub content: MessageContent,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
#[serde(rename_all = "lowercase")]
pub enum Role {
User,
Assistant,
System,
Tool,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(untagged)]
pub enum MessageContent {
Text(String),
Blocks(Vec<ContentBlock>),
}
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(tag = "type", rename_all = "snake_case")]
pub enum ContentBlock {
Text {
text: String,
},
Image {
source: ImageSource,
},
ToolUse {
id: String,
name: String,
input: serde_json::Value,
},
ToolResult {
tool_use_id: String,
content: String,
is_error: Option<bool>,
},
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ImageSource {
#[serde(rename = "type")]
pub source_type: String,
pub media_type: String,
pub data: String,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct Tool {
pub name: String,
pub description: String,
pub input_schema: serde_json::Value,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct CompletionResponse {
pub id: String,
pub model: String,
pub content: Vec<ContentBlock>,
pub stop_reason: Option<StopReason>,
pub usage: TokenUsage,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
pub enum StopReason {
EndTurn,
MaxTokens,
StopSequence,
ToolUse,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct StreamingChunk {
pub chunk_type: ChunkType,
pub delta: Option<String>,
pub index: Option<usize>,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
pub enum ChunkType {
MessageStart,
ContentBlockStart,
ContentBlockDelta,
ContentBlockStop,
MessageDelta,
MessageStop,
}
#[async_trait]
pub trait Provider: Send + Sync {
fn name(&self) -> &str;
async fn list_models(&self) -> Result<Vec<String>, ProviderError>;
async fn complete(
&self,
request: CompletionRequest,
) -> Result<CompletionResponse, ProviderError>;
async fn complete_stream(
&self,
request: CompletionRequest,
) -> Result<
Pin<Box<dyn futures::Stream<Item = Result<StreamingChunk, ProviderError>> + Send>>,
ProviderError,
>;
}