use crate::error::CompletionError;
use crate::providers::provider_type::ProviderType;
use crate::providers::zai::{CompletionRequest, Message, ToolDefinition};
use crate::providers::{
AnthropicClient, CerebrasClient, LmStudioClient, MinimaxClient, MlxLmClient, OllamaClient,
OpenAIClient, OpenRouterClient, ZaiClient,
};
use schemars::JsonSchema;
use serde::{de::DeserializeOwned, Serialize};
use thiserror::Error;
#[derive(Error, Debug)]
pub enum ExtractorError {
#[error("Provider error: {0}")]
Provider(#[from] CompletionError),
#[error("Serialization error: {0}")]
Serialization(String),
#[error("No tool call or valid JSON found in response")]
NoToolCall,
}
pub enum ExtractorWrapper {
OpenRouter(OpenRouterClient),
OpenAI(OpenAIClient),
Anthropic(AnthropicClient),
Minimax(MinimaxClient),
Cerebras(CerebrasClient),
Ollama(OllamaClient),
Zai(ZaiClient),
MlxLm(MlxLmClient),
LmStudio(LmStudioClient),
}
impl ExtractorWrapper {
pub async fn extract<T>(
&self,
model: &str,
preamble: &str,
prompt: &str,
) -> Result<T, ExtractorError>
where
T: JsonSchema + DeserializeOwned + Serialize + Send + Sync + 'static,
{
let schema = schemars::schema_for!(T);
let schema_json = serde_json::to_value(&schema).map_err(|e| {
ExtractorError::Serialization(format!("Failed to serialize schema: {}", e))
})?;
let tool = ToolDefinition {
name: "extract".to_string(),
description: "Extract structured data from the provided text".to_string(),
parameters: schema_json,
};
let request = CompletionRequest {
preamble: Some(preamble.to_string()),
messages: vec![Message::user(prompt)],
tools: vec![tool],
temperature: Some(0.0), max_tokens: None,
additional_params: None,
};
let message = match self {
ExtractorWrapper::OpenRouter(client) => {
let completion_model = client.completion_model(model);
let response = completion_model
.completion(request)
.await
.map_err(|e| ExtractorError::Serialization(e.to_string()))?;
response.message
}
ExtractorWrapper::OpenAI(client) => {
let completion_model = client.completion_model(model);
let response = completion_model
.completion(request)
.await
.map_err(|e| ExtractorError::Serialization(e.to_string()))?;
response.message
}
ExtractorWrapper::Anthropic(client) => {
let completion_model = client.completion_model(model);
let response = completion_model
.completion(request)
.await
.map_err(|e| ExtractorError::Serialization(e.to_string()))?;
response.message
}
ExtractorWrapper::Minimax(client) => {
let completion_model = client.completion_model(model);
let response = completion_model
.completion(request)
.await
.map_err(|e| ExtractorError::Serialization(e.to_string()))?;
response.message
}
ExtractorWrapper::Cerebras(client) => {
let completion_model = client.completion_model(model);
let response = completion_model
.completion(request)
.await
.map_err(|e| ExtractorError::Serialization(e.to_string()))?;
response.message
}
ExtractorWrapper::Ollama(client) => {
let completion_model = client.completion_model(model);
let response = completion_model
.completion(request)
.await
.map_err(|e| ExtractorError::Serialization(e.to_string()))?;
response.message
}
ExtractorWrapper::Zai(client) => {
let completion_model = client.completion_model(model);
let response = completion_model
.completion(request)
.await
.map_err(|e| ExtractorError::Serialization(e.to_string()))?;
response.message
}
ExtractorWrapper::MlxLm(client) => {
let completion_model = client.completion_model(model);
let response = completion_model
.completion(request)
.await
.map_err(|e| ExtractorError::Serialization(e.to_string()))?;
response.message
}
ExtractorWrapper::LmStudio(client) => {
let completion_model = client.completion_model(model);
let response = completion_model
.completion(request)
.await
.map_err(|e| ExtractorError::Serialization(e.to_string()))?;
response.message
}
};
let arguments = Self::extract_tool_arguments(&message)?;
serde_json::from_str(&arguments).map_err(|e| {
ExtractorError::Serialization(format!(
"Failed to deserialize extracted data: {}. Raw arguments: {}",
e, arguments
))
})
}
fn extract_tool_arguments(message: &Message) -> Result<String, ExtractorError> {
if let Some(tool_call) = message.tool_calls.as_ref().and_then(|calls| calls.first()) {
return Ok(tool_call.arguments.clone());
}
if !message.content.is_empty() {
if let Some(json) = Self::extract_json_from_content(&message.content) {
return Ok(json);
}
}
Err(ExtractorError::NoToolCall)
}
fn extract_json_from_content(content: &str) -> Option<String> {
if serde_json::from_str::<serde_json::Value>(content).is_ok() {
return Some(content.to_string());
}
let trimmed = content.trim();
if trimmed.starts_with("```json") {
if let Some(end) = trimmed.rfind("```") {
let json_start = trimmed.find('\n').map(|i| i + 1).unwrap_or(7);
if json_start < end {
let json_str = trimmed[json_start..end].trim();
if serde_json::from_str::<serde_json::Value>(json_str).is_ok() {
return Some(json_str.to_string());
}
}
}
}
if let Some(start) = content.find('{') {
if let Some(end) = content.rfind('}') {
if start < end {
let json_str = &content[start..=end];
if serde_json::from_str::<serde_json::Value>(json_str).is_ok() {
return Some(json_str.to_string());
}
}
}
}
None
}
pub fn provider_type(&self) -> ProviderType {
match self {
ExtractorWrapper::OpenRouter(_) => ProviderType::OpenRouter,
ExtractorWrapper::OpenAI(_) => ProviderType::OpenAI,
ExtractorWrapper::Anthropic(_) => ProviderType::Anthropic,
ExtractorWrapper::Minimax(_) => ProviderType::Minimax,
ExtractorWrapper::Cerebras(_) => ProviderType::Cerebras,
ExtractorWrapper::Ollama(_) => ProviderType::Ollama,
ExtractorWrapper::Zai(_) => ProviderType::Zai,
ExtractorWrapper::MlxLm(_) => ProviderType::MlxLm,
ExtractorWrapper::LmStudio(_) => ProviderType::LmStudio,
}
}
}