Fix reasoning model directive response identity

This commit is contained in:
fawney19
2026-05-11 22:48:44 +08:00
parent 7ae38b6c43
commit 0fa97595bf
9 changed files with 465 additions and 25 deletions

View File

@@ -98,6 +98,21 @@ pub fn model_directive_base_model(model: &str) -> Option<String> {
parse_model_directive(model).map(|directive| directive.base_model)
}
pub(crate) fn model_directive_display_model(model: &str) -> Option<String> {
let model = model.trim();
parse_model_directive(model)?;
Some(model.to_string())
}
pub(crate) fn model_directive_display_model_from_report_context(
report_context: &Value,
) -> Option<String> {
report_context
.get("model")
.and_then(Value::as_str)
.and_then(model_directive_display_model)
}
pub fn normalize_model_directive_model(model: &str) -> String {
parse_model_directive(model)
.map(|directive| directive.base_model)

View File

@@ -1,5 +1,7 @@
use serde_json::{json, Map, Value};
use crate::formats::shared::model_directives::model_directive_display_model_from_report_context;
pub use aether_ai_formats::protocol::stream::{
CanonicalContentPart, CanonicalStreamEvent, CanonicalStreamFrame, CanonicalUsage,
};
@@ -27,6 +29,9 @@ pub fn resolve_identity(
.filter(|value| !value.is_empty())
.unwrap_or(default_id)
.to_string();
if let Some(display_model) = model_directive_display_model_from_report_context(report_context) {
return (id, display_model);
}
let model = model
.filter(|value| !value.is_empty())
.or_else(|| report_context.get("mapped_model").and_then(Value::as_str))

View File

@@ -1,6 +1,7 @@
use serde_json::Value;
use crate::formats::openai::image::stream::OpenAiImageStreamState;
use crate::formats::shared::model_directives::model_directive_display_model_from_report_context;
use crate::formats::shared::stream_core::StreamingStandardFormatMatrix;
use crate::formats::shared::AiSurfaceFinalizeError;
use crate::provider_compat::kiro_stream::KiroToClaudeCliStreamState;
@@ -12,6 +13,7 @@ use crate::provider_compat::surfaces::{
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum FinalizeStreamRewriteMode {
EnvelopeUnwrap,
ModelDirectiveDisplay,
OpenAiImage,
Standard,
KiroToClaudeCli,
@@ -73,6 +75,17 @@ pub fn resolve_finalize_stream_rewrite_mode(
.then_some(FinalizeStreamRewriteMode::KiroToClaudeCli);
}
if model_directive_display_model_from_report_context(report_context).is_some()
&& provider_api_format == client_api_format
&& is_standard_provider_api_format(provider_api_format.as_str())
&& !provider_adaptation_should_unwrap_stream_envelope(
envelope_name.as_str(),
provider_api_format.as_str(),
)
{
return Some(FinalizeStreamRewriteMode::ModelDirectiveDisplay);
}
(provider_api_format == client_api_format
&& provider_adaptation_should_unwrap_stream_envelope(
envelope_name.as_str(),
@@ -83,6 +96,7 @@ pub fn resolve_finalize_stream_rewrite_mode(
enum AiSurfaceStreamRewriteState {
EnvelopeUnwrap,
ModelDirectiveDisplay,
OpenAiImage(Box<OpenAiImageStreamState>),
Standard(Box<StreamingStandardFormatMatrix>),
KiroToClaudeCli(Box<KiroToClaudeCliStreamState>),
@@ -104,6 +118,9 @@ pub fn maybe_build_ai_surface_stream_rewriter<'a>(
let report_context = report_context?;
let state = match resolve_finalize_stream_rewrite_mode(report_context)? {
FinalizeStreamRewriteMode::EnvelopeUnwrap => AiSurfaceStreamRewriteState::EnvelopeUnwrap,
FinalizeStreamRewriteMode::ModelDirectiveDisplay => {
AiSurfaceStreamRewriteState::ModelDirectiveDisplay
}
FinalizeStreamRewriteMode::OpenAiImage => {
AiSurfaceStreamRewriteState::OpenAiImage(Box::<OpenAiImageStreamState>::default())
}
@@ -142,6 +159,7 @@ impl AiSurfaceStreamRewriter<'_> {
transform_standard_bytes(standard, self.report_context, claude_bytes)
}
AiSurfaceStreamRewriteState::EnvelopeUnwrap
| AiSurfaceStreamRewriteState::ModelDirectiveDisplay
| AiSurfaceStreamRewriteState::Standard(_) => {
self.buffered.extend_from_slice(chunk);
let mut output = Vec::new();
@@ -170,6 +188,7 @@ impl AiSurfaceStreamRewriter<'_> {
Ok(output)
}
AiSurfaceStreamRewriteState::EnvelopeUnwrap
| AiSurfaceStreamRewriteState::ModelDirectiveDisplay
| AiSurfaceStreamRewriteState::Standard(_) => {
if self.buffered.is_empty() {
if let AiSurfaceStreamRewriteState::Standard(state) = &mut self.state {
@@ -190,8 +209,12 @@ impl AiSurfaceStreamRewriter<'_> {
fn transform_line(&mut self, line: Vec<u8>) -> Result<Vec<u8>, AiSurfaceFinalizeError> {
match &mut self.state {
AiSurfaceStreamRewriteState::EnvelopeUnwrap => {
transform_provider_private_stream_line(self.report_context, line)
.map_err(AiSurfaceFinalizeError::from)
let output = transform_provider_private_stream_line(self.report_context, line)
.map_err(AiSurfaceFinalizeError::from)?;
rewrite_model_directive_stream_line(self.report_context, output)
}
AiSurfaceStreamRewriteState::ModelDirectiveDisplay => {
rewrite_model_directive_stream_line(self.report_context, line)
}
AiSurfaceStreamRewriteState::Standard(state) => {
transform_standard_line(state, self.report_context, line)
@@ -203,6 +226,63 @@ impl AiSurfaceStreamRewriter<'_> {
}
}
fn rewrite_model_directive_stream_line(
report_context: &Value,
line: Vec<u8>,
) -> Result<Vec<u8>, AiSurfaceFinalizeError> {
let Some(display_model) = model_directive_display_model_from_report_context(report_context)
else {
return Ok(line);
};
let text = match std::str::from_utf8(&line) {
Ok(text) => text,
Err(_) => return Ok(line),
};
let trimmed_line_end = text.trim_end_matches(['\r', '\n']);
let trailing = &text[trimmed_line_end.len()..];
let Some((prefix, payload)) = trimmed_line_end.split_once(':') else {
return Ok(line);
};
if prefix.trim() != "data" {
return Ok(line);
}
let payload = payload.trim_start();
if payload.is_empty() || payload == "[DONE]" {
return Ok(line);
}
let mut value = match serde_json::from_str::<Value>(payload) {
Ok(value) => value,
Err(_) => return Ok(line),
};
if !rewrite_stream_payload_model(&mut value, &display_model) {
return Ok(line);
}
let mut output = Vec::new();
output.extend_from_slice(b"data: ");
output.extend(serde_json::to_vec(&value)?);
output.extend_from_slice(trailing.as_bytes());
Ok(output)
}
fn rewrite_stream_payload_model(value: &mut Value, display_model: &str) -> bool {
let Some(object) = value.as_object_mut() else {
return false;
};
let mut changed = false;
for key in ["model", "modelVersion"] {
if object.get(key).and_then(Value::as_str).is_some() {
object.insert(key.to_string(), Value::String(display_model.to_string()));
changed = true;
}
}
for key in ["response", "message"] {
if let Some(nested) = object.get_mut(key) {
changed |= rewrite_stream_payload_model(nested, display_model);
}
}
changed
}
fn transform_standard_bytes(
standard: &mut StreamingStandardFormatMatrix,
report_context: &Value,
@@ -288,7 +368,10 @@ fn is_standard_cli_client_api_format(api_format: &str) -> bool {
mod tests {
use serde_json::json;
use super::{resolve_finalize_stream_rewrite_mode, FinalizeStreamRewriteMode};
use super::{
maybe_build_ai_surface_stream_rewriter, resolve_finalize_stream_rewrite_mode,
FinalizeStreamRewriteMode,
};
#[test]
fn resolves_standard_mode_for_cross_format_standard_streams() {
@@ -341,6 +424,84 @@ mod tests {
assert_eq!(resolve_finalize_stream_rewrite_mode(&report_context), None);
}
#[test]
fn resolves_model_directive_display_mode_for_same_format_standard_streams() {
let report_context = json!({
"provider_api_format": "openai:responses",
"client_api_format": "openai:responses",
"model": "gpt-5.5-xhigh",
"mapped_model": "gpt-5.5",
"needs_conversion": false,
});
assert_eq!(
resolve_finalize_stream_rewrite_mode(&report_context),
Some(FinalizeStreamRewriteMode::ModelDirectiveDisplay)
);
}
#[test]
fn model_directive_display_mode_does_not_displace_kiro_stream_bridge() {
let report_context = json!({
"provider_api_format": "claude:messages",
"client_api_format": "claude:messages",
"envelope_name": "kiro:generateAssistantResponse",
"model": "claude-sonnet-4.5-high",
"mapped_model": "claude-sonnet-4.5",
"needs_conversion": false,
});
assert_eq!(
resolve_finalize_stream_rewrite_mode(&report_context),
Some(FinalizeStreamRewriteMode::KiroToClaudeCli)
);
}
#[test]
fn envelope_unwrap_rewriter_restores_model_directive_display_model() {
let report_context = json!({
"provider_api_format": "gemini:generate_content",
"client_api_format": "gemini:generate_content",
"envelope_name": "gemini_cli:v1internal",
"model": "gemini-2.5-pro-high",
"mapped_model": "gemini-2.5-pro",
"needs_conversion": false,
});
let mut rewriter = maybe_build_ai_surface_stream_rewriter(Some(&report_context))
.expect("rewriter should exist");
let output = rewriter
.push_chunk(
b"data: {\"response\":{\"modelVersion\":\"gemini-2.5-pro\",\"candidates\":[]}}\n\n",
)
.expect("rewrite should succeed");
let output = String::from_utf8(output).expect("output should be utf8");
assert!(output.contains("\"modelVersion\":\"gemini-2.5-pro-high\""));
assert!(!output.contains("\"modelVersion\":\"gemini-2.5-pro\""));
}
#[test]
fn model_directive_display_rewriter_restores_response_model() {
let report_context = json!({
"provider_api_format": "openai:responses",
"client_api_format": "openai:responses",
"model": "gpt-5.5-xhigh",
"mapped_model": "gpt-5.5",
"needs_conversion": false,
});
let mut rewriter = maybe_build_ai_surface_stream_rewriter(Some(&report_context))
.expect("rewriter should exist");
let output = rewriter
.push_chunk(
b"event: response.created\n\
data: {\"type\":\"response.created\",\"response\":{\"id\":\"resp_123\",\"object\":\"response\",\"model\":\"gpt-5.5\",\"status\":\"in_progress\"}}\n\n",
)
.expect("rewrite should succeed");
let output = String::from_utf8(output).expect("output should be utf8");
assert!(output.contains("event: response.created"));
assert!(output.contains("\"model\":\"gpt-5.5-xhigh\""));
assert!(!output.contains("\"model\":\"gpt-5.5\""));
}
#[test]
fn resolves_openai_image_mode_for_same_format_image_streams() {
let report_context = json!({

View File

@@ -19,6 +19,7 @@ use serde_json::{json, Map, Value};
use super::AiSurfaceFinalizeError;
use crate::formats::gemini::generate_content::stream::GeminiProviderState;
use crate::formats::shared::model_directives::model_directive_display_model_from_report_context;
use crate::formats::shared::response::remove_empty_pages_from_tool_arguments;
use crate::formats::shared::stream_core::common::{
map_openai_finish_reason_to_gemini, parse_json_arguments_value, CanonicalContentPart,
@@ -382,7 +383,11 @@ fn maybe_build_standard_same_format_sync_body(
return None;
}
Some(body_json.clone())
Some(client_body_with_report_context_model(
body_json.clone(),
report_context,
&client_api_format,
))
}
fn maybe_build_standard_same_format_stream_sync_body(
@@ -431,10 +436,11 @@ fn maybe_build_standard_same_format_stream_sync_body(
return Ok(None);
};
let body_bytes = base64::engine::general_purpose::STANDARD.decode(body_base64)?;
Ok(aggregate_same_format_stream_sync_response(
expected_api_format,
&body_bytes,
))
Ok(
aggregate_same_format_stream_sync_response(expected_api_format, &body_bytes).map(|body| {
client_body_with_report_context_model(body, report_context, &client_api_format)
}),
)
}
fn maybe_build_openai_responses_same_family_sync_body(
@@ -480,7 +486,11 @@ fn maybe_build_openai_responses_same_family_sync_body(
return None;
}
Some(body_json.clone())
Some(client_body_with_report_context_model(
body_json.clone(),
report_context,
&client_api_format,
))
}
fn maybe_build_openai_responses_same_family_stream_sync_body(
@@ -528,7 +538,11 @@ fn maybe_build_openai_responses_same_family_stream_sync_body(
return Ok(None);
};
let body_bytes = base64::engine::general_purpose::STANDARD.decode(body_base64)?;
Ok(aggregate_openai_responses_stream_sync_response(&body_bytes))
Ok(
aggregate_openai_responses_stream_sync_response(&body_bytes).map(|body| {
client_body_with_report_context_model(body, report_context, &client_api_format)
}),
)
}
fn maybe_build_openai_cross_format_provider_body_from_normalized_payload(
@@ -625,6 +639,8 @@ pub fn maybe_build_standard_cross_format_sync_product(
} else {
return None;
};
let client_body_json =
client_body_with_report_context_model(client_body_json, report_context, &client_api_format);
Some(StandardCrossFormatSyncProduct {
client_body_json,
@@ -797,6 +813,30 @@ fn format_context_from_report_context(report_context: &Value) -> FormatContext {
context
}
fn client_body_with_report_context_model(
mut body_json: Value,
report_context: &Value,
client_api_format: &str,
) -> Value {
let Some(display_model) = model_directive_display_model_from_report_context(report_context)
else {
return body_json;
};
let Some(object) = body_json.as_object_mut() else {
return body_json;
};
match normalize_openai_responses_family_api_format(client_api_format).as_str() {
"openai:chat" | "openai:responses" | "openai:responses:compact" | "claude:messages" => {
object.insert("model".to_string(), Value::String(display_model));
}
"gemini:generate_content" => {
object.insert("modelVersion".to_string(), Value::String(display_model));
}
_ => {}
}
body_json
}
fn convert_openai_chat_canonical_chat_response(
body_json: &Value,
client_api_format: &str,
@@ -1195,6 +1235,10 @@ fn gemini_response_can_use_single_response_canonical(body_json: &Value) -> bool
}
fn apply_report_context_model_fallback(model: &mut String, report_context: &Value) {
if let Some(display_model) = model_directive_display_model_from_report_context(report_context) {
*model = display_model;
return;
}
if model != "unknown" && !model.trim().is_empty() {
return;
}
@@ -2881,6 +2925,7 @@ mod tests {
maybe_build_openai_chat_cross_format_sync_product_from_normalized_payload,
maybe_build_openai_responses_cross_format_sync_product_from_normalized_payload,
maybe_build_openai_responses_same_family_sync_body_from_normalized_payload,
maybe_build_standard_cross_format_sync_product,
maybe_build_standard_cross_format_sync_product_from_normalized_payload,
maybe_build_standard_same_format_sync_body_from_normalized_payload,
maybe_build_standard_sync_finalize_product_from_normalized_payload,
@@ -3222,6 +3267,75 @@ mod tests {
assert_eq!(body_json, provider_body_json);
}
#[test]
fn same_format_sync_response_restores_model_directive_display_model() {
let report_context = json!({
"provider_api_format": "openai:responses",
"client_api_format": "openai:responses",
"model": "gpt-5.5-xhigh",
"mapped_model": "gpt-5.5",
"needs_conversion": false,
});
let provider_body_json = json!({
"id": "resp_123",
"object": "response",
"model": "gpt-5.5",
"status": "completed",
"output": []
});
let product = maybe_build_standard_sync_finalize_product_from_normalized_payload(
"openai_responses_sync_finalize",
200,
Some(&report_context),
Some(&provider_body_json),
None,
)
.expect("same-format sync body should succeed")
.expect("body should exist");
let StandardSyncFinalizeNormalizedProduct::SuccessBody(body_json) = product else {
panic!("same-format response should be a success body");
};
assert_eq!(body_json["model"], "gpt-5.5-xhigh");
}
#[test]
fn cross_format_sync_response_restores_model_directive_display_model() {
let report_context = json!({
"provider_api_format": "openai:responses",
"client_api_format": "claude:messages",
"model": "gpt-5.5-xhigh",
"mapped_model": "gpt-5.5",
"needs_conversion": true,
});
let provider_body_json = json!({
"id": "resp_123",
"object": "response",
"model": "gpt-5.5",
"status": "completed",
"output": [{
"type": "message",
"id": "msg_123",
"role": "assistant",
"status": "completed",
"content": [{ "type": "output_text", "text": "done" }]
}]
});
let product = maybe_build_standard_cross_format_sync_product(
"claude_cli_sync_finalize",
"openai:responses",
"claude:messages",
&report_context,
provider_body_json,
)
.expect("cross-format product should exist");
assert_eq!(product.client_body_json["model"], "gpt-5.5-xhigh");
assert_eq!(product.provider_body_json["model"], "gpt-5.5");
}
#[test]
fn rejects_standard_same_format_when_needs_conversion_is_true() {
let report_context = json!({

View File

@@ -1,6 +1,7 @@
use serde_json::Value;
use uuid::Uuid;
use crate::formats::shared::model_directives::model_directive_display_model_from_report_context;
use crate::formats::shared::AiSurfaceFinalizeError;
use crate::provider_compat::kiro_stream::{
build_kiro_initial_sse_events, build_kiro_stream_error_sse_events, encode_kiro_sse_events,
@@ -63,18 +64,22 @@ impl KiroToClaudeCliStreamState {
impl KiroClaudeStreamState {
pub(super) fn new(report_context: &Value) -> Self {
let model = report_context
.get("mapped_model")
.and_then(Value::as_str)
.filter(|value| !value.is_empty())
let model = model_directive_display_model_from_report_context(report_context)
.or_else(|| {
report_context
.get("mapped_model")
.and_then(Value::as_str)
.filter(|value| !value.is_empty())
.map(ToOwned::to_owned)
})
.or_else(|| {
report_context
.get("model")
.and_then(Value::as_str)
.filter(|value| !value.is_empty())
.map(ToOwned::to_owned)
})
.unwrap_or("unknown")
.to_string();
.unwrap_or_else(|| "unknown".to_string());
let thinking_enabled = report_context
.get("original_request_body")
.and_then(Value::as_object)

View File

@@ -86,6 +86,32 @@ fn kiro_stream_rewriter_converts_text_events_to_claude_sse() {
assert!(text.contains("\"input_tokens\":2000"));
}
#[test]
fn kiro_stream_rewriter_restores_model_directive_display_model() {
let report_context = json!({
"provider_api_format": "claude:messages",
"client_api_format": "claude:messages",
"envelope_name": "kiro:generateAssistantResponse",
"model": "claude-sonnet-4.5-high",
"mapped_model": "claude-sonnet-4.5"
});
let mut rewriter = KiroToClaudeCliStreamState::new(&report_context);
let first = rewriter
.push_chunk(
&report_context,
&encode_event_frame(
"event",
Some("assistantResponseEvent"),
&json!({"content": "Hello"}),
),
)
.expect("rewrite should succeed");
let text = String::from_utf8(first).expect("utf8 should decode");
assert!(text.contains("\"model\":\"claude-sonnet-4.5-high\""));
assert!(!text.contains("\"model\":\"claude-sonnet-4.5\""));
}
#[test]
fn kiro_stream_rewriter_converts_tool_use_to_claude_events() {
let report_context = kiro_report_context(false);