Stream OpenAI-compatible provider responses

This commit is contained in:
Gerard-Devlin 2026-06-08 02:31:59 +08:00
parent 55daad36e0
commit 782024a502
3 changed files with 479 additions and 21 deletions

View file

@ -1981,6 +1981,26 @@ struct DirectToolOutput {
is_error: bool,
}
struct DirectChatResponse {
message: serde_json::Value,
content: String,
tool_calls: Vec<DirectToolCall>,
usage: serde_json::Value,
streamed_text: bool,
}
#[derive(Default)]
struct DirectStreamingToolCall {
id: Option<String>,
name: String,
arguments: String,
}
struct DirectStreamFailure {
message: String,
can_retry_non_streaming: bool,
}
fn emit_direct_output(window: &WebviewWindow, tab_id: &str, event: &serde_json::Value) {
let _ = window.emit(
"claude-output",
@ -2736,6 +2756,17 @@ async fn execute_direct_provider_tool(
}
}
fn parse_direct_tool_arguments(arguments: &serde_json::Value) -> serde_json::Value {
if let Some(arguments) = arguments.as_str() {
return serde_json::from_str::<serde_json::Value>(arguments)
.unwrap_or_else(|_| json!({ "_raw_arguments": arguments }));
}
if arguments.is_object() {
return arguments.clone();
}
json!({})
}
fn parse_direct_tool_calls(response: &serde_json::Value) -> Vec<DirectToolCall> {
response
.pointer("/choices/0/message/tool_calls")
@ -2751,12 +2782,10 @@ fn parse_direct_tool_calls(response: &serde_json::Value) -> Vec<DirectToolCall>
.and_then(|v| v.as_str())
.map(|value| value.to_string())
.unwrap_or_else(|| format!("toolu_{}", idx + 1));
let arguments = call
let input = call
.pointer("/function/arguments")
.and_then(|v| v.as_str())
.unwrap_or("{}");
let input = serde_json::from_str::<serde_json::Value>(arguments)
.unwrap_or_else(|_| json!({ "_raw_arguments": arguments }));
.map(parse_direct_tool_arguments)
.unwrap_or_else(|| json!({}));
Some(DirectToolCall { id, name, input })
})
.collect()
@ -2785,11 +2814,153 @@ fn direct_assistant_content(response: &serde_json::Value) -> String {
String::new()
}
fn direct_chat_response_from_value(response: serde_json::Value) -> DirectChatResponse {
let content = direct_assistant_content(&response);
let tool_calls = parse_direct_tool_calls(&response);
let usage = json_usage(&response);
let message = response
.pointer("/choices/0/message")
.cloned()
.unwrap_or_else(|| json!({ "role": "assistant", "content": content.clone() }));
DirectChatResponse {
message,
content,
tool_calls,
usage,
streamed_text: false,
}
}
fn direct_stream_text_delta(delta: &serde_json::Value) -> String {
let Some(content) = delta.get("content") else {
return String::new();
};
if let Some(text) = content.as_str() {
return text.to_string();
}
if let Some(parts) = content.as_array() {
return parts
.iter()
.filter_map(|part| {
part.get("text")
.and_then(|v| v.as_str())
.or_else(|| part.get("content").and_then(|v| v.as_str()))
})
.collect::<Vec<_>>()
.join("");
}
String::new()
}
fn emit_direct_streaming_delta(window: &WebviewWindow, tab_id: &str, text: &str) {
if text.is_empty() {
return;
}
emit_direct_output(
window,
tab_id,
&json!({
"type": "assistant",
"subtype": "streaming_delta",
"message": {
"content": [{ "type": "text", "text": text }],
},
}),
);
}
fn direct_tool_calls_from_stream(
calls: HashMap<usize, DirectStreamingToolCall>,
) -> Vec<DirectToolCall> {
let mut indexed = calls.into_iter().collect::<Vec<_>>();
indexed.sort_by_key(|(idx, _)| *idx);
indexed
.into_iter()
.filter_map(|(idx, call)| {
if call.name.trim().is_empty() {
return None;
}
let id = call.id.unwrap_or_else(|| format!("toolu_stream_{}", idx + 1));
let input = serde_json::from_str::<serde_json::Value>(&call.arguments)
.unwrap_or_else(|_| json!({ "_raw_arguments": call.arguments }));
Some(DirectToolCall {
id,
name: call.name,
input,
})
})
.collect()
}
fn direct_message_from_parts(content: &str, tool_calls: &[DirectToolCall]) -> serde_json::Value {
let mut message = json!({
"role": "assistant",
"content": if content.trim().is_empty() {
serde_json::Value::Null
} else {
json!(content)
},
});
if !tool_calls.is_empty() {
if let Some(object) = message.as_object_mut() {
object.insert(
"tool_calls".to_string(),
json!(
tool_calls
.iter()
.map(|tool_call| json!({
"id": tool_call.id,
"type": "function",
"function": {
"name": tool_call.name,
"arguments": tool_call.input.to_string(),
},
}))
.collect::<Vec<_>>()
),
);
}
}
message
}
async fn send_openai_compatible_chat_request(
client: &reqwest::Client,
credential: &StoredOpenAiCompatibleCredential,
messages: &[serde_json::Value],
) -> Result<serde_json::Value, String> {
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,
)
.await
{
Ok(response) => Ok(response),
Err(err) if err.can_retry_non_streaming => {
eprintln!(
"[direct-provider] streaming request failed, retrying non-streaming: {}",
err.message
);
send_openai_compatible_non_streaming_chat_request(client, credential, messages).await
}
Err(err) => Err(err.message),
}
}
async fn send_openai_compatible_non_streaming_chat_request(
client: &reqwest::Client,
credential: &StoredOpenAiCompatibleCredential,
messages: &[serde_json::Value],
) -> Result<DirectChatResponse, String> {
let mut request_body = json!({
"model": credential.model.clone(),
"messages": messages,
@ -2819,8 +2990,142 @@ async fn send_openai_compatible_chat_request(
return Err(format!("Provider returned HTTP {}: {}", status, response_text));
}
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| format!("Provider returned invalid JSON: {}", err))?;
Ok(direct_chat_response_from_value(response))
}
async fn send_openai_compatible_streaming_chat_request(
client: &reqwest::Client,
credential: &StoredOpenAiCompatibleCredential,
messages: &[serde_json::Value],
window: &WebviewWindow,
tab_id: &str,
state: &ClaudeProcessState,
process_key: &str,
) -> Result<DirectChatResponse, DirectStreamFailure> {
let mut request_body = json!({
"model": credential.model.clone(),
"messages": 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"));
}
let mut response = client
.post(openai_chat_completions_url(&credential.base_url))
.bearer_auth(&credential.api_key)
.header("Content-Type", "application/json")
.body(request_body.to_string())
.send()
.await
.map_err(|err| DirectStreamFailure {
message: format!("Provider request failed: {}", err),
can_retry_non_streaming: true,
})?;
let status = response.status();
if !status.is_success() {
let response_text = response.text().await.unwrap_or_default();
return Err(DirectStreamFailure {
message: format!("Provider returned HTTP {}: {}", status, response_text),
can_retry_non_streaming: true,
});
}
let mut buffer = String::new();
let mut content = String::new();
let mut tool_calls: HashMap<usize, DirectStreamingToolCall> = HashMap::new();
let mut usage = json!({ "input_tokens": 0, "output_tokens": 0 });
let mut streamed_text = false;
while let Some(chunk) = response.chunk().await.map_err(|err| DirectStreamFailure {
message: format!("Failed to read provider stream: {}", err),
can_retry_non_streaming: false,
})? {
if direct_provider_cancelled(state, process_key).await {
return Err(DirectStreamFailure {
message: "Direct provider request cancelled".to_string(),
can_retry_non_streaming: false,
});
}
buffer.push_str(&String::from_utf8_lossy(&chunk));
while let Some(pos) = buffer.find('\n') {
let line = buffer[..pos].trim_end_matches('\r').trim().to_string();
buffer = buffer[pos + 1..].to_string();
if line.is_empty() || !line.starts_with("data:") {
continue;
}
let data = line.trim_start_matches("data:").trim();
if data == "[DONE]" {
let tool_calls = direct_tool_calls_from_stream(tool_calls);
let message = direct_message_from_parts(&content, &tool_calls);
return Ok(DirectChatResponse {
message,
content,
tool_calls,
usage,
streamed_text,
});
}
let value = serde_json::from_str::<serde_json::Value>(data).map_err(|err| {
DirectStreamFailure {
message: format!("Provider returned invalid stream JSON: {}", err),
can_retry_non_streaming: false,
}
})?;
if value.get("usage").is_some() {
usage = json_usage(&value);
}
let Some(delta) = value.pointer("/choices/0/delta") else {
continue;
};
let text_delta = direct_stream_text_delta(delta);
if !text_delta.is_empty() {
content.push_str(&text_delta);
streamed_text = true;
emit_direct_streaming_delta(window, tab_id, &text_delta);
}
if let Some(calls) = delta.get("tool_calls").and_then(|v| v.as_array()) {
for (fallback_idx, call) in calls.iter().enumerate() {
let idx = call
.get("index")
.and_then(|v| v.as_u64())
.and_then(|v| usize::try_from(v).ok())
.unwrap_or(fallback_idx);
let entry = tool_calls.entry(idx).or_default();
if let Some(id) = call.get("id").and_then(|v| v.as_str()) {
if !id.is_empty() {
entry.id = Some(id.to_string());
}
}
if let Some(name) = call.pointer("/function/name").and_then(|v| v.as_str()) {
entry.name.push_str(name);
}
if let Some(arguments) =
call.pointer("/function/arguments").and_then(|v| v.as_str())
{
entry.arguments.push_str(arguments);
}
}
}
}
}
let tool_calls = direct_tool_calls_from_stream(tool_calls);
let message = direct_message_from_parts(&content, &tool_calls);
Ok(DirectChatResponse {
message,
content,
tool_calls,
usage,
streamed_text,
})
}
async fn execute_openai_compatible_provider(
@ -2916,11 +3221,18 @@ async fn execute_openai_compatible_provider(
&client,
&credential,
&request_messages,
&window,
&tab_id,
&state,
&process_key,
)
.await
{
Ok(value) => value,
Err(err) => {
if direct_provider_cancelled(&state, &process_key).await {
return Ok(());
}
emit_direct_error(&window, &tab_id, err);
emit_direct_complete(&window, &tab_id, false);
return Ok(());
@ -2931,9 +3243,9 @@ async fn execute_openai_compatible_provider(
return Ok(());
}
let content = direct_assistant_content(&response);
let tool_calls = parse_direct_tool_calls(&response);
let usage = json_usage(&response);
let content = response.content;
let tool_calls = response.tool_calls;
let usage = response.usage;
final_usage = usage.clone();
let mut content_blocks = Vec::new();
@ -2956,13 +3268,18 @@ async fn execute_openai_compatible_provider(
return Ok(());
}
let assistant_event = json!({
let mut assistant_event = json!({
"type": "assistant",
"message": {
"content": content_blocks,
"usage": usage.clone(),
},
});
if response.streamed_text {
if let Some(object) = assistant_event.as_object_mut() {
object.insert("subtype".to_string(), json!("streaming_final"));
}
}
emit_and_persist_direct_output(
&window,
&project_path,
@ -2971,15 +3288,7 @@ async fn execute_openai_compatible_provider(
&assistant_event,
);
let assistant_message = response
.pointer("/choices/0/message")
.cloned()
.unwrap_or_else(|| {
json!({
"role": "assistant",
"content": content.clone(),
})
});
let assistant_message = response.message;
request_messages.push(assistant_message.clone());
messages.push(assistant_message);
@ -3922,6 +4231,63 @@ mod tests {
assert!(!glob.content.contains("notes.md"));
}
#[test]
fn test_direct_provider_parses_tool_arguments_string_and_object() {
let response = json!({
"choices": [{
"message": {
"role": "assistant",
"content": null,
"tool_calls": [
{
"id": "call-1",
"type": "function",
"function": {
"name": "Read",
"arguments": "{\"file_path\":\"main.tex\"}"
}
},
{
"id": "call-2",
"type": "function",
"function": {
"name": "Grep",
"arguments": { "pattern": "Intro", "glob": "*.tex" }
}
}
]
}
}]
});
let calls = parse_direct_tool_calls(&response);
assert_eq!(calls.len(), 2);
assert_eq!(calls[0].input["file_path"], "main.tex");
assert_eq!(calls[1].input["pattern"], "Intro");
}
#[test]
fn test_direct_provider_reconstructs_streamed_tool_calls() {
let mut calls = HashMap::new();
calls.insert(
0,
DirectStreamingToolCall {
id: Some("call-1".to_string()),
name: "Read".to_string(),
arguments: "{\"file_path\":\"main.tex\"}".to_string(),
},
);
let calls = direct_tool_calls_from_stream(calls);
let message = direct_message_from_parts("Checking the file", &calls);
assert_eq!(calls.len(), 1);
assert_eq!(calls[0].name, "Read");
assert_eq!(calls[0].input["file_path"], "main.tex");
assert_eq!(message["content"], "Checking the file");
assert_eq!(message["tool_calls"][0]["function"]["name"], "Read");
}
// --- claude_required_dirs ---
#[cfg(not(target_os = "windows"))]

View file

@ -271,6 +271,56 @@ describe("Multi-tab merge triggers", () => {
expect(state.messages).toHaveLength(0);
});
it("_appendMessage merges direct provider streaming deltas", () => {
const chat = useClaudeChatStore.getState();
chat._appendMessage("tab-default", {
type: "assistant",
subtype: "streaming_delta",
message: { content: [{ type: "text", text: "Hello" }] },
});
chat._appendMessage("tab-default", {
type: "assistant",
subtype: "streaming_delta",
message: { content: [{ type: "text", text: " world" }] },
});
const messages = useClaudeChatStore.getState().messages;
expect(messages).toHaveLength(1);
expect(messages[0].message?.content?.[0].text).toBe("Hello world");
});
it("_appendMessage replaces streaming deltas with final direct provider message", () => {
const chat = useClaudeChatStore.getState();
chat._appendMessage("tab-default", {
type: "assistant",
subtype: "streaming_delta",
message: { content: [{ type: "text", text: "Draft" }] },
});
chat._appendMessage("tab-default", {
type: "assistant",
subtype: "streaming_final",
message: {
content: [
{ type: "text", text: "Final" },
{
type: "tool_use",
id: "tool-1",
name: "Read",
input: { file_path: "main.tex" },
},
],
},
});
const messages = useClaudeChatStore.getState().messages;
expect(messages).toHaveLength(1);
expect(messages[0].subtype).toBe("streaming_final");
expect(messages[0].message?.content?.[0].text).toBe("Final");
expect(messages[0].message?.content?.[1].type).toBe("tool_use");
});
it("_setSessionId routes to the specified tab", () => {
const chat = useClaudeChatStore.getState();
const tabB = chat.createTab();

View file

@ -589,6 +589,48 @@ export const useClaudeChatStore = create<ClaudeChatState>()((set, get) => ({
const tab = state.tabs.find((t) => t.id === tabId);
if (!tab) return {};
if (msg.type === "assistant" && msg.subtype === "streaming_delta") {
const last = tab.messages[tab.messages.length - 1];
if (last?.type === "assistant" && last.subtype === "streaming_delta") {
const existing = last.message?.content ?? [];
const incoming = msg.message?.content ?? [];
const existingText = existing.find((block) => block.type === "text");
const incomingText = incoming.find((block) => block.type === "text");
if (existingText && incomingText?.text) {
const merged: ClaudeStreamMessage = {
...last,
message: {
...last.message,
content: existing.map((block) =>
block === existingText
? {
...block,
text: `${block.text ?? ""}${incomingText.text}`,
}
: block,
),
},
};
return applyTabUpdate(state, tabId, {
messages: [...tab.messages.slice(0, -1), merged],
totalInputTokens: tab.totalInputTokens + inputDelta,
totalOutputTokens: tab.totalOutputTokens + outputDelta,
});
}
}
}
if (msg.type === "assistant" && msg.subtype === "streaming_final") {
const last = tab.messages[tab.messages.length - 1];
if (last?.type === "assistant" && last.subtype === "streaming_delta") {
return applyTabUpdate(state, tabId, {
messages: [...tab.messages.slice(0, -1), msg],
totalInputTokens: tab.totalInputTokens + inputDelta,
totalOutputTokens: tab.totalOutputTokens + outputDelta,
});
}
}
return applyTabUpdate(state, tabId, {
messages: [...tab.messages, msg],
totalInputTokens: tab.totalInputTokens + inputDelta,