Team Ai
Apppublic

devin15/cursor2api-rust

sourceHugging Faceupdated 2y agoView on Hugging Face
0likes
stream.rs210 linesDownload Raw Back to chat
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