use crate::{
adapters::{AnthropicAdapter, GeminiAdapter, OpenAIResponsesAdapter},
adapter::{AdapterError, ChatAdapter},
stream::*,
types::*,
};
use async_trait::async_trait;
use futures_util::StreamExt;
use std::collections::HashMap;
use tokio_util::sync::CancellationToken;
pub struct OpenRouterAdapter;
#[async_trait]
impl ChatAdapter for OpenRouterAdapter {
fn provider_kind(&self) -> ProviderKind {
ProviderKind::OpenRouter
}
async fn discover_models(
&self,
provider_name: &str,
endpoint: &ProviderEndpoint,
) -> Result<Vec<DiscoveredModel>, AdapterError> {
let client = reqwest::Client::new();
let url = format!("{}/v1/models/user", endpoint.base_url);
let mut request = client.get(&url);
if let Some(timeout) = endpoint.timeout {
request = request.timeout(std::time::Duration::from_millis(timeout));
}
if let Some(api_key) = &endpoint.api_key {
request = request.header("Authorization", format!("Bearer {}", api_key));
}
for (key, value) in &endpoint.extra_headers {
request = request.header(key, value);
}
let resp = request
.send()
.await
.map_err(|e| AdapterError::Http(format!("Failed to fetch models: {}", e)))?;
if !resp.status().is_success() {
let status = resp.status();
let text = resp
.text()
.await
.unwrap_or_else(|_| "Unknown error".to_string());
return Err(AdapterError::Provider {
code: status.as_u16().to_string(),
message: text,
});
}
let models_response: OpenRouterModelsResponse = resp
.json()
.await
.map_err(|e| AdapterError::Http(format!("Failed to parse models response: {}", e)))?;
let discovered_models: Vec<DiscoveredModel> = models_response
.data
.into_iter()
.map(|model| {
let capabilities = self.parse_model_capabilities_internal(&model);
DiscoveredModel {
id: format!("{}/{}", provider_name.to_lowercase(), model.id),
name: model.name,
provider_name: provider_name.to_string(),
provider_kind: ProviderKind::OpenRouter,
input_modalities: capabilities.input_modalities,
output_modalities: capabilities.output_modalities,
capabilities: capabilities.capabilities,
context_length: capabilities.context_length,
max_tokens: capabilities.max_tokens,
}
})
.collect();
Ok(discovered_models)
}
async fn execute_chat(
&self,
ir: ChatRequestIR,
cancel: CancellationToken,
) -> Result<Box<dyn futures_util::Stream<Item = StreamEvent> + Send + Unpin>, AdapterError>
{
let payload = self.build_openai_request(&ir)?;
let client = reqwest::Client::new();
let url = format!(
"{}/v1/chat/completions",
ir.model.provider.endpoint.base_url
);
let mut request = client.post(&url).json(&payload);
if let Some(timeout) = ir.model.provider.endpoint.timeout {
request = request.timeout(std::time::Duration::from_millis(timeout));
}
if let Some(api_key) = &ir.model.provider.endpoint.api_key {
request = request.header("Authorization", format!("Bearer {}", api_key));
}
for (key, value) in &ir.model.provider.endpoint.extra_headers {
request = request.header(key, value);
}
let mut resp = request
.send()
.await
.map_err(|e| AdapterError::Http(format!("Failed to send request: {}", e)))?;
if !resp.status().is_success() {
let status = resp.status();
let text = resp
.text()
.await
.unwrap_or_else(|_| "Unknown error".to_string());
if let Ok(error_response) = serde_json::from_str::<OpenAIErrorResponse>(&text) {
return Err(AdapterError::Provider {
code: error_response
.error
.code
.unwrap_or_else(|| status.as_u16().to_string()),
message: error_response.error.message,
});
}
return Err(AdapterError::Provider {
code: status.as_u16().to_string(),
message: text,
});
}
if ir.stream {
let s = async_stream::try_stream! {
use crate::sse::SseParser;
let mut tool_calls_buffer: HashMap<u32, OpenAIToolCall> = HashMap::new();
let mut sse_parser = SseParser::new();
while let Some(chunk) = resp.chunk().await
.map_err(|e| AdapterError::Http(format!("Failed to read chunk: {}", e)))?
{
if cancel.is_cancelled() {
yield StreamEvent::Error {
code: "cancelled".to_string(),
message: "Request was cancelled".to_string(),
};
break;
}
let chunk_str = String::from_utf8_lossy(&chunk);
let events = sse_parser.feed(&chunk_str);
for sse_event in events {
let json_str = &sse_event.data;
if json_str == "[DONE]" {
for tool_call in tool_calls_buffer.values() {
let args_json = serde_json::from_str(&tool_call.function.arguments)
.unwrap_or(serde_json::json!({}));
yield StreamEvent::ToolCallEnd {
id: tool_call.id.clone(),
args_json,
};
}
yield StreamEvent::Done;
return;
}
if let Ok(response) = serde_json::from_str::<OpenAIChatResponse>(json_str) {
if let Some(choice) = response.choices.first() {
if let Some(delta) = &choice.delta {
if let Some(content) = &delta.content {
yield StreamEvent::TextDelta {
content: content.clone(),
};
}
if let Some(tool_calls) = &delta.tool_calls {
for tool_call_delta in tool_calls {
let index = tool_call_delta.index;
if let Some(id) = &tool_call_delta.id {
tool_calls_buffer.insert(index, OpenAIToolCall {
id: id.clone(),
r#type: tool_call_delta.r#type.clone().unwrap_or_else(|| "function".to_string()),
function: OpenAIFunctionCall {
name: tool_call_delta.function.as_ref().and_then(|f| f.name.clone()).unwrap_or_default(),
arguments: String::new(),
},
});
yield StreamEvent::ToolCallStart {
id: id.clone(),
name: tool_call_delta.function.as_ref().and_then(|f| f.name.clone()).unwrap_or_default(),
args_json: serde_json::Value::Object(serde_json::Map::new()),
};
}
if let Some(tool_call) = tool_calls_buffer.get_mut(&index) {
if let Some(function) = &tool_call_delta.function {
if let Some(args_delta) = &function.arguments {
tool_call.function.arguments.push_str(args_delta);
yield StreamEvent::ToolCallDelta {
id: tool_call.id.clone(),
args_delta_json: serde_json::Value::String(args_delta.clone()),
};
}
}
}
}
}
}
}
if let Some(usage) = response.usage {
yield StreamEvent::Tokens {
input: usage.prompt_tokens,
output: usage.completion_tokens,
};
}
}
}
}
for tool_call in tool_calls_buffer.values() {
let args_json = serde_json::from_str(&tool_call.function.arguments)
.unwrap_or(serde_json::json!({}));
yield StreamEvent::ToolCallEnd {
id: tool_call.id.clone(),
args_json,
};
}
yield StreamEvent::Done;
};
Ok(Box::new(Box::pin(s.map(
|r: Result<StreamEvent, AdapterError>| match r {
Ok(ev) => ev,
Err(e) => StreamEvent::Error {
code: "stream_error".to_string(),
message: e.to_string(),
},
},
))))
} else {
let response: OpenAIChatResponse = resp
.json()
.await
.map_err(|e| AdapterError::Http(format!("Failed to parse response: {}", e)))?;
let s = async_stream::try_stream! {
if let Some(choice) = response.choices.first() {
if let Some(message) = &choice.message {
if let Some(content) = &message.content {
yield StreamEvent::TextDelta {
content: content.clone(),
};
}
if let Some(tool_calls) = &message.tool_calls {
for tool_call in tool_calls {
yield StreamEvent::ToolCallStart {
id: tool_call.id.clone(),
name: tool_call.function.name.clone(),
args_json: serde_json::Value::Object(serde_json::Map::new()),
};
yield StreamEvent::ToolCallDelta {
id: tool_call.id.clone(),
args_delta_json: serde_json::Value::String(tool_call.function.arguments.clone()),
};
let args_json = serde_json::from_str(&tool_call.function.arguments)
.unwrap_or(serde_json::json!({}));
yield StreamEvent::ToolCallEnd {
id: tool_call.id.clone(),
args_json,
};
}
}
}
let (prompt_details, completion_details) = if let Some(ref usage) = response.usage {
yield StreamEvent::Tokens {
input: usage.prompt_tokens,
output: usage.completion_tokens,
};
(usage.prompt_tokens_details.clone(), usage.completion_tokens_details.clone())
} else {
(None, None)
};
yield StreamEvent::OpenAIMetadata {
system_fingerprint: response.system_fingerprint,
service_tier: response.service_tier,
prompt_tokens_details: prompt_details,
completion_tokens_details: completion_details,
};
yield StreamEvent::Done;
}
};
Ok(Box::new(Box::pin(s.map(
|r: Result<StreamEvent, AdapterError>| match r {
Ok(ev) => ev,
Err(e) => StreamEvent::Error {
code: "response_error".to_string(),
message: e.to_string(),
},
},
))))
}
}
}
impl OpenRouterAdapter {
fn normalize_messages(messages: &[Message]) -> Vec<Message> {
let mut normalized: Vec<Message> = Vec::new();
let mut pending_tools: Vec<(usize, Message)> = Vec::new();
for (idx, msg) in messages.iter().enumerate() {
if msg.role == Role::Tool {
pending_tools.push((idx, msg.clone()));
} else if msg.role == Role::Assistant {
if !normalized.is_empty() {
let prev_role = &normalized.last().unwrap().role;
if prev_role == &Role::User
|| prev_role == &Role::System
|| prev_role == &Role::Developer
{
pending_tools.sort_by_key(|(i, _)| *i);
for (_, tool_msg) in pending_tools.drain(..) {
normalized.push(tool_msg);
}
}
}
normalized.push(msg.clone());
if !pending_tools.is_empty() {
pending_tools.sort_by_key(|(i, _)| *i);
for (_, tool_msg) in pending_tools.drain(..) {
normalized.push(tool_msg);
}
}
} else {
if !pending_tools.is_empty() {
if let Some(last) = normalized.last() {
if last.role == Role::Assistant {
pending_tools.sort_by_key(|(i, _)| *i);
for (_, tool_msg) in pending_tools.drain(..) {
normalized.push(tool_msg);
}
}
}
}
normalized.push(msg.clone());
}
}
if !pending_tools.is_empty() {
pending_tools.sort_by_key(|(i, _)| *i);
for (_, tool_msg) in pending_tools.drain(..) {
normalized.push(tool_msg);
}
}
normalized
}
fn build_openai_request(&self, ir: &ChatRequestIR) -> Result<OpenAIChatRequest, AdapterError> {
let normalized_messages = Self::normalize_messages(&ir.messages);
let mut required_tool_calls: HashMap<usize, Vec<OpenAIToolCall>> = HashMap::new();
for (idx, msg) in normalized_messages.iter().enumerate() {
if msg.role == Role::Tool {
let mut assistant_idx = None;
for i in (0..idx).rev() {
if normalized_messages[i].role == Role::Assistant {
assistant_idx = Some(i);
break;
}
}
if let Some(a_idx) = assistant_idx {
let name_field = msg.name.clone().unwrap_or_default();
let (tool_name, tool_call_id) = if let Some(colon_pos) = name_field.rfind(':') {
(
name_field[..colon_pos].to_string(),
name_field[colon_pos + 1..].to_string(),
)
} else {
("unknown_tool".to_string(), name_field)
};
let assistant_msg = &normalized_messages[a_idx];
let has_tool_call = assistant_msg.parts.iter().any(|p| match p {
ContentPart::ToolCall { id, .. } => id == &tool_call_id,
_ => false,
});
if !has_tool_call {
required_tool_calls
.entry(a_idx)
.or_default()
.push(OpenAIToolCall {
id: tool_call_id,
r#type: "function".to_string(),
function: OpenAIFunctionCall {
name: tool_name,
arguments: "{}".to_string(), },
});
}
}
}
}
let messages: Vec<OpenAIMessage> = normalized_messages
.iter()
.enumerate()
.map(|(idx, msg)| {
let mut text_content = String::new();
let mut has_images = false;
let mut content_parts: Vec<OpenAIContentPart> = Vec::new();
let mut tool_calls_out = Vec::new();
for part in &msg.parts {
match part {
ContentPart::Text(text) => {
text_content.push_str(text);
content_parts.push(OpenAIContentPart {
kind: "text".to_string(),
text: Some(text.clone()),
image_url: None,
audio: None,
file: None,
});
}
ContentPart::ImageUrl { url, mime: _ } => {
has_images = true;
content_parts.push(OpenAIContentPart {
kind: "image_url".to_string(),
text: None,
image_url: Some(crate::OpenAIImageUrl::Obj {
url: url.clone(),
detail: Some("auto".to_string()),
}),
audio: None,
file: None,
});
}
ContentPart::BlobRef { .. } => {}
ContentPart::Audio { .. } => {}
ContentPart::File { .. } => {}
ContentPart::ToolCall {
id,
name,
arguments,
} => {
tool_calls_out.push(OpenAIToolCall {
id: id.clone(),
r#type: "function".to_string(),
function: OpenAIFunctionCall {
name: name.clone(),
arguments: arguments.clone(),
},
});
}
}
}
if let Some(missing_tools) = required_tool_calls.get(&idx) {
tool_calls_out.extend(missing_tools.clone());
}
let content = if has_images {
crate::OpenAIMessageContent::Parts(content_parts)
} else if !text_content.is_empty() {
crate::OpenAIMessageContent::Text(text_content)
} else {
crate::OpenAIMessageContent::Text(String::new())
};
let role = match msg.role {
Role::System => "system",
Role::User => "user",
Role::Assistant => "assistant",
Role::Tool => "tool",
Role::Developer => "system",
};
let tool_call_id = if msg.role == Role::Tool {
let name_field = msg.name.clone().unwrap_or_default();
if let Some(colon_pos) = name_field.rfind(':') {
Some(name_field[colon_pos + 1..].to_string())
} else {
Some(name_field)
}
} else {
None
};
OpenAIMessage {
role: role.to_string(),
content,
name: msg.name.clone(),
tool_calls: if tool_calls_out.is_empty() {
None
} else {
Some(tool_calls_out)
},
tool_call_id,
}
})
.collect();
let tools = if ir.tools.is_empty() {
None
} else {
Some(
ir.tools
.iter()
.map(|tool| match tool {
ToolSpec::JsonSchema {
name,
description,
schema,
strict: _,
} => OpenAITool {
r#type: "function".to_string(),
function: OpenAIFunction {
name: name.clone(),
description: description.clone(),
parameters: schema.clone(),
},
},
})
.collect(),
)
};
let _tool_choice = match &ir.tool_choice {
ToolChoice::Auto => Some(serde_json::json!("auto")),
ToolChoice::None => Some(serde_json::json!("none")),
ToolChoice::Required => Some(serde_json::json!("required")),
ToolChoice::Named(name) => Some(serde_json::json!({
"type": "function",
"function": { "name": name }
})),
ToolChoice::Allowed { .. } => Some(serde_json::json!("auto")), };
Ok(OpenAIChatRequest {
model: self.resolve_adapter_model_id(&ir.model.model_id, &ir.model.provider.name),
messages,
temperature: ir.sampling.temperature,
top_p: ir.sampling.top_p,
max_tokens: None,
max_completion_tokens: ir.sampling.max_tokens,
stream: Some(ir.stream),
stop: if ir.sampling.stop.is_empty() {
None
} else {
Some(crate::OpenAIStop::Many(ir.sampling.stop.clone()))
},
presence_penalty: ir.sampling.presence_penalty,
frequency_penalty: ir.sampling.frequency_penalty,
tools: tools.clone(),
tool_choice: None, functions: None,
function_call: None,
response_format: None,
logit_bias: None,
logprobs: None,
top_logprobs: None,
n: None,
seed: None,
user: None,
stream_options: None,
modalities: None,
audio: None,
parallel_tool_calls: if tools.is_some() && !ir.tools.is_empty() {
Some(ir.sampling.parallel_tool_calls.unwrap_or(true))
} else {
None
},
store: None,
metadata: None,
prediction: None,
service_tier: None,
reasoning_effort: ir.reasoning.as_ref().and_then(|r| {
r.effort.as_ref().and_then(|e| match e.as_str() {
"minimal" => Some(crate::types::OpenAIReasoningEffort::Minimal),
"low" => Some(crate::types::OpenAIReasoningEffort::Low),
"medium" => Some(crate::types::OpenAIReasoningEffort::Medium),
"high" => Some(crate::types::OpenAIReasoningEffort::High),
_ => None,
})
}),
verbosity: None,
web_search_options: None,
prompt_cache_key: None,
safety_identifier: None,
})
}
fn parse_reasoning(&self, model_id: &str) -> Vec<ModelCapabilities> {
let (provider, inner_model_id) = if let Some(pos) = model_id.find('/') {
(&model_id[..pos], &model_id[pos + 1..])
} else {
("", model_id)
};
match provider {
"openai" => OpenAIResponsesAdapter.parse_reasoning(inner_model_id),
"anthropic" => AnthropicAdapter.parse_reasoning(inner_model_id),
"google" => GeminiAdapter.parse_reasoning(inner_model_id),
_ => Vec::new(),
}
}
fn parse_model_capabilities(&self, model_info: &str) -> ModelCapabilitiesWithModalities {
if let Ok(model) = serde_json::from_str::<OpenRouterModel>(model_info) {
self.parse_model_capabilities_internal(&model)
} else {
ModelCapabilitiesWithModalities {
context_length: None,
max_tokens: None,
capabilities: Vec::new(),
input_modalities: vec![Modality::Text],
output_modalities: vec![Modality::Text],
}
}
}
}
impl OpenRouterAdapter {
fn parse_model_capabilities_internal(
&self,
model: &OpenRouterModel,
) -> ModelCapabilitiesWithModalities {
let mut capabilities = ModelCapabilitiesWithModalities {
context_length: model.context_length,
max_tokens: model
.top_provider
.as_ref()
.and_then(|tp| tp.max_completion_tokens),
capabilities: vec![],
input_modalities: vec![],
output_modalities: vec![],
};
let arch = &model.architecture;
for input_modality in &arch.input_modalities {
match input_modality.as_str() {
"text" => capabilities.input_modalities.push(Modality::Text),
"image" => capabilities.input_modalities.push(Modality::Image),
"audio" => capabilities.input_modalities.push(Modality::Audio),
"video" => capabilities.input_modalities.push(Modality::Video),
_ => {}
}
}
for output_modality in &arch.output_modalities {
match output_modality.as_str() {
"text" => capabilities.output_modalities.push(Modality::Text),
"image" => capabilities.output_modalities.push(Modality::Image),
"audio" => capabilities.output_modalities.push(Modality::Audio),
"embeddings" => capabilities.output_modalities.push(Modality::Embeddings),
_ => {}
}
}
if model
.supported_parameters
.iter()
.any(|p| p == "tools" || p == "tool_choice")
{
capabilities.capabilities.push(ModelCapabilities::Tools);
}
let inferred_reasoning = self.parse_reasoning(&model.id);
if !inferred_reasoning.is_empty() {
capabilities.capabilities.extend(inferred_reasoning);
} else if model
.supported_parameters
.iter()
.any(|p| p == "reasoning")
{
capabilities.capabilities.extend([
ModelCapabilities::ReasoningEffortNone,
ModelCapabilities::ReasoningEffortMinimal,
ModelCapabilities::ReasoningEffortLow,
ModelCapabilities::ReasoningEffortMedium,
ModelCapabilities::ReasoningEffortHigh,
ModelCapabilities::ReasoningEffortXHigh,
]);
}
capabilities
}
}