use std::collections::HashMap;
use async_trait::async_trait;
use mermaid_domain::ChatRequest;
use mermaid_model::models::adapters::ModelLimits;
use mermaid_model::models::adapters::openai_compat::OpenAICompatAdapter;
use mermaid_model::models::{Model, ModelConfig, ModelError, ProviderProfile, Result};
use super::super::ctx::{FinalResponse, StreamContext, StreamEvent};
use super::{
ContextSizing, ModelProvider, RejectionCache, learn_output_cap, output_cap_from_error,
resolve_limits_cached, retry_cap,
};
use mermaid_model::models::ModelCapabilities;
pub struct OpenAICompatProvider {
adapter: OpenAICompatAdapter,
capabilities: ModelCapabilities,
rejections: RejectionCache,
}
impl OpenAICompatProvider {
pub fn new(
profile: &'static ProviderProfile,
base_url: String,
api_key: Option<String>,
model_name: String,
extra_headers: HashMap<String, String>,
) -> Result<Self> {
let adapter =
OpenAICompatAdapter::new(profile, base_url, api_key, model_name, extra_headers)?;
let capabilities = adapter.capabilities().clone();
Ok(Self {
adapter,
capabilities,
rejections: RejectionCache::default(),
})
}
}
#[async_trait]
impl ModelProvider for OpenAICompatProvider {
fn capabilities(&self) -> &ModelCapabilities {
&self.capabilities
}
async fn resolve_context_window(&self, request: &ChatRequest) -> ContextSizing {
let _ = request;
let provider = self.adapter.provider_name().to_string();
let model = Model::name(&self.adapter).to_string();
let limits = resolve_limits_cached(&provider, &model, || async {
let listings = self.adapter.list_models_for_limits().await?;
let found = listings.into_iter().find(|m| m.id == model);
Ok(ModelLimits {
max_context_tokens: found.as_ref().and_then(|m| m.max_context_tokens),
max_output_tokens: found.as_ref().and_then(|m| m.max_output_tokens),
})
})
.await;
let window = limits.as_ref().and_then(|l| l.max_context_tokens);
ContextSizing {
model_max: window,
effective: window,
source: None,
max_output: limits.as_ref().and_then(|l| l.max_output_tokens),
compacts_natively: false,
}
}
async fn supports_vision(&self) -> Option<bool> {
Some(self.capabilities.supports_vision)
}
async fn chat(&self, request: ChatRequest, ctx: StreamContext) -> Result<FinalResponse> {
let config = ModelConfig::from(&request);
let chat_fut = async {
match self
.adapter
.chat(&request.messages, &config, Some(ctx.sink.clone()))
.await
{
Ok(response) => Ok(response),
Err(err) => {
let Some(cap) = output_cap_from_error(&err) else {
return Err(err);
};
let Some(clamped) = retry_cap(config.max_tokens, cap) else {
return Err(err);
};
let provider = self.adapter.provider_name().to_string();
let model = Model::name(&self.adapter).to_string();
learn_output_cap(provider, model.clone(), cap).await;
let _ = ctx.sink.send(StreamEvent::Status(format!(
"{model} rejected the output budget; learned its {cap}-token cap and retrying"
))).await;
let retry_config = ModelConfig {
max_tokens: clamped,
..config.clone()
};
self.adapter
.chat(&request.messages, &retry_config, Some(ctx.sink.clone()))
.await
},
}
};
let learned_model = Model::name(&self.adapter).to_string();
let memory = self.adapter.param_memory();
self.rejections
.seed(self.adapter.provider_name(), &learned_model, memory)
.await;
let response = tokio::select! {
biased;
_ = ctx.token.cancelled() => {
return Err(ModelError::Cancelled);
},
r = chat_fut => {
self.rejections.persist(self.adapter.provider_name(), &learned_model, memory).await;
r?
},
};
let usage = response.usage.clone();
let stop_reason = response.stop_reason.clone();
let _ = ctx
.sink
.send(StreamEvent::Done {
usage: usage.clone(),
provider_continuation: None,
stop_reason: stop_reason.clone(),
})
.await;
Ok(FinalResponse {
usage,
provider_continuation: None,
tool_calls: response.tool_calls.unwrap_or_default(),
stop_reason,
})
}
}