use std::collections::HashMap;
use std::marker::PhantomData;
use std::sync::{Arc, Mutex};
use reqwest::{Client as HttpClient, RequestBuilder};
use serde_json::{Map, Value};
use vtcode_config::core::{ModelConfig, PromptCachingConfig};
use crate::provider::{LLMError, LLMRequest, LLMResponse, LLMStream, ToolDefinition};
use super::common::{
chat_completions_url, ensure_model, extract_prompt_cache_settings_default, float_to_json_number, override_base_url,
parse_json_response, parse_response_openai_format, resolve_model, send_chat_completions,
serialize_messages_openai_format, serialize_tools_openai_format, spawn_openai_compatible_stream,
validate_request_common, validate_supported_models,
};
use super::error_handling::handle_openai_http_error;
use super::shared::OpenAiDeltaOrder;
pub(crate) type ReasoningExtractor = fn(&Value, &Value) -> Option<String>;
#[derive(Clone, Copy, PartialEq, Eq)]
pub(crate) enum SystemPromptPlacement {
Omitted,
FirstMessage,
TopLevelField,
}
pub(crate) trait OpenAiCompatSpec: Sized + Send + Sync + 'static {
const NAME: &'static str;
const KEY: &'static str;
const API_KEY_ENV: &'static str;
const DEFAULT_MODEL: &'static str;
const DEFAULT_BASE_URL: &'static str;
const BASE_URL_ENV: Option<&'static str>;
const LISTED_MODELS: &'static [&'static str];
const VALIDATION_ALLOWLIST: Option<&'static [&'static str]>;
const MAX_TOKENS_KEY: &'static str = "max_tokens";
const SYSTEM_PROMPT: SystemPromptPlacement = SystemPromptPlacement::FirstMessage;
const INCLUDE_TOP_P: bool = true;
const SUPPRESS_SAMPLING_WHEN_REASONING: bool = true;
const STREAM_OPTIONS_INCLUDE_USAGE: bool = false;
const INCLUDE_USER_ID: bool = false;
const VALIDATE_ON_GENERATE: bool = false;
const STREAM_REASONING_FIELDS: &'static [&'static str] = &["reasoning_content"];
const DELTA_ORDER: OpenAiDeltaOrder = OpenAiDeltaOrder::ReasoningFirst;
const RESPONSE_REASONING_EXTRACTOR: Option<ReasoningExtractor> = None;
fn resolve_api_key(api_key: Option<String>) -> String {
api_key.unwrap_or_default()
}
fn resolve_base_url(_api_key: &str, base_url: Option<String>) -> String {
override_base_url(Self::DEFAULT_BASE_URL, base_url, Self::BASE_URL_ENV)
}
fn normalize_model(model: String) -> String {
model
}
fn prompt_cache_enabled(prompt_cache: Option<&PromptCachingConfig>) -> bool {
extract_prompt_cache_settings_default(prompt_cache.cloned(), Self::KEY).0
}
fn response_cache_metrics(_core: &OpenAiCompatCore<Self>) -> bool {
false
}
fn stream_cache_metrics(_core: &OpenAiCompatCore<Self>) -> bool {
false
}
fn reasoning_enabled(_core: &OpenAiCompatCore<Self>, request: &LLMRequest) -> bool {
request
.reasoning_effort
.is_some_and(|effort| effort != vtcode_config::types::ReasoningEffortLevel::None)
}
fn float_number(value: f32) -> Result<serde_json::Number, LLMError> {
float_to_json_number(value)
}
fn insert_tool_choice(_core: &OpenAiCompatCore<Self>, request: &LLMRequest, payload: &mut Map<String, Value>) {
if let Some(choice) = &request.tool_choice {
payload.insert("tool_choice".to_owned(), choice.to_provider_format(Self::KEY));
}
}
fn insert_reasoning(
_core: &OpenAiCompatCore<Self>,
_request: &LLMRequest,
_payload: &mut Map<String, Value>,
) -> Result<(), LLMError> {
Ok(())
}
fn finish_payload(
_core: &OpenAiCompatCore<Self>,
_request: &LLMRequest,
_payload: &mut Map<String, Value>,
) -> Result<(), LLMError> {
Ok(())
}
fn apply_auth(core: &OpenAiCompatCore<Self>, builder: RequestBuilder) -> RequestBuilder {
builder.bearer_auth(&core.api_key)
}
fn api_key_env(_core: &OpenAiCompatCore<Self>) -> &'static str {
Self::API_KEY_ENV
}
fn listed_models(_core: &OpenAiCompatCore<Self>) -> &'static [&'static str] {
Self::LISTED_MODELS
}
fn validate(_core: &OpenAiCompatCore<Self>, request: &LLMRequest) -> Result<(), LLMError> {
match Self::VALIDATION_ALLOWLIST {
Some(models) => validate_supported_models(request, Self::NAME, Self::KEY, models),
None => validate_request_common(request, Self::NAME, Self::KEY, None),
}
}
}
pub(crate) struct OpenAiCompatCore<S: OpenAiCompatSpec> {
pub(crate) api_key: String,
pub(crate) http_client: HttpClient,
pub(crate) base_url: String,
pub(crate) model: String,
pub(crate) prompt_cache_enabled: bool,
pub(crate) model_behavior: Option<ModelConfig>,
spec: PhantomData<S>,
tools_cache: Mutex<ToolsCache>,
}
type ToolsCache = HashMap<usize, (Arc<Vec<ToolDefinition>>, Arc<Vec<Value>>)>;
impl<S: OpenAiCompatSpec> OpenAiCompatCore<S> {
pub(crate) fn direct(api_key: String, model: String) -> Self {
Self::assemble(api_key, model, None, None, None)
}
pub(crate) fn from_config(
api_key: Option<String>,
model: Option<String>,
base_url: Option<String>,
prompt_cache: Option<PromptCachingConfig>,
timeouts: Option<vtcode_config::TimeoutsConfig>,
model_behavior: Option<ModelConfig>,
) -> Self {
let api_key = S::resolve_api_key(api_key);
let model = resolve_model(model, S::DEFAULT_MODEL);
let mut core = Self::assemble(api_key, model, base_url, timeouts, model_behavior);
core.prompt_cache_enabled = S::prompt_cache_enabled(prompt_cache.as_ref());
core
}
fn assemble(
api_key: String,
model: String,
base_url: Option<String>,
timeouts: Option<vtcode_config::TimeoutsConfig>,
model_behavior: Option<ModelConfig>,
) -> Self {
use crate::http_client::HttpClientFactory;
let timeouts = timeouts.unwrap_or_default();
let base_url = S::resolve_base_url(&api_key, base_url);
Self {
api_key,
http_client: HttpClientFactory::for_llm(&timeouts),
base_url,
model: S::normalize_model(model),
prompt_cache_enabled: false,
model_behavior,
spec: PhantomData,
tools_cache: Mutex::new(HashMap::new()),
}
}
pub(crate) fn from_parts(api_key: String, model: String, http_client: HttpClient, base_url: String) -> Self {
Self {
api_key,
http_client,
base_url,
model: S::normalize_model(model),
prompt_cache_enabled: false,
model_behavior: None,
spec: PhantomData,
tools_cache: Mutex::new(HashMap::new()),
}
}
pub(crate) fn prepare(&self, request: &mut LLMRequest) {
ensure_model(request, &self.model);
request.model = S::normalize_model(std::mem::take(&mut request.model));
}
pub(crate) fn convert_request(&self, request: &LLMRequest) -> Result<Value, LLMError> {
let mut payload = Map::new();
payload.insert("model".to_owned(), Value::String(request.model.clone()));
let mut messages = serialize_messages_openai_format(request, S::KEY)?;
if S::SYSTEM_PROMPT == SystemPromptPlacement::FirstMessage
&& let Some(system) = &request.system_prompt
{
let trimmed = system.trim();
if !trimmed.is_empty() {
messages.insert(0, serde_json::json!({"role": "system", "content": trimmed}));
}
}
payload.insert("messages".to_owned(), Value::Array(messages));
if S::SYSTEM_PROMPT == SystemPromptPlacement::TopLevelField
&& let Some(system) = &request.system_prompt
{
let trimmed = system.trim();
if !trimmed.is_empty() {
payload.insert("system".to_owned(), Value::String(trimmed.to_owned()));
}
}
if let Some(max_tokens) = request.max_tokens {
payload
.insert(S::MAX_TOKENS_KEY.to_owned(), Value::Number(serde_json::Number::from(u64::from(max_tokens))));
}
let suppress_sampling = S::SUPPRESS_SAMPLING_WHEN_REASONING && S::reasoning_enabled(self, request);
if !suppress_sampling {
if let Some(temperature) = request.temperature {
payload.insert("temperature".to_owned(), Value::Number(S::float_number(temperature)?));
}
if S::INCLUDE_TOP_P
&& let Some(top_p) = request.top_p
{
payload.insert("top_p".to_owned(), Value::Number(S::float_number(top_p)?));
}
}
if request.stream {
payload.insert("stream".to_owned(), Value::Bool(true));
if S::STREAM_OPTIONS_INCLUDE_USAGE {
payload.insert("stream_options".to_owned(), serde_json::json!({"include_usage": true}));
}
}
if let Some(tools) = &request.tools {
let key = Arc::as_ptr(tools) as usize;
let serialized_tools = {
let mut cache = self.tools_cache.lock().unwrap_or_else(|poisoned| poisoned.into_inner());
match cache.get(&key) {
Some((arc, serialized)) if Arc::ptr_eq(arc, tools) => Arc::clone(serialized),
_ => {
let serialized = serialize_tools_openai_format(tools)
.map(Arc::new)
.unwrap_or_else(|| Arc::new(Vec::new()));
cache.insert(key, (Arc::clone(tools), Arc::clone(&serialized)));
serialized
}
}
};
if !serialized_tools.is_empty() {
payload.insert("tools".to_owned(), Value::Array(serialized_tools.as_ref().clone()));
}
}
S::insert_tool_choice(self, request, &mut payload);
S::insert_reasoning(self, request, &mut payload)?;
if S::INCLUDE_USER_ID
&& let Some(user_id) = request
.metadata
.as_ref()
.and_then(|metadata| metadata.get("user_id"))
.and_then(Value::as_str)
{
payload.insert("user_id".to_owned(), Value::String(user_id.to_owned()));
}
S::finish_payload(self, request, &mut payload)?;
Ok(Value::Object(payload))
}
pub(crate) async fn dispatch(&self, request: &LLMRequest) -> Result<reqwest::Response, LLMError> {
let payload = self.convert_request(request)?;
let url = chat_completions_url(&self.base_url);
let builder = S::apply_auth(self, self.http_client.post(&url));
let response = send_chat_completions(builder, &payload, S::NAME).await?;
handle_openai_http_error(response, S::NAME, S::api_key_env(self)).await
}
pub(crate) async fn generate_prepared(&self, request: LLMRequest) -> Result<LLMResponse, LLMError> {
let model = request.model.clone();
let response = self.dispatch(&request).await?;
let response_json = parse_json_response(response, S::NAME).await?;
parse_response_openai_format::<ReasoningExtractor>(
response_json,
S::NAME,
model,
S::response_cache_metrics(self),
S::RESPONSE_REASONING_EXTRACTOR,
)
}
pub(crate) async fn stream_prepared(&self, request: LLMRequest) -> Result<LLMStream, LLMError> {
let model = request.model.clone();
let response = self.dispatch(&request).await?;
Ok(spawn_openai_compatible_stream(
response,
S::NAME,
model,
S::STREAM_REASONING_FIELDS,
S::DELTA_ORDER,
S::stream_cache_metrics(self),
))
}
pub(crate) fn supported_models(&self) -> Vec<String> {
S::listed_models(self).iter().map(|model| model.to_string()).collect()
}
pub(crate) fn validate(&self, request: &LLMRequest) -> Result<(), LLMError> {
S::validate(self, request)
}
}
macro_rules! impl_openai_compat_provider {
($provider:ident, $spec:ty $(, { $($extra:item)* })?) => {
pub struct $provider {
core: crate::providers::openai_compat::OpenAiCompatCore<$spec>,
}
impl $provider {
pub fn new(api_key: String) -> Self {
Self::with_model(
api_key,
<$spec as crate::providers::openai_compat::OpenAiCompatSpec>::DEFAULT_MODEL
.to_string(),
)
}
pub fn with_model(api_key: String, model: String) -> Self {
Self {
core: crate::providers::openai_compat::OpenAiCompatCore::direct(
api_key, model,
),
}
}
pub fn new_with_client(
api_key: String,
model: String,
http_client: reqwest::Client,
base_url: String,
_timeouts: vtcode_config::TimeoutsConfig,
) -> Self {
Self {
core: crate::providers::openai_compat::OpenAiCompatCore::from_parts(
api_key,
model,
http_client,
base_url,
),
}
}
pub fn from_config(
api_key: Option<String>,
model: Option<String>,
base_url: Option<String>,
prompt_cache: Option<vtcode_config::core::PromptCachingConfig>,
timeouts: Option<vtcode_config::TimeoutsConfig>,
_anthropic: Option<vtcode_config::core::AnthropicConfig>,
model_behavior: Option<vtcode_config::core::ModelConfig>,
) -> Self {
Self {
core: crate::providers::openai_compat::OpenAiCompatCore::from_config(
api_key,
model,
base_url,
prompt_cache,
timeouts,
model_behavior,
),
}
}
}
#[async_trait::async_trait]
impl crate::provider::LLMProvider for $provider {
fn name(&self) -> &str {
<$spec as crate::providers::openai_compat::OpenAiCompatSpec>::KEY
}
async fn generate(
&self,
mut request: crate::provider::LLMRequest,
) -> Result<crate::provider::LLMResponse, crate::provider::LLMError> {
self.core.prepare(&mut request);
if <$spec as crate::providers::openai_compat::OpenAiCompatSpec>::VALIDATE_ON_GENERATE
{
crate::provider::LLMProvider::validate_request(self, &request)?;
}
self.core.generate_prepared(request).await
}
async fn stream(
&self,
mut request: crate::provider::LLMRequest,
) -> Result<crate::provider::LLMStream, crate::provider::LLMError> {
self.core.prepare(&mut request);
crate::provider::LLMProvider::validate_request(self, &request)?;
request.stream = true;
self.core.stream_prepared(request).await
}
fn supported_models(&self) -> Vec<String> {
self.core.supported_models()
}
fn validate_request(
&self,
request: &crate::provider::LLMRequest,
) -> Result<(), crate::provider::LLMError> {
self.core.validate(request)
}
$($($extra)*)?
}
#[async_trait::async_trait]
impl crate::client::LLMClient for $provider {
async fn generate(
&mut self,
prompt: &str,
) -> Result<crate::provider::LLMResponse, crate::provider::LLMError> {
let request =
crate::providers::common::make_default_request(prompt, &self.core.model);
Ok(<$provider as crate::provider::LLMProvider>::generate(self, request).await?)
}
fn model_id(&self) -> &str {
&self.core.model
}
}
};
}
pub(crate) use impl_openai_compat_provider;