devin15/cursor2api-rust
0
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 