refactor: 抽离 AI pipeline 与调度共享能力逻辑

This commit is contained in:
fawney19
2026-04-10 01:46:14 +08:00
parent b901a6ffc7
commit 5014e2f5fd
255 changed files with 15057 additions and 3115 deletions

View File

@@ -145,6 +145,25 @@ impl ClaudeProviderState {
},
});
}
"thinking_delta" => {
let Some(piece) = delta
.get("thinking")
.and_then(Value::as_str)
.or_else(|| delta.get("text").and_then(Value::as_str))
else {
return Ok(out);
};
if piece.is_empty() {
return Ok(out);
}
self.ensure_started(report_context, &mut out);
let (id, model) = self.identity(report_context);
out.push(CanonicalStreamFrame {
id,
model,
event: CanonicalStreamEvent::ReasoningDelta(piece.to_string()),
});
}
_ => {}
}
}
@@ -162,6 +181,26 @@ impl ClaudeProviderState {
.get("type")
.and_then(Value::as_str)
.unwrap_or_default();
if block_type == "thinking" {
let Some(piece) = block
.get("thinking")
.and_then(Value::as_str)
.or_else(|| block.get("text").and_then(Value::as_str))
else {
return Ok(out);
};
if piece.is_empty() {
return Ok(out);
}
self.ensure_started(report_context, &mut out);
let (id, model) = self.identity(report_context);
out.push(CanonicalStreamFrame {
id,
model,
event: CanonicalStreamEvent::ReasoningDelta(piece.to_string()),
});
return Ok(out);
}
if block_type == "text" {
let Some(text) = block.get("text").and_then(Value::as_str) else {
return Ok(out);
@@ -273,12 +312,21 @@ enum ClaudeOpenBlock {
Text {
block_index: usize,
},
Thinking {
block_index: usize,
},
Tool {
tool_index: usize,
block_index: usize,
},
}
#[derive(Default)]
struct ClaudeClientToolState {
call_id: String,
name: String,
}
#[derive(Default)]
pub struct ClaudeClientEmitter {
message_id: Option<String>,
@@ -288,6 +336,7 @@ pub struct ClaudeClientEmitter {
next_block_index: usize,
open_block: Option<ClaudeOpenBlock>,
tool_block_indices: BTreeMap<usize, usize>,
tool_states: BTreeMap<usize, ClaudeClientToolState>,
}
impl ClaudeClientEmitter {
@@ -313,6 +362,10 @@ impl ClaudeClientEmitter {
"content": [],
"stop_reason": Value::Null,
"stop_sequence": Value::Null,
"usage": {
"input_tokens": 0,
"output_tokens": 0,
},
}
}),
)
@@ -324,6 +377,7 @@ impl ClaudeClientEmitter {
};
let block_index = match open_block {
ClaudeOpenBlock::Text { block_index } => block_index,
ClaudeOpenBlock::Thinking { block_index } => block_index,
ClaudeOpenBlock::Tool { block_index, .. } => block_index,
};
encode_json_sse(
@@ -358,6 +412,29 @@ impl ClaudeClientEmitter {
Ok(out)
}
fn ensure_thinking_block(&mut self) -> Result<Vec<u8>, PipelineFinalizeError> {
let mut out = Vec::new();
if let Some(ClaudeOpenBlock::Thinking { .. }) = self.open_block {
return Ok(out);
}
out.extend(self.close_open_block()?);
let block_index = self.next_block_index;
self.next_block_index += 1;
self.open_block = Some(ClaudeOpenBlock::Thinking { block_index });
out.extend(encode_json_sse(
Some("content_block_start"),
&json!({
"type": "content_block_start",
"index": block_index,
"content_block": {
"type": "thinking",
"thinking": "",
}
}),
)?);
Ok(out)
}
fn ensure_tool_block(
&mut self,
tool_index: usize,
@@ -429,19 +506,52 @@ impl ClaudeClientEmitter {
)?);
Ok(out)
}
CanonicalStreamEvent::ReasoningDelta(text) => {
let mut out = self.ensure_started()?;
out.extend(self.ensure_thinking_block()?);
let block_index = match self.open_block {
Some(ClaudeOpenBlock::Thinking { block_index }) => block_index,
_ => return Ok(out),
};
out.extend(encode_json_sse(
Some("content_block_delta"),
&json!({
"type": "content_block_delta",
"index": block_index,
"delta": {
"type": "thinking_delta",
"thinking": text,
}
}),
)?);
Ok(out)
}
CanonicalStreamEvent::ToolCallStart {
index,
call_id,
name,
} => {
let mut out = self.ensure_started()?;
let state = self.tool_states.entry(index).or_default();
state.call_id = call_id.clone();
state.name = name.clone();
out.extend(self.ensure_tool_block(index, &call_id, &name)?);
Ok(out)
}
CanonicalStreamEvent::ToolCallArgumentsDelta { index, arguments } => {
let mut out = self.ensure_started()?;
let call_id = format!("tool_{index}");
out.extend(self.ensure_tool_block(index, &call_id, "unknown")?);
let state = self.tool_states.entry(index).or_default();
let call_id = if state.call_id.is_empty() {
format!("tool_{index}")
} else {
state.call_id.clone()
};
let name = if state.name.is_empty() {
"unknown".to_string()
} else {
state.name.clone()
};
out.extend(self.ensure_tool_block(index, &call_id, &name)?);
let block_index = match self.open_block {
Some(ClaudeOpenBlock::Tool { block_index, .. }) => block_index,
_ => return Ok(out),
@@ -482,15 +592,14 @@ impl ClaudeClientEmitter {
"stop_sequence": Value::Null,
}),
);
if let Some(usage) = usage {
payload.insert(
"usage".to_string(),
json!({
"input_tokens": usage.input_tokens,
"output_tokens": usage.output_tokens,
}),
);
}
let usage = usage.unwrap_or_default();
payload.insert(
"usage".to_string(),
json!({
"input_tokens": usage.input_tokens,
"output_tokens": usage.output_tokens,
}),
);
out.extend(encode_json_sse(
Some("message_delta"),
&Value::Object(payload),
@@ -524,3 +633,124 @@ impl ClaudeClientEmitter {
})
}
}
#[cfg(test)]
mod tests {
use super::*;
use serde_json::json;
fn data_line(value: Value) -> Vec<u8> {
format!("data: {}\n", value).into_bytes()
}
#[test]
fn claude_provider_state_parses_thinking_deltas() {
let mut state = ClaudeProviderState::default();
let report_context = json!({});
let _ = state
.push_line(
&report_context,
data_line(json!({
"type": "message_start",
"message": {
"id": "msg_123",
"model": "claude-sonnet-4-5"
}
})),
)
.expect("message_start should parse");
let frames = state
.push_line(
&report_context,
data_line(json!({
"type": "content_block_delta",
"index": 0,
"delta": {
"type": "thinking_delta",
"thinking": "step by step"
}
})),
)
.expect("thinking delta should parse");
assert!(frames.iter().any(|frame| matches!(
frame.event,
CanonicalStreamEvent::ReasoningDelta(ref text) if text == "step by step"
)));
}
#[test]
fn claude_client_emitter_preserves_tool_identity_and_emits_thinking_blocks() {
let mut emitter = ClaudeClientEmitter::default();
let mut bytes = emitter
.emit(CanonicalStreamFrame {
id: "msg_123".to_string(),
model: "claude-sonnet-4-5".to_string(),
event: CanonicalStreamEvent::Start,
})
.expect("start should encode");
bytes.extend(
emitter
.emit(CanonicalStreamFrame {
id: "msg_123".to_string(),
model: "claude-sonnet-4-5".to_string(),
event: CanonicalStreamEvent::ReasoningDelta("step by step".to_string()),
})
.expect("reasoning should encode"),
);
bytes.extend(
emitter
.emit(CanonicalStreamFrame {
id: "msg_123".to_string(),
model: "claude-sonnet-4-5".to_string(),
event: CanonicalStreamEvent::ToolCallStart {
index: 0,
call_id: "toolu_1".to_string(),
name: "lookup".to_string(),
},
})
.expect("tool start should encode"),
);
bytes.extend(
emitter
.emit(CanonicalStreamFrame {
id: "msg_123".to_string(),
model: "claude-sonnet-4-5".to_string(),
event: CanonicalStreamEvent::ToolCallArgumentsDelta {
index: 0,
arguments: "{\"city\":\"Shanghai\"}".to_string(),
},
})
.expect("tool delta should encode"),
);
let sse = String::from_utf8(bytes).expect("sse should be utf8");
assert!(sse.contains("\"type\":\"thinking\""));
assert!(sse.contains("\"type\":\"thinking_delta\""));
assert!(sse.contains("\"id\":\"toolu_1\""));
assert!(sse.contains("\"name\":\"lookup\""));
assert!(sse.contains("\"partial_json\":\"{\\\"city\\\":\\\"Shanghai\\\"}\""));
assert!(sse.contains("\"usage\":{\"input_tokens\":0,\"output_tokens\":0}"));
}
#[test]
fn claude_client_emitter_injects_default_usage_into_finish_events() {
let mut emitter = ClaudeClientEmitter::default();
let bytes = emitter
.emit(CanonicalStreamFrame {
id: "msg_456".to_string(),
model: "gpt-5.4".to_string(),
event: CanonicalStreamEvent::Finish {
finish_reason: Some("stop".to_string()),
usage: None,
},
})
.expect("finish should encode");
let sse = String::from_utf8(bytes).expect("sse should be utf8");
assert!(sse.contains("event: message_start"));
assert!(sse.contains("event: message_delta"));
assert!(sse.contains("\"stop_reason\":\"end_turn\""));
assert!(sse.contains("\"usage\":{\"input_tokens\":0,\"output_tokens\":0}"));
}
}

View File

@@ -22,6 +22,7 @@ pub struct GeminiProviderState {
started: bool,
finished: bool,
text_parts: BTreeMap<usize, String>,
reasoning_parts: BTreeMap<usize, String>,
tool_calls: BTreeMap<usize, GeminiProviderToolState>,
}
@@ -97,8 +98,16 @@ impl GeminiProviderState {
let Some(part_object) = part.as_object() else {
continue;
};
if let Some(text) = part_object.get("text").and_then(Value::as_str) {
let previous = self.text_parts.entry(index).or_default();
if let Some(text) = render_gemini_part_as_text(part_object) {
let is_reasoning = part_object
.get("thought")
.and_then(Value::as_bool)
.unwrap_or(false);
let previous = if is_reasoning {
self.reasoning_parts.entry(index).or_default()
} else {
self.text_parts.entry(index).or_default()
};
let delta = if text.starts_with(previous.as_str()) {
text[previous.len()..].to_string()
} else if previous.as_str() == text {
@@ -106,12 +115,16 @@ impl GeminiProviderState {
} else {
text.to_string()
};
*previous = text.to_string();
*previous = text;
if !delta.is_empty() {
out.push(CanonicalStreamFrame {
id: id.clone(),
model: model.clone(),
event: CanonicalStreamEvent::TextDelta(delta),
event: if is_reasoning {
CanonicalStreamEvent::ReasoningDelta(delta)
} else {
CanonicalStreamEvent::TextDelta(delta)
},
});
}
continue;
@@ -179,7 +192,8 @@ impl GeminiProviderState {
let mut finish_reason = normalize_openai_finish_reason(match finish_reason {
"STOP" => Some("stop"),
"MAX_TOKENS" => Some("length"),
"SAFETY" => Some("content_filter"),
"SAFETY" | "RECITATION" | "BLOCKLIST" | "PROHIBITED_CONTENT" | "SPII"
| "OTHER" => Some("content_filter"),
other => Some(other),
});
if has_tool_calls && finish_reason.as_deref().is_none_or(|value| value == "stop") {
@@ -332,6 +346,9 @@ impl GeminiClientEmitter {
CanonicalStreamEvent::TextDelta(text) => {
self.emit_candidate(vec![json!({ "text": text })], None, None)
}
CanonicalStreamEvent::ReasoningDelta(text) => {
self.emit_candidate(vec![json!({ "text": text, "thought": true })], None, None)
}
CanonicalStreamEvent::ToolCallStart {
index,
call_id,
@@ -399,3 +416,112 @@ impl GeminiClientEmitter {
Ok(out)
}
}
fn render_gemini_part_as_text(part: &Map<String, Value>) -> Option<String> {
if let Some(text) = part.get("text").and_then(Value::as_str) {
return Some(text.to_string());
}
if let Some(code) = part.get("executableCode").and_then(Value::as_object) {
let language = code
.get("language")
.and_then(Value::as_str)
.unwrap_or_default();
let source = code.get("code").and_then(Value::as_str).unwrap_or_default();
return Some(format!("```{language}\n{source}\n```"));
}
if let Some(result) = part.get("codeExecutionResult").and_then(Value::as_object) {
let output = result
.get("output")
.and_then(Value::as_str)
.unwrap_or_default();
return Some(format!("```output\n{output}\n```"));
}
None
}
#[cfg(test)]
mod tests {
use super::*;
use serde_json::json;
fn data_line(value: Value) -> Vec<u8> {
format!("data: {}\n", value).into_bytes()
}
#[test]
fn gemini_provider_state_parses_thoughts_code_and_content_filter_finish() {
let mut state = GeminiProviderState::default();
let report_context = json!({});
let frames = state
.push_line(
&report_context,
data_line(json!({
"responseId": "resp_123",
"modelVersion": "gemini-2.5-pro",
"candidates": [{
"index": 0,
"finishReason": "RECITATION",
"content": {
"parts": [
{ "text": "reason", "thought": true },
{ "executableCode": { "language": "python", "code": "print(1)" } }
]
}
}],
"usageMetadata": {
"promptTokenCount": 1,
"candidatesTokenCount": 2,
"totalTokenCount": 3
}
})),
)
.expect("chunk should parse");
assert!(matches!(frames[0].event, CanonicalStreamEvent::Start));
assert!(frames.iter().any(|frame| matches!(
frame.event,
CanonicalStreamEvent::ReasoningDelta(ref text) if text == "reason"
)));
assert!(frames.iter().any(|frame| matches!(
frame.event,
CanonicalStreamEvent::TextDelta(ref text) if text == "```python\nprint(1)\n```"
)));
assert!(frames.iter().any(|frame| matches!(
frame.event,
CanonicalStreamEvent::Finish { ref finish_reason, .. }
if finish_reason.as_deref() == Some("content_filter")
)));
}
#[test]
fn gemini_client_emitter_marks_reasoning_parts_as_thoughts() {
let mut emitter = GeminiClientEmitter::default();
let mut bytes = emitter
.emit(CanonicalStreamFrame {
id: "resp_123".to_string(),
model: "gemini-2.5-pro".to_string(),
event: CanonicalStreamEvent::ReasoningDelta("reason".to_string()),
})
.expect("reasoning should encode");
bytes.extend(
emitter
.emit(CanonicalStreamFrame {
id: "resp_123".to_string(),
model: "gemini-2.5-pro".to_string(),
event: CanonicalStreamEvent::Finish {
finish_reason: Some("stop".to_string()),
usage: Some(CanonicalUsage {
input_tokens: 1,
output_tokens: 2,
total_tokens: 3,
}),
},
})
.expect("finish should encode"),
);
let sse = String::from_utf8(bytes).expect("sse should be utf8");
assert!(sse.contains("\"thought\":true"));
assert!(sse.contains("\"finishReason\":\"STOP\""));
}
}

View File

@@ -11,6 +11,7 @@ pub struct CanonicalUsage {
pub enum CanonicalStreamEvent {
Start,
TextDelta(String),
ReasoningDelta(String),
ToolCallStart {
index: usize,
call_id: String,
@@ -216,3 +217,23 @@ pub fn build_openai_chat_finish_chunk(id: &str, model: &str, finish_reason: Opti
}]
})
}
pub fn build_openai_chat_usage_chunk(
id: &str,
model: &str,
prompt_tokens: u64,
completion_tokens: u64,
total_tokens: u64,
) -> Value {
json!({
"id": id,
"object": "chat.completion.chunk",
"model": model,
"choices": [],
"usage": {
"prompt_tokens": prompt_tokens,
"completion_tokens": completion_tokens,
"total_tokens": total_tokens,
}
})
}

View File

@@ -1,18 +1,21 @@
use serde_json::Value;
use crate::conversion::{build_core_error_body_for_client_format, LocalCoreSyncErrorKind};
use crate::finalize::sse::encode_json_sse;
use crate::finalize::standard::claude::stream::{ClaudeClientEmitter, ClaudeProviderState};
use crate::finalize::standard::gemini::stream::{GeminiClientEmitter, GeminiProviderState};
use crate::finalize::standard::openai::stream::{
OpenAIChatClientEmitter, OpenAIChatProviderState, OpenAICliClientEmitter,
OpenAICliProviderState,
};
use crate::finalize::standard::stream_core::common::CanonicalStreamFrame;
use crate::finalize::standard::stream_core::common::{decode_json_data_line, CanonicalStreamFrame};
use crate::finalize::PipelineFinalizeError;
#[derive(Default)]
pub struct StreamingStandardFormatMatrix {
provider: Option<ProviderStreamParser>,
client: Option<ClientStreamEmitter>,
terminated: bool,
}
impl StreamingStandardFormatMatrix {
@@ -21,7 +24,14 @@ impl StreamingStandardFormatMatrix {
report_context: &Value,
line: Vec<u8>,
) -> Result<Vec<u8>, PipelineFinalizeError> {
if self.terminated {
return Ok(Vec::new());
}
self.ensure_initialized(report_context);
if let Some(error_body) = build_client_error_body_for_line(report_context, &line) {
self.terminated = true;
return self.emit_error(error_body);
}
let Some(provider) = self.provider.as_mut() else {
return Ok(Vec::new());
};
@@ -30,6 +40,9 @@ impl StreamingStandardFormatMatrix {
}
pub fn finish(&mut self, report_context: &Value) -> Result<Vec<u8>, PipelineFinalizeError> {
if self.terminated {
return Ok(Vec::new());
}
self.ensure_initialized(report_context);
let Some(provider) = self.provider.as_mut() else {
return Ok(Vec::new());
@@ -77,6 +90,13 @@ impl StreamingStandardFormatMatrix {
}
Ok(out)
}
fn emit_error(&mut self, error_body: Value) -> Result<Vec<u8>, PipelineFinalizeError> {
let Some(client) = self.client.as_mut() else {
return Ok(Vec::new());
};
client.emit_error(error_body)
}
}
enum ProviderStreamParser {
@@ -158,4 +178,399 @@ impl ClientStreamEmitter {
ClientStreamEmitter::Gemini(state) => state.finish(),
}
}
fn emit_error(&mut self, error_body: Value) -> Result<Vec<u8>, PipelineFinalizeError> {
match self {
ClientStreamEmitter::OpenAICli(state) => state.emit_error(error_body),
ClientStreamEmitter::Claude(_) => {
let event = error_body.get("type").and_then(Value::as_str);
encode_json_sse(event, &error_body)
}
ClientStreamEmitter::OpenAIChat(_) | ClientStreamEmitter::Gemini(_) => {
encode_json_sse(None, &error_body)
}
}
}
}
fn build_client_error_body_for_line(report_context: &Value, line: &[u8]) -> Option<Value> {
let value = decode_json_data_line(line)?;
let provider_api_format = report_context
.get("provider_api_format")
.and_then(Value::as_str)
.unwrap_or_default()
.trim()
.to_ascii_lowercase();
let client_api_format = report_context
.get("client_api_format")
.and_then(Value::as_str)
.unwrap_or_default()
.trim()
.to_ascii_lowercase();
let (message, code, kind) = parse_provider_error(&provider_api_format, &value)?;
build_core_error_body_for_client_format(&client_api_format, &message, code.as_deref(), kind)
}
fn parse_provider_error(
provider_api_format: &str,
payload: &Value,
) -> Option<(String, Option<String>, LocalCoreSyncErrorKind)> {
match provider_api_format {
"openai:chat" | "openai:cli" | "openai:compact" => parse_openai_error(payload),
"claude:chat" | "claude:cli" => parse_claude_error(payload),
"gemini:chat" | "gemini:cli" => parse_gemini_error(payload),
_ => None,
}
}
fn parse_openai_error(payload: &Value) -> Option<(String, Option<String>, LocalCoreSyncErrorKind)> {
let error = payload.get("error")?.as_object()?;
let message = error.get("message").and_then(Value::as_str)?.to_string();
let code = error
.get("code")
.and_then(Value::as_str)
.map(ToOwned::to_owned);
let kind = match error
.get("type")
.and_then(Value::as_str)
.unwrap_or_default()
{
"invalid_request_error" => LocalCoreSyncErrorKind::InvalidRequest,
"authentication_error" => LocalCoreSyncErrorKind::Authentication,
"permission_error" => LocalCoreSyncErrorKind::PermissionDenied,
"not_found_error" => LocalCoreSyncErrorKind::NotFound,
"rate_limit_error" => LocalCoreSyncErrorKind::RateLimit,
"context_length_exceeded" => LocalCoreSyncErrorKind::ContextLengthExceeded,
"overloaded_error" => LocalCoreSyncErrorKind::Overloaded,
_ => LocalCoreSyncErrorKind::ServerError,
};
Some((message, code, kind))
}
fn parse_claude_error(payload: &Value) -> Option<(String, Option<String>, LocalCoreSyncErrorKind)> {
let error = payload.get("error")?.as_object()?;
let message = error.get("message").and_then(Value::as_str)?.to_string();
let code = error
.get("code")
.and_then(Value::as_str)
.map(ToOwned::to_owned);
let kind = match error
.get("type")
.and_then(Value::as_str)
.unwrap_or_default()
{
"invalid_request_error" => LocalCoreSyncErrorKind::InvalidRequest,
"authentication_error" => LocalCoreSyncErrorKind::Authentication,
"permission_error" => LocalCoreSyncErrorKind::PermissionDenied,
"not_found_error" => LocalCoreSyncErrorKind::NotFound,
"rate_limit_error" => LocalCoreSyncErrorKind::RateLimit,
"overloaded_error" => LocalCoreSyncErrorKind::Overloaded,
_ => LocalCoreSyncErrorKind::ServerError,
};
Some((message, code, kind))
}
fn parse_gemini_error(payload: &Value) -> Option<(String, Option<String>, LocalCoreSyncErrorKind)> {
let error = payload.get("error")?.as_object()?;
let message = error.get("message").and_then(Value::as_str)?.to_string();
let code = error.get("code").map(|value| match value {
Value::String(text) => text.clone(),
Value::Number(number) => number.to_string(),
_ => String::new(),
});
let kind = match error
.get("status")
.and_then(Value::as_str)
.unwrap_or_default()
{
"INVALID_ARGUMENT" => LocalCoreSyncErrorKind::InvalidRequest,
"UNAUTHENTICATED" => LocalCoreSyncErrorKind::Authentication,
"PERMISSION_DENIED" => LocalCoreSyncErrorKind::PermissionDenied,
"NOT_FOUND" => LocalCoreSyncErrorKind::NotFound,
"RESOURCE_EXHAUSTED" => LocalCoreSyncErrorKind::RateLimit,
"UNAVAILABLE" => LocalCoreSyncErrorKind::Overloaded,
_ => LocalCoreSyncErrorKind::ServerError,
};
let code = code.filter(|value| !value.is_empty());
Some((message, code, kind))
}
#[cfg(test)]
mod tests {
use super::StreamingStandardFormatMatrix;
use serde_json::{json, Value};
fn report_context(provider_api_format: &str, client_api_format: &str) -> Value {
json!({
"provider_api_format": provider_api_format,
"client_api_format": client_api_format,
"mapped_model": "test-model",
})
}
fn data_line(value: Value) -> Vec<u8> {
format!("data: {}\n", value).into_bytes()
}
#[test]
fn transforms_provider_errors_to_openai_chat_error_bodies() {
let cases = [
(
"openai:chat",
data_line(json!({
"error": {
"message": "bad request",
"type": "invalid_request_error",
"code": "invalid_request",
}
})),
"\"message\":\"bad request\"",
"\"type\":\"invalid_request_error\"",
"\"code\":\"invalid_request\"",
),
(
"claude:chat",
data_line(json!({
"type": "error",
"error": {
"message": "slow down",
"type": "rate_limit_error",
"code": "rate_limit",
}
})),
"\"message\":\"slow down\"",
"\"type\":\"rate_limit_error\"",
"\"code\":\"rate_limit\"",
),
(
"gemini:cli",
data_line(json!({
"error": {
"code": 429,
"message": "quota exceeded",
"status": "RESOURCE_EXHAUSTED",
}
})),
"\"message\":\"quota exceeded\"",
"\"type\":\"rate_limit_error\"",
"\"code\":\"429\"",
),
];
for (provider_api_format, line, message, err_type, code) in cases {
let report_context = report_context(provider_api_format, "openai:chat");
let mut matrix = StreamingStandardFormatMatrix::default();
let output = matrix
.transform_line(&report_context, line)
.expect("error should convert");
let sse = String::from_utf8(output).expect("sse should be utf8");
assert!(sse.starts_with("data: {\"error\":"));
assert!(!sse.contains("event: "));
assert!(sse.contains(message));
assert!(sse.contains(err_type));
assert!(sse.contains(code));
assert!(matrix
.finish(&report_context)
.expect("finish should succeed")
.is_empty());
}
}
#[test]
fn transforms_provider_errors_to_claude_error_events() {
let cases = [
(
"openai:chat",
data_line(json!({
"error": {
"message": "bad request",
"type": "invalid_request_error",
"code": "invalid_request",
}
})),
"\"message\":\"bad request\"",
"\"type\":\"invalid_request_error\"",
"\"code\":\"invalid_request\"",
),
(
"claude:chat",
data_line(json!({
"type": "error",
"error": {
"message": "slow down",
"type": "rate_limit_error",
"code": "rate_limit",
}
})),
"\"message\":\"slow down\"",
"\"type\":\"rate_limit_error\"",
"\"code\":\"rate_limit\"",
),
(
"gemini:cli",
data_line(json!({
"error": {
"code": 429,
"message": "quota exceeded",
"status": "RESOURCE_EXHAUSTED",
}
})),
"\"message\":\"quota exceeded\"",
"\"type\":\"rate_limit_error\"",
"\"code\":\"429\"",
),
];
for (provider_api_format, line, message, err_type, code) in cases {
let report_context = report_context(provider_api_format, "claude:chat");
let mut matrix = StreamingStandardFormatMatrix::default();
let output = matrix
.transform_line(&report_context, line)
.expect("error should convert");
let sse = String::from_utf8(output).expect("sse should be utf8");
assert!(sse.starts_with("event: error\n"));
assert!(sse.contains("data: {"));
assert!(sse.contains("\"type\":\"error\""));
assert!(sse.contains("\"error\":{"));
assert!(sse.contains(message));
assert!(sse.contains(err_type));
assert!(sse.contains(code));
assert!(matrix
.finish(&report_context)
.expect("finish should succeed")
.is_empty());
}
}
#[test]
fn transforms_provider_errors_to_gemini_error_bodies() {
let cases = [
(
"openai:chat",
data_line(json!({
"error": {
"message": "bad request",
"type": "invalid_request_error",
"code": "invalid_request",
}
})),
"\"message\":\"bad request\"",
"\"code\":400",
"\"status\":\"INVALID_ARGUMENT\"",
),
(
"claude:chat",
data_line(json!({
"type": "error",
"error": {
"message": "slow down",
"type": "rate_limit_error",
"code": "rate_limit",
}
})),
"\"message\":\"slow down\"",
"\"code\":429",
"\"status\":\"RESOURCE_EXHAUSTED\"",
),
(
"gemini:cli",
data_line(json!({
"error": {
"code": 429,
"message": "quota exceeded",
"status": "RESOURCE_EXHAUSTED",
}
})),
"\"message\":\"quota exceeded\"",
"\"code\":429",
"\"status\":\"RESOURCE_EXHAUSTED\"",
),
];
for (provider_api_format, line, message, code, status) in cases {
let report_context = report_context(provider_api_format, "gemini:chat");
let mut matrix = StreamingStandardFormatMatrix::default();
let output = matrix
.transform_line(&report_context, line)
.expect("error should convert");
let sse = String::from_utf8(output).expect("sse should be utf8");
assert!(sse.starts_with("data: {\"error\":"));
assert!(!sse.contains("event: "));
assert!(sse.contains(message));
assert!(sse.contains(code));
assert!(sse.contains(status));
assert!(matrix
.finish(&report_context)
.expect("finish should succeed")
.is_empty());
}
}
#[test]
fn transforms_provider_errors_to_openai_cli_failed_events() {
let cases = [
(
"openai:chat",
data_line(json!({
"error": {
"message": "bad request",
"type": "invalid_request_error",
"code": "invalid_request",
}
})),
"\"message\":\"bad request\"",
"\"type\":\"invalid_request_error\"",
"\"code\":\"invalid_request\"",
),
(
"claude:chat",
data_line(json!({
"type": "error",
"error": {
"message": "slow down",
"type": "rate_limit_error",
"code": "rate_limit",
}
})),
"\"message\":\"slow down\"",
"\"type\":\"rate_limit_error\"",
"\"code\":\"rate_limit\"",
),
(
"gemini:cli",
data_line(json!({
"error": {
"code": 429,
"message": "quota exceeded",
"status": "RESOURCE_EXHAUSTED",
}
})),
"\"message\":\"quota exceeded\"",
"\"type\":\"rate_limit_error\"",
"\"code\":\"429\"",
),
];
for (provider_api_format, line, message, err_type, code) in cases {
let report_context = report_context(provider_api_format, "openai:cli");
let mut matrix = StreamingStandardFormatMatrix::default();
let output = matrix
.transform_line(&report_context, line)
.expect("error should convert");
let sse = String::from_utf8(output).expect("sse should be utf8");
assert!(sse.starts_with("event: response.failed\n"));
assert!(sse.contains("\"sequence_number\":1"));
assert!(sse.contains(message));
assert!(sse.contains(err_type));
assert!(sse.contains(code));
assert!(matrix
.finish(&report_context)
.expect("finish should succeed")
.is_empty());
}
}
}

View File

@@ -170,7 +170,7 @@ pub fn maybe_build_openai_chat_cross_format_sync_product_from_normalized_payload
"gemini:chat" | "gemini:cli" => {
convert_gemini_chat_response_to_openai_chat(&provider_body_json, report_context)
}
"openai:cli" | "openai:compact" => {
"openai:cli" => {
convert_openai_cli_response_to_openai_chat(&provider_body_json, report_context)
}
_ => None,
@@ -211,7 +211,7 @@ pub fn maybe_build_openai_cli_cross_format_sync_product_from_normalized_payload(
.trim()
.to_ascii_lowercase();
if !is_openai_cli_family_api_format(&client_api_format)
if client_api_format != "openai:cli"
|| sync_cli_response_conversion_kind(&provider_api_format, &client_api_format).is_none()
{
return Ok(None);
@@ -228,7 +228,7 @@ pub fn maybe_build_openai_cli_cross_format_sync_product_from_normalized_payload(
};
let Some(client_body_json) = (match provider_api_format.as_str() {
"openai:cli" | "openai:compact" => Some(provider_body_json.clone()),
"openai:cli" => Some(provider_body_json.clone()),
"claude:chat" | "claude:cli" => {
convert_claude_cli_response_to_openai_cli(&provider_body_json, report_context)
}
@@ -351,7 +351,12 @@ fn maybe_build_standard_same_format_sync_body(
return None;
}
body_json.cloned()
let body_json = body_json?;
if is_error_like_sync_body(body_json) {
return None;
}
Some(body_json.clone())
}
fn maybe_build_standard_same_format_stream_sync_body(
@@ -429,14 +434,25 @@ fn maybe_build_openai_cli_same_family_sync_body(
.unwrap_or_default()
.trim()
.to_ascii_lowercase();
let needs_conversion = report_context
.get("needs_conversion")
.and_then(Value::as_bool)
.unwrap_or(false);
if !is_openai_cli_family_api_format(&provider_api_format)
|| !is_openai_cli_family_api_format(&client_api_format)
|| provider_api_format != client_api_format
|| needs_conversion
{
return None;
}
body_json.cloned()
let body_json = body_json?;
if is_error_like_sync_body(body_json) {
return None;
}
Some(body_json.clone())
}
fn maybe_build_openai_cli_same_family_stream_sync_body(
@@ -472,7 +488,8 @@ fn maybe_build_openai_cli_same_family_stream_sync_body(
if !is_openai_cli_family_api_format(&provider_api_format)
|| !is_openai_cli_family_api_format(&client_api_format)
|| (provider_api_format == client_api_format && needs_conversion)
|| provider_api_format != client_api_format
|| needs_conversion
{
return Ok(None);
}
@@ -507,6 +524,32 @@ fn maybe_build_openai_cross_format_provider_body_from_normalized_payload(
Ok(aggregated_stream_body.or_else(|| body_json.cloned()))
}
fn is_error_like_sync_body(value: &Value) -> bool {
let Some(object) = value.as_object() else {
return false;
};
object.contains_key("error")
|| object
.get("type")
.and_then(Value::as_str)
.is_some_and(|value| value == "error")
|| object
.get("chunks")
.and_then(Value::as_array)
.is_some_and(|chunks| {
chunks.iter().any(|chunk| {
chunk.as_object().is_some_and(|chunk_object| {
chunk_object.contains_key("error")
|| chunk_object
.get("type")
.and_then(Value::as_str)
.is_some_and(|value| value == "error")
})
})
})
}
pub fn maybe_build_standard_cross_format_sync_product(
report_kind: &str,
provider_api_format: &str,
@@ -1402,6 +1445,33 @@ mod tests {
assert!(body_json.is_none());
}
#[test]
fn rejects_standard_same_format_error_body_json() {
let report_context = json!({
"provider_api_format": "claude:chat",
"client_api_format": "claude:chat",
"needs_conversion": false,
});
let provider_body_json = json!({
"type": "error",
"error": {
"type": "rate_limit_error",
"message": "slow down"
}
});
let body_json = maybe_build_standard_same_format_sync_body_from_normalized_payload(
"claude_chat_sync_finalize",
200,
Some(&report_context),
Some(&provider_body_json),
None,
)
.expect("same-format error guard should not error");
assert!(body_json.is_none());
}
#[test]
fn builds_openai_cli_same_family_body_from_stream_payload() {
let body = concat!(
@@ -1412,12 +1482,12 @@ mod tests {
);
let report_context = json!({
"provider_api_format": "openai:compact",
"client_api_format": "openai:cli",
"needs_conversion": true,
"client_api_format": "openai:compact",
"needs_conversion": false,
});
let body_json = maybe_build_openai_cli_same_family_sync_body_from_normalized_payload(
"openai_cli_sync_finalize",
"openai_compact_sync_finalize",
200,
Some(&report_context),
None,
@@ -1458,7 +1528,7 @@ mod tests {
fn falls_back_to_body_json_for_openai_cli_same_family_sync_payload() {
let report_context = json!({
"provider_api_format": "openai:compact",
"client_api_format": "openai:cli",
"client_api_format": "openai:compact",
"needs_conversion": false,
});
let provider_body_json = json!({
@@ -1481,6 +1551,32 @@ mod tests {
assert_eq!(body_json, provider_body_json);
}
#[test]
fn rejects_openai_cli_same_family_error_body_json() {
let report_context = json!({
"provider_api_format": "openai:cli",
"client_api_format": "openai:cli",
"needs_conversion": false,
});
let provider_body_json = json!({
"error": {
"message": "quota reached",
"type": "rate_limit_error"
}
});
let body_json = maybe_build_openai_cli_same_family_sync_body_from_normalized_payload(
"openai_cli_sync_finalize",
200,
Some(&report_context),
Some(&provider_body_json),
None,
)
.expect("openai-cli same-family error guard should not error");
assert!(body_json.is_none());
}
#[test]
fn builds_openai_chat_cross_format_sync_product_from_claude_body_json() {
let report_context = json!({
@@ -1640,54 +1736,6 @@ mod tests {
);
}
#[test]
fn builds_openai_compact_cross_format_sync_product_for_function_call_case() {
let report_context = json!({
"provider_api_format": "gemini:cli",
"client_api_format": "openai:compact",
"model": "gpt-5",
});
let provider_body_json = json!({
"responseId": "resp_cli_tool_123",
"candidates": [{
"content": {
"parts": [
{"text": "Need a tool."},
{"functionCall": {"name": "get_weather", "args": {"location": "Tokyo"}}}
],
"role": "model"
},
"finishReason": "STOP",
"index": 0
}],
"modelVersion": "gemini-cli-upstream",
"usageMetadata": {
"promptTokenCount": 3,
"candidatesTokenCount": 5,
"thoughtsTokenCount": 2,
"totalTokenCount": 10
}
});
let product = maybe_build_openai_cli_cross_format_sync_product_from_normalized_payload(
"openai_compact_sync_finalize",
200,
Some(&report_context),
Some(&provider_body_json),
None,
)
.expect("openai-compact cross-format should succeed")
.expect("product should exist");
assert_eq!(product.provider_body_json, provider_body_json);
assert_eq!(product.client_body_json["object"], "response");
assert_eq!(
product.client_body_json["output"][1]["type"],
"function_call"
);
assert_eq!(product.client_body_json["output"][1]["name"], "get_weather");
}
#[test]
fn rejects_openai_chat_cross_format_for_unsupported_matrix() {
let report_context = json!({
@@ -1741,8 +1789,8 @@ mod tests {
fn standard_sync_finalize_product_handles_openai_cli_same_family_body() {
let report_context = json!({
"provider_api_format": "openai:compact",
"client_api_format": "openai:cli",
"needs_conversion": true,
"client_api_format": "openai:compact",
"needs_conversion": false,
});
let provider_body_json = json!({
"id": "resp_123",
@@ -1752,7 +1800,7 @@ mod tests {
});
let product = maybe_build_standard_sync_finalize_product_from_normalized_payload(
"openai_cli_sync_finalize",
"openai_compact_sync_finalize",
200,
Some(&report_context),
Some(&provider_body_json),