ai-lib 0.4.0

A unified AI SDK for Rust providing a single interface for multiple AI providers with hybrid architecture
Documentation
use crate::api::ChatProvider;
use crate::types::AiLibError;
use crate::types::{ChatCompletionRequest, ChatCompletionResponse, Message, Role};
use async_trait::async_trait;
use futures::Stream;
use reqwest::Client;
use serde::{Deserialize, Serialize};

/// AI21 API adapter
///
/// AI21 provides Jurassic series models with a custom API format.
/// Documentation: <https://docs.ai21.com/reference/introduction>
pub struct AI21Adapter {
    client: Client,
    api_key: String,
}

impl AI21Adapter {
    pub fn new() -> Result<Self, AiLibError> {
        let api_key = std::env::var("AI21_API_KEY").map_err(|_| {
            AiLibError::ConfigurationError("AI21_API_KEY environment variable not set".to_string())
        })?;

        Ok(Self {
            client: Client::new(),
            api_key,
        })
    }

    pub async fn chat_completion(
        &self,
        request: ChatCompletionRequest,
    ) -> Result<ChatCompletionResponse, AiLibError> {
        let ai21_request = self.convert_request(&request)?;

        let response = self
            .client
            .post("https://api.ai21.com/studio/v1/chat/completions")
            .header("Authorization", format!("Bearer {}", self.api_key))
            .header("Content-Type", "application/json")
            .json(&ai21_request)
            .send()
            .await
            .map_err(|e| AiLibError::NetworkError(format!("AI21 API request failed: {}", e)))?;

        if !response.status().is_success() {
            let status = response.status();
            let error_text = response
                .text()
                .await
                .unwrap_or_else(|_| "Unknown error".to_string());
            return Err(AiLibError::ProviderError(format!(
                "AI21 API error {}: {}",
                status, error_text
            )));
        }

        let ai21_response: AI21Response = response.json().await.map_err(|e| {
            AiLibError::DeserializationError(format!("Failed to parse AI21 response: {}", e))
        })?;

        self.convert_response(ai21_response)
    }

    pub async fn chat_completion_stream(
        &self,
        request: ChatCompletionRequest,
    ) -> Result<
        Box<
            dyn futures::Stream<Item = Result<crate::api::ChatCompletionChunk, AiLibError>>
                + Send
                + Unpin,
        >,
        AiLibError,
    > {
        let mut ai21_request = self.convert_request(&request)?;
        ai21_request.stream = Some(true);

        let response = self
            .client
            .post("https://api.ai21.com/studio/v1/chat/completions")
            .header("Authorization", format!("Bearer {}", self.api_key))
            .header("Content-Type", "application/json")
            .json(&ai21_request)
            .send()
            .await
            .map_err(|e| AiLibError::NetworkError(format!("AI21 API request failed: {}", e)))?;

        if !response.status().is_success() {
            let status = response.status();
            let error_text = response
                .text()
                .await
                .unwrap_or_else(|_| "Unknown error".to_string());
            return Err(AiLibError::ProviderError(format!(
                "AI21 API error {}: {}",
                status, error_text
            )));
        }

        // For now, convert streaming request to non-streaming and return a single chunk
        let response = self.chat_completion(request.clone()).await?;

        // Create a single chunk from the response
        let chunk = crate::api::ChatCompletionChunk {
            id: response.id.clone(),
            object: "chat.completion.chunk".to_string(),
            created: response.created,
            model: response.model.clone(),
            choices: response
                .choices
                .into_iter()
                .map(|choice| crate::api::ChoiceDelta {
                    index: choice.index,
                    delta: crate::api::MessageDelta {
                        role: Some(choice.message.role),
                        content: Some(match &choice.message.content {
                            crate::Content::Text(text) => text.clone(),
                            _ => "".to_string(),
                        }),
                    },
                    finish_reason: choice.finish_reason,
                })
                .collect(),
        };

        let stream = futures::stream::once(async move { Ok(chunk) });
        Ok(Box::new(Box::pin(stream)))
    }

