Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
4 changes: 3 additions & 1 deletion providers.json
Original file line number Diff line number Diff line change
Expand Up @@ -9,8 +9,9 @@
"api_key_required": true,
"base_url_env": "OPENAI_BASE_URL",
"model_env": "OPENAI_MODEL",
"default_model": "gpt-4o",
"default_model": "gpt-5-mini",
"description": "OpenAI GPT models (direct API)",
"unsupported_params": ["temperature"],
"setup": {
"kind": "api_key",
"secret_name": "llm_openai_api_key",
Expand Down Expand Up @@ -86,6 +87,7 @@
"model_env": "TINFOIL_MODEL",
"default_model": "kimi-k2-5",
"description": "Tinfoil private inference (hardware-attested TEE)",
"unsupported_params": ["temperature"],
"setup": {
"kind": "api_key",
"secret_name": "llm_tinfoil_api_key",
Expand Down
10 changes: 10 additions & 0 deletions src/config/llm.rs
Original file line number Diff line number Diff line change
Expand Up @@ -209,6 +209,7 @@ impl LlmConfig {
extra_headers_env,
api_key_required,
base_url_required,
unsupported_params,
) = if let Some(def) = def {
(
def.id.as_str(),
Expand All @@ -221,6 +222,7 @@ impl LlmConfig {
def.extra_headers_env.as_deref(),
def.api_key_required,
def.base_url_required,
def.unsupported_params.clone(),
)
} else {
// Absolute fallback: treat as generic openai_completions
Expand All @@ -235,6 +237,7 @@ impl LlmConfig {
Some("LLM_EXTRA_HEADERS"),
false,
true,
Vec::new(),
)
};

Expand Down Expand Up @@ -338,6 +341,7 @@ impl LlmConfig {
extra_headers,
oauth_token,
cache_retention,
unsupported_params,
})
}
}
Expand Down Expand Up @@ -624,6 +628,12 @@ mod tests {
let provider = cfg.provider.expect("provider config should be present");
assert_eq!(provider.base_url, "https://inference.tinfoil.sh/v1");
assert_eq!(provider.model, "kimi-k2-5");
assert!(
provider
.unsupported_params
.contains(&"temperature".to_string()),
"tinfoil should propagate unsupported_params from registry"
);
}

#[test]
Expand Down
44 changes: 40 additions & 4 deletions src/llm/anthropic_oauth.rs
Original file line number Diff line number Diff line change
Expand Up @@ -6,6 +6,8 @@
//!
//! Pattern follows `nearai_chat.rs`: direct HTTP calls via `reqwest::Client`.

use std::collections::HashSet;

use async_trait::async_trait;
use reqwest::Client;
use rust_decimal::Decimal;
Expand Down Expand Up @@ -35,6 +37,8 @@ pub struct AnthropicOAuthProvider {
model: String,
base_url: Option<String>,
active_model: std::sync::RwLock<String>,
/// Parameter names that this provider does not support.
unsupported_params: HashSet<String>,
}

impl AnthropicOAuthProvider {
Expand All @@ -61,15 +65,45 @@ impl AnthropicOAuthProvider {
Some(config.base_url.clone())
};

let unsupported_params: HashSet<String> =
config.unsupported_params.iter().cloned().collect();

Ok(Self {
client,
token,
model: config.model.clone(),
base_url,
active_model,
unsupported_params,
})
}

/// Strip unsupported fields from a `CompletionRequest` in place.
fn strip_unsupported_completion_params(&self, req: &mut CompletionRequest) {
if self.unsupported_params.is_empty() {
return;
}
if self.unsupported_params.contains("temperature") {
req.temperature = None;
}
if self.unsupported_params.contains("max_tokens") {
req.max_tokens = None;
}
}

/// Strip unsupported fields from a `ToolCompletionRequest` in place.
fn strip_unsupported_tool_params(&self, req: &mut ToolCompletionRequest) {
if self.unsupported_params.is_empty() {
return;
}
if self.unsupported_params.contains("temperature") {
req.temperature = None;
}
if self.unsupported_params.contains("max_tokens") {
req.max_tokens = None;
}
}

fn api_url(&self) -> String {
if let Some(ref base) = self.base_url {
let base = base.trim_end_matches('/');
Expand Down Expand Up @@ -197,8 +231,9 @@ impl AnthropicOAuthProvider {

#[async_trait]
impl LlmProvider for AnthropicOAuthProvider {
async fn complete(&self, req: CompletionRequest) -> Result<CompletionResponse, LlmError> {
let model = req.model.unwrap_or_else(|| self.active_model_name());
async fn complete(&self, mut req: CompletionRequest) -> Result<CompletionResponse, LlmError> {
let model = req.model.take().unwrap_or_else(|| self.active_model_name());
self.strip_unsupported_completion_params(&mut req);
let (system, messages) = convert_messages(req.messages);
Comment on lines +234 to 237

Copilot AI Mar 10, 2026

Copy link

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Same as RigAdapter: unsupported-parameter stripping happens after outer wrappers compute cache keys from the original CompletionRequest. If CachedProvider is enabled, different temperatures/max_tokens that are later stripped can still fragment the cache. Consider applying request normalization before caching (e.g., as an outer wrapper) so the cache key reflects the effective on-wire request.

Copilot uses AI. Check for mistakes.

Copy link
Copy Markdown
Member Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Same as above — not an issue. The stripped values are consistent per-provider defaults, not user-varying inputs, so cache keys remain stable.


let request = AnthropicRequest {
Expand Down Expand Up @@ -233,9 +268,10 @@ impl LlmProvider for AnthropicOAuthProvider {

async fn complete_with_tools(
&self,
req: ToolCompletionRequest,
mut req: ToolCompletionRequest,
) -> Result<ToolCompletionResponse, LlmError> {
let model = req.model.unwrap_or_else(|| self.active_model_name());
let model = req.model.take().unwrap_or_else(|| self.active_model_name());
self.strip_unsupported_tool_params(&mut req);
let (system, messages) = convert_messages(req.messages);

let tools: Vec<AnthropicTool> = req
Expand Down
4 changes: 4 additions & 0 deletions src/llm/config.rs
Original file line number Diff line number Diff line change
Expand Up @@ -87,6 +87,10 @@ pub struct RegistryProviderConfig {
pub oauth_token: Option<SecretString>,
/// Prompt cache retention (Anthropic-specific).
pub cache_retention: CacheRetention,
/// Parameter names that this provider does not support (e.g., `["temperature"]`).
/// Supported keys: `"temperature"`, `"max_tokens"`, `"stop_sequences"`.
/// Listed parameters are stripped from requests before sending to avoid 400 errors.
pub unsupported_params: Vec<String>,
}

/// Configuration for AWS Bedrock (native Converse API).
Expand Down
12 changes: 9 additions & 3 deletions src/llm/mod.rs
Original file line number Diff line number Diff line change
Expand Up @@ -228,7 +228,9 @@ fn create_openai_compat_from_registry(
"Using OpenAI-compatible provider"
);

Ok(Arc::new(RigAdapter::new(model, &config.model)))
let adapter = RigAdapter::new(model, &config.model)
.with_unsupported_params(config.unsupported_params.clone());
Ok(Arc::new(adapter))
}

fn create_anthropic_from_registry(
Expand Down Expand Up @@ -296,7 +298,9 @@ fn create_anthropic_from_registry(
);

Ok(Arc::new(
RigAdapter::new(model, &config.model).with_cache_retention(cache_retention),
RigAdapter::new(model, &config.model)
.with_cache_retention(cache_retention)
.with_unsupported_params(config.unsupported_params.clone()),
))
}

Expand Down Expand Up @@ -324,7 +328,9 @@ fn create_ollama_from_registry(
"Using Ollama provider"
);

Ok(Arc::new(RigAdapter::new(model, &config.model)))
let adapter = RigAdapter::new(model, &config.model)
.with_unsupported_params(config.unsupported_params.clone());
Ok(Arc::new(adapter))
}

/// Create a cheap/fast LLM provider for lightweight tasks (heartbeat, routing, evaluation).
Expand Down
56 changes: 56 additions & 0 deletions src/llm/registry.rs
Original file line number Diff line number Diff line change
Expand Up @@ -152,6 +152,11 @@ pub struct ProviderDefinition {
/// Setup wizard hints.
#[serde(default)]
pub setup: Option<SetupHint>,
/// Parameter names that this provider does not support (e.g., `["temperature"]`).
/// Supported keys: `"temperature"`, `"max_tokens"`, `"stop_sequences"`.
/// Listed parameters are stripped from requests before sending to avoid 400 errors.
#[serde(default)]
pub unsupported_params: Vec<String>,
}

/// Registry of known LLM providers.
Expand Down Expand Up @@ -378,6 +383,7 @@ mod tests {
description: "Custom tinfoil".to_string(),
extra_headers_env: None,
setup: None,
unsupported_params: vec![],
});
let registry = ProviderRegistry::new(all);
let tf = registry.find("tinfoil").expect("tinfoil should exist");
Expand Down Expand Up @@ -517,6 +523,7 @@ mod tests {
description: "No setup".to_string(),
extra_headers_env: None,
setup: None, // no setup hint
unsupported_params: vec![],
}];

let registry = ProviderRegistry::new(providers.clone());
Expand Down Expand Up @@ -546,6 +553,7 @@ mod tests {
can_list_models: false,
models_filter: None,
}),
unsupported_params: vec![],
});

let registry = ProviderRegistry::new(providers);
Expand Down Expand Up @@ -587,6 +595,7 @@ mod tests {
can_list_models: false,
models_filter: None,
}),
unsupported_params: vec![],
},
// User override removes setup
ProviderDefinition {
Expand All @@ -603,6 +612,7 @@ mod tests {
description: "No setup now".to_string(),
extra_headers_env: None,
setup: None,
unsupported_params: vec![],
},
];

Expand Down Expand Up @@ -640,6 +650,7 @@ mod tests {
display_name: "A".to_string(),
can_list_models: false,
}),
unsupported_params: vec![],
},
ProviderDefinition {
id: "bbb".to_string(),
Expand All @@ -658,6 +669,7 @@ mod tests {
display_name: "B".to_string(),
can_list_models: false,
}),
unsupported_params: vec![],
},
ProviderDefinition {
id: "ccc".to_string(),
Expand All @@ -676,6 +688,7 @@ mod tests {
display_name: "C".to_string(),
can_list_models: false,
}),
unsupported_params: vec![],
},
// User override for B
ProviderDefinition {
Expand All @@ -695,6 +708,7 @@ mod tests {
display_name: "B".to_string(),
can_list_models: false,
}),
unsupported_params: vec![],
},
];

Expand All @@ -708,6 +722,48 @@ mod tests {
);
}

#[test]
fn test_unsupported_params_deserialized() {
let providers: Vec<ProviderDefinition> =
serde_json::from_str(include_str!("../../providers.json")).unwrap();

// Tinfoil should have temperature in unsupported_params
let tinfoil = providers.iter().find(|p| p.id == "tinfoil").unwrap();
assert!(
tinfoil
.unsupported_params
.contains(&"temperature".to_string()),
"tinfoil should have 'temperature' in unsupported_params"
);

// OpenAI should also have temperature in unsupported_params
let openai = providers.iter().find(|p| p.id == "openai").unwrap();
assert!(
openai
.unsupported_params
.contains(&"temperature".to_string()),
"openai should have 'temperature' in unsupported_params"
);

// Providers without the field in JSON should deserialize to empty vec
let groq = providers.iter().find(|p| p.id == "groq").unwrap();
assert!(
groq.unsupported_params.is_empty(),
"groq should have empty unsupported_params (field absent in JSON)"
);

// Every non-empty entry should contain valid param names
for def in &providers {
for param in &def.unsupported_params {
assert!(
!param.is_empty(),
"{}: unsupported_params contains empty string",
def.id
);
}
}
}

#[test]
fn test_all_builtin_api_key_providers_have_api_key_env() {
// Every built-in provider with SetupHint::ApiKey must have api_key_env
Expand Down
Loading