Team Ai
Apppublic

devin15/cursor2api-rust

sourceHugging Faceupdated 2y agoView on Hugging Face
0likes
service.rs711 linesDownload Raw Back to chat
1use crate::{2    app::{3        constant::{4            AUTHORIZATION_BEARER_PREFIX, FINISH_REASON_STOP, OBJECT_CHAT_COMPLETION,5            OBJECT_CHAT_COMPLETION_CHUNK, STATUS_FAILED, STATUS_PENDING, STATUS_SUCCESS,6        },7        lazy::{AUTH_TOKEN, SHARED_AUTH_TOKEN, USE_SHARE},8        model::{AppConfig, AppState, ChatRequest, RequestLog, TimingInfo, TokenInfo},9    },10    chat::{11        constant::{AVAILABLE_MODELS, USAGE_CHECK_MODELS},12        error::StreamError,13        model::{14            ChatResponse, Choice, Delta, Message, MessageContent, ModelsResponse, Role, Usage,15        },16        stream::{parse_stream_data, StreamMessage},17    },18    common::{19        client::build_client,20        models::{error::ChatError, userinfo::MembershipType, ErrorResponse},21        utils::{format_time_ms, generate_checksum_with_repair, get_token_profile, validate_token_and_checksum},22    },23};24use axum::{25    body::Body,26    extract::State,27    http::{28        header::{AUTHORIZATION, CONTENT_TYPE},29        HeaderMap, StatusCode,30    },31    response::Response,32    Json,33};34use bytes::Bytes;35use futures::{Stream, StreamExt};36use std::{37    convert::Infallible,38    sync::{atomic::AtomicBool, Arc},39};40use std::{41    pin::Pin,42    sync::atomic::{AtomicUsize, Ordering},43};44use tokio::sync::Mutex;45use uuid::Uuid;46 47const REQUEST_LOGS_LIMIT: usize = 1000;48 49// 模型列表处理50pub async fn handle_models() -> Json<ModelsResponse> {51    Json(ModelsResponse {52        object: "list",53        data: &AVAILABLE_MODELS,54    })55}56 57// 聊天处理函数的签名58pub async fn handle_chat(59    State(state): State<Arc<Mutex<AppState>>>,60    headers: HeaderMap,61    Json(request): Json<ChatRequest>,62) -> Result<Response<Body>, (StatusCode, Json<ErrorResponse>)> {63    let allow_claude = AppConfig::get_allow_claude();64    // 验证模型是否支持并获取模型信息65    let model = AVAILABLE_MODELS.iter().find(|m| m.id == request.model);66    let model_supported = model.is_some();67 68    if !(model_supported || allow_claude && request.model.starts_with("claude")) {69        return Err((70            StatusCode::BAD_REQUEST,71            Json(ChatError::ModelNotSupported(request.model).to_json()),72        ));73    }74 75    let request_time = chrono::Local::now();76 77    // 验证请求78    if request.messages.is_empty() {79        return Err((80            StatusCode::BAD_REQUEST,81            Json(ChatError::EmptyMessages.to_json()),82        ));83    }84 85    // 获取并处理认证令牌86    let auth_header = headers87        .get(AUTHORIZATION)88        .and_then(|h| h.to_str().ok())89        .and_then(|h| h.strip_prefix(AUTHORIZATION_BEARER_PREFIX))90        .ok_or((91            StatusCode::UNAUTHORIZED,92            Json(ChatError::Unauthorized.to_json()),93        ))?;94 95    // 验证认证token并获取token信息96    let (auth_token, checksum) = match auth_header {97        // 管理员Token验证逻辑98        token if token == AUTH_TOKEN.as_str() || (*USE_SHARE && token == SHARED_AUTH_TOKEN.as_str()) => {99            static CURRENT_KEY_INDEX: AtomicUsize = AtomicUsize::new(0);100            let state_guard = state.lock().await;101            let token_infos = &state_guard.token_infos;102 103            // 检查是否存在可用的token104            if token_infos.is_empty() {105                return Err((106                    StatusCode::SERVICE_UNAVAILABLE, 107                    Json(ChatError::NoTokens.to_json()),108                ));109            }110 111            // 轮询选择token112            let index = CURRENT_KEY_INDEX.fetch_add(1, Ordering::SeqCst) % token_infos.len();113            let token_info = &token_infos[index];114            (token_info.token.clone(), token_info.checksum.clone())115        },116 117        // 普通用户Token验证逻辑118        token => validate_token_and_checksum(token).ok_or((119            StatusCode::UNAUTHORIZED,120            Json(ChatError::Unauthorized.to_json()),121        ))?,122    };123 124    let current_id: u64;125 126    // 更新请求日志127    {128        let state_clone = state.clone();129        let mut state = state.lock().await;130        state.total_requests += 1;131        state.active_requests += 1;132 133        // 查找最新的相同token的日志,检查使用情况134        let need_profile_check = state135            .request_logs136            .iter()137            .rev()138            .find(|log| log.token_info.token == auth_token && log.token_info.profile.is_some())139            .and_then(|log| log.token_info.profile.as_ref())140            .map(|profile| {141                if profile.stripe.membership_type != MembershipType::Free {142                    return false;143                }144 145                let is_premium = USAGE_CHECK_MODELS.contains(&request.model.as_str());146                let standard = &profile.usage.standard;147                let premium = &profile.usage.premium;148 149                if is_premium {150                    premium151                        .max_requests152                        .map_or(false, |max| premium.num_requests >= max)153                } else {154                    standard155                        .max_requests156                        .map_or(false, |max| standard.num_requests >= max)157                }158            })159            .unwrap_or(false);160 161        // 如果达到限制,直接返回未授权错误162        if need_profile_check {163            state.active_requests -= 1;164            state.error_requests += 1;165            return Err((166                StatusCode::UNAUTHORIZED,167                Json(ChatError::Unauthorized.to_json()),168            ));169        }170 171        let next_id = state.request_logs.last().map_or(1, |log| log.id + 1);172        current_id = next_id;173 174        // 如果需要获取用户使用情况,创建后台任务获取profile175        if model.map(|m| m.is_usage_check()).unwrap_or(false) {176            let auth_token_clone = auth_token.clone();177            let state_clone = state_clone.clone();178            let log_id = next_id;179 180            tokio::spawn(async move {181                let profile = get_token_profile(&auth_token_clone).await;182                let mut state = state_clone.lock().await;183                // 根据id查找对应的日志184                if let Some(log) = state185                    .request_logs186                    .iter_mut()187                    .rev()188                    .find(|log| log.id == log_id)189                {190                    log.token_info.profile = profile;191                }192            });193        }194 195        state.request_logs.push(RequestLog {196            id: next_id,197            timestamp: request_time,198            model: request.model.clone(),199            token_info: TokenInfo {200                token: auth_token.clone(),201                checksum: checksum.clone(),202                profile: None,203            },204            prompt: None,205            timing: TimingInfo {206                total: 0.0,207                first: None,208            },209            stream: request.stream,210            status: STATUS_PENDING,211            error: None,212        });213 214        if state.request_logs.len() > REQUEST_LOGS_LIMIT {215            state.request_logs.remove(0);216        }217    }218 219    // 将消息转换为hex格式220    let hex_data = match super::adapter::encode_chat_message(request.messages, &request.model).await221    {222        Ok(data) => data,223        Err(e) => {224            let mut state = state.lock().await;225            if let Some(log) = state226                .request_logs227                .iter_mut()228                .rev()229                .find(|log| log.id == current_id)230            {231                log.status = STATUS_FAILED;232                log.error = Some(e.to_string());233            }234            state.active_requests -= 1;235            state.error_requests += 1;236            return Err((237                StatusCode::INTERNAL_SERVER_ERROR,238                Json(239                    ChatError::RequestFailed("Failed to encode chat message".to_string()).to_json(),240                ),241            ));242        }243    };244 245    // 构建请求客户端246    let client = build_client(&auth_token, &generate_checksum_with_repair(&checksum));247    let response = client.body(hex_data).send().await;248 249    // 处理请求结果250    let response = match response {251        Ok(resp) => {252            // 更新请求日志为成功253            {254                let mut state = state.lock().await;255                if let Some(log) = state256                    .request_logs257                    .iter_mut()258                    .rev()259                    .find(|log| log.id == current_id)260                {261                    log.status = STATUS_SUCCESS;262                }263            }264            resp265        }266        Err(e) => {267            // 更新请求日志为失败268            {269                let mut state = state.lock().await;270                if let Some(log) = state271                    .request_logs272                    .iter_mut()273                    .rev()274                    .find(|log| log.id == current_id)275                {276                    log.status = STATUS_FAILED;277                    log.error = Some(e.to_string());278                }279                state.active_requests -= 1;280                state.error_requests += 1;281            }282            return Err((283                StatusCode::INTERNAL_SERVER_ERROR,284                Json(ChatError::RequestFailed(e.to_string()).to_json()),285            ));286        }287    };288 289    // 释放活动请求计数290    {291        let mut state = state.lock().await;292        state.active_requests -= 1;293    }294 295    if request.stream {296        let response_id = format!("chatcmpl-{}", Uuid::new_v4().simple());297        let full_text = Arc::new(Mutex::new(String::with_capacity(1024)));298        let is_start = Arc::new(AtomicBool::new(true));299        let start_time = std::time::Instant::now();300        let first_chunk_time = Arc::new(Mutex::new(None));301 302        let stream = {303            // 创建新的 stream304            let mut stream = response.bytes_stream();305 306            let enable_stream_check = AppConfig::get_stream_check();307 308            if enable_stream_check {309                // 检查第一个 chunk310                match stream.next().await {311                    Some(first_chunk) => {312                        let chunk = first_chunk.map_err(|e| {313                            let error_message = format!("Failed to read response chunk: {}", e);314                            // 理论上,若程序正常,必定成功,因为前面判断过了315                            (316                                StatusCode::INTERNAL_SERVER_ERROR,317                                Json(ChatError::RequestFailed(error_message).to_json()),318                            )319                        })?;320 321                        match parse_stream_data(&chunk) {322                            Err(StreamError::ChatError(error)) => {323                                let error_respone = error.to_error_response();324                                // 更新请求日志为失败325                                {326                                    let mut state = state.lock().await;327                                    if let Some(log) = state328                                        .request_logs329                                        .iter_mut()330                                        .rev()331                                        .find(|log| log.id == current_id)332                                    {333                                        log.status = STATUS_FAILED;334                                        log.error = Some(error_respone.native_code());335                                        log.timing.total =336                                            format_time_ms(start_time.elapsed().as_secs_f64());337                                        state.error_requests += 1;338                                    }339                                }340                                return Err((341                                    error_respone.status_code(),342                                    Json(error_respone.to_common()),343                                ));344                            }345                            Ok(_) | Err(_) => {346                                // 创建一个包含第一个 chunk 的 stream347                                Box::pin(348                                    futures::stream::once(async move { Ok(chunk) }).chain(stream),349                                )350                                    as Pin<351                                        Box<352                                            dyn Stream<Item = Result<Bytes, reqwest::Error>> + Send,353                                        >,354                                    >355                            }356                        }357                    }358                    None => {359                        // Box::pin(stream)360                        //     as Pin<Box<dyn Stream<Item = Result<Bytes, reqwest::Error>> + Send>>361                        // 更新请求日志为失败362                        {363                            let mut state = state.lock().await;364                            if let Some(log) = state365                                .request_logs366                                .iter_mut()367                                .rev()368                                .find(|log| log.id == current_id)369                            {370                                log.status = STATUS_FAILED;371                                log.error = Some("Empty stream response".to_string());372                                state.error_requests += 1;373                            }374                        }375                        return Err((376                            StatusCode::INTERNAL_SERVER_ERROR,377                            Json(378                                ChatError::RequestFailed("Empty stream response".to_string())379                                    .to_json(),380                            ),381                        ));382                    }383                }384            } else {385                Box::pin(stream)386                    as Pin<Box<dyn Stream<Item = Result<Bytes, reqwest::Error>> + Send>>387            }388        }389        .then({390            let buffer = Arc::new(Mutex::new(Vec::new()));391            let first_chunk_time = first_chunk_time.clone();392            let state = state.clone();393 394            move |chunk| {395                let buffer = buffer.clone();396                let response_id = response_id.clone();397                let model = request.model.clone();398                let is_start = is_start.clone();399                let full_text = full_text.clone();400                let first_chunk_time = first_chunk_time.clone();401                let state = state.clone();402 403                async move {404                    let chunk = chunk.unwrap_or_default();405                    let mut buffer_guard = buffer.lock().await;406                    buffer_guard.extend_from_slice(&chunk);407 408                    match parse_stream_data(&buffer_guard) {409                        Ok(StreamMessage::Content(texts)) => {410                            buffer_guard.clear();411                            let mut response_data = String::new();412 413                            // 记录首字时间(如果还未记录)414                            if let Ok(mut first_time) = first_chunk_time.try_lock() {415                                if first_time.is_none() {416                                    *first_time =417                                        Some(format_time_ms(start_time.elapsed().as_secs_f64()));418                                }419                            }420 421                            // 处理文本内容422                            for text in texts {423                                let mut text_guard = full_text.lock().await;424                                text_guard.push_str(&text);425                                let is_first = is_start.load(Ordering::SeqCst);426 427                                let response = ChatResponse {428                                    id: response_id.clone(),429                                    object: OBJECT_CHAT_COMPLETION_CHUNK.to_string(),430                                    created: chrono::Utc::now().timestamp(),431                                    model: if is_first { Some(model.clone()) } else { None },432                                    choices: vec![Choice {433                                        index: 0,434                                        message: None,435                                        delta: Some(Delta {436                                            role: if is_first {437                                                is_start.store(false, Ordering::SeqCst);438                                                Some(Role::Assistant)439                                            } else {440                                                None441                                            },442                                            content: Some(text),443                                        }),444                                        finish_reason: None,445                                    }],446                                    usage: None,447                                };448 449                                response_data.push_str(&format!(450                                    "data: {}\n\n",451                                    serde_json::to_string(&response).unwrap()452                                ));453                            }454 455                            Ok::<_, Infallible>(Bytes::from(response_data))456                        }457                        Ok(StreamMessage::StreamStart) => {458                            buffer_guard.clear();459                            // 发送初始响应,包含模型信息460                            let response = ChatResponse {461                                id: response_id.clone(),462                                object: OBJECT_CHAT_COMPLETION_CHUNK.to_string(),463                                created: chrono::Utc::now().timestamp(),464                                model: {465                                    is_start.store(true, Ordering::SeqCst);466                                    Some(model.clone())467                                },468                                choices: vec![Choice {469                                    index: 0,470                                    message: None,471                                    delta: Some(Delta {472                                        role: Some(Role::Assistant),473                                        content: Some(String::new()),474                                    }),475                                    finish_reason: None,476                                }],477                                usage: None,478                            };479 480                            Ok(Bytes::from(format!(481                                "data: {}\n\n",482                                serde_json::to_string(&response).unwrap()483                            )))484                        }485                        Ok(StreamMessage::StreamEnd) => {486                            buffer_guard.clear();487                            // 根据配置决定是否发送最后的 finish_reason488                            let include_finish_reason = AppConfig::get_stop_stream();489 490                            // 计算总时间和首次片段时间491                            let total_time = format_time_ms(start_time.elapsed().as_secs_f64());492                            let first_time = first_chunk_time.lock().await.unwrap_or(total_time);493 494                            {495                                let mut state = state.lock().await;496                                if let Some(log) = state497                                    .request_logs498                                    .iter_mut()499                                    .rev()500                                    .find(|log| log.id == current_id)501                                {502                                    log.timing.total = total_time;503                                    log.timing.first = Some(first_time);504                                }505                            }506 507                            if include_finish_reason {508                                let response = ChatResponse {509                                    id: response_id.clone(),510                                    object: OBJECT_CHAT_COMPLETION_CHUNK.to_string(),511                                    created: chrono::Utc::now().timestamp(),512                                    model: None,513                                    choices: vec![Choice {514                                        index: 0,515                                        message: None,516                                        delta: Some(Delta {517                                            role: None,518                                            content: None,519                                        }),520                                        finish_reason: Some(FINISH_REASON_STOP.to_string()),521                                    }],522                                    usage: None,523                                };524                                Ok(Bytes::from(format!(525                                    "data: {}\n\ndata: [DONE]\n\n",526                                    serde_json::to_string(&response).unwrap()527                                )))528                            } else {529                                Ok(Bytes::from("data: [DONE]\n\n"))530                            }531                        }532                        Ok(StreamMessage::Incomplete) => {533                            // 保持buffer中的数据以待下一个chunk534                            Ok(Bytes::new())535                        }536                        Ok(StreamMessage::Debug(debug_prompt)) => {537                            buffer_guard.clear();538                            if let Ok(mut state) = state.try_lock() {539                                if let Some(last_log) = state.request_logs.last_mut() {540                                    last_log.prompt = Some(debug_prompt.clone());541                                }542                            }543                            Ok(Bytes::new())544                        }545                        Err(e) => {546                            buffer_guard.clear();547                            eprintln!("[警告] Stream error: {}", e);548                            Ok(Bytes::new())549                        }550                    }551                }552            }553        });554 555        Ok(Response::builder()556            .header("Cache-Control", "no-cache")557            .header("Connection", "keep-alive")558            .header(CONTENT_TYPE, "text/event-stream")559            .body(Body::from_stream(stream))560            .unwrap())561    } else {562        // 非流式响应563        let start_time = std::time::Instant::now();564        let mut first_chunk_received = false;565        let mut first_chunk_time = 0.0;566        let mut full_text = String::with_capacity(1024);567        let mut stream = response.bytes_stream();568        let mut prompt = None;569 570        let mut buffer = Vec::new();571        while let Some(chunk) = stream.next().await {572            let chunk = chunk.map_err(|e| {573                // 更新请求日志为失败574                if let Ok(mut state) = state.try_lock() {575                    if let Some(log) = state576                        .request_logs577                        .iter_mut()578                        .rev()579                        .find(|log| log.id == current_id)580                    {581                        log.status = STATUS_FAILED;582                        log.error = Some(format!("Failed to read response chunk: {}", e));583                        state.error_requests += 1;584                    }585                }586                (587                    StatusCode::INTERNAL_SERVER_ERROR,588                    Json(589                        ChatError::RequestFailed(format!("Failed to read response chunk: {}", e))590                            .to_json(),591                    ),592                )593            })?;594 595            buffer.extend_from_slice(&chunk);596 597            match parse_stream_data(&buffer) {598                Ok(StreamMessage::Content(texts)) => {599                    if !first_chunk_received {600                        first_chunk_time = format_time_ms(start_time.elapsed().as_secs_f64());601                        first_chunk_received = true;602                    }603                    for text in texts {604                        full_text.push_str(&text);605                    }606                    buffer.clear();607                }608                Ok(StreamMessage::Incomplete) => continue,609                Ok(StreamMessage::Debug(debug_prompt)) => {610                    prompt = Some(debug_prompt);611                    buffer.clear();612                }613                Ok(StreamMessage::StreamStart) | Ok(StreamMessage::StreamEnd) => {614                    buffer.clear();615                }616                Err(StreamError::ChatError(error)) => {617                    let error = error.to_error_response();618                    // 更新请求日志为失败619                    {620                        let mut state = state.lock().await;621                        if let Some(log) = state622                            .request_logs623                            .iter_mut()624                            .rev()625                            .find(|log| log.id == current_id)626                        {627                            log.status = STATUS_FAILED;628                            log.error = Some(error.native_code());629                            log.timing.total = format_time_ms(start_time.elapsed().as_secs_f64());630                            state.error_requests += 1;631                        }632                    }633                    return Err((error.status_code(), Json(error.to_common())));634                }635                Err(_) => {636                    buffer.clear();637                    continue;638                }639            }640        }641 642        // 检查响应是否为空643        if full_text.is_empty() {644            // 更新请求日志为失败645            {646                let mut state = state.lock().await;647                if let Some(log) = state648                    .request_logs649                    .iter_mut()650                    .rev()651                    .find(|log| log.id == current_id)652                {653                    log.status = STATUS_FAILED;654                    log.error = Some("Empty response received".to_string());655                    if let Some(p) = prompt {656                        log.prompt = Some(p);657                    }658                    state.error_requests += 1;659                }660            }661            return Err((662                StatusCode::INTERNAL_SERVER_ERROR,663                Json(ChatError::RequestFailed("Empty response received".to_string()).to_json()),664            ));665        }666 667        let response_data = ChatResponse {668            id: format!("chatcmpl-{}", Uuid::new_v4().simple()),669            object: OBJECT_CHAT_COMPLETION.to_string(),670            created: chrono::Utc::now().timestamp(),671            model: Some(request.model),672            choices: vec![Choice {673                index: 0,674                message: Some(Message {675                    role: Role::Assistant,676                    content: MessageContent::Text(full_text),677                }),678                delta: None,679                finish_reason: Some(FINISH_REASON_STOP.to_string()),680            }],681            usage: Some(Usage {682                prompt_tokens: 0,683                completion_tokens: 0,684                total_tokens: 0,685            }),686        };687 688        {689            // 更新请求日志时间信息和状态690            let total_time = format_time_ms(start_time.elapsed().as_secs_f64());691            let mut state = state.lock().await;692            if let Some(log) = state693                .request_logs694                .iter_mut()695                .rev()696                .find(|log| log.id == current_id)697            {698                log.timing.total = total_time;699                log.timing.first = Some(first_chunk_time);700                log.prompt = prompt;701                log.status = STATUS_SUCCESS;702            }703        }704 705        Ok(Response::builder()706            .header(CONTENT_TYPE, "application/json")707            .body(Body::from(serde_json::to_string(&response_data).unwrap()))708            .unwrap())709    }710}711