mirror of
https://github.com/delibae/claude-prism.git
synced 2026-10-07 02:58:00 +00:00
Fallback when compatible providers reject tools
This commit is contained in:
parent
782024a502
commit
b059e3bb07
1 changed files with 268 additions and 24 deletions
|
|
@ -1816,6 +1816,18 @@ fn direct_provider_system_prompt() -> String {
|
|||
.join("\n")
|
||||
}
|
||||
|
||||
fn direct_provider_no_tools_system_prompt() -> String {
|
||||
[
|
||||
"You are an AI assistant integrated into ClaudePrism, a LaTeX document editor.",
|
||||
"Help the user write, revise, and reason about academic documents.",
|
||||
"Preserve existing LaTeX structure unless the user asks for a rewrite.",
|
||||
"Use proper LaTeX sectioning, citations, labels, references, and bibliography conventions.",
|
||||
"This provider endpoint does not support tool calls in ClaudePrism. Do not claim to read, edit, or run files directly.",
|
||||
"When file changes are needed, provide precise patches, replacement snippets, or step-by-step commands for the user.",
|
||||
]
|
||||
.join("\n")
|
||||
}
|
||||
|
||||
fn direct_provider_tools() -> serde_json::Value {
|
||||
json!([
|
||||
{
|
||||
|
|
@ -1999,6 +2011,7 @@ struct DirectStreamingToolCall {
|
|||
struct DirectStreamFailure {
|
||||
message: String,
|
||||
can_retry_non_streaming: bool,
|
||||
can_retry_without_tools: bool,
|
||||
}
|
||||
|
||||
fn emit_direct_output(window: &WebviewWindow, tab_id: &str, event: &serde_json::Value) {
|
||||
|
|
@ -2924,6 +2937,113 @@ fn direct_message_from_parts(content: &str, tool_calls: &[DirectToolCall]) -> se
|
|||
message
|
||||
}
|
||||
|
||||
fn provider_error_allows_toolless_retry(status: reqwest::StatusCode, body: &str) -> bool {
|
||||
if !(status == reqwest::StatusCode::BAD_REQUEST
|
||||
|| status == reqwest::StatusCode::UNPROCESSABLE_ENTITY
|
||||
|| status == reqwest::StatusCode::NOT_FOUND
|
||||
|| status == reqwest::StatusCode::METHOD_NOT_ALLOWED)
|
||||
{
|
||||
return false;
|
||||
}
|
||||
let body = body.to_ascii_lowercase();
|
||||
body.contains("tool")
|
||||
|| body.contains("tools")
|
||||
|| body.contains("function")
|
||||
|| body.contains("tool_choice")
|
||||
|| body.contains("unsupported parameter")
|
||||
|| body.contains("unknown parameter")
|
||||
}
|
||||
|
||||
fn add_direct_request_tooling(request_body: &mut serde_json::Value, use_tools: bool) {
|
||||
if !use_tools {
|
||||
return;
|
||||
}
|
||||
if let Some(object) = request_body.as_object_mut() {
|
||||
object.insert("tools".to_string(), direct_provider_tools());
|
||||
object.insert("tool_choice".to_string(), json!("auto"));
|
||||
}
|
||||
}
|
||||
|
||||
fn direct_messages_for_tool_capability(
|
||||
messages: &[serde_json::Value],
|
||||
use_tools: bool,
|
||||
) -> Vec<serde_json::Value> {
|
||||
if use_tools {
|
||||
return messages.to_vec();
|
||||
}
|
||||
|
||||
let mut messages = messages.to_vec();
|
||||
for message in &mut messages {
|
||||
let role = message.get("role").and_then(|v| v.as_str());
|
||||
if role == Some("system") {
|
||||
if let Some(object) = message.as_object_mut() {
|
||||
object.insert(
|
||||
"content".to_string(),
|
||||
json!(direct_provider_no_tools_system_prompt()),
|
||||
);
|
||||
}
|
||||
break;
|
||||
}
|
||||
}
|
||||
|
||||
if !messages
|
||||
.iter()
|
||||
.any(|message| message.get("role").and_then(|v| v.as_str()) == Some("system"))
|
||||
{
|
||||
messages.insert(
|
||||
0,
|
||||
json!({
|
||||
"role": "system",
|
||||
"content": direct_provider_no_tools_system_prompt(),
|
||||
}),
|
||||
);
|
||||
}
|
||||
|
||||
let mut sanitized = Vec::new();
|
||||
for mut message in messages {
|
||||
match message.get("role").and_then(|v| v.as_str()) {
|
||||
Some("assistant") => {
|
||||
if let Some(object) = message.as_object_mut() {
|
||||
object.remove("tool_calls");
|
||||
if object.get("content").map(|v| v.is_null()).unwrap_or(false) {
|
||||
object.insert(
|
||||
"content".to_string(),
|
||||
json!("[Assistant requested a tool call, but this provider endpoint does not support tool calls.]"),
|
||||
);
|
||||
}
|
||||
}
|
||||
sanitized.push(message);
|
||||
}
|
||||
Some("tool") => {
|
||||
let tool_call_id = message
|
||||
.get("tool_call_id")
|
||||
.and_then(|v| v.as_str())
|
||||
.unwrap_or("unknown");
|
||||
let content = message
|
||||
.get("content")
|
||||
.and_then(|v| v.as_str())
|
||||
.map(|v| v.to_string())
|
||||
.unwrap_or_else(|| {
|
||||
message
|
||||
.get("content")
|
||||
.map(|v| v.to_string())
|
||||
.unwrap_or_default()
|
||||
});
|
||||
sanitized.push(json!({
|
||||
"role": "user",
|
||||
"content": format!(
|
||||
"[Previous tool result {}]\n{}",
|
||||
tool_call_id,
|
||||
content
|
||||
),
|
||||
}));
|
||||
}
|
||||
_ => sanitized.push(message),
|
||||
}
|
||||
}
|
||||
sanitized
|
||||
}
|
||||
|
||||
async fn send_openai_compatible_chat_request(
|
||||
client: &reqwest::Client,
|
||||
credential: &StoredOpenAiCompatibleCredential,
|
||||
|
|
@ -2934,13 +3054,7 @@ async fn send_openai_compatible_chat_request(
|
|||
process_key: &str,
|
||||
) -> Result<DirectChatResponse, String> {
|
||||
match send_openai_compatible_streaming_chat_request(
|
||||
client,
|
||||
credential,
|
||||
messages,
|
||||
window,
|
||||
tab_id,
|
||||
state,
|
||||
process_key,
|
||||
client, credential, messages, window, tab_id, state, process_key, true,
|
||||
)
|
||||
.await
|
||||
{
|
||||
|
|
@ -2950,26 +3064,81 @@ async fn send_openai_compatible_chat_request(
|
|||
"[direct-provider] streaming request failed, retrying non-streaming: {}",
|
||||
err.message
|
||||
);
|
||||
send_openai_compatible_non_streaming_chat_request(client, credential, messages).await
|
||||
match send_openai_compatible_non_streaming_chat_request(
|
||||
client, credential, messages, true,
|
||||
)
|
||||
.await
|
||||
{
|
||||
Ok(response) => Ok(response),
|
||||
Err(err) if err.can_retry_without_tools => {
|
||||
eprintln!(
|
||||
"[direct-provider] provider rejected tools, retrying without tools: {}",
|
||||
err.message
|
||||
);
|
||||
send_openai_compatible_without_tools(
|
||||
client, credential, messages, window, tab_id, state, process_key,
|
||||
)
|
||||
.await
|
||||
}
|
||||
Err(err) => Err(err.message),
|
||||
}
|
||||
}
|
||||
Err(err) if err.can_retry_without_tools => {
|
||||
eprintln!(
|
||||
"[direct-provider] provider rejected streaming tools, retrying without tools: {}",
|
||||
err.message
|
||||
);
|
||||
send_openai_compatible_without_tools(
|
||||
client, credential, messages, window, tab_id, state, process_key,
|
||||
)
|
||||
.await
|
||||
}
|
||||
Err(err) => Err(err.message),
|
||||
}
|
||||
}
|
||||
|
||||
async fn send_openai_compatible_without_tools(
|
||||
client: &reqwest::Client,
|
||||
credential: &StoredOpenAiCompatibleCredential,
|
||||
messages: &[serde_json::Value],
|
||||
window: &WebviewWindow,
|
||||
tab_id: &str,
|
||||
state: &ClaudeProcessState,
|
||||
process_key: &str,
|
||||
) -> Result<DirectChatResponse, String> {
|
||||
match send_openai_compatible_streaming_chat_request(
|
||||
client, credential, messages, window, tab_id, state, process_key, false,
|
||||
)
|
||||
.await
|
||||
{
|
||||
Ok(response) => Ok(response),
|
||||
Err(err) if err.can_retry_non_streaming => {
|
||||
send_openai_compatible_non_streaming_chat_request(client, credential, messages, false)
|
||||
.await
|
||||
.map_err(|err| err.message)
|
||||
}
|
||||
Err(err) => Err(err.message),
|
||||
}
|
||||
}
|
||||
|
||||
struct DirectProviderRequestFailure {
|
||||
message: String,
|
||||
can_retry_without_tools: bool,
|
||||
}
|
||||
|
||||
async fn send_openai_compatible_non_streaming_chat_request(
|
||||
client: &reqwest::Client,
|
||||
credential: &StoredOpenAiCompatibleCredential,
|
||||
messages: &[serde_json::Value],
|
||||
) -> Result<DirectChatResponse, String> {
|
||||
use_tools: bool,
|
||||
) -> Result<DirectChatResponse, DirectProviderRequestFailure> {
|
||||
let request_messages = direct_messages_for_tool_capability(messages, use_tools);
|
||||
let mut request_body = json!({
|
||||
"model": credential.model.clone(),
|
||||
"messages": messages,
|
||||
"messages": request_messages,
|
||||
"stream": false,
|
||||
});
|
||||
if let Some(object) = request_body.as_object_mut() {
|
||||
object.insert("tools".to_string(), direct_provider_tools());
|
||||
object.insert("tool_choice".to_string(), json!("auto"));
|
||||
}
|
||||
add_direct_request_tooling(&mut request_body, use_tools);
|
||||
|
||||
let response = client
|
||||
.post(openai_chat_completions_url(&credential.base_url))
|
||||
|
|
@ -2978,20 +3147,33 @@ async fn send_openai_compatible_non_streaming_chat_request(
|
|||
.body(request_body.to_string())
|
||||
.send()
|
||||
.await
|
||||
.map_err(|err| format!("Provider request failed: {}", err))?;
|
||||
.map_err(|err| DirectProviderRequestFailure {
|
||||
message: format!("Provider request failed: {}", err),
|
||||
can_retry_without_tools: false,
|
||||
})?;
|
||||
|
||||
let status = response.status();
|
||||
let response_text = response
|
||||
.text()
|
||||
.await
|
||||
.map_err(|err| format!("Failed to read provider response: {}", err))?;
|
||||
.map_err(|err| DirectProviderRequestFailure {
|
||||
message: format!("Failed to read provider response: {}", err),
|
||||
can_retry_without_tools: false,
|
||||
})?;
|
||||
|
||||
if !status.is_success() {
|
||||
return Err(format!("Provider returned HTTP {}: {}", status, response_text));
|
||||
return Err(DirectProviderRequestFailure {
|
||||
message: format!("Provider returned HTTP {}: {}", status, response_text),
|
||||
can_retry_without_tools: provider_error_allows_toolless_retry(status, &response_text),
|
||||
});
|
||||
}
|
||||
|
||||
let response = serde_json::from_str(&response_text)
|
||||
.map_err(|err| format!("Provider returned invalid JSON: {}", err))?;
|
||||
let response = serde_json::from_str(&response_text).map_err(|err| {
|
||||
DirectProviderRequestFailure {
|
||||
message: format!("Provider returned invalid JSON: {}", err),
|
||||
can_retry_without_tools: false,
|
||||
}
|
||||
})?;
|
||||
Ok(direct_chat_response_from_value(response))
|
||||
}
|
||||
|
||||
|
|
@ -3003,16 +3185,15 @@ async fn send_openai_compatible_streaming_chat_request(
|
|||
tab_id: &str,
|
||||
state: &ClaudeProcessState,
|
||||
process_key: &str,
|
||||
use_tools: bool,
|
||||
) -> Result<DirectChatResponse, DirectStreamFailure> {
|
||||
let request_messages = direct_messages_for_tool_capability(messages, use_tools);
|
||||
let mut request_body = json!({
|
||||
"model": credential.model.clone(),
|
||||
"messages": messages,
|
||||
"messages": request_messages,
|
||||
"stream": true,
|
||||
});
|
||||
if let Some(object) = request_body.as_object_mut() {
|
||||
object.insert("tools".to_string(), direct_provider_tools());
|
||||
object.insert("tool_choice".to_string(), json!("auto"));
|
||||
}
|
||||
add_direct_request_tooling(&mut request_body, use_tools);
|
||||
|
||||
let mut response = client
|
||||
.post(openai_chat_completions_url(&credential.base_url))
|
||||
|
|
@ -3024,6 +3205,7 @@ async fn send_openai_compatible_streaming_chat_request(
|
|||
.map_err(|err| DirectStreamFailure {
|
||||
message: format!("Provider request failed: {}", err),
|
||||
can_retry_non_streaming: true,
|
||||
can_retry_without_tools: false,
|
||||
})?;
|
||||
|
||||
let status = response.status();
|
||||
|
|
@ -3032,6 +3214,7 @@ async fn send_openai_compatible_streaming_chat_request(
|
|||
return Err(DirectStreamFailure {
|
||||
message: format!("Provider returned HTTP {}: {}", status, response_text),
|
||||
can_retry_non_streaming: true,
|
||||
can_retry_without_tools: provider_error_allows_toolless_retry(status, &response_text),
|
||||
});
|
||||
}
|
||||
|
||||
|
|
@ -3044,11 +3227,13 @@ async fn send_openai_compatible_streaming_chat_request(
|
|||
while let Some(chunk) = response.chunk().await.map_err(|err| DirectStreamFailure {
|
||||
message: format!("Failed to read provider stream: {}", err),
|
||||
can_retry_non_streaming: false,
|
||||
can_retry_without_tools: false,
|
||||
})? {
|
||||
if direct_provider_cancelled(state, process_key).await {
|
||||
return Err(DirectStreamFailure {
|
||||
message: "Direct provider request cancelled".to_string(),
|
||||
can_retry_non_streaming: false,
|
||||
can_retry_without_tools: false,
|
||||
});
|
||||
}
|
||||
buffer.push_str(&String::from_utf8_lossy(&chunk));
|
||||
|
|
@ -3075,6 +3260,7 @@ async fn send_openai_compatible_streaming_chat_request(
|
|||
DirectStreamFailure {
|
||||
message: format!("Provider returned invalid stream JSON: {}", err),
|
||||
can_retry_non_streaming: false,
|
||||
can_retry_without_tools: false,
|
||||
}
|
||||
})?;
|
||||
if value.get("usage").is_some() {
|
||||
|
|
@ -4288,6 +4474,64 @@ mod tests {
|
|||
assert_eq!(message["tool_calls"][0]["function"]["name"], "Read");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_direct_provider_detects_tool_unsupported_errors() {
|
||||
assert!(provider_error_allows_toolless_retry(
|
||||
reqwest::StatusCode::BAD_REQUEST,
|
||||
"unknown parameter: tools"
|
||||
));
|
||||
assert!(provider_error_allows_toolless_retry(
|
||||
reqwest::StatusCode::UNPROCESSABLE_ENTITY,
|
||||
"tool_choice is not supported"
|
||||
));
|
||||
assert!(!provider_error_allows_toolless_retry(
|
||||
reqwest::StatusCode::UNAUTHORIZED,
|
||||
"invalid api key"
|
||||
));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_direct_provider_no_tools_prompt_replaces_system_message() {
|
||||
let messages = vec![
|
||||
json!({ "role": "system", "content": direct_provider_system_prompt() }),
|
||||
json!({ "role": "user", "content": "Please edit main.tex" }),
|
||||
];
|
||||
|
||||
let no_tools = direct_messages_for_tool_capability(&messages, false);
|
||||
let content = no_tools[0]["content"].as_str().unwrap();
|
||||
|
||||
assert!(content.contains("does not support tool calls"));
|
||||
assert!(!content.contains("You can inspect and edit files through the provided tools"));
|
||||
assert_eq!(no_tools[1]["content"], "Please edit main.tex");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_direct_provider_no_tools_sanitizes_tool_history() {
|
||||
let messages = vec![
|
||||
json!({ "role": "system", "content": direct_provider_system_prompt() }),
|
||||
json!({
|
||||
"role": "assistant",
|
||||
"content": null,
|
||||
"tool_calls": [{
|
||||
"id": "call-1",
|
||||
"type": "function",
|
||||
"function": { "name": "Read", "arguments": "{\"file_path\":\"main.tex\"}" }
|
||||
}]
|
||||
}),
|
||||
json!({ "role": "tool", "tool_call_id": "call-1", "content": "main.tex contents" }),
|
||||
];
|
||||
|
||||
let no_tools = direct_messages_for_tool_capability(&messages, false);
|
||||
|
||||
assert!(no_tools[1].get("tool_calls").is_none());
|
||||
assert_eq!(no_tools[1]["role"], "assistant");
|
||||
assert_eq!(no_tools[2]["role"], "user");
|
||||
assert!(no_tools[2]["content"]
|
||||
.as_str()
|
||||
.unwrap()
|
||||
.contains("Previous tool result call-1"));
|
||||
}
|
||||
|
||||
// --- claude_required_dirs ---
|
||||
|
||||
#[cfg(not(target_os = "windows"))]
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue