mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-02 09:20:22 +08:00
refactor: fold ai surfaces into formats
This commit is contained in:
180
crates/aether-ai-formats/src/provider_compat/kiro_stream.rs
Normal file
180
crates/aether-ai-formats/src/provider_compat/kiro_stream.rs
Normal file
@@ -0,0 +1,180 @@
|
||||
use serde_json::{json, Value};
|
||||
|
||||
pub use self::state::KiroToClaudeCliStreamState;
|
||||
|
||||
mod state;
|
||||
|
||||
pub const KIRO_CONTEXT_WINDOW_TOKENS: f64 = 200_000.0;
|
||||
pub const KIRO_MAX_THINKING_BUFFER: usize = 1024 * 1024;
|
||||
|
||||
const KIRO_QUOTE_CHARS: &str = "`\"'\\#!@$%^&*()-_=+[]{};:<>,.?/";
|
||||
|
||||
pub fn encode_kiro_sse_events(events: Vec<Value>) -> Result<Vec<u8>, serde_json::Error> {
|
||||
let mut output = Vec::new();
|
||||
for event in events {
|
||||
output.extend(encode_kiro_sse_event(&event)?);
|
||||
}
|
||||
Ok(output)
|
||||
}
|
||||
|
||||
pub fn encode_kiro_sse_event(event: &Value) -> Result<Vec<u8>, serde_json::Error> {
|
||||
let encoded = serde_json::to_string(event)?;
|
||||
if let Some(event_type) = event.get("type").and_then(Value::as_str) {
|
||||
Ok(format!("event: {event_type}\ndata: {encoded}\n\n").into_bytes())
|
||||
} else {
|
||||
Ok(format!("data: {encoded}\n\n").into_bytes())
|
||||
}
|
||||
}
|
||||
|
||||
pub fn build_kiro_initial_sse_events(
|
||||
message_id: &str,
|
||||
model: &str,
|
||||
estimated_input_tokens: usize,
|
||||
) -> Vec<Value> {
|
||||
vec![json!({
|
||||
"type": "message_start",
|
||||
"message": {
|
||||
"id": message_id,
|
||||
"type": "message",
|
||||
"role": "assistant",
|
||||
"content": [],
|
||||
"model": model,
|
||||
"stop_reason": Value::Null,
|
||||
"stop_sequence": Value::Null,
|
||||
"usage": {
|
||||
"input_tokens": estimated_input_tokens as u64,
|
||||
"output_tokens": 1,
|
||||
},
|
||||
}
|
||||
})]
|
||||
}
|
||||
|
||||
pub fn build_kiro_stream_error_sse_events(error_type: &str, message: &str) -> Vec<Value> {
|
||||
vec![json!({
|
||||
"type": "error",
|
||||
"error": {
|
||||
"type": error_type,
|
||||
"message": message,
|
||||
}
|
||||
})]
|
||||
}
|
||||
|
||||
pub fn build_kiro_final_message_sse_events(
|
||||
stop_reason: &str,
|
||||
input_tokens: usize,
|
||||
output_tokens: usize,
|
||||
) -> Vec<Value> {
|
||||
vec![
|
||||
json!({
|
||||
"type": "message_delta",
|
||||
"delta": {
|
||||
"stop_reason": stop_reason,
|
||||
"stop_sequence": Value::Null,
|
||||
},
|
||||
"usage": {
|
||||
"input_tokens": input_tokens as u64,
|
||||
"output_tokens": output_tokens as u64,
|
||||
}
|
||||
}),
|
||||
json!({"type": "message_stop"}),
|
||||
]
|
||||
}
|
||||
|
||||
pub fn calculate_kiro_context_input_tokens(percentage: f64) -> usize {
|
||||
((percentage * KIRO_CONTEXT_WINDOW_TOKENS) / 100.0) as usize
|
||||
}
|
||||
|
||||
pub fn estimate_kiro_tokens(text: &str) -> usize {
|
||||
if text.is_empty() {
|
||||
return 0;
|
||||
}
|
||||
let mut chinese = 0usize;
|
||||
let mut other = 0usize;
|
||||
for ch in text.chars() {
|
||||
if ('\u{4e00}'..='\u{9fff}').contains(&ch) {
|
||||
chinese += 1;
|
||||
} else {
|
||||
other += 1;
|
||||
}
|
||||
}
|
||||
let chinese_tokens = (chinese * 2).div_ceil(3);
|
||||
let other_tokens = other.div_ceil(4);
|
||||
(chinese_tokens + other_tokens).max(1)
|
||||
}
|
||||
|
||||
pub fn find_kiro_real_thinking_start_tag(buffer: &str) -> Option<usize> {
|
||||
let tag = "<thinking>";
|
||||
let mut search = 0usize;
|
||||
loop {
|
||||
let pos = buffer[search..].find(tag).map(|value| value + search)?;
|
||||
let has_before = pos > 0 && is_kiro_quote_char(buffer, pos - 1);
|
||||
let after_pos = pos + tag.len();
|
||||
let has_after = is_kiro_quote_char(buffer, after_pos);
|
||||
if !has_before && !has_after {
|
||||
return Some(pos);
|
||||
}
|
||||
search = pos + 1;
|
||||
}
|
||||
}
|
||||
|
||||
pub fn find_kiro_real_thinking_end_tag(buffer: &str) -> Option<usize> {
|
||||
let tag = "</thinking>";
|
||||
let mut search = 0usize;
|
||||
loop {
|
||||
let pos = buffer[search..].find(tag).map(|value| value + search)?;
|
||||
let has_before = pos > 0 && is_kiro_quote_char(buffer, pos - 1);
|
||||
let after_pos = pos + tag.len();
|
||||
let has_after = is_kiro_quote_char(buffer, after_pos);
|
||||
if has_before || has_after {
|
||||
search = pos + 1;
|
||||
continue;
|
||||
}
|
||||
let after = &buffer[after_pos..];
|
||||
if after.len() < 2 {
|
||||
return None;
|
||||
}
|
||||
if after.starts_with("\n\n") {
|
||||
return Some(pos);
|
||||
}
|
||||
search = pos + 1;
|
||||
}
|
||||
}
|
||||
|
||||
pub fn find_kiro_real_thinking_end_tag_at_buffer_end(buffer: &str) -> Option<usize> {
|
||||
let tag = "</thinking>";
|
||||
let mut search = 0usize;
|
||||
loop {
|
||||
let pos = buffer[search..].find(tag).map(|value| value + search)?;
|
||||
let has_before = pos > 0 && is_kiro_quote_char(buffer, pos - 1);
|
||||
let after_pos = pos + tag.len();
|
||||
let has_after = is_kiro_quote_char(buffer, after_pos);
|
||||
if has_before || has_after {
|
||||
search = pos + 1;
|
||||
continue;
|
||||
}
|
||||
if buffer[after_pos..].trim().is_empty() {
|
||||
return Some(pos);
|
||||
}
|
||||
search = pos + 1;
|
||||
}
|
||||
}
|
||||
|
||||
pub fn kiro_crc32(data: &[u8]) -> u32 {
|
||||
let mut crc = 0xffff_ffffu32;
|
||||
for &byte in data {
|
||||
crc ^= byte as u32;
|
||||
for _ in 0..8 {
|
||||
let mask = if crc & 1 == 1 { 0xedb8_8320 } else { 0 };
|
||||
crc = (crc >> 1) ^ mask;
|
||||
}
|
||||
}
|
||||
!crc
|
||||
}
|
||||
|
||||
fn is_kiro_quote_char(buffer: &str, pos: usize) -> bool {
|
||||
buffer
|
||||
.as_bytes()
|
||||
.get(pos)
|
||||
.map(|byte| KIRO_QUOTE_CHARS.as_bytes().contains(byte))
|
||||
.unwrap_or(false)
|
||||
}
|
||||
@@ -0,0 +1,71 @@
|
||||
use std::collections::BTreeMap;
|
||||
|
||||
const MAX_MESSAGE_SIZE: usize = 16 * 1024 * 1024;
|
||||
const MAX_BUFFER_SIZE: usize = MAX_MESSAGE_SIZE;
|
||||
const MAX_ERRORS: usize = 5;
|
||||
|
||||
#[derive(Default)]
|
||||
pub struct KiroToClaudeCliStreamState {
|
||||
decoder: EventStreamDecoder,
|
||||
state: KiroClaudeStreamState,
|
||||
started: bool,
|
||||
}
|
||||
|
||||
#[derive(Default)]
|
||||
struct KiroClaudeStreamState {
|
||||
model: String,
|
||||
thinking_enabled: bool,
|
||||
estimated_input_tokens: usize,
|
||||
message_id: String,
|
||||
output_tokens: usize,
|
||||
context_input_tokens: Option<usize>,
|
||||
next_block_index: usize,
|
||||
open_blocks: BTreeMap<usize, String>,
|
||||
text_block_index: Option<usize>,
|
||||
thinking_block_index: Option<usize>,
|
||||
tool_block_indices: BTreeMap<String, usize>,
|
||||
thinking_buffer: String,
|
||||
in_thinking_block: bool,
|
||||
thinking_extracted: bool,
|
||||
strip_thinking_leading_newline: bool,
|
||||
has_tool_use: bool,
|
||||
stop_reason_override: Option<String>,
|
||||
had_error: bool,
|
||||
last_content: String,
|
||||
}
|
||||
|
||||
#[derive(Default)]
|
||||
struct EventStreamDecoder {
|
||||
buffer: Vec<u8>,
|
||||
error_count: usize,
|
||||
stopped: bool,
|
||||
}
|
||||
|
||||
#[derive(Default)]
|
||||
struct AwsHeaders {
|
||||
values: BTreeMap<String, AwsHeaderValue>,
|
||||
}
|
||||
|
||||
enum AwsHeaderValue {
|
||||
Ignored,
|
||||
String(String),
|
||||
}
|
||||
|
||||
struct AwsEventFrame {
|
||||
headers: AwsHeaders,
|
||||
payload: Vec<u8>,
|
||||
}
|
||||
|
||||
enum FrameParseError {
|
||||
Incomplete,
|
||||
Invalid(String),
|
||||
}
|
||||
|
||||
#[path = "stream/decoder.rs"]
|
||||
mod decoder;
|
||||
#[path = "stream/state.rs"]
|
||||
mod stream_state;
|
||||
|
||||
#[cfg(test)]
|
||||
#[path = "stream/tests.rs"]
|
||||
mod tests;
|
||||
@@ -0,0 +1,222 @@
|
||||
use std::collections::BTreeMap;
|
||||
|
||||
use super::{
|
||||
AwsEventFrame, AwsHeaderValue, AwsHeaders, EventStreamDecoder, FrameParseError,
|
||||
MAX_BUFFER_SIZE, MAX_ERRORS, MAX_MESSAGE_SIZE,
|
||||
};
|
||||
use crate::provider_compat::kiro_stream::kiro_crc32 as crc32;
|
||||
|
||||
impl EventStreamDecoder {
|
||||
pub(super) fn feed(&mut self, data: &[u8]) -> Result<(), String> {
|
||||
if self.stopped || data.is_empty() {
|
||||
return Ok(());
|
||||
}
|
||||
let new_size = self.buffer.len() + data.len();
|
||||
if new_size > MAX_BUFFER_SIZE {
|
||||
self.stopped = true;
|
||||
return Err(format!(
|
||||
"buffer overflow: size={new_size} max={MAX_BUFFER_SIZE}"
|
||||
));
|
||||
}
|
||||
self.buffer.extend_from_slice(data);
|
||||
Ok(())
|
||||
}
|
||||
|
||||
pub(super) fn decode_available(&mut self) -> Result<Vec<AwsEventFrame>, String> {
|
||||
let mut out = Vec::new();
|
||||
if self.stopped {
|
||||
return Ok(out);
|
||||
}
|
||||
|
||||
loop {
|
||||
match parse_frame(&self.buffer) {
|
||||
Ok(Some((frame, consumed))) => {
|
||||
if consumed == 0 {
|
||||
break;
|
||||
}
|
||||
out.push(frame);
|
||||
self.buffer.drain(..consumed);
|
||||
self.error_count = 0;
|
||||
}
|
||||
Ok(None) => break,
|
||||
Err(FrameParseError::Incomplete) => break,
|
||||
Err(FrameParseError::Invalid(message)) => {
|
||||
self.error_count += 1;
|
||||
if self.error_count >= MAX_ERRORS {
|
||||
self.stopped = true;
|
||||
return Err(message);
|
||||
}
|
||||
if self.buffer.is_empty() {
|
||||
break;
|
||||
}
|
||||
self.buffer.drain(..1);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
Ok(out)
|
||||
}
|
||||
}
|
||||
|
||||
impl AwsHeaders {
|
||||
fn get_string(&self, name: &str) -> Option<&str> {
|
||||
match self.values.get(name) {
|
||||
Some(AwsHeaderValue::String(value)) => Some(value.as_str()),
|
||||
_ => None,
|
||||
}
|
||||
}
|
||||
|
||||
pub(super) fn message_type(&self) -> Option<&str> {
|
||||
self.get_string(":message-type")
|
||||
}
|
||||
|
||||
pub(super) fn event_type(&self) -> Option<&str> {
|
||||
self.get_string(":event-type")
|
||||
}
|
||||
|
||||
pub(super) fn exception_type(&self) -> Option<&str> {
|
||||
self.get_string(":exception-type")
|
||||
}
|
||||
|
||||
pub(super) fn error_code(&self) -> Option<&str> {
|
||||
self.get_string(":error-code")
|
||||
}
|
||||
}
|
||||
|
||||
fn parse_frame(buffer: &[u8]) -> Result<Option<(AwsEventFrame, usize)>, FrameParseError> {
|
||||
if buffer.len() < 12 {
|
||||
return Ok(None);
|
||||
}
|
||||
let total_length = u32::from_be_bytes(buffer[0..4].try_into().expect("slice size")) as usize;
|
||||
let header_length = u32::from_be_bytes(buffer[4..8].try_into().expect("slice size")) as usize;
|
||||
let prelude_crc = u32::from_be_bytes(buffer[8..12].try_into().expect("slice size"));
|
||||
|
||||
if total_length < 16 {
|
||||
return Err(FrameParseError::Invalid(format!(
|
||||
"message too small: length={total_length}"
|
||||
)));
|
||||
}
|
||||
if total_length > MAX_MESSAGE_SIZE {
|
||||
return Err(FrameParseError::Invalid(format!(
|
||||
"message too large: length={total_length}"
|
||||
)));
|
||||
}
|
||||
if buffer.len() < total_length {
|
||||
return Ok(None);
|
||||
}
|
||||
if crc32(&buffer[0..8]) != prelude_crc {
|
||||
return Err(FrameParseError::Invalid("prelude crc mismatch".to_string()));
|
||||
}
|
||||
let message_crc = u32::from_be_bytes(
|
||||
buffer[total_length - 4..total_length]
|
||||
.try_into()
|
||||
.expect("slice size"),
|
||||
);
|
||||
if crc32(&buffer[..total_length - 4]) != message_crc {
|
||||
return Err(FrameParseError::Invalid("message crc mismatch".to_string()));
|
||||
}
|
||||
|
||||
let headers_start = 12;
|
||||
let headers_end = headers_start + header_length;
|
||||
if headers_end > total_length - 4 {
|
||||
return Err(FrameParseError::Invalid(
|
||||
"header length exceeds frame boundary".to_string(),
|
||||
));
|
||||
}
|
||||
let headers = parse_headers(&buffer[headers_start..headers_end], header_length)?;
|
||||
let payload = buffer[headers_end..total_length - 4].to_vec();
|
||||
Ok(Some((AwsEventFrame { headers, payload }, total_length)))
|
||||
}
|
||||
|
||||
fn parse_headers(data: &[u8], header_length: usize) -> Result<AwsHeaders, FrameParseError> {
|
||||
if data.len() < header_length {
|
||||
return Err(FrameParseError::Incomplete);
|
||||
}
|
||||
let mut values = BTreeMap::new();
|
||||
let mut offset = 0usize;
|
||||
while offset < header_length {
|
||||
ensure_header_bytes(data, offset, 1)?;
|
||||
let name_len = data[offset] as usize;
|
||||
offset += 1;
|
||||
if name_len == 0 {
|
||||
return Err(FrameParseError::Invalid(
|
||||
"header name length cannot be 0".to_string(),
|
||||
));
|
||||
}
|
||||
ensure_header_bytes(data, offset, name_len)?;
|
||||
let name = String::from_utf8_lossy(&data[offset..offset + name_len]).to_string();
|
||||
offset += name_len;
|
||||
|
||||
ensure_header_bytes(data, offset, 1)?;
|
||||
let value_type = data[offset];
|
||||
offset += 1;
|
||||
|
||||
let value = match value_type {
|
||||
0 => AwsHeaderValue::Ignored,
|
||||
1 => AwsHeaderValue::Ignored,
|
||||
2 => {
|
||||
ensure_header_bytes(data, offset, 1)?;
|
||||
let _ = i8::from_be_bytes([data[offset]]);
|
||||
offset += 1;
|
||||
AwsHeaderValue::Ignored
|
||||
}
|
||||
3 => {
|
||||
ensure_header_bytes(data, offset, 2)?;
|
||||
let _ = i16::from_be_bytes(data[offset..offset + 2].try_into().expect("slice"));
|
||||
offset += 2;
|
||||
AwsHeaderValue::Ignored
|
||||
}
|
||||
4 => {
|
||||
ensure_header_bytes(data, offset, 4)?;
|
||||
let _ = i32::from_be_bytes(data[offset..offset + 4].try_into().expect("slice"));
|
||||
offset += 4;
|
||||
AwsHeaderValue::Ignored
|
||||
}
|
||||
5 | 8 => {
|
||||
ensure_header_bytes(data, offset, 8)?;
|
||||
let _ = i64::from_be_bytes(data[offset..offset + 8].try_into().expect("slice"));
|
||||
offset += 8;
|
||||
AwsHeaderValue::Ignored
|
||||
}
|
||||
6 => {
|
||||
ensure_header_bytes(data, offset, 2)?;
|
||||
let length = u16::from_be_bytes(data[offset..offset + 2].try_into().expect("slice"))
|
||||
as usize;
|
||||
offset += 2;
|
||||
ensure_header_bytes(data, offset, length)?;
|
||||
offset += length;
|
||||
AwsHeaderValue::Ignored
|
||||
}
|
||||
7 => {
|
||||
ensure_header_bytes(data, offset, 2)?;
|
||||
let length = u16::from_be_bytes(data[offset..offset + 2].try_into().expect("slice"))
|
||||
as usize;
|
||||
offset += 2;
|
||||
ensure_header_bytes(data, offset, length)?;
|
||||
let out = String::from_utf8_lossy(&data[offset..offset + length]).to_string();
|
||||
offset += length;
|
||||
AwsHeaderValue::String(out)
|
||||
}
|
||||
9 => {
|
||||
ensure_header_bytes(data, offset, 16)?;
|
||||
offset += 16;
|
||||
AwsHeaderValue::Ignored
|
||||
}
|
||||
other => {
|
||||
return Err(FrameParseError::Invalid(format!(
|
||||
"invalid header type: {other}"
|
||||
)));
|
||||
}
|
||||
};
|
||||
values.insert(name, value);
|
||||
}
|
||||
Ok(AwsHeaders { values })
|
||||
}
|
||||
|
||||
fn ensure_header_bytes(data: &[u8], offset: usize, needed: usize) -> Result<(), FrameParseError> {
|
||||
let available = data.len().saturating_sub(offset);
|
||||
if available < needed {
|
||||
return Err(FrameParseError::Incomplete);
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
@@ -0,0 +1,8 @@
|
||||
#[path = "state/blocks.rs"]
|
||||
mod blocks;
|
||||
#[path = "state/events.rs"]
|
||||
mod events;
|
||||
#[path = "state/finalize.rs"]
|
||||
mod finalize;
|
||||
#[path = "state/lifecycle.rs"]
|
||||
mod lifecycle;
|
||||
@@ -0,0 +1,97 @@
|
||||
use serde_json::{json, Value};
|
||||
|
||||
use super::super::KiroClaudeStreamState;
|
||||
|
||||
impl KiroClaudeStreamState {
|
||||
pub(super) fn ensure_text_block_open(&mut self) -> Vec<Value> {
|
||||
if let Some(idx) = self.text_block_index {
|
||||
if self
|
||||
.open_blocks
|
||||
.get(&idx)
|
||||
.map(|value| value == "text")
|
||||
.unwrap_or(false)
|
||||
{
|
||||
return Vec::new();
|
||||
}
|
||||
}
|
||||
let idx = self.next_block_index;
|
||||
self.next_block_index += 1;
|
||||
self.text_block_index = Some(idx);
|
||||
self.open_blocks.insert(idx, "text".to_string());
|
||||
vec![json!({
|
||||
"type": "content_block_start",
|
||||
"index": idx,
|
||||
"content_block": {"type": "text", "text": ""}
|
||||
})]
|
||||
}
|
||||
|
||||
pub(super) fn ensure_thinking_block_open(&mut self) -> Vec<Value> {
|
||||
if let Some(idx) = self.thinking_block_index {
|
||||
if self
|
||||
.open_blocks
|
||||
.get(&idx)
|
||||
.map(|value| value == "thinking")
|
||||
.unwrap_or(false)
|
||||
{
|
||||
return Vec::new();
|
||||
}
|
||||
}
|
||||
let idx = self.next_block_index;
|
||||
self.next_block_index += 1;
|
||||
self.thinking_block_index = Some(idx);
|
||||
self.open_blocks.insert(idx, "thinking".to_string());
|
||||
vec![json!({
|
||||
"type": "content_block_start",
|
||||
"index": idx,
|
||||
"content_block": {"type": "thinking", "thinking": ""}
|
||||
})]
|
||||
}
|
||||
|
||||
pub(super) fn close_block(&mut self, idx: usize) -> Vec<Value> {
|
||||
if self.open_blocks.remove(&idx).is_none() {
|
||||
return Vec::new();
|
||||
}
|
||||
vec![json!({"type": "content_block_stop", "index": idx})]
|
||||
}
|
||||
|
||||
pub(super) fn emit_text_delta(&mut self, text: &str) -> Vec<Value> {
|
||||
if text.is_empty() {
|
||||
return Vec::new();
|
||||
}
|
||||
let mut events = self.ensure_text_block_open();
|
||||
let idx = self.text_block_index.unwrap_or_default();
|
||||
events.push(json!({
|
||||
"type": "content_block_delta",
|
||||
"index": idx,
|
||||
"delta": {"type": "text_delta", "text": text}
|
||||
}));
|
||||
events
|
||||
}
|
||||
|
||||
pub(super) fn emit_thinking_delta(&mut self, thinking: &str) -> Vec<Value> {
|
||||
if thinking.is_empty() {
|
||||
return Vec::new();
|
||||
}
|
||||
let mut events = self.ensure_thinking_block_open();
|
||||
let idx = self.thinking_block_index.unwrap_or_default();
|
||||
events.push(json!({
|
||||
"type": "content_block_delta",
|
||||
"index": idx,
|
||||
"delta": {"type": "thinking_delta", "thinking": thinking}
|
||||
}));
|
||||
events
|
||||
}
|
||||
|
||||
pub(super) fn close_thinking_block(&mut self) -> Vec<Value> {
|
||||
let Some(idx) = self.thinking_block_index else {
|
||||
return Vec::new();
|
||||
};
|
||||
let mut events = vec![json!({
|
||||
"type": "content_block_delta",
|
||||
"index": idx,
|
||||
"delta": {"type": "thinking_delta", "thinking": ""}
|
||||
})];
|
||||
events.extend(self.close_block(idx));
|
||||
events
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,331 @@
|
||||
use serde_json::{json, Value};
|
||||
|
||||
use crate::provider_compat::kiro_stream::{
|
||||
calculate_kiro_context_input_tokens, encode_kiro_sse_events, estimate_kiro_tokens,
|
||||
find_kiro_real_thinking_end_tag, find_kiro_real_thinking_end_tag_at_buffer_end,
|
||||
find_kiro_real_thinking_start_tag, KIRO_MAX_THINKING_BUFFER,
|
||||
};
|
||||
|
||||
use crate::response::AiSurfaceFinalizeError;
|
||||
|
||||
use super::super::AwsEventFrame;
|
||||
use super::super::KiroClaudeStreamState;
|
||||
|
||||
fn floor_char_boundary(text: &str, index: usize) -> usize {
|
||||
let mut boundary = index.min(text.len());
|
||||
while boundary > 0 && !text.is_char_boundary(boundary) {
|
||||
boundary -= 1;
|
||||
}
|
||||
boundary
|
||||
}
|
||||
|
||||
fn split_preserving_trailing_bytes(
|
||||
buffer: &str,
|
||||
trailing_bytes: usize,
|
||||
) -> Option<(String, String)> {
|
||||
if buffer.len() <= trailing_bytes {
|
||||
return None;
|
||||
}
|
||||
|
||||
let split = floor_char_boundary(buffer, buffer.len() - trailing_bytes);
|
||||
if split == 0 {
|
||||
return None;
|
||||
}
|
||||
|
||||
Some((buffer[..split].to_string(), buffer[split..].to_string()))
|
||||
}
|
||||
|
||||
impl KiroClaudeStreamState {
|
||||
pub(super) fn process_frame(
|
||||
&mut self,
|
||||
frame: AwsEventFrame,
|
||||
) -> Result<Vec<u8>, AiSurfaceFinalizeError> {
|
||||
let message_type = frame.headers.message_type().unwrap_or("event");
|
||||
match message_type {
|
||||
"event" => self.process_event_frame(frame),
|
||||
"exception" => self.process_exception_frame(frame),
|
||||
"error" => self.process_error_frame(frame),
|
||||
_ => Ok(Vec::new()),
|
||||
}
|
||||
}
|
||||
|
||||
pub(super) fn process_event_frame(
|
||||
&mut self,
|
||||
frame: AwsEventFrame,
|
||||
) -> Result<Vec<u8>, AiSurfaceFinalizeError> {
|
||||
let event_type = frame.headers.event_type().unwrap_or_default();
|
||||
let payload: Value = if frame.payload.is_empty() {
|
||||
json!({})
|
||||
} else {
|
||||
serde_json::from_slice(&frame.payload).unwrap_or_else(|_| json!({}))
|
||||
};
|
||||
let payload_object = payload.as_object();
|
||||
let mut events = Vec::new();
|
||||
match event_type {
|
||||
"assistantResponseEvent" => {
|
||||
if let Some(content) = payload_object
|
||||
.and_then(|value| value.get("content"))
|
||||
.and_then(Value::as_str)
|
||||
{
|
||||
events.extend(self.process_assistant_response(content));
|
||||
}
|
||||
}
|
||||
"toolUseEvent" => {
|
||||
if let Some(payload_object) = payload_object {
|
||||
let name = payload_object
|
||||
.get("name")
|
||||
.and_then(Value::as_str)
|
||||
.unwrap_or_default();
|
||||
let tool_use_id = payload_object
|
||||
.get("toolUseId")
|
||||
.or_else(|| payload_object.get("tool_use_id"))
|
||||
.and_then(Value::as_str)
|
||||
.unwrap_or_default();
|
||||
let input_json = match payload_object.get("input") {
|
||||
None | Some(Value::Null) => String::new(),
|
||||
Some(Value::String(text)) => text.clone(),
|
||||
Some(other) => {
|
||||
serde_json::to_string(other).map_err(AiSurfaceFinalizeError::from)?
|
||||
}
|
||||
};
|
||||
let stop = payload_object
|
||||
.get("stop")
|
||||
.and_then(Value::as_bool)
|
||||
.unwrap_or(false);
|
||||
events.extend(self.process_tool_use(name, tool_use_id, &input_json, stop));
|
||||
}
|
||||
}
|
||||
"contextUsageEvent" => {
|
||||
if let Some(percentage) = payload_object
|
||||
.and_then(|value| value.get("contextUsagePercentage"))
|
||||
.and_then(Value::as_f64)
|
||||
{
|
||||
self.context_input_tokens =
|
||||
Some(calculate_kiro_context_input_tokens(percentage));
|
||||
}
|
||||
}
|
||||
_ => {}
|
||||
}
|
||||
encode_kiro_sse_events(events).map_err(AiSurfaceFinalizeError::from)
|
||||
}
|
||||
|
||||
pub(super) fn process_exception_frame(
|
||||
&mut self,
|
||||
frame: AwsEventFrame,
|
||||
) -> Result<Vec<u8>, AiSurfaceFinalizeError> {
|
||||
let exception_type = frame
|
||||
.headers
|
||||
.exception_type()
|
||||
.unwrap_or("UnknownException")
|
||||
.to_string();
|
||||
if exception_type == "ContentLengthExceededException" {
|
||||
self.stop_reason_override = Some("max_tokens".to_string());
|
||||
return Ok(Vec::new());
|
||||
}
|
||||
self.emit_stream_error("upstream_exception", &exception_type)
|
||||
}
|
||||
|
||||
pub(super) fn process_error_frame(
|
||||
&mut self,
|
||||
frame: AwsEventFrame,
|
||||
) -> Result<Vec<u8>, AiSurfaceFinalizeError> {
|
||||
let error_code = frame
|
||||
.headers
|
||||
.error_code()
|
||||
.unwrap_or("UnknownError")
|
||||
.to_string();
|
||||
self.emit_stream_error("upstream_error", &error_code)
|
||||
}
|
||||
|
||||
pub(super) fn process_assistant_response(&mut self, content: &str) -> Vec<Value> {
|
||||
if content.is_empty() || content == self.last_content {
|
||||
return Vec::new();
|
||||
}
|
||||
self.last_content = content.to_string();
|
||||
self.output_tokens += estimate_kiro_tokens(content);
|
||||
|
||||
if !self.thinking_enabled {
|
||||
return self.emit_text_delta(content);
|
||||
}
|
||||
|
||||
self.thinking_buffer.push_str(content);
|
||||
if self.thinking_buffer.len() > KIRO_MAX_THINKING_BUFFER {
|
||||
let overflow = std::mem::take(&mut self.thinking_buffer);
|
||||
if self.in_thinking_block {
|
||||
let mut events = self.emit_thinking_delta(&overflow);
|
||||
events.extend(self.close_thinking_block());
|
||||
self.in_thinking_block = false;
|
||||
self.thinking_extracted = true;
|
||||
return events;
|
||||
}
|
||||
return self.emit_text_delta(&overflow);
|
||||
}
|
||||
|
||||
let mut events = Vec::new();
|
||||
loop {
|
||||
if !self.in_thinking_block && !self.thinking_extracted {
|
||||
if let Some(start_pos) = find_kiro_real_thinking_start_tag(&self.thinking_buffer) {
|
||||
let before = self.thinking_buffer[..start_pos].to_string();
|
||||
if !before.trim().is_empty() {
|
||||
events.extend(self.emit_text_delta(&before));
|
||||
}
|
||||
self.in_thinking_block = true;
|
||||
self.strip_thinking_leading_newline = true;
|
||||
self.thinking_buffer =
|
||||
self.thinking_buffer[start_pos + "<thinking>".len()..].to_string();
|
||||
events.extend(self.ensure_thinking_block_open());
|
||||
continue;
|
||||
}
|
||||
|
||||
let keep = "<thinking>".len();
|
||||
if let Some((safe, remaining)) =
|
||||
split_preserving_trailing_bytes(&self.thinking_buffer, keep)
|
||||
{
|
||||
if !safe.trim().is_empty() {
|
||||
events.extend(self.emit_text_delta(&safe));
|
||||
self.thinking_buffer = remaining;
|
||||
}
|
||||
}
|
||||
break;
|
||||
}
|
||||
|
||||
if self.in_thinking_block {
|
||||
if self.strip_thinking_leading_newline {
|
||||
if self.thinking_buffer.starts_with('\n') {
|
||||
self.thinking_buffer.remove(0);
|
||||
self.strip_thinking_leading_newline = false;
|
||||
} else if !self.thinking_buffer.is_empty() {
|
||||
self.strip_thinking_leading_newline = false;
|
||||
}
|
||||
}
|
||||
|
||||
if let Some(end_pos) = find_kiro_real_thinking_end_tag(&self.thinking_buffer) {
|
||||
let thinking_text = self.thinking_buffer[..end_pos].to_string();
|
||||
if !thinking_text.is_empty() {
|
||||
events.extend(self.emit_thinking_delta(&thinking_text));
|
||||
}
|
||||
events.extend(self.close_thinking_block());
|
||||
self.in_thinking_block = false;
|
||||
self.thinking_extracted = true;
|
||||
self.thinking_buffer =
|
||||
self.thinking_buffer[end_pos + "</thinking>".len()..].to_string();
|
||||
continue;
|
||||
}
|
||||
|
||||
let keep = "</thinking>".len();
|
||||
if let Some((safe, remaining)) =
|
||||
split_preserving_trailing_bytes(&self.thinking_buffer, keep)
|
||||
{
|
||||
if !safe.is_empty() {
|
||||
events.extend(self.emit_thinking_delta(&safe));
|
||||
self.thinking_buffer = remaining;
|
||||
}
|
||||
}
|
||||
break;
|
||||
}
|
||||
|
||||
if !self.thinking_buffer.is_empty() {
|
||||
let remaining = std::mem::take(&mut self.thinking_buffer);
|
||||
events.extend(self.emit_text_delta(&remaining));
|
||||
}
|
||||
break;
|
||||
}
|
||||
|
||||
events
|
||||
}
|
||||
|
||||
pub(super) fn process_tool_use(
|
||||
&mut self,
|
||||
name: &str,
|
||||
tool_use_id: &str,
|
||||
input_json: &str,
|
||||
stop: bool,
|
||||
) -> Vec<Value> {
|
||||
if tool_use_id.is_empty() {
|
||||
return Vec::new();
|
||||
}
|
||||
|
||||
self.has_tool_use = true;
|
||||
let mut events = Vec::new();
|
||||
|
||||
if self.thinking_enabled && self.in_thinking_block && !self.thinking_buffer.is_empty() {
|
||||
if let Some(end_pos) =
|
||||
find_kiro_real_thinking_end_tag_at_buffer_end(&self.thinking_buffer)
|
||||
{
|
||||
let thinking_text = self.thinking_buffer[..end_pos].to_string();
|
||||
if !thinking_text.is_empty() {
|
||||
events.extend(self.emit_thinking_delta(&thinking_text));
|
||||
}
|
||||
events.extend(self.close_thinking_block());
|
||||
let remaining = self.thinking_buffer[end_pos + "</thinking>".len()..].to_string();
|
||||
self.thinking_buffer.clear();
|
||||
self.in_thinking_block = false;
|
||||
self.thinking_extracted = true;
|
||||
if !remaining.is_empty() {
|
||||
events.extend(self.emit_text_delta(&remaining));
|
||||
}
|
||||
} else {
|
||||
let thinking = std::mem::take(&mut self.thinking_buffer);
|
||||
events.extend(self.emit_thinking_delta(&thinking));
|
||||
events.extend(self.close_thinking_block());
|
||||
self.in_thinking_block = false;
|
||||
self.thinking_extracted = true;
|
||||
}
|
||||
}
|
||||
|
||||
if self.thinking_enabled
|
||||
&& !self.in_thinking_block
|
||||
&& !self.thinking_extracted
|
||||
&& !self.thinking_buffer.is_empty()
|
||||
{
|
||||
let buffered = std::mem::take(&mut self.thinking_buffer);
|
||||
events.extend(self.emit_text_delta(&buffered));
|
||||
}
|
||||
|
||||
if let Some(idx) = self.text_block_index.take() {
|
||||
events.extend(self.close_block(idx));
|
||||
}
|
||||
|
||||
let block_index = if let Some(block_index) = self.tool_block_indices.get(tool_use_id) {
|
||||
*block_index
|
||||
} else {
|
||||
let block_index = self.next_block_index;
|
||||
self.next_block_index += 1;
|
||||
self.tool_block_indices
|
||||
.insert(tool_use_id.to_string(), block_index);
|
||||
block_index
|
||||
};
|
||||
|
||||
if let std::collections::btree_map::Entry::Vacant(e) = self.open_blocks.entry(block_index) {
|
||||
e.insert("tool_use".to_string());
|
||||
events.push(json!({
|
||||
"type": "content_block_start",
|
||||
"index": block_index,
|
||||
"content_block": {
|
||||
"type": "tool_use",
|
||||
"id": tool_use_id,
|
||||
"name": name,
|
||||
"input": {},
|
||||
}
|
||||
}));
|
||||
}
|
||||
|
||||
if !input_json.is_empty() {
|
||||
self.output_tokens += estimate_kiro_tokens(input_json);
|
||||
events.push(json!({
|
||||
"type": "content_block_delta",
|
||||
"index": block_index,
|
||||
"delta": {
|
||||
"type": "input_json_delta",
|
||||
"partial_json": input_json,
|
||||
}
|
||||
}));
|
||||
}
|
||||
|
||||
if stop {
|
||||
events.extend(self.close_block(block_index));
|
||||
}
|
||||
|
||||
events
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,96 @@
|
||||
use crate::provider_compat::kiro_stream::{
|
||||
build_kiro_final_message_sse_events, encode_kiro_sse_events,
|
||||
find_kiro_real_thinking_end_tag_at_buffer_end,
|
||||
};
|
||||
|
||||
use crate::response::AiSurfaceFinalizeError;
|
||||
|
||||
use super::super::KiroClaudeStreamState;
|
||||
|
||||
impl KiroClaudeStreamState {
|
||||
pub(super) fn finalize(&mut self) -> Result<Vec<u8>, AiSurfaceFinalizeError> {
|
||||
if self.thinking_enabled && !self.thinking_buffer.is_empty() {
|
||||
let flush_events = if self.in_thinking_block {
|
||||
if let Some(end_pos) =
|
||||
find_kiro_real_thinking_end_tag_at_buffer_end(&self.thinking_buffer)
|
||||
{
|
||||
let thinking_text = self.thinking_buffer[..end_pos].to_string();
|
||||
let mut events = Vec::new();
|
||||
if !thinking_text.is_empty() {
|
||||
events.extend(self.emit_thinking_delta(&thinking_text));
|
||||
}
|
||||
events.extend(self.close_thinking_block());
|
||||
let remaining =
|
||||
self.thinking_buffer[end_pos + "</thinking>".len()..].to_string();
|
||||
if !remaining.is_empty() {
|
||||
events.extend(self.emit_text_delta(&remaining));
|
||||
}
|
||||
events
|
||||
} else {
|
||||
let mut events = self.emit_thinking_delta(&self.thinking_buffer.clone());
|
||||
events.extend(self.close_thinking_block());
|
||||
events
|
||||
}
|
||||
} else {
|
||||
self.emit_text_delta(&self.thinking_buffer.clone())
|
||||
};
|
||||
self.thinking_buffer.clear();
|
||||
self.in_thinking_block = false;
|
||||
self.thinking_extracted = true;
|
||||
let mut output =
|
||||
encode_kiro_sse_events(flush_events).map_err(AiSurfaceFinalizeError::from)?;
|
||||
for idx in self
|
||||
.open_blocks
|
||||
.keys()
|
||||
.cloned()
|
||||
.collect::<Vec<_>>()
|
||||
.into_iter()
|
||||
.rev()
|
||||
{
|
||||
output.extend(
|
||||
encode_kiro_sse_events(self.close_block(idx))
|
||||
.map_err(AiSurfaceFinalizeError::from)?,
|
||||
);
|
||||
}
|
||||
output.extend(self.final_message_bytes()?);
|
||||
return Ok(output);
|
||||
}
|
||||
|
||||
let mut output = Vec::new();
|
||||
for idx in self
|
||||
.open_blocks
|
||||
.keys()
|
||||
.cloned()
|
||||
.collect::<Vec<_>>()
|
||||
.into_iter()
|
||||
.rev()
|
||||
{
|
||||
output.extend(
|
||||
encode_kiro_sse_events(self.close_block(idx))
|
||||
.map_err(AiSurfaceFinalizeError::from)?,
|
||||
);
|
||||
}
|
||||
output.extend(self.final_message_bytes()?);
|
||||
Ok(output)
|
||||
}
|
||||
|
||||
pub(super) fn final_message_bytes(&self) -> Result<Vec<u8>, AiSurfaceFinalizeError> {
|
||||
let stop_reason = self.stop_reason_override.clone().unwrap_or_else(|| {
|
||||
if self.has_tool_use {
|
||||
"tool_use"
|
||||
} else {
|
||||
"end_turn"
|
||||
}
|
||||
.to_string()
|
||||
});
|
||||
let input_tokens = self
|
||||
.context_input_tokens
|
||||
.unwrap_or(self.estimated_input_tokens) as u64;
|
||||
encode_kiro_sse_events(build_kiro_final_message_sse_events(
|
||||
&stop_reason,
|
||||
input_tokens as usize,
|
||||
self.output_tokens,
|
||||
))
|
||||
.map_err(AiSurfaceFinalizeError::from)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,129 @@
|
||||
use serde_json::Value;
|
||||
use uuid::Uuid;
|
||||
|
||||
use crate::provider_compat::kiro_stream::{
|
||||
build_kiro_initial_sse_events, build_kiro_stream_error_sse_events, encode_kiro_sse_events,
|
||||
};
|
||||
use crate::response::AiSurfaceFinalizeError;
|
||||
|
||||
use super::super::{EventStreamDecoder, KiroClaudeStreamState, KiroToClaudeCliStreamState};
|
||||
|
||||
impl KiroToClaudeCliStreamState {
|
||||
pub fn new(report_context: &Value) -> Self {
|
||||
Self {
|
||||
decoder: EventStreamDecoder::default(),
|
||||
state: KiroClaudeStreamState::new(report_context),
|
||||
started: false,
|
||||
}
|
||||
}
|
||||
|
||||
pub fn push_chunk(
|
||||
&mut self,
|
||||
_report_context: &Value,
|
||||
chunk: &[u8],
|
||||
) -> Result<Vec<u8>, AiSurfaceFinalizeError> {
|
||||
let mut output = Vec::new();
|
||||
if !self.started {
|
||||
self.started = true;
|
||||
output.extend(self.state.generate_initial_bytes()?);
|
||||
}
|
||||
|
||||
if let Err(err) = self.decoder.feed(chunk) {
|
||||
output.extend(
|
||||
self.state
|
||||
.emit_stream_error("upstream_stream_error", &err)?,
|
||||
);
|
||||
return Ok(output);
|
||||
}
|
||||
|
||||
match self.decoder.decode_available() {
|
||||
Ok(frames) => {
|
||||
for frame in frames {
|
||||
output.extend(self.state.process_frame(frame)?);
|
||||
}
|
||||
}
|
||||
Err(err) => {
|
||||
output.extend(
|
||||
self.state
|
||||
.emit_stream_error("upstream_stream_error", &err)?,
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
Ok(output)
|
||||
}
|
||||
|
||||
pub fn finish(&mut self, _report_context: &Value) -> Result<Vec<u8>, AiSurfaceFinalizeError> {
|
||||
if !self.started || self.state.had_error {
|
||||
return Ok(Vec::new());
|
||||
}
|
||||
self.state.finalize()
|
||||
}
|
||||
}
|
||||
|
||||
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())
|
||||
.or_else(|| {
|
||||
report_context
|
||||
.get("model")
|
||||
.and_then(Value::as_str)
|
||||
.filter(|value| !value.is_empty())
|
||||
})
|
||||
.unwrap_or("unknown")
|
||||
.to_string();
|
||||
let thinking_enabled = report_context
|
||||
.get("original_request_body")
|
||||
.and_then(Value::as_object)
|
||||
.and_then(|body| body.get("thinking"))
|
||||
.and_then(Value::as_object)
|
||||
.and_then(|thinking| thinking.get("type"))
|
||||
.and_then(Value::as_str)
|
||||
.map(|value| {
|
||||
value.trim().eq_ignore_ascii_case("enabled")
|
||||
|| value.trim().eq_ignore_ascii_case("adaptive")
|
||||
})
|
||||
.unwrap_or(false);
|
||||
let estimated_input_tokens = report_context
|
||||
.get("input_tokens")
|
||||
.and_then(Value::as_u64)
|
||||
.map(|value| value as usize)
|
||||
.unwrap_or(0);
|
||||
Self {
|
||||
model,
|
||||
thinking_enabled,
|
||||
estimated_input_tokens,
|
||||
message_id: format!("msg_{}", Uuid::new_v4().simple()),
|
||||
..Self::default()
|
||||
}
|
||||
}
|
||||
|
||||
pub(super) fn generate_initial_bytes(&mut self) -> Result<Vec<u8>, AiSurfaceFinalizeError> {
|
||||
let events = build_kiro_initial_sse_events(
|
||||
&self.message_id,
|
||||
&self.model,
|
||||
self.estimated_input_tokens,
|
||||
);
|
||||
let mut events = events;
|
||||
if !self.thinking_enabled {
|
||||
events.extend(self.ensure_text_block_open());
|
||||
}
|
||||
encode_kiro_sse_events(events).map_err(AiSurfaceFinalizeError::from)
|
||||
}
|
||||
|
||||
pub(super) fn emit_stream_error(
|
||||
&mut self,
|
||||
error_type: &str,
|
||||
message: &str,
|
||||
) -> Result<Vec<u8>, AiSurfaceFinalizeError> {
|
||||
if self.had_error {
|
||||
return Ok(Vec::new());
|
||||
}
|
||||
self.had_error = true;
|
||||
encode_kiro_sse_events(build_kiro_stream_error_sse_events(error_type, message))
|
||||
.map_err(AiSurfaceFinalizeError::from)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,166 @@
|
||||
use crate::provider_compat::kiro_stream::kiro_crc32 as crc32;
|
||||
use serde_json::{json, Value};
|
||||
|
||||
use super::KiroToClaudeCliStreamState;
|
||||
|
||||
fn encode_string_header(name: &str, value: &str) -> Vec<u8> {
|
||||
let mut out = Vec::new();
|
||||
out.push(name.len() as u8);
|
||||
out.extend_from_slice(name.as_bytes());
|
||||
out.push(7);
|
||||
out.extend_from_slice(&(value.len() as u16).to_be_bytes());
|
||||
out.extend_from_slice(value.as_bytes());
|
||||
out
|
||||
}
|
||||
|
||||
fn encode_event_frame(message_type: &str, event_type: Option<&str>, payload: &Value) -> Vec<u8> {
|
||||
let mut headers = encode_string_header(":message-type", message_type);
|
||||
if let Some(event_type) = event_type {
|
||||
headers.extend_from_slice(&encode_string_header(":event-type", event_type));
|
||||
}
|
||||
let payload_bytes = serde_json::to_vec(payload).expect("payload should encode");
|
||||
encode_frame(headers, payload_bytes)
|
||||
}
|
||||
|
||||
fn encode_frame(headers: Vec<u8>, payload: Vec<u8>) -> Vec<u8> {
|
||||
let total_len = 12 + headers.len() + payload.len() + 4;
|
||||
let header_len = headers.len();
|
||||
let mut out = Vec::with_capacity(total_len);
|
||||
out.extend_from_slice(&(total_len as u32).to_be_bytes());
|
||||
out.extend_from_slice(&(header_len as u32).to_be_bytes());
|
||||
let prelude_crc = crc32(&out[..8]);
|
||||
out.extend_from_slice(&prelude_crc.to_be_bytes());
|
||||
out.extend_from_slice(&headers);
|
||||
out.extend_from_slice(&payload);
|
||||
let message_crc = crc32(&out);
|
||||
out.extend_from_slice(&message_crc.to_be_bytes());
|
||||
out
|
||||
}
|
||||
|
||||
fn kiro_report_context(thinking_enabled: bool) -> Value {
|
||||
let mut context = json!({
|
||||
"provider_api_format": "claude:messages",
|
||||
"client_api_format": "claude:messages",
|
||||
"envelope_name": "kiro:generateAssistantResponse",
|
||||
"mapped_model": "claude-sonnet-4.5"
|
||||
});
|
||||
if thinking_enabled {
|
||||
context["original_request_body"] = json!({
|
||||
"thinking": {
|
||||
"type": "enabled"
|
||||
}
|
||||
});
|
||||
}
|
||||
context
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn kiro_stream_rewriter_converts_text_events_to_claude_sse() {
|
||||
let report_context = kiro_report_context(false);
|
||||
let mut rewriter = KiroToClaudeCliStreamState::new(&report_context);
|
||||
let chunk = [
|
||||
encode_event_frame(
|
||||
"event",
|
||||
Some("assistantResponseEvent"),
|
||||
&json!({"content": "Hello from Kiro"}),
|
||||
),
|
||||
encode_event_frame(
|
||||
"event",
|
||||
Some("contextUsageEvent"),
|
||||
&json!({"contextUsagePercentage": 1.0}),
|
||||
),
|
||||
]
|
||||
.concat();
|
||||
|
||||
let first = rewriter
|
||||
.push_chunk(&report_context, &chunk)
|
||||
.expect("rewrite should succeed");
|
||||
let rest = rewriter
|
||||
.finish(&report_context)
|
||||
.expect("finish should succeed");
|
||||
let text = String::from_utf8([first, rest].concat()).expect("utf8 should decode");
|
||||
assert!(text.contains("event: message_start"));
|
||||
assert!(text.contains("\"type\":\"content_block_delta\""));
|
||||
assert!(text.contains("Hello from Kiro"));
|
||||
assert!(text.contains("\"stop_reason\":\"end_turn\""));
|
||||
assert!(text.contains("\"input_tokens\":2000"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn kiro_stream_rewriter_converts_tool_use_to_claude_events() {
|
||||
let report_context = kiro_report_context(false);
|
||||
let mut rewriter = KiroToClaudeCliStreamState::new(&report_context);
|
||||
let chunk = [
|
||||
encode_event_frame(
|
||||
"event",
|
||||
Some("assistantResponseEvent"),
|
||||
&json!({"content": "Need a tool."}),
|
||||
),
|
||||
encode_event_frame(
|
||||
"event",
|
||||
Some("toolUseEvent"),
|
||||
&json!({
|
||||
"name": "get_weather",
|
||||
"toolUseId": "tool_123",
|
||||
"input": {"city": "SF"},
|
||||
"stop": true
|
||||
}),
|
||||
),
|
||||
]
|
||||
.concat();
|
||||
|
||||
let first = rewriter
|
||||
.push_chunk(&report_context, &chunk)
|
||||
.expect("rewrite should succeed");
|
||||
let rest = rewriter
|
||||
.finish(&report_context)
|
||||
.expect("finish should succeed");
|
||||
let text = String::from_utf8([first, rest].concat()).expect("utf8 should decode");
|
||||
assert!(text.contains("\"type\":\"tool_use\""));
|
||||
assert!(text.contains("\"id\":\"tool_123\""));
|
||||
assert!(text.contains("\"name\":\"get_weather\""));
|
||||
assert!(text.contains("\"partial_json\":\"{\\\"city\\\":\\\"SF\\\"}\""));
|
||||
assert!(text.contains("\"stop_reason\":\"tool_use\""));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn kiro_stream_rewriter_handles_multibyte_text_without_thinking_tag() {
|
||||
let report_context = kiro_report_context(true);
|
||||
let mut rewriter = KiroToClaudeCliStreamState::new(&report_context);
|
||||
let chunk = encode_event_frame(
|
||||
"event",
|
||||
Some("assistantResponseEvent"),
|
||||
&json!({"content": "\n\n你好!有"}),
|
||||
);
|
||||
|
||||
let first = rewriter
|
||||
.push_chunk(&report_context, &chunk)
|
||||
.expect("rewrite should succeed");
|
||||
let rest = rewriter
|
||||
.finish(&report_context)
|
||||
.expect("finish should succeed");
|
||||
let text = String::from_utf8([first, rest].concat()).expect("utf8 should decode");
|
||||
assert!(text.contains("\"type\":\"text_delta\""));
|
||||
assert!(text.contains("你好!有"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn kiro_stream_rewriter_handles_multibyte_text_inside_thinking_block() {
|
||||
let report_context = kiro_report_context(true);
|
||||
let mut rewriter = KiroToClaudeCliStreamState::new(&report_context);
|
||||
let chunk = encode_event_frame(
|
||||
"event",
|
||||
Some("assistantResponseEvent"),
|
||||
&json!({"content": "<thinking>\n\n你好!有"}),
|
||||
);
|
||||
|
||||
let first = rewriter
|
||||
.push_chunk(&report_context, &chunk)
|
||||
.expect("rewrite should succeed");
|
||||
let rest = rewriter
|
||||
.finish(&report_context)
|
||||
.expect("finish should succeed");
|
||||
let text = String::from_utf8([first, rest].concat()).expect("utf8 should decode");
|
||||
assert!(text.contains("\"type\":\"thinking_delta\""));
|
||||
assert!(text.contains("你好!有"));
|
||||
}
|
||||
4
crates/aether-ai-formats/src/provider_compat/mod.rs
Normal file
4
crates/aether-ai-formats/src/provider_compat/mod.rs
Normal file
@@ -0,0 +1,4 @@
|
||||
pub mod kiro_stream;
|
||||
pub mod private_envelope;
|
||||
pub mod proxy;
|
||||
pub mod surfaces;
|
||||
529
crates/aether-ai-formats/src/provider_compat/private_envelope.rs
Normal file
529
crates/aether-ai-formats/src/provider_compat/private_envelope.rs
Normal file
@@ -0,0 +1,529 @@
|
||||
use std::collections::BTreeMap;
|
||||
|
||||
use serde_json::Value;
|
||||
|
||||
use crate::provider_compat::kiro_stream::KiroToClaudeCliStreamState;
|
||||
use crate::response::AiSurfaceFinalizeError;
|
||||
|
||||
use super::surfaces::{
|
||||
provider_adaptation_allows_sync_finalize_envelope, provider_adaptation_descriptor_for_envelope,
|
||||
provider_adaptation_should_unwrap_stream_envelope, ANTIGRAVITY_V1INTERNAL_ENVELOPE_NAME,
|
||||
GEMINI_CLI_V1INTERNAL_ENVELOPE_NAME, KIRO_ENVELOPE_NAME,
|
||||
};
|
||||
|
||||
pub fn provider_private_response_allows_sync_finalize(report_context: &Value) -> bool {
|
||||
let has_envelope = report_context
|
||||
.get("has_envelope")
|
||||
.and_then(Value::as_bool)
|
||||
.unwrap_or(false);
|
||||
if !has_envelope {
|
||||
return true;
|
||||
}
|
||||
let envelope_name = report_context
|
||||
.get("envelope_name")
|
||||
.and_then(Value::as_str)
|
||||
.unwrap_or_default();
|
||||
let provider_api_format = report_context
|
||||
.get("provider_api_format")
|
||||
.and_then(Value::as_str)
|
||||
.unwrap_or_default();
|
||||
provider_adaptation_allows_sync_finalize_envelope(envelope_name, provider_api_format)
|
||||
}
|
||||
|
||||
pub fn normalize_provider_private_report_context(report_context: Option<&Value>) -> Option<Value> {
|
||||
let report_context = report_context?;
|
||||
if !report_context
|
||||
.get("has_envelope")
|
||||
.and_then(Value::as_bool)
|
||||
.unwrap_or(false)
|
||||
{
|
||||
return Some(report_context.clone());
|
||||
}
|
||||
let envelope_name = report_context
|
||||
.get("envelope_name")
|
||||
.and_then(Value::as_str)
|
||||
.unwrap_or_default();
|
||||
let provider_api_format = report_context
|
||||
.get("provider_api_format")
|
||||
.and_then(Value::as_str)
|
||||
.unwrap_or_default();
|
||||
if provider_adaptation_descriptor_for_envelope(envelope_name, provider_api_format).is_none() {
|
||||
return Some(report_context.clone());
|
||||
}
|
||||
Some(clear_private_envelope_context(report_context))
|
||||
}
|
||||
|
||||
pub fn normalize_provider_private_response_value(
|
||||
data: Value,
|
||||
report_context: &Value,
|
||||
) -> Option<Value> {
|
||||
if !report_context
|
||||
.get("has_envelope")
|
||||
.and_then(Value::as_bool)
|
||||
.unwrap_or(false)
|
||||
{
|
||||
return Some(data);
|
||||
}
|
||||
let mut unwrapped = match report_context.get("envelope_name").and_then(Value::as_str) {
|
||||
Some(KIRO_ENVELOPE_NAME) => data,
|
||||
Some(GEMINI_CLI_V1INTERNAL_ENVELOPE_NAME) => {
|
||||
if let Some(response) = data
|
||||
.get("response")
|
||||
.and_then(Value::as_object)
|
||||
.filter(|response| !response.contains_key("response"))
|
||||
{
|
||||
Value::Object(response.clone())
|
||||
} else {
|
||||
data
|
||||
}
|
||||
}
|
||||
Some(ANTIGRAVITY_V1INTERNAL_ENVELOPE_NAME) => {
|
||||
if let Some(response) = data
|
||||
.get("response")
|
||||
.and_then(Value::as_object)
|
||||
.filter(|response| !response.contains_key("response"))
|
||||
{
|
||||
let mut unwrapped = response.clone();
|
||||
if let Some(response_id) = data.get("responseId").cloned() {
|
||||
unwrapped.insert("_v1internal_response_id".to_string(), response_id);
|
||||
}
|
||||
Value::Object(unwrapped)
|
||||
} else {
|
||||
data
|
||||
}
|
||||
}
|
||||
_ => return None,
|
||||
};
|
||||
postprocess_private_response_value(&mut unwrapped, report_context);
|
||||
Some(unwrapped)
|
||||
}
|
||||
|
||||
pub fn transform_provider_private_stream_line(
|
||||
report_context: &Value,
|
||||
line: Vec<u8>,
|
||||
) -> Result<Vec<u8>, serde_json::Error> {
|
||||
let Ok(text) = std::str::from_utf8(&line) else {
|
||||
return Ok(line);
|
||||
};
|
||||
let trimmed = text.trim_matches('\r').trim();
|
||||
if trimmed.is_empty() || trimmed.starts_with(':') || trimmed.starts_with("event:") {
|
||||
return Ok(Vec::new());
|
||||
}
|
||||
let Some(data_line) = trimmed.strip_prefix("data:") else {
|
||||
return Ok(line);
|
||||
};
|
||||
let data_line = data_line.trim();
|
||||
if data_line.is_empty() || data_line == "[DONE]" {
|
||||
return Ok(line);
|
||||
}
|
||||
|
||||
let body: Value = match serde_json::from_str(data_line) {
|
||||
Ok(value) => value,
|
||||
Err(_) => return Ok(line),
|
||||
};
|
||||
|
||||
let envelope_name = report_context
|
||||
.get("envelope_name")
|
||||
.and_then(Value::as_str)
|
||||
.unwrap_or_default();
|
||||
let provider_api_format = report_context
|
||||
.get("provider_api_format")
|
||||
.and_then(Value::as_str)
|
||||
.unwrap_or_default();
|
||||
if !provider_adaptation_should_unwrap_stream_envelope(envelope_name, provider_api_format) {
|
||||
return Ok(line);
|
||||
}
|
||||
let unwrapped = match envelope_name {
|
||||
GEMINI_CLI_V1INTERNAL_ENVELOPE_NAME => body.get("response").cloned().unwrap_or(body),
|
||||
ANTIGRAVITY_V1INTERNAL_ENVELOPE_NAME => {
|
||||
let mut response = body.get("response").cloned().unwrap_or(body.clone());
|
||||
if let Some(response_id) = body.get("responseId").cloned() {
|
||||
if let Some(object) = response.as_object_mut() {
|
||||
object
|
||||
.entry("_v1internal_response_id".to_string())
|
||||
.or_insert(response_id);
|
||||
}
|
||||
}
|
||||
inject_antigravity_stream_tool_ids(&mut response);
|
||||
response
|
||||
}
|
||||
_ => body,
|
||||
};
|
||||
|
||||
let mut out = b"data: ".to_vec();
|
||||
out.extend(serde_json::to_vec(&unwrapped)?);
|
||||
out.extend_from_slice(b"\n\n");
|
||||
Ok(out)
|
||||
}
|
||||
|
||||
enum ProviderPrivateStreamNormalizeMode {
|
||||
EnvelopeUnwrap,
|
||||
KiroToClaudeCli(Box<KiroToClaudeCliStreamState>),
|
||||
}
|
||||
|
||||
pub struct ProviderPrivateStreamNormalizer<'a> {
|
||||
report_context: &'a Value,
|
||||
buffered: Vec<u8>,
|
||||
mode: ProviderPrivateStreamNormalizeMode,
|
||||
}
|
||||
|
||||
pub fn maybe_build_provider_private_stream_normalizer<'a>(
|
||||
report_context: Option<&'a Value>,
|
||||
) -> Option<ProviderPrivateStreamNormalizer<'a>> {
|
||||
let report_context = report_context?;
|
||||
if !report_context
|
||||
.get("has_envelope")
|
||||
.and_then(Value::as_bool)
|
||||
.unwrap_or(false)
|
||||
{
|
||||
return None;
|
||||
}
|
||||
let envelope_name = report_context
|
||||
.get("envelope_name")
|
||||
.and_then(Value::as_str)
|
||||
.unwrap_or_default();
|
||||
let provider_api_format = report_context
|
||||
.get("provider_api_format")
|
||||
.and_then(Value::as_str)
|
||||
.unwrap_or_default();
|
||||
let descriptor =
|
||||
provider_adaptation_descriptor_for_envelope(envelope_name, provider_api_format)?;
|
||||
let mode = if descriptor
|
||||
.envelope_name
|
||||
.eq_ignore_ascii_case(KIRO_ENVELOPE_NAME)
|
||||
{
|
||||
ProviderPrivateStreamNormalizeMode::KiroToClaudeCli(Box::new(
|
||||
KiroToClaudeCliStreamState::new(report_context),
|
||||
))
|
||||
} else if descriptor.unwraps_response_envelope {
|
||||
ProviderPrivateStreamNormalizeMode::EnvelopeUnwrap
|
||||
} else {
|
||||
return None;
|
||||
};
|
||||
Some(ProviderPrivateStreamNormalizer {
|
||||
report_context,
|
||||
buffered: Vec::new(),
|
||||
mode,
|
||||
})
|
||||
}
|
||||
|
||||
impl ProviderPrivateStreamNormalizer<'_> {
|
||||
pub fn push_chunk(&mut self, chunk: &[u8]) -> Result<Vec<u8>, AiSurfaceFinalizeError> {
|
||||
match &mut self.mode {
|
||||
ProviderPrivateStreamNormalizeMode::KiroToClaudeCli(state) => {
|
||||
state.push_chunk(self.report_context, chunk)
|
||||
}
|
||||
ProviderPrivateStreamNormalizeMode::EnvelopeUnwrap => {
|
||||
self.buffered.extend_from_slice(chunk);
|
||||
let mut output = Vec::new();
|
||||
while let Some(line_end) = self.buffered.iter().position(|byte| *byte == b'\n') {
|
||||
let line = self.buffered.drain(..=line_end).collect::<Vec<_>>();
|
||||
output.extend(
|
||||
transform_provider_private_stream_line(self.report_context, line)
|
||||
.map_err(AiSurfaceFinalizeError::from)?,
|
||||
);
|
||||
}
|
||||
Ok(output)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
pub fn finish(&mut self) -> Result<Vec<u8>, AiSurfaceFinalizeError> {
|
||||
match &mut self.mode {
|
||||
ProviderPrivateStreamNormalizeMode::KiroToClaudeCli(state) => {
|
||||
state.finish(self.report_context)
|
||||
}
|
||||
ProviderPrivateStreamNormalizeMode::EnvelopeUnwrap => {
|
||||
if self.buffered.is_empty() {
|
||||
return Ok(Vec::new());
|
||||
}
|
||||
let line = std::mem::take(&mut self.buffered);
|
||||
transform_provider_private_stream_line(self.report_context, line)
|
||||
.map_err(AiSurfaceFinalizeError::from)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
pub fn stream_body_contains_error_event(body: &[u8]) -> bool {
|
||||
let Ok(text) = std::str::from_utf8(body) else {
|
||||
return false;
|
||||
};
|
||||
let mut current_event_type: Option<String> = None;
|
||||
for raw_line in text.lines() {
|
||||
let line = raw_line.trim_matches('\r').trim();
|
||||
if line.is_empty() || line.starts_with(':') {
|
||||
continue;
|
||||
}
|
||||
if let Some(event_name) = line.strip_prefix("event:") {
|
||||
current_event_type = Some(event_name.trim().to_string());
|
||||
continue;
|
||||
}
|
||||
let data_line = if let Some(rest) = line.strip_prefix("data:") {
|
||||
rest.trim()
|
||||
} else {
|
||||
line
|
||||
};
|
||||
if data_line.is_empty() || data_line == "[DONE]" {
|
||||
continue;
|
||||
}
|
||||
let Ok(mut event) = serde_json::from_str::<Value>(data_line) else {
|
||||
continue;
|
||||
};
|
||||
if let Some(event_object) = event.as_object_mut() {
|
||||
if !event_object.contains_key("type") {
|
||||
if let Some(event_name) = current_event_type.take() {
|
||||
event_object.insert("type".to_string(), Value::String(event_name));
|
||||
}
|
||||
}
|
||||
}
|
||||
if event
|
||||
.get("type")
|
||||
.and_then(Value::as_str)
|
||||
.is_some_and(|value| value.eq_ignore_ascii_case("error"))
|
||||
{
|
||||
return true;
|
||||
}
|
||||
current_event_type = None;
|
||||
}
|
||||
false
|
||||
}
|
||||
|
||||
fn clear_private_envelope_context(report_context: &Value) -> Value {
|
||||
let mut normalized = report_context.clone();
|
||||
if let Some(object) = normalized.as_object_mut() {
|
||||
object.insert("has_envelope".to_string(), Value::Bool(false));
|
||||
object.remove("envelope_name");
|
||||
}
|
||||
normalized
|
||||
}
|
||||
|
||||
fn local_finalize_response_model(report_context: &Value) -> &str {
|
||||
report_context
|
||||
.get("mapped_model")
|
||||
.and_then(Value::as_str)
|
||||
.or_else(|| report_context.get("model").and_then(Value::as_str))
|
||||
.unwrap_or_default()
|
||||
}
|
||||
|
||||
fn inject_antigravity_stream_tool_ids(value: &mut Value) {
|
||||
let Some(candidates) = value.get_mut("candidates").and_then(Value::as_array_mut) else {
|
||||
return;
|
||||
};
|
||||
|
||||
for candidate in candidates {
|
||||
let Some(parts) = candidate
|
||||
.get_mut("content")
|
||||
.and_then(Value::as_object_mut)
|
||||
.and_then(|content| content.get_mut("parts"))
|
||||
.and_then(Value::as_array_mut)
|
||||
else {
|
||||
continue;
|
||||
};
|
||||
|
||||
let mut counters: BTreeMap<String, usize> = BTreeMap::new();
|
||||
for part in parts {
|
||||
let Some(function_call) = part.get_mut("functionCall").and_then(Value::as_object_mut)
|
||||
else {
|
||||
continue;
|
||||
};
|
||||
let has_id = function_call
|
||||
.get("id")
|
||||
.and_then(Value::as_str)
|
||||
.is_some_and(|value| !value.is_empty());
|
||||
if has_id {
|
||||
continue;
|
||||
}
|
||||
let name = function_call
|
||||
.get("name")
|
||||
.and_then(Value::as_str)
|
||||
.filter(|value| !value.is_empty())
|
||||
.unwrap_or("unknown")
|
||||
.to_string();
|
||||
let index = counters.entry(name.clone()).or_insert(0);
|
||||
function_call.insert(
|
||||
"id".to_string(),
|
||||
Value::String(format!("call_{name}_{index}")),
|
||||
);
|
||||
*index += 1;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
fn inject_antigravity_sync_tool_ids(response: &mut Value, model: &str) {
|
||||
if !model.to_ascii_lowercase().contains("claude") {
|
||||
return;
|
||||
}
|
||||
|
||||
let Some(candidates) = response.get_mut("candidates").and_then(Value::as_array_mut) else {
|
||||
return;
|
||||
};
|
||||
|
||||
for candidate in candidates {
|
||||
let Some(parts) = candidate
|
||||
.get_mut("content")
|
||||
.and_then(Value::as_object_mut)
|
||||
.and_then(|content| content.get_mut("parts"))
|
||||
.and_then(Value::as_array_mut)
|
||||
else {
|
||||
continue;
|
||||
};
|
||||
|
||||
let mut name_counters: BTreeMap<String, usize> = BTreeMap::new();
|
||||
for part in parts {
|
||||
let function_call = if let Some(function_call) =
|
||||
part.get_mut("functionCall").and_then(Value::as_object_mut)
|
||||
{
|
||||
function_call
|
||||
} else if let Some(function_call) =
|
||||
part.get_mut("function_call").and_then(Value::as_object_mut)
|
||||
{
|
||||
function_call
|
||||
} else {
|
||||
continue;
|
||||
};
|
||||
let has_id = function_call
|
||||
.get("id")
|
||||
.and_then(Value::as_str)
|
||||
.is_some_and(|value| !value.is_empty());
|
||||
if has_id {
|
||||
continue;
|
||||
}
|
||||
let function_name = function_call
|
||||
.get("name")
|
||||
.and_then(Value::as_str)
|
||||
.filter(|value| !value.is_empty())
|
||||
.unwrap_or("unknown")
|
||||
.to_string();
|
||||
let count = name_counters.entry(function_name.clone()).or_insert(0);
|
||||
function_call.insert(
|
||||
"id".to_string(),
|
||||
Value::String(format!("call_{function_name}_{count}")),
|
||||
);
|
||||
*count += 1;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
fn postprocess_private_response_value(data: &mut Value, report_context: &Value) {
|
||||
if !matches!(
|
||||
report_context.get("envelope_name").and_then(Value::as_str),
|
||||
Some(ANTIGRAVITY_V1INTERNAL_ENVELOPE_NAME)
|
||||
) {
|
||||
return;
|
||||
}
|
||||
if let Some(object) = data.as_object_mut() {
|
||||
if !object.contains_key("_v1internal_response_id") {
|
||||
if let Some(response_id) = object.remove("responseId") {
|
||||
object.insert("_v1internal_response_id".to_string(), response_id);
|
||||
}
|
||||
}
|
||||
}
|
||||
inject_antigravity_sync_tool_ids(data, local_finalize_response_model(report_context));
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use serde_json::json;
|
||||
|
||||
use super::{
|
||||
maybe_build_provider_private_stream_normalizer, normalize_provider_private_report_context,
|
||||
normalize_provider_private_response_value, stream_body_contains_error_event,
|
||||
transform_provider_private_stream_line,
|
||||
};
|
||||
|
||||
#[test]
|
||||
fn normalizes_supported_private_report_context() {
|
||||
let report_context = json!({
|
||||
"has_envelope": true,
|
||||
"envelope_name": "antigravity:v1internal",
|
||||
"provider_api_format": "gemini:generate_content",
|
||||
});
|
||||
let normalized = normalize_provider_private_report_context(Some(&report_context))
|
||||
.expect("context should normalize");
|
||||
assert_eq!(normalized["has_envelope"], json!(false));
|
||||
assert!(normalized.get("envelope_name").is_none());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn unwraps_antigravity_sync_response_and_injects_ids() {
|
||||
let report_context = json!({
|
||||
"has_envelope": true,
|
||||
"provider_api_format": "gemini:generate_content",
|
||||
"envelope_name": "antigravity:v1internal",
|
||||
"mapped_model": "claude-sonnet-4-5",
|
||||
});
|
||||
let body = json!({
|
||||
"response": {
|
||||
"candidates": [{
|
||||
"content": {
|
||||
"parts": [{
|
||||
"functionCall": {
|
||||
"name": "get_weather",
|
||||
"args": {"city": "SF"}
|
||||
}
|
||||
}]
|
||||
}
|
||||
}]
|
||||
},
|
||||
"responseId": "resp_123"
|
||||
});
|
||||
|
||||
let normalized = normalize_provider_private_response_value(body, &report_context)
|
||||
.expect("body should normalize");
|
||||
assert_eq!(normalized["_v1internal_response_id"], json!("resp_123"));
|
||||
assert_eq!(
|
||||
normalized["candidates"][0]["content"]["parts"][0]["functionCall"]["id"],
|
||||
json!("call_get_weather_0")
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn unwraps_antigravity_stream_line_and_injects_ids() {
|
||||
let report_context = json!({
|
||||
"has_envelope": true,
|
||||
"provider_api_format": "gemini:generate_content",
|
||||
"client_api_format": "gemini:generate_content",
|
||||
"envelope_name": "antigravity:v1internal",
|
||||
"mapped_model": "claude-sonnet-4-5",
|
||||
});
|
||||
let output = transform_provider_private_stream_line(
|
||||
&report_context,
|
||||
b"data: {\"response\":{\"candidates\":[{\"content\":{\"parts\":[{\"functionCall\":{\"name\":\"get_weather\",\"args\":{\"city\":\"SF\"}}}],\"role\":\"model\"},\"index\":0}],\"modelVersion\":\"claude-sonnet-4-5\"},\"responseId\":\"resp_123\"}\n\n".to_vec(),
|
||||
)
|
||||
.expect("unwrap should succeed");
|
||||
let output_text = String::from_utf8(output).expect("text should decode");
|
||||
assert!(output_text.contains("\"_v1internal_response_id\":\"resp_123\""));
|
||||
assert!(output_text.contains("\"id\":\"call_get_weather_0\""));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn private_stream_normalizer_unwraps_antigravity_stream() {
|
||||
let report_context = json!({
|
||||
"has_envelope": true,
|
||||
"provider_api_format": "gemini:generate_content",
|
||||
"client_api_format": "gemini:generate_content",
|
||||
"envelope_name": "antigravity:v1internal",
|
||||
"mapped_model": "claude-sonnet-4-5",
|
||||
});
|
||||
let mut normalizer = maybe_build_provider_private_stream_normalizer(Some(&report_context))
|
||||
.expect("normalizer should exist");
|
||||
let output = normalizer
|
||||
.push_chunk(
|
||||
b"data: {\"response\":{\"candidates\":[{\"content\":{\"parts\":[{\"functionCall\":{\"name\":\"get_weather\",\"args\":{\"city\":\"SF\"}}}],\"role\":\"model\"},\"index\":0}],\"modelVersion\":\"claude-sonnet-4-5\"},\"responseId\":\"resp_123\"}\n\n",
|
||||
)
|
||||
.expect("unwrap should succeed");
|
||||
let output_text = String::from_utf8(output).expect("text should decode");
|
||||
assert!(output_text.contains("\"_v1internal_response_id\":\"resp_123\""));
|
||||
assert!(output_text.contains("\"id\":\"call_get_weather_0\""));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn detects_sse_error_events_without_explicit_type_field() {
|
||||
let body = br#"event: error
|
||||
data: {"message":"bad"}
|
||||
|
||||
"#;
|
||||
assert!(stream_body_contains_error_event(body));
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,6 @@
|
||||
pub mod rules;
|
||||
|
||||
pub use rules::{
|
||||
apply_local_body_rules, apply_local_header_rules, body_rules_are_locally_supported,
|
||||
body_rules_handle_path, header_rules_are_locally_supported,
|
||||
};
|
||||
1616
crates/aether-ai-formats/src/provider_compat/proxy/rules.rs
Normal file
1616
crates/aether-ai-formats/src/provider_compat/proxy/rules.rs
Normal file
File diff suppressed because it is too large
Load Diff
188
crates/aether-ai-formats/src/provider_compat/surfaces.rs
Normal file
188
crates/aether-ai-formats/src/provider_compat/surfaces.rs
Normal file
@@ -0,0 +1,188 @@
|
||||
pub const ANTIGRAVITY_PROVIDER_TYPE: &str = "antigravity";
|
||||
pub const KIRO_PROVIDER_TYPE: &str = "kiro";
|
||||
pub const KIRO_ENVELOPE_NAME: &str = "kiro:generateAssistantResponse";
|
||||
pub const ANTIGRAVITY_V1INTERNAL_ENVELOPE_NAME: &str = "antigravity:v1internal";
|
||||
pub const GEMINI_CLI_V1INTERNAL_ENVELOPE_NAME: &str = "gemini_cli:v1internal";
|
||||
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
||||
pub enum ProviderAdaptationSurface {
|
||||
AntigravityGeminiChat,
|
||||
AntigravityGeminiCli,
|
||||
GeminiCliV1Internal,
|
||||
KiroClaudeCli,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Copy)]
|
||||
pub struct ProviderAdaptationDescriptor {
|
||||
pub surface: ProviderAdaptationSurface,
|
||||
pub provider_type: Option<&'static str>,
|
||||
pub envelope_name: &'static str,
|
||||
pub anchor_api_format: &'static str,
|
||||
pub supports_request_bridge: bool,
|
||||
pub supports_sync_finalize_bridge: bool,
|
||||
pub supports_stream_bridge: bool,
|
||||
pub requires_eventstream_accept: bool,
|
||||
pub unwraps_response_envelope: bool,
|
||||
}
|
||||
|
||||
const PROVIDER_ADAPTATION_SURFACES: &[ProviderAdaptationDescriptor] = &[
|
||||
ProviderAdaptationDescriptor {
|
||||
surface: ProviderAdaptationSurface::AntigravityGeminiChat,
|
||||
provider_type: Some(ANTIGRAVITY_PROVIDER_TYPE),
|
||||
envelope_name: ANTIGRAVITY_V1INTERNAL_ENVELOPE_NAME,
|
||||
anchor_api_format: "gemini:generate_content",
|
||||
supports_request_bridge: true,
|
||||
supports_sync_finalize_bridge: true,
|
||||
supports_stream_bridge: true,
|
||||
requires_eventstream_accept: false,
|
||||
unwraps_response_envelope: true,
|
||||
},
|
||||
ProviderAdaptationDescriptor {
|
||||
surface: ProviderAdaptationSurface::AntigravityGeminiCli,
|
||||
provider_type: Some(ANTIGRAVITY_PROVIDER_TYPE),
|
||||
envelope_name: ANTIGRAVITY_V1INTERNAL_ENVELOPE_NAME,
|
||||
anchor_api_format: "gemini:generate_content",
|
||||
supports_request_bridge: true,
|
||||
supports_sync_finalize_bridge: true,
|
||||
supports_stream_bridge: true,
|
||||
requires_eventstream_accept: false,
|
||||
unwraps_response_envelope: true,
|
||||
},
|
||||
ProviderAdaptationDescriptor {
|
||||
surface: ProviderAdaptationSurface::GeminiCliV1Internal,
|
||||
provider_type: None,
|
||||
envelope_name: GEMINI_CLI_V1INTERNAL_ENVELOPE_NAME,
|
||||
anchor_api_format: "gemini:generate_content",
|
||||
supports_request_bridge: false,
|
||||
supports_sync_finalize_bridge: true,
|
||||
supports_stream_bridge: true,
|
||||
requires_eventstream_accept: false,
|
||||
unwraps_response_envelope: true,
|
||||
},
|
||||
ProviderAdaptationDescriptor {
|
||||
surface: ProviderAdaptationSurface::KiroClaudeCli,
|
||||
provider_type: Some(KIRO_PROVIDER_TYPE),
|
||||
envelope_name: KIRO_ENVELOPE_NAME,
|
||||
anchor_api_format: "claude:messages",
|
||||
supports_request_bridge: true,
|
||||
supports_sync_finalize_bridge: true,
|
||||
supports_stream_bridge: true,
|
||||
requires_eventstream_accept: true,
|
||||
unwraps_response_envelope: false,
|
||||
},
|
||||
];
|
||||
|
||||
pub fn provider_adaptation_descriptor_for_envelope(
|
||||
envelope_name: &str,
|
||||
provider_api_format: &str,
|
||||
) -> Option<&'static ProviderAdaptationDescriptor> {
|
||||
let envelope_name = envelope_name.trim();
|
||||
let provider_api_format = provider_api_format.trim().to_ascii_lowercase();
|
||||
PROVIDER_ADAPTATION_SURFACES.iter().find(|descriptor| {
|
||||
descriptor.envelope_name.eq_ignore_ascii_case(envelope_name)
|
||||
&& descriptor
|
||||
.anchor_api_format
|
||||
.eq_ignore_ascii_case(provider_api_format.as_str())
|
||||
})
|
||||
}
|
||||
|
||||
pub fn provider_adaptation_descriptor_for_provider_type(
|
||||
provider_type: &str,
|
||||
provider_api_format: &str,
|
||||
) -> Option<&'static ProviderAdaptationDescriptor> {
|
||||
let provider_type = provider_type.trim();
|
||||
let provider_api_format = provider_api_format.trim().to_ascii_lowercase();
|
||||
PROVIDER_ADAPTATION_SURFACES.iter().find(|descriptor| {
|
||||
descriptor
|
||||
.provider_type
|
||||
.is_some_and(|value| value.eq_ignore_ascii_case(provider_type))
|
||||
&& descriptor
|
||||
.anchor_api_format
|
||||
.eq_ignore_ascii_case(provider_api_format.as_str())
|
||||
})
|
||||
}
|
||||
|
||||
pub fn provider_adaptation_anchor_api_format(
|
||||
envelope_name: &str,
|
||||
provider_api_format: &str,
|
||||
) -> Option<&'static str> {
|
||||
provider_adaptation_descriptor_for_envelope(envelope_name, provider_api_format)
|
||||
.map(|descriptor| descriptor.anchor_api_format)
|
||||
}
|
||||
|
||||
pub fn provider_adaptation_allows_sync_finalize_envelope(
|
||||
envelope_name: &str,
|
||||
provider_api_format: &str,
|
||||
) -> bool {
|
||||
provider_adaptation_descriptor_for_envelope(envelope_name, provider_api_format)
|
||||
.is_some_and(|descriptor| descriptor.supports_sync_finalize_bridge)
|
||||
}
|
||||
|
||||
pub fn provider_adaptation_requires_eventstream_accept(
|
||||
envelope_name: Option<&str>,
|
||||
provider_api_format: &str,
|
||||
) -> bool {
|
||||
envelope_name
|
||||
.and_then(|value| provider_adaptation_descriptor_for_envelope(value, provider_api_format))
|
||||
.is_some_and(|descriptor| descriptor.requires_eventstream_accept)
|
||||
}
|
||||
|
||||
pub fn provider_adaptation_should_unwrap_stream_envelope(
|
||||
envelope_name: &str,
|
||||
provider_api_format: &str,
|
||||
) -> bool {
|
||||
provider_adaptation_descriptor_for_envelope(envelope_name, provider_api_format)
|
||||
.is_some_and(|descriptor| descriptor.unwraps_response_envelope)
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::{
|
||||
provider_adaptation_allows_sync_finalize_envelope, provider_adaptation_anchor_api_format,
|
||||
provider_adaptation_requires_eventstream_accept,
|
||||
provider_adaptation_should_unwrap_stream_envelope, ANTIGRAVITY_V1INTERNAL_ENVELOPE_NAME,
|
||||
GEMINI_CLI_V1INTERNAL_ENVELOPE_NAME, KIRO_ENVELOPE_NAME,
|
||||
};
|
||||
|
||||
#[test]
|
||||
fn resolves_private_surface_anchor_contracts() {
|
||||
assert_eq!(
|
||||
provider_adaptation_anchor_api_format(
|
||||
ANTIGRAVITY_V1INTERNAL_ENVELOPE_NAME,
|
||||
"gemini:generate_content"
|
||||
),
|
||||
Some("gemini:generate_content")
|
||||
);
|
||||
assert_eq!(
|
||||
provider_adaptation_anchor_api_format(
|
||||
GEMINI_CLI_V1INTERNAL_ENVELOPE_NAME,
|
||||
"gemini:generate_content"
|
||||
),
|
||||
Some("gemini:generate_content")
|
||||
);
|
||||
assert_eq!(
|
||||
provider_adaptation_anchor_api_format(KIRO_ENVELOPE_NAME, "claude:messages"),
|
||||
Some("claude:messages")
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn exposes_private_surface_capabilities() {
|
||||
assert!(provider_adaptation_allows_sync_finalize_envelope(
|
||||
ANTIGRAVITY_V1INTERNAL_ENVELOPE_NAME,
|
||||
"gemini:generate_content"
|
||||
));
|
||||
assert!(provider_adaptation_should_unwrap_stream_envelope(
|
||||
GEMINI_CLI_V1INTERNAL_ENVELOPE_NAME,
|
||||
"gemini:generate_content"
|
||||
));
|
||||
assert!(provider_adaptation_requires_eventstream_accept(
|
||||
Some(KIRO_ENVELOPE_NAME),
|
||||
"claude:messages"
|
||||
));
|
||||
assert!(!provider_adaptation_requires_eventstream_accept(
|
||||
Some(ANTIGRAVITY_V1INTERNAL_ENVELOPE_NAME),
|
||||
"gemini:generate_content"
|
||||
));
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user