    fn convert_request(&self, request: &ChatCompletionRequest) -> Result<AI21Request, AiLibError> {
        // Convert messages to AI21 format
        let messages = request
            .messages
            .iter()
            .map(|msg| AI21Message {
                role: match msg.role {
                    Role::System => "system".to_string(),
                    Role::User => "user".to_string(),
                    Role::Assistant => "assistant".to_string(),
                },
                content: match &msg.content {
                    crate::Content::Text(text) => text.clone(),
                    _ => "Unsupported content type".to_string(),
                },
            })
            .collect();

        Ok(AI21Request {
            model: request.model.clone(),
            messages,
            max_tokens: request.max_tokens,
            temperature: request.temperature,
            top_p: request.top_p,
            stream: Some(false),
            extensions: request.extensions.clone(),
        })
    }

    fn convert_response(
        &self,
        response: AI21Response,
    ) -> Result<ChatCompletionResponse, AiLibError> {
        let choice = response.choices.first().ok_or_else(|| {
            AiLibError::InvalidModelResponse("No choices in AI21 response".to_string())
        })?;

        let message = Message {
            role: match choice.message.role.as_str() {
                "assistant" => Role::Assistant,
                "user" => Role::User,
                "system" => Role::System,
                _ => Role::Assistant,
            },
            content: crate::Content::Text(choice.message.content.clone().unwrap_or_default()),
            function_call: None,
        };

        Ok(ChatCompletionResponse {
            id: response.id,
            object: "chat.completion".to_string(),
            created: response.created,
            model: response.model,
            choices: vec![crate::types::Choice {
                index: 0,
                message,
                finish_reason: choice.finish_reason.clone(),
            }],
            usage: response
                .usage
                .map(|u| crate::types::Usage {
                    prompt_tokens: u.prompt_tokens,
                    completion_tokens: u.completion_tokens,
                    total_tokens: u.total_tokens,
                })
                .unwrap_or_else(|| crate::types::Usage {
                    prompt_tokens: 0,
                    completion_tokens: 0,
                    total_tokens: 0,
                }),
            usage_status: crate::types::response::UsageStatus::Finalized,
        })
    }
}

#[async_trait]
impl ChatProvider for AI21Adapter {
    fn name(&self) -> &str {
        "AI21"
    }

    async fn chat(
        &self,
        request: ChatCompletionRequest,
    ) -> Result<ChatCompletionResponse, AiLibError> {
        self.chat_completion(request).await
    }

    async fn stream(
        &self,
        request: ChatCompletionRequest,
    ) -> Result<
        Box<dyn Stream<Item = Result<crate::api::ChatCompletionChunk, AiLibError>> + Send + Unpin>,
        AiLibError,
    > {
        self.chat_completion_stream(request).await
    }

    async fn list_models(&self) -> Result<Vec<String>, AiLibError> {
        // Return default models for AI21
        Ok(vec![
            "j2-ultra".to_string(),
            "j2-mid".to_string(),
            "j2-light".to_string(),
        ])
    }

    async fn get_model_info(&self, model_id: &str) -> Result<crate::api::ModelInfo, AiLibError> {
        Ok(crate::api::ModelInfo {
            id: model_id.to_string(),
            object: "model".to_string(),
            created: 0,
            owned_by: "ai21".to_string(),
            permission: vec![],
        })
    }
}

#[derive(Serialize)]
struct AI21Request {
    model: String,
    messages: Vec<AI21Message>,
    max_tokens: Option<u32>,
    temperature: Option<f32>,
    top_p: Option<f32>,
    stream: Option<bool>,
    #[serde(flatten)]
    extensions: Option<serde_json::Map<String, serde_json::Value>>,
}

#[derive(Serialize)]
struct AI21Message {
    role: String,
    content: String,
}

#[derive(Deserialize)]
struct AI21Response {
    id: String,
    #[allow(dead_code)]
    object: String,
    created: u64,
    model: String,
    choices: Vec<AI21Choice>,
    usage: Option<AI21Usage>,
}

#[derive(Deserialize)]
struct AI21Choice {
    message: AI21MessageResponse,
    finish_reason: Option<String>,
}

#[derive(Deserialize)]
struct AI21MessageResponse {
    role: String,
    content: Option<String>,
}

#[derive(Deserialize)]
struct AI21Usage {
    prompt_tokens: u32,
    completion_tokens: u32,
    total_tokens: u32,
}