devin15/cursor2api-rust
0
1use super::aiserver::v1::StreamChatResponse;2use flate2::read::GzDecoder;3use prost::Message;4use std::io::Read;5 6use super::error::{ChatError, StreamError};7 8// 解压gzip数据9fn decompress_gzip(data: &[u8]) -> Option<Vec<u8>> {10 let mut decoder = GzDecoder::new(data);11 let mut decompressed = Vec::new();12 13 match decoder.read_to_end(&mut decompressed) {14 Ok(_) => Some(decompressed),15 Err(_) => {16 // println!("gzip解压失败: {}", e);17 None18 }19 }20}21 22pub enum StreamMessage {23 // 未完成24 Incomplete,25 // 调试26 Debug(String),27 // 流开始标志 b"\0\0\0\0\0"28 StreamStart,29 // 消息内容30 Content(Vec<String>),31 // 流结束标志 b"\x02\0\0\0\x02{}"32 StreamEnd,33}34 35pub fn parse_stream_data(data: &[u8]) -> Result<StreamMessage, StreamError> {36 if data.len() < 5 {37 return Err(StreamError::DataLengthLessThan5);38 }39 40 // 检查是否为流开始标志41 // if data == b"\0\0\0\0\0" {42 // return Ok(StreamMessage::StreamStart);43 // }44 45 // 检查是否为流结束标志46 // if data == b"\x02\0\0\0\x02{}" {47 // return Ok(StreamMessage::StreamEnd);48 // }49 50 let mut messages = Vec::new();51 let mut offset = 0;52 53 while offset + 5 <= data.len() {54 // 获取消息类型和长度55 let msg_type = data[offset];56 let msg_len = u32::from_be_bytes([57 data[offset + 1],58 data[offset + 2],59 data[offset + 3],60 data[offset + 4],61 ]) as usize;62 63 // 流开始64 if msg_type == 0 && msg_len == 0 {65 return Ok(StreamMessage::StreamStart);66 }67 68 // 检查剩余数据长度是否足够69 if offset + 5 + msg_len > data.len() {70 return Ok(StreamMessage::Incomplete);71 }72 73 let msg_data = &data[offset + 5..offset + 5 + msg_len];74 75 match msg_type {76 // 文本消息77 0 => {78 if let Ok(response) = StreamChatResponse::decode(msg_data) {79 // crate::debug_println!("[text] StreamChatResponse: {:?}", response);80 if !response.text.is_empty() {81 messages.push(response.text);82 } else {83 // println!("[text] StreamChatResponse: {:?}", response);84 return Ok(StreamMessage::Debug(85 response.filled_prompt.unwrap_or_default(),86 // response.is_using_slow_request,87 ));88 }89 }90 }91 // gzip压缩消息92 1 => {93 if let Some(text) = decompress_gzip(msg_data) {94 let response = StreamChatResponse::decode(&text[..]).unwrap_or_default();95 // crate::debug_println!("[gzip] StreamChatResponse: {:?}", response);96 if !response.text.is_empty() {97 messages.push(response.text);98 } else {99 // println!("[gzip] StreamChatResponse: {:?}", response);100 return Ok(StreamMessage::Debug(101 response.filled_prompt.unwrap_or_default(),102 // response.is_using_slow_request,103 ));104 }105 }106 }107 // JSON字符串108 2 => {109 if msg_len == 2 {110 return Ok(StreamMessage::StreamEnd);111 }112 if let Ok(text) = String::from_utf8(msg_data.to_vec()) {113 // println!("JSON消息: {}", text);114 if let Ok(error) = serde_json::from_str::<ChatError>(&text) {115 return Err(StreamError::ChatError(error));116 }117 // 未预计118 // messages.push(text);119 }120 }121 // 其他类型暂不处理122 t => eprintln!("收到未知消息类型: {},请尝试联系开发者以获取支持", t),123 }124 125 offset += 5 + msg_len;126 }127 128 if messages.is_empty() {129 Err(StreamError::EmptyMessage)130 } else {131 Ok(StreamMessage::Content(messages))132 }133}134 135#[test]136fn test_parse_stream_data() {137 // 使用include_str!加载测试数据文件138 let stream_data = include_str!("../../tests/data/stream_data.txt");139 140 // 将整个字符串按每两个字符分割成字节141 let bytes: Vec<u8> = stream_data142 .as_bytes()143 .chunks(2)144 .map(|chunk| {145 let hex_str = std::str::from_utf8(chunk).unwrap();146 u8::from_str_radix(hex_str, 16).unwrap()147 })148 .collect();149 150 // 辅助函数:找到下一个消息边界151 fn find_next_message_boundary(bytes: &[u8]) -> usize {152 if bytes.len() < 5 {153 return bytes.len();154 }155 let msg_len = u32::from_be_bytes([bytes[1], bytes[2], bytes[3], bytes[4]]) as usize;156 5 + msg_len157 }158 159 // 辅助函数:将字节转换为hex字符串160 fn bytes_to_hex(bytes: &[u8]) -> String {161 bytes.iter()162 .map(|b| format!("{:02X}", b))163 .collect::<Vec<String>>()164 .join("")165 }166 167 // 多次解析数据168 let mut offset = 0;169 while offset < bytes.len() {170 let remaining_bytes = &bytes[offset..];171 let msg_boundary = find_next_message_boundary(remaining_bytes);172 let current_msg_bytes = &remaining_bytes[..msg_boundary];173 let hex_str = bytes_to_hex(current_msg_bytes);174 175 match parse_stream_data(current_msg_bytes) {176 Ok(message) => {177 match message {178 StreamMessage::Content(messages) => {179 print!("消息内容 [hex: {}]:", hex_str);180 for msg in messages {181 println!(" {}", msg);182 }183 offset += msg_boundary;184 }185 StreamMessage::Debug(_) => {186 // println!("调试信息 [hex: {}]: {}", hex_str, prompt);187 offset += msg_boundary;188 }189 StreamMessage::StreamEnd => {190 println!("流结束 [hex: {}]", hex_str);191 break;192 }193 StreamMessage::StreamStart => {194 println!("流开始 [hex: {}]", hex_str);195 offset += msg_boundary;196 }197 StreamMessage::Incomplete => {198 println!("数据不完整 [hex: {}]", hex_str);199 break;200 }201 }202 }203 Err(e) => {204 println!("解析错误 [hex: {}]: {}", hex_str, e);205 break;206 }207 }208 }209}210 