Files
Genarrative/server-rs/crates/api-server/src/llm/mod.rs
T
lhk229 9aa6f5efea
Project CI / AI game creator shell Rust shard 1/4 (push) Successful in 4m44s
Project CI / AI game creator shell Rust shard 2/4 (push) Successful in 5m9s
Project CI / AI game creator shell Rust shard 3/4 (push) Successful in 4m33s
Project CI / AI game creator shell Rust smoke (push) Successful in 1m50s
Project CI / AI game creator shell Rust shard 4/4 (push) Successful in 3m53s
Project CI / AI game creator shell Rust crates (push) Successful in 2m51s
Project CI / Frontend tests (push) Successful in 4m58s
Project CI / Repository checks (push) Successful in 3m18s
Project CI / Native shell tests (push) Successful in 6m10s
Project CI / Backend tests (push) Successful in 7m7s
Project CI / AI game creator shell web tests (push) Successful in 2m16s
Project CI / AI game creator shell Rust shard 2/4 (pull_request) Has been cancelled
Project CI / AI game creator shell Rust shard 3/4 (pull_request) Has been cancelled
Project CI / AI game creator shell Rust shard 4/4 (pull_request) Has been cancelled
Project CI / AI game creator shell Rust smoke (pull_request) Has been cancelled
Project CI / AI game creator shell Rust crates (pull_request) Has been cancelled
Project CI / Backend tests (pull_request) Has been cancelled
Project CI / Native shell tests (pull_request) Has been cancelled
Project CI / Frontend tests (pull_request) Has been cancelled
Project CI / Repository checks (pull_request) Has been cancelled
Project CI / AI game creator shell web tests (pull_request) Has been cancelled
Project CI / AI game creator shell Rust shard 1/4 (pull_request) Has been cancelled
解决platform-llm和策划agent reasoning问题,并修复若干bug (#350)
Reviewed-on: #350
2026-09-15 00:50:36 +08:00

1723 lines
62 KiB
Rust

use axum::{
Json,
body::{Body, Bytes},
extract::{Extension, State},
http::{HeaderMap, HeaderValue, StatusCode},
response::{
IntoResponse, Response,
sse::{Event, Sse},
},
};
use futures_util::StreamExt;
use platform_llm::{LlmApiKind, LlmMessage, LlmMessageRole, LlmRunRequest};
use serde_json::{Value, json};
use shared_contracts::llm::{
LlmChatCompletionRequest, LlmChatCompletionResponse, LlmChatMessagePayload, LlmChatMessageRole,
LlmModelSummary, LlmModelsResponse,
};
use spacetime_client::SpacetimeClientError;
use std::convert::Infallible;
#[cfg(test)]
use std::collections::HashMap;
#[cfg(test)]
use std::sync::{Mutex, OnceLock};
use crate::{
api_response::json_success_body, auth::AuthenticatedAccessToken, http_error::AppError,
platform_errors::map_llm_error, request_context::RequestContext, state::AppState,
};
pub(crate) const LLM_REQUEST_MAX_BODY_BYTES: usize = 32 * 1024 * 1024;
pub(crate) mod icon_specs;
#[cfg(test)]
mod model_catalog_tests {
use super::*;
#[test]
fn public_catalog_only_exposes_alias_and_stable_id() {
let mut catalog = module_runtime::AgcModelCatalog::default();
catalog.revision = 7;
catalog.models[1].enabled = false;
let payload = serde_json::to_value(public_model_catalog(catalog)).unwrap();
assert_eq!(
payload["models"],
json!([{"id": "quality", "displayName": "高质量"}])
);
assert_eq!(payload["defaultModelId"], "quality");
assert_eq!(payload["revision"], json!(7));
assert!(!payload.to_string().contains("gpt-"));
}
}
#[cfg(test)]
#[derive(Clone, Debug)]
struct TestProvisionedRouterCredential {
base_url: String,
api_key: String,
key_id: String,
}
#[cfg(test)]
static TEST_PROVISIONED_ROUTER_CREDENTIALS: OnceLock<
Mutex<HashMap<String, TestProvisionedRouterCredential>>,
> = OnceLock::new();
#[cfg(test)]
static TEST_LLM_ROUTER_WALLET_BALANCES: OnceLock<Mutex<HashMap<String, u64>>> = OnceLock::new();
#[cfg(test)]
fn test_provisioned_router_credentials()
-> &'static Mutex<HashMap<String, TestProvisionedRouterCredential>> {
TEST_PROVISIONED_ROUTER_CREDENTIALS.get_or_init(|| Mutex::new(HashMap::new()))
}
#[cfg(test)]
fn test_llm_router_wallet_balances() -> &'static Mutex<HashMap<String, u64>> {
TEST_LLM_ROUTER_WALLET_BALANCES.get_or_init(|| Mutex::new(HashMap::new()))
}
pub async fn proxy_llm_chat_completions(
State(state): State<AppState>,
Extension(request_context): Extension<RequestContext>,
Extension(authenticated): Extension<AuthenticatedAccessToken>,
Json(payload): Json<LlmChatCompletionRequest>,
) -> Result<Response, Response> {
if let Err(error) =
ensure_llm_router_user_can_start_conversation(&state, authenticated.claims().user_id())
.await
{
return Err(llm_error_response(&request_context, error));
}
let (llm_client, key_id) =
match resolve_llm_router_client(&state, authenticated.claims().user_id()).await {
Ok((client, key_id)) => (client, key_id),
Err(error) => {
return Err(llm_error_response(
&request_context,
AppError::from_status(StatusCode::SERVICE_UNAVAILABLE).with_message(error),
));
}
};
prepare_llm_router_billing(&state, authenticated.claims().user_id())
.await
.map_err(|error| llm_error_response(&request_context, error))?;
let api_kind = LlmApiKind::OpenAiResponses;
let request = LlmRunRequest {
model: None,
api_kind,
messages: payload
.messages
.into_iter()
.map(map_chat_message)
.collect::<Vec<_>>(),
responses_input: None,
max_output_tokens: None,
enable_web_search: false,
request_timeout_ms: None,
response_reasoning_effort: None,
response_text_verbosity: None,
capture_reasoning: false,
function_tools: Vec::new(),
tool_choice: None,
};
if payload.stream {
return Ok(stream_llm_chat_completions(
llm_client.clone(),
request,
state.clone(),
authenticated.claims().user_id().to_string(),
key_id,
request_context.request_id().to_string(),
)
.into_response());
}
let response = match llm_client.run(request).await {
Ok(response) => response,
Err(error) => {
revoke_llm_router_key_after_auth_failure(
&state,
authenticated.claims().user_id(),
key_id.as_str(),
&error,
request_context.request_id(),
)
.await;
return Err(llm_error_response(&request_context, map_llm_error(error)));
}
};
if let Err(error) = settle_llm_router_usage(&state, authenticated.claims().user_id()).await {
tracing::error!(
request_id = request_context.request_id(),
user_id = %authenticated.claims().user_id(),
error = %error,
"LLM Router 响应成功但累计额度同步未完成"
);
}
Ok(json_success_body(
Some(&request_context),
LlmChatCompletionResponse {
id: response.response_id,
model: response.model,
content: response.text,
finish_reason: response.finish_reason,
},
)
.into_response())
}
/// Only platform model identifiers and aliases are exposed to the client.
pub async fn list_llm_models(
State(state): State<AppState>,
Extension(request_context): Extension<RequestContext>,
Extension(authenticated): Extension<AuthenticatedAccessToken>,
) -> Result<Response, Response> {
let catalog = load_llm_catalog(&state, authenticated.claims().user_id())
.await
.map_err(|error| llm_error_response(&request_context, error))?;
Ok(json_success_body(Some(&request_context), public_model_catalog(catalog)).into_response())
}
fn public_model_catalog(catalog: module_runtime::AgcModelCatalog) -> LlmModelsResponse {
LlmModelsResponse {
default_model_id: catalog.default_model_id,
models: catalog
.models
.into_iter()
.filter(|model| model.enabled)
.map(|model| LlmModelSummary {
id: model.id,
display_name: model.alias,
})
.collect(),
revision: catalog.revision,
}
}
async fn load_llm_catalog(
state: &AppState,
owner: &str,
) -> Result<module_runtime::AgcModelCatalog, AppError> {
#[cfg(test)]
if test_provisioned_router_credentials()
.lock()
.expect("fixture lock")
.contains_key(owner)
{
return Ok(module_runtime::AgcModelCatalog::default());
}
let _ = owner;
crate::agc_models::load_catalog(state).await
}
/// Proxies the OpenAI-compatible Responses protocol for the LLM Router.
///
/// The caller only presents the platform access token. The Router credential
/// is resolved from the authenticated account inside api-server and is never
/// returned to the client or placed in the request payload.
pub async fn proxy_llm_responses(
State(state): State<AppState>,
Extension(request_context): Extension<RequestContext>,
Extension(authenticated): Extension<AuthenticatedAccessToken>,
headers: HeaderMap,
body: Bytes,
) -> Result<Response, Response> {
if body.len() > LLM_REQUEST_MAX_BODY_BYTES {
return Err(llm_error_response(
&request_context,
AppError::from_status(StatusCode::PAYLOAD_TOO_LARGE)
.with_message("LLM Responses 请求体超过大小限制"),
));
}
if let Err(error) =
ensure_llm_router_user_can_start_conversation(&state, authenticated.claims().user_id())
.await
{
return Err(llm_error_response(&request_context, error));
}
let mut payload = serde_json::from_slice::<Value>(&body).map_err(|_| {
llm_error_response(
&request_context,
AppError::from_status(StatusCode::BAD_REQUEST)
.with_message("LLM Responses 请求体必须是合法 JSON"),
)
})?;
let object = payload.as_object_mut().ok_or_else(|| {
llm_error_response(
&request_context,
AppError::from_status(StatusCode::BAD_REQUEST)
.with_message("LLM Responses 请求体必须是 JSON 对象"),
)
})?;
let requested_model = object
.get("model")
.and_then(Value::as_str)
.map(str::trim)
.filter(|value| !value.is_empty())
.map(str::to_string);
// The LLM Router is an account-owned route. Ignore legacy client/provider controls;
// they must not reach Router even when an older desktop build still sends
// them. Runtime controls such as `stream`, `input`, `tools` and `metadata`
// remain part of the Responses contract.
for field in [
"apiKey",
"api_key",
"baseUrl",
"base_url",
"provider",
"apiKind",
"api_kind",
"agentLlm",
"agent_llm",
"agentMode",
"agent_mode",
] {
object.remove(field);
}
// The AGC client may select a model from the server-provided Router
// directory. Older callers without the reserved marker remain pinned to
// the official default model.
let agc_client = headers
.get("x-genarrative-client")
.and_then(|value| value.to_str().ok())
.is_some_and(|value| value == "agc");
let catalog = load_llm_catalog(&state, authenticated.claims().user_id())
.await
.map_err(|error| llm_error_response(&request_context, error))?;
let selected_id = if agc_client {
requested_model
.as_deref()
.filter(|id| *id != "platform-default")
} else {
None
}
.unwrap_or(&catalog.default_model_id);
let selected_model = catalog
.resolve(selected_id)
.map_err(|message| {
llm_error_response(
&request_context,
AppError::from_status(StatusCode::UNPROCESSABLE_ENTITY).with_message(message),
)
})?
.to_string();
object.insert("model".to_string(), Value::String(selected_model));
let (base_url, api_key, key_id) =
resolve_llm_router_credentials(&state, authenticated.claims().user_id())
.await
.map_err(|error| {
llm_error_response(
&request_context,
AppError::from_status(StatusCode::SERVICE_UNAVAILABLE).with_message(error),
)
})?;
prepare_llm_router_billing(&state, authenticated.claims().user_id())
.await
.map_err(|error| llm_error_response(&request_context, error))?;
let client = reqwest::Client::builder()
.connect_timeout(std::time::Duration::from_secs(15))
.read_timeout(std::time::Duration::from_secs(180))
.redirect(reqwest::redirect::Policy::none())
.build()
.map_err(|error| {
llm_error_response(
&request_context,
AppError::from_status(StatusCode::INTERNAL_SERVER_ERROR)
.with_message(format!("创建 LLM Router 请求客户端失败:{error}")),
)
})?;
let upstream_url = format!("{}/responses", base_url.trim_end_matches('/'));
let mut request = client
.post(upstream_url)
.bearer_auth(api_key)
.header("content-type", "application/json");
if let Some(accept) = headers.get("accept") {
request = request.header("accept", accept);
}
let upstream = request
.body(serde_json::to_vec(&payload).map_err(|error| {
llm_error_response(
&request_context,
AppError::from_status(StatusCode::INTERNAL_SERVER_ERROR)
.with_message(format!("序列化 LLM Responses 请求失败:{error}")),
)
})?)
.send()
.await
.map_err(|error| {
llm_error_response(
&request_context,
AppError::from_status(StatusCode::BAD_GATEWAY)
.with_message(format!("LLM Router 暂时不可用:{error}")),
)
})?;
let status = upstream.status();
if matches!(status, StatusCode::UNAUTHORIZED | StatusCode::FORBIDDEN) {
if let Err(error) = crate::external_api_keys::revoke_llm_router_account(
&state,
authenticated.claims().user_id(),
&key_id,
)
.await
{
tracing::warn!(
request_id = request_context.request_id(),
user_id = %authenticated.claims().user_id(),
key_id = %key_id,
error = %error,
"LLM Router 返回确定鉴权失败,但本地账号 Key 失效标记未完成"
);
}
}
let upstream_headers = upstream.headers().clone();
let is_stream = payload
.get("stream")
.and_then(Value::as_bool)
.unwrap_or(false);
if !is_stream {
let body = upstream.bytes().await.map_err(|error| {
llm_error_response(
&request_context,
AppError::from_status(StatusCode::BAD_GATEWAY)
.with_message(format!("读取 LLM Router 响应失败:{error}")),
)
})?;
if status.is_success() {
if let Err(error) =
settle_llm_router_usage(&state, authenticated.claims().user_id()).await
{
tracing::error!(
request_id = request_context.request_id(),
user_id = %authenticated.claims().user_id(),
error = %error,
"LLM Router Responses 成功但累计额度同步未完成"
);
}
}
return build_upstream_response(
status,
&upstream_headers,
Body::from(body),
&request_context,
);
}
if !status.is_success() {
let stream = upstream.bytes_stream().map(|chunk| {
chunk.map_err(|error| {
std::io::Error::other(format!("LLM Router 响应流读取失败:{error}"))
})
});
return build_upstream_response(
status,
&upstream_headers,
Body::from_stream(stream),
&request_context,
);
}
let stream = stream_responses_with_billing(
upstream.bytes_stream(),
state.clone(),
authenticated.claims().user_id().to_string(),
request_context.request_id().to_string(),
);
build_upstream_response(
status,
&upstream_headers,
Body::from_stream(stream),
&request_context,
)
}
fn build_upstream_response(
status: StatusCode,
upstream_headers: &HeaderMap,
body: Body,
request_context: &RequestContext,
) -> Result<Response, Response> {
let mut response = Response::builder().status(status);
if let Some(response_headers) = response.headers_mut() {
for (name, value) in upstream_headers {
if matches!(
name.as_str().to_ascii_lowercase().as_str(),
"connection"
| "keep-alive"
| "proxy-authenticate"
| "proxy-authorization"
| "te"
| "trailer"
| "transfer-encoding"
| "upgrade"
| "host"
| "content-length"
) {
continue;
}
response_headers.append(name.clone(), value.clone());
}
if !response_headers.contains_key("content-type") {
response_headers.insert(
"content-type",
HeaderValue::from_static("application/json; charset=utf-8"),
);
}
}
response.body(body).map_err(|_| {
llm_error_response(
request_context,
AppError::from_status(StatusCode::BAD_GATEWAY).with_message("LLM Router 响应无法建立"),
)
})
}
async fn ensure_llm_router_user_can_start_conversation(
state: &AppState,
owner_user_id: &str,
) -> Result<(), AppError> {
#[cfg(test)]
if let Some(wallet_balance) = test_llm_router_wallet_balances()
.lock()
.expect("test LLM wallet fixture lock should not poison")
.get(owner_user_id)
.copied()
{
return if wallet_balance == 0 {
Err(insufficient_mud_points_error())
} else {
Ok(())
};
}
#[cfg(test)]
if test_provisioned_router_credentials()
.lock()
.expect("test Router credential fixture lock should not poison")
.contains_key(owner_user_id)
{
return Ok(());
}
let dashboard = state
.spacetime_client()
.get_profile_dashboard(owner_user_id.to_string())
.await
.map_err(|error| {
tracing::warn!(
user_id = %owner_user_id,
error = %error,
"读取用户泥点余额失败,已拒绝发起 LLM 对话"
);
AppError::from_status(StatusCode::SERVICE_UNAVAILABLE)
.with_message("泥点余额暂时不可用,请稍后重试")
})?;
if dashboard.wallet_balance == 0 {
return Err(insufficient_mud_points_error());
}
Ok(())
}
fn insufficient_mud_points_error() -> AppError {
AppError::from_status(StatusCode::CONFLICT)
.with_code("MUD_POINTS_INSUFFICIENT")
.with_message("泥点余额不足")
.with_details(json!({
"reason": "insufficient-mud-points",
}))
}
fn map_llm_router_billing_error(_error: SpacetimeClientError) -> AppError {
AppError::from_status(StatusCode::SERVICE_UNAVAILABLE)
.with_code("LLM_BILLING_FAILED")
.with_message("LLM 额度同步失败,请稍后重试")
}
async fn settle_llm_router_usage(state: &AppState, owner_user_id: &str) -> Result<(), AppError> {
#[cfg(test)]
if test_provisioned_router_credentials()
.lock()
.expect("test Router credential fixture lock should not poison")
.contains_key(owner_user_id)
{
// Transport fixtures have no database; settlement arithmetic is tested in module-runtime.
return Ok(());
}
sync_llm_router_quota(state, owner_user_id)
.await
.map(|_| ())
}
async fn prepare_llm_router_billing(state: &AppState, owner_user_id: &str) -> Result<(), AppError> {
#[cfg(test)]
if test_provisioned_router_credentials()
.lock()
.expect("test Router credential fixture lock should not poison")
.contains_key(owner_user_id)
{
return ensure_llm_router_user_can_start_conversation(state, owner_user_id).await;
}
if sync_llm_router_quota(state, owner_user_id).await? == 0 {
return Err(insufficient_mud_points_error());
}
Ok(())
}
async fn sync_llm_router_quota(state: &AppState, owner_user_id: &str) -> Result<u64, AppError> {
let unavailable = || {
AppError::from_status(StatusCode::SERVICE_UNAVAILABLE)
.with_message("LLM 额度同步失败,请稍后重试")
};
let route = state.config.llm_router_base_url.trim_end_matches('/');
let account = state
.spacetime_client()
.get_llm_router_account(owner_user_id.to_string(), route.to_string())
.await
.map_err(|error| {
tracing::warn!(user_id = %owner_user_id, error = %error, "读取 LLM Router 账号映射失败");
map_llm_router_billing_error(error)
})?
.ok_or_else(unavailable)?;
let user_id =
crate::external_api_keys::router_user_id_from_account(&account).ok_or_else(unavailable)?;
let token = state
.config
.llm_router_admin_token
.as_deref()
.filter(|token| !token.trim().is_empty())
.ok_or_else(unavailable)?;
let origin = route.strip_suffix("/v1").unwrap_or(route);
let client = reqwest::Client::builder()
.redirect(reqwest::redirect::Policy::none())
.connect_timeout(std::time::Duration::from_secs(5))
.timeout(std::time::Duration::from_secs(10))
.build()
.map_err(|_| unavailable())?;
let used_quota =
platform_llm::router_billing::read_router_used_quota(&client, origin, token, user_id)
.await
.map_err(|error| {
tracing::warn!(user_id = %owner_user_id, error = %error, "LLM 累计额度读取失败");
unavailable()
})?;
state
.spacetime_client()
.settle_llm_router_quota(
owner_user_id.to_string(),
route.to_string(),
user_id,
used_quota,
)
.await
.map(|result| result.spendable_points)
.map_err(|error| {
tracing::warn!(user_id = %owner_user_id, error = %error, "提交 LLM 累计额度结算失败");
map_llm_router_billing_error(error)
})
}
fn llm_router_done_sse_event() -> Bytes {
Bytes::from_static(b"data: [DONE]\n\n")
}
fn is_responses_done_sse_event(event: &str) -> bool {
event
.lines()
.any(|line| matches!(line.trim(), "data: [DONE]" | "[DONE]"))
}
fn finalize_llm_router_terminal_event(
terminal_event: &mut Option<String>,
billing_result: Option<&Result<(), AppError>>,
request_id: &str,
owner_user_id: &str,
) -> Option<Bytes> {
let billing_result = billing_result?;
match billing_result {
Ok(()) => terminal_event.take().map(Bytes::from),
Err(error) => {
tracing::error!(
request_id = %request_id,
user_id = %owner_user_id,
error = %error,
"LLM Router 流式响应已完成但累计额度同步未完成"
);
terminal_event.take().map(Bytes::from)
}
}
}
fn stream_responses_with_billing(
mut upstream: impl futures_util::Stream<Item = Result<Bytes, reqwest::Error>> + Unpin,
state: AppState,
owner_user_id: String,
request_id: String,
) -> impl futures_util::Stream<Item = Result<Bytes, std::io::Error>> {
async_stream::stream! {
let mut pending = String::new();
let mut utf8_pending = Vec::new();
let mut terminal_event = None;
let mut completed_seen = false;
let mut billing_result = None;
let mut done_seen = false;
let mut pending_done_event = None;
while let Some(chunk) = upstream.next().await {
match chunk {
Ok(bytes) => {
append_utf8_chunk(&mut pending, &mut utf8_pending, bytes.as_ref());
pending = pending.replace("\r\n", "\n");
while let Some(separator) = pending.find("\n\n") {
let raw_event = pending[..separator + 2].to_string();
let event = pending[..separator].to_string();
pending.drain(..separator + 2);
if is_responses_terminal_sse_event(&event) {
completed_seen = true;
terminal_event = Some(raw_event);
if billing_result.is_none() {
billing_result = Some(settle_llm_router_usage(
&state,
owner_user_id.as_str(),
).await);
}
continue;
}
if is_responses_done_sse_event(&event) {
done_seen = true;
if completed_seen && billing_result.is_none() {
billing_result = Some(settle_llm_router_usage(
&state,
owner_user_id.as_str(),
).await);
}
if terminal_event.is_some() {
if let Some(final_event) = finalize_llm_router_terminal_event(
&mut terminal_event,
billing_result.as_ref(),
request_id.as_str(),
owner_user_id.as_str(),
) {
yield Ok(final_event);
}
}
yield Ok(Bytes::from(raw_event));
continue;
}
yield Ok(Bytes::from(raw_event));
}
}
Err(error) => {
if completed_seen && billing_result.is_none() {
billing_result = Some(settle_llm_router_usage(
&state,
owner_user_id.as_str(),
).await);
}
if terminal_event.is_some() {
if let Some(final_event) = finalize_llm_router_terminal_event(
&mut terminal_event,
billing_result.as_ref(),
request_id.as_str(),
owner_user_id.as_str(),
) {
yield Ok(final_event);
if !done_seen {
yield Ok(llm_router_done_sse_event());
}
}
}
yield Err(std::io::Error::other(format!("LLM Router 响应流读取失败:{error}")));
return;
}
}
}
pending = pending.replace("\r\n", "\n");
if !pending.trim().is_empty() {
let event = pending.trim().to_string();
if is_responses_terminal_sse_event(&event) {
completed_seen = true;
terminal_event = Some(pending.clone());
} else if is_responses_done_sse_event(&event) {
done_seen = true;
pending_done_event = Some(pending.clone());
} else {
yield Ok(Bytes::from(pending.clone()));
}
}
if completed_seen && billing_result.is_none() {
billing_result = Some(settle_llm_router_usage(
&state,
owner_user_id.as_str(),
).await);
} else if !completed_seen {
tracing::warn!(
request_id = %request_id,
user_id = %owner_user_id,
"LLM Router 流式响应未收到 response.completed/incomplete,跳过泥点扣费"
);
}
if terminal_event.is_some() {
if let Some(final_event) = finalize_llm_router_terminal_event(
&mut terminal_event,
billing_result.as_ref(),
request_id.as_str(),
owner_user_id.as_str(),
) {
yield Ok(final_event);
}
}
if let Some(done_event) = pending_done_event {
yield Ok(Bytes::from(done_event));
} else if completed_seen && !done_seen {
yield Ok(llm_router_done_sse_event());
}
}
}
fn append_utf8_chunk(pending: &mut String, carry: &mut Vec<u8>, bytes: &[u8]) {
carry.extend_from_slice(bytes);
loop {
match std::str::from_utf8(carry.as_slice()) {
Ok(text) => {
pending.push_str(text);
carry.clear();
break;
}
Err(error) => {
let valid = error.valid_up_to();
if valid > 0 {
pending.push_str(std::str::from_utf8(&carry[..valid]).unwrap_or_default());
carry.drain(..valid);
}
if let Some(error_len) = error.error_len() {
pending.push('\u{FFFD}');
carry.drain(..error_len.min(carry.len()));
continue;
}
break;
}
}
}
}
fn is_responses_terminal_sse_event(event: &str) -> bool {
if event.lines().any(|line| {
line.strip_prefix("event:")
.map(str::trim)
.is_some_and(|name| matches!(name, "response.completed" | "response.incomplete"))
}) {
return true;
}
let data = event
.lines()
.filter_map(|line| line.strip_prefix("data:").map(str::trim_start))
.collect::<Vec<_>>()
.join("\n");
serde_json::from_str::<Value>(&data)
.ok()
.and_then(|value| {
value
.get("type")
.and_then(Value::as_str)
.map(str::to_string)
})
.is_some_and(|kind| matches!(kind.as_str(), "response.completed" | "response.incomplete"))
}
const SSE_DONE_MARKER: &[u8] = b"data: [DONE]";
fn find_sse_done_marker(bytes: &[u8]) -> Option<usize> {
bytes
.windows(SSE_DONE_MARKER.len())
.enumerate()
.find_map(|(position, window)| {
(window == SSE_DONE_MARKER && (position == 0 || bytes[position - 1] == b'\n'))
.then_some(position)
})
}
fn find_sse_event_end(bytes: &[u8], event_start: usize) -> Option<usize> {
let event = &bytes[event_start..];
if let Some(offset) = event.windows(2).position(|window| window == b"\n\n") {
return Some(event_start + offset + 2);
}
event
.windows(4)
.position(|window| window == b"\r\n\r\n")
.map(|offset| event_start + offset + 4)
}
fn sse_done_marker_suffix_len(bytes: &[u8]) -> usize {
(1..SSE_DONE_MARKER.len())
.rev()
.find(|&length| bytes.ends_with(&SSE_DONE_MARKER[..length]))
.unwrap_or(0)
}
async fn resolve_llm_router_client(
state: &AppState,
owner_user_id: &str,
) -> Result<(platform_llm::LlmClient, String), String> {
let (base_url, api_key, key_id) = resolve_llm_router_credentials(state, owner_user_id).await?;
let catalog = load_llm_catalog(state, owner_user_id)
.await
.map_err(|_| "模型目录暂不可用".to_string())?;
let model = catalog.resolve(&catalog.default_model_id)?;
let config = platform_llm::LlmConfig::new(
platform_llm::LlmProvider::OpenAiCompatible,
base_url.to_string(),
api_key,
model.to_string(),
state.config.llm_request_timeout_ms,
state.config.llm_max_retries,
state.config.llm_retry_backoff_ms,
)
.map_err(|error| error.to_string())?;
platform_llm::LlmClient::new(config)
.map(|client| (client, key_id))
.map_err(|error| error.to_string())
}
async fn resolve_llm_router_credentials(
state: &AppState,
owner_user_id: &str,
) -> Result<(String, String, String), String> {
#[cfg(test)]
if let Some(fixture) = test_provisioned_router_credentials()
.lock()
.expect("test Router credential fixture lock should not poison")
.get(owner_user_id)
.cloned()
{
if !test_router_fixture_base_url_is_loopback(fixture.base_url.as_str()) {
return Err("LLM Router 测试账号只允许配合 loopback 地址使用".to_string());
}
if fixture.base_url.trim_end_matches('/')
!= state.config.llm_router_base_url.trim_end_matches('/')
{
return Err("LLM Router 测试账号路由与当前配置不一致".to_string());
}
if fixture.api_key.trim().is_empty() || fixture.key_id.trim().is_empty() {
return Err("LLM Router 测试账号凭据不完整".to_string());
}
// This is an explicit, owner-scoped pre-provisioned account fixture,
// not a fallback credential. Release binaries do not compile this
// branch and always require a validated SpacetimeDB row.
return Ok((
fixture.base_url.trim_end_matches('/').to_string(),
fixture.api_key.trim().to_string(),
fixture.key_id,
));
}
// No fixture means the request must follow the same provisioning/read path
// as production. The dedicated llm_router_account row is authoritative.
if let Some(credentials) =
crate::external_api_keys::read_active_llm_router_credentials(state, owner_user_id).await?
{
return Ok(credentials);
}
crate::external_api_keys::ensure_llm_router_account(state, owner_user_id).await?;
crate::external_api_keys::read_active_llm_router_credentials(state, owner_user_id)
.await?
.ok_or_else(|| "LLM Router 账号密钥缺失,请重新登录后重试".to_string())
}
#[cfg(test)]
fn test_router_fixture_base_url_is_loopback(value: &str) -> bool {
reqwest::Url::parse(value)
.ok()
.and_then(|url| url.host_str().map(str::to_string))
.is_some_and(|host| {
host.eq_ignore_ascii_case("localhost")
|| host
.parse::<std::net::IpAddr>()
.is_ok_and(|address| address.is_loopback())
})
}
fn stream_llm_chat_completions(
llm_client: platform_llm::LlmClient,
request: LlmRunRequest,
state: AppState,
owner_user_id: String,
key_id: String,
request_id: String,
) -> Sse<impl tokio_stream::Stream<Item = Result<Event, Infallible>>> {
let stream = async_stream::stream! {
let (delta_tx, mut delta_rx) = tokio::sync::mpsc::unbounded_channel::<Value>();
let llm_stream = llm_client.stream_run(request, move |delta| {
let _ = delta_tx.send(json!({
"delta": delta.delta_text,
"content": delta.accumulated_text,
"finishReason": delta.finish_reason,
}));
});
tokio::pin!(llm_stream);
let llm_result = loop {
// `platform-llm` 负责上游 SSE 解析;这里尽快把增量转成 API 层 SSE 事件。
tokio::select! {
result = &mut llm_stream => break result,
maybe_delta = delta_rx.recv() => {
if let Some(delta) = maybe_delta {
yield Ok::<Event, Infallible>(llm_sse_json_event_or_error("delta", delta));
}
}
}
};
while let Some(delta) = delta_rx.recv().await {
yield Ok::<Event, Infallible>(llm_sse_json_event_or_error("delta", delta));
}
match llm_result {
Ok(response) => {
if let Err(error) = settle_llm_router_usage(
&state,
owner_user_id.as_str(),
)
.await
{
tracing::error!(
request_id = %request_id,
user_id = %owner_user_id,
error = %error,
"LLM Router Chat 流式响应已完成但累计额度同步未完成"
);
}
yield Ok::<Event, Infallible>(llm_sse_json_event_or_error(
"complete",
json!(LlmChatCompletionResponse {
id: response.response_id,
model: response.model,
content: response.text,
finish_reason: response.finish_reason,
}),
));
}
Err(error) => {
revoke_llm_router_key_after_auth_failure(
&state,
owner_user_id.as_str(),
key_id.as_str(),
&error,
"stream",
)
.await;
let app_error = map_llm_error(error);
yield Ok::<Event, Infallible>(llm_sse_json_event_or_error(
"error",
json!({
"code": app_error.code(),
"message": app_error.message(),
}),
));
}
}
yield Ok::<Event, Infallible>(Event::default().data("[DONE]"));
};
Sse::new(stream)
}
async fn revoke_llm_router_key_after_auth_failure(
state: &AppState,
owner_user_id: &str,
key_id: &str,
error: &platform_llm::LlmError,
request_id: &str,
) {
if !matches!(
error,
platform_llm::LlmError::Upstream {
status_code: 401 | 403,
..
}
) {
return;
}
if let Err(revoke_error) =
crate::external_api_keys::revoke_llm_router_account(state, owner_user_id, key_id).await
{
tracing::warn!(
request_id,
user_id = %owner_user_id,
key_id,
error = %revoke_error,
"LLM Router 返回确定鉴权失败,但账号 Key 失效标记未完成"
);
}
}
fn llm_sse_json_event_or_error(event_name: &str, payload: Value) -> Event {
match serde_json::to_string(&payload) {
Ok(payload_text) => Event::default().event(event_name).data(payload_text),
Err(_) => Event::default()
.event("error")
.data("{\"code\":\"INTERNAL_SERVER_ERROR\",\"message\":\"SSE payload 序列化失败\"}"),
}
}
fn map_chat_message(message: LlmChatMessagePayload) -> LlmMessage {
let role = match message.role {
LlmChatMessageRole::System => LlmMessageRole::System,
LlmChatMessageRole::User => LlmMessageRole::User,
LlmChatMessageRole::Assistant => LlmMessageRole::Assistant,
};
LlmMessage::new(role, message.content)
}
fn llm_error_response(request_context: &RequestContext, error: AppError) -> Response {
error.into_response_with_context(Some(request_context))
}
#[cfg(test)]
mod tests {
use super::*;
use std::{
io::{Read, Write},
net::TcpListener,
sync::atomic::{AtomicUsize, Ordering},
sync::{Arc, Mutex},
thread,
time::Duration as StdDuration,
};
use axum::{
body::Body,
http::{Request, StatusCode},
};
use http_body_util::BodyExt;
use platform_auth::{
AccessTokenClaims, AccessTokenClaimsInput, AuthProvider, BindingStatus, sign_access_token,
};
use serde_json::{Value, json};
use time::OffsetDateTime;
use tower::ServiceExt;
use crate::{app::build_router, config::AppConfig, state::AppState};
struct MockResponse {
status_line: &'static str,
content_type: &'static str,
body: String,
extra_headers: Vec<(&'static str, &'static str)>,
}
static NEXT_TEST_PHONE: AtomicUsize = AtomicUsize::new(1);
#[tokio::test]
async fn llm_chat_completions_returns_non_stream_run_payload() {
let server_url = spawn_mock_server(vec![MockResponse {
status_line: "200 OK",
content_type: "application/json; charset=utf-8",
body: r#"{"id":"resp_api_server_01","model":"ark-router-test","status":"completed","output":[{"type":"message","content":[{"type":"output_text","text":"代理成功"}]}]}"#.to_string(),
extra_headers: Vec::new(),
}]);
let (state, user_id) = seed_authenticated_state(AppConfig {
llm_router_base_url: server_url.clone(),
..AppConfig::default()
})
.await;
install_test_provisioned_router_credential(&user_id, server_url, "test-key");
let token = issue_access_token(&state, &user_id);
let app = build_router(state);
let response = app
.oneshot(
Request::builder()
.method("POST")
.uri("/api/llm/chat/completions")
.header("authorization", format!("Bearer {token}"))
.header("content-type", "application/json")
.header("x-genarrative-response-envelope", "v1")
.body(Body::from(
json!({
"messages": [
{ "role": "system", "content": "系统" },
{ "role": "user", "content": "用户" }
]
})
.to_string(),
))
.expect("request should build"),
)
.await
.expect("request should succeed");
assert_eq!(response.status(), StatusCode::OK);
let body = response
.into_body()
.collect()
.await
.expect("body should collect")
.to_bytes();
let payload: Value =
serde_json::from_slice(&body).expect("response body should be valid json");
assert_eq!(payload["ok"], Value::Bool(true));
assert_eq!(
payload["data"]["id"],
Value::String("resp_api_server_01".to_string())
);
assert_eq!(
payload["data"]["model"],
Value::String("ark-router-test".to_string())
);
assert_eq!(
payload["data"]["content"],
Value::String("代理成功".to_string())
);
assert_eq!(
payload["data"]["finishReason"],
Value::String("completed".to_string())
);
}
#[tokio::test]
async fn llm_routes_accept_large_context_bodies_beyond_axum_default() {
let large_input = "x".repeat(2 * 1024 * 1024 + 1024);
let (state, user_id) = seed_authenticated_state(AppConfig::default()).await;
install_test_provisioned_router_credential(
&user_id,
"http://127.0.0.1:1".to_string(),
"fixture-key",
);
let token = issue_access_token(&state, &user_id);
let app = build_router(state);
let response = app
.clone()
.oneshot(
Request::builder()
.method("POST")
.uri("/api/llm/responses")
.header("authorization", format!("Bearer {token}"))
.header("content-type", "application/json")
.body(Body::from(json!({ "input": large_input }).to_string()))
.expect("request should build"),
)
.await
.expect("response should not hit the default body limit");
assert_ne!(response.status(), StatusCode::PAYLOAD_TOO_LARGE);
let response = app
.oneshot(
Request::builder()
.method("POST")
.uri("/api/llm/chat/completions")
.header("authorization", format!("Bearer {token}"))
.header("content-type", "application/json")
.body(Body::from(
json!({
"messages": [
{ "role": "user", "content": large_input }
]
})
.to_string(),
))
.expect("request should build"),
)
.await
.expect("response should not hit the default body limit");
assert_ne!(response.status(), StatusCode::PAYLOAD_TOO_LARGE);
}
#[tokio::test]
async fn llm_responses_rejects_bodies_above_explicit_limit() {
let (state, user_id) = seed_authenticated_state(AppConfig::default()).await;
let token = issue_access_token(&state, &user_id);
let app = build_router(state);
let response = app
.oneshot(
Request::builder()
.method("POST")
.uri("/api/llm/responses")
.header("authorization", format!("Bearer {token}"))
.header("content-type", "application/json")
.body(Body::from(vec![b'x'; LLM_REQUEST_MAX_BODY_BYTES + 1]))
.expect("request should build"),
)
.await
.expect("oversized response should be returned");
assert_eq!(response.status(), StatusCode::PAYLOAD_TOO_LARGE);
}
#[tokio::test]
async fn llm_chat_completions_streams_sse_payload() {
let server_url = spawn_mock_server(vec![MockResponse {
status_line: "200 OK",
content_type: "text/event-stream; charset=utf-8",
body: concat!(
"data: {\"type\":\"response.output_text.delta\",\"delta\":\"你\"}\n\n",
"data: {\"type\":\"response.output_text.delta\",\"delta\":\"好\"}\n\n",
"data: {\"type\":\"response.completed\"}\n\n"
)
.to_string(),
extra_headers: vec![("x-request-id", "req_llm_stream_01")],
}]);
let (state, user_id) = seed_authenticated_state(AppConfig {
llm_router_base_url: server_url.clone(),
..AppConfig::default()
})
.await;
install_test_provisioned_router_credential(&user_id, server_url, "test-key");
let token = issue_access_token(&state, &user_id);
let app = build_router(state);
let response = app
.oneshot(
Request::builder()
.method("POST")
.uri("/api/llm/chat/completions")
.header("authorization", format!("Bearer {token}"))
.header("content-type", "application/json")
.body(Body::from(
json!({
"stream": true,
"messages": [
{ "role": "user", "content": "用户" }
]
})
.to_string(),
))
.expect("request should build"),
)
.await
.expect("request should succeed");
assert_eq!(response.status(), StatusCode::OK);
assert_eq!(
response
.headers()
.get("content-type")
.and_then(|value| value.to_str().ok()),
Some("text/event-stream")
);
let body = response
.into_body()
.collect()
.await
.expect("body should collect")
.to_bytes();
let body_text = String::from_utf8(body.to_vec()).expect("body should be utf8");
assert!(body_text.contains("event: delta"));
assert!(body_text.contains(r#""delta":"你""#));
assert!(body_text.contains(r#""content":"你好""#));
assert!(body_text.contains("event: complete"));
assert!(body_text.contains(r#""id":"req_llm_stream_01""#));
assert!(body_text.contains(r#""finishReason":"completed""#));
assert!(body_text.contains("data: [DONE]"));
}
#[tokio::test]
async fn llm_responses_proxy_forces_official_model_and_keeps_router_key_server_side() {
let (server_url, captured_request) = spawn_capturing_mock_server(MockResponse {
status_line: "200 OK",
content_type: "application/json; charset=utf-8",
body: r#"{"id":"resp_proxy_01","model":"gpt-6-astra","output":[]}"#.to_string(),
extra_headers: Vec::new(),
});
let (state, user_id) = seed_authenticated_state(AppConfig {
llm_router_base_url: server_url.clone(),
llm_router_api_key_encryption_secret: Some("fixture-encryption-secret".to_string()),
..AppConfig::default()
})
.await;
install_test_provisioned_router_credential(&user_id, server_url, "fixture-router-key");
let token = issue_access_token(&state, &user_id);
let app = build_router(state);
let response = app
.oneshot(
Request::builder()
.method("POST")
.uri("/api/llm/responses")
.header("authorization", format!("Bearer {token}"))
.header("content-type", "application/json")
.body(Body::from(
json!({
"model": "client-must-not-control",
"input": "hello",
"stream": false
})
.to_string(),
))
.expect("request should build"),
)
.await
.expect("request should succeed");
let response_status = response.status();
let body = response
.into_body()
.collect()
.await
.expect("body should collect")
.to_bytes();
assert_eq!(response_status, StatusCode::OK);
let payload: Value = serde_json::from_slice(&body).expect("response body should be json");
assert_eq!(payload["id"], "resp_proxy_01");
let upstream_request = captured_request
.lock()
.expect("captured request lock")
.clone()
.expect("mock server should capture upstream request");
assert!(upstream_request.starts_with("POST /responses HTTP/1.1"));
assert!(
upstream_request
.lines()
.any(|line| line.eq_ignore_ascii_case("authorization: Bearer fixture-router-key")),
"upstream must receive the server-side Router key"
);
assert!(
!upstream_request.contains(token.as_str()),
"platform access token must not be forwarded to Router"
);
let (_, upstream_body) = upstream_request
.split_once("\r\n\r\n")
.expect("upstream request body");
let upstream_payload: Value =
serde_json::from_str(upstream_body).expect("upstream body should be json");
assert_eq!(upstream_payload["model"], "gpt-6-astra");
assert_ne!(upstream_payload["model"], "client-must-not-control");
}
#[tokio::test]
async fn llm_responses_rejects_upstream_names_and_unknown_catalog_ids() {
let (state, user_id) = seed_authenticated_state(AppConfig::default()).await;
install_test_provisioned_router_credential(
&user_id,
"http://127.0.0.1:1".into(),
"fixture-key",
);
let token = issue_access_token(&state, &user_id);
let app = build_router(state);
for model in ["gpt-6-astra", "unlisted"] {
let response = app
.clone()
.oneshot(
Request::builder()
.method("POST")
.uri("/api/llm/responses")
.header("authorization", format!("Bearer {token}"))
.header("x-genarrative-client", "agc")
.header("content-type", "application/json")
.body(Body::from(
json!({"model": model, "input": "hello"}).to_string(),
))
.unwrap(),
)
.await
.unwrap();
assert_eq!(response.status(), StatusCode::UNPROCESSABLE_ENTITY);
}
}
#[tokio::test]
async fn llm_responses_never_uses_legacy_llm_api_key_without_router_provisioning() {
let (state, user_id) = seed_authenticated_state(AppConfig {
llm_router_base_url: "https://127.0.0.1:1/v1".to_string(),
llm_router_admin_token: None,
llm_router_api_key_encryption_secret: Some("fixture-encryption-secret".to_string()),
llm_base_url: "http://127.0.0.1:1/v1".to_string(),
llm_api_key: Some("legacy-key-must-not-be-used".to_string()),
llm_model: "legacy-model-must-not-be-used".to_string(),
..AppConfig::default()
})
.await;
let token = issue_access_token(&state, &user_id);
let app = build_router(state);
let response = app
.oneshot(
Request::builder()
.method("POST")
.uri("/api/llm/responses")
.header("authorization", format!("Bearer {token}"))
.header("content-type", "application/json")
.body(Body::from(json!({"input":"hello"}).to_string()))
.expect("request should build"),
)
.await
.expect("provisioning response should be returned");
assert_eq!(response.status(), StatusCode::SERVICE_UNAVAILABLE);
}
#[tokio::test]
async fn llm_responses_rejects_zero_wallet_before_router_provisioning() {
let (state, user_id) = seed_authenticated_state(AppConfig {
llm_router_base_url: "https://127.0.0.1:1/v1".to_string(),
llm_router_api_key_encryption_secret: Some("fixture-encryption-secret".to_string()),
..AppConfig::default()
})
.await;
test_llm_router_wallet_balances()
.lock()
.expect("test LLM wallet fixture lock should not poison")
.insert(user_id.clone(), 0);
let token = issue_access_token(&state, &user_id);
let app = build_router(state);
let response = app
.oneshot(
Request::builder()
.method("POST")
.uri("/api/llm/responses")
.header("authorization", format!("Bearer {token}"))
.header("content-type", "application/json")
.body(Body::from(json!({"input":"hello"}).to_string()))
.expect("request should build"),
)
.await
.expect("request should succeed");
assert_eq!(response.status(), StatusCode::CONFLICT);
let body = response
.into_body()
.collect()
.await
.expect("body should collect")
.to_bytes();
let payload: Value =
serde_json::from_slice(&body).expect("response body should be valid json");
assert_eq!(
payload["error"]["code"],
Value::String("MUD_POINTS_INSUFFICIENT".to_string())
);
assert_eq!(
payload["error"]["message"],
Value::String("泥点余额不足".to_string())
);
}
async fn seed_authenticated_state(config: AppConfig) -> (AppState, String) {
let state = AppState::new(config).expect("state should build");
let phone = format!(
"1380013{:04}",
NEXT_TEST_PHONE.fetch_add(1, Ordering::Relaxed)
);
let user_id = state
.seed_test_phone_user_with_password(phone.as_str(), "secret123")
.await
.id;
(state, user_id)
}
fn install_test_provisioned_router_credential(
owner_user_id: &str,
base_url: String,
api_key: &str,
) {
test_provisioned_router_credentials()
.lock()
.expect("test Router credential fixture lock should not poison")
.insert(
owner_user_id.to_string(),
TestProvisionedRouterCredential {
base_url,
api_key: api_key.to_string(),
key_id: format!("test-llm-router-key-{owner_user_id}"),
},
);
}
fn issue_access_token(state: &AppState, user_id: &str) -> String {
let claims = AccessTokenClaims::from_input(
AccessTokenClaimsInput {
user_id: user_id.to_string(),
session_id: state.seed_test_refresh_session_for_user_id(user_id, "sess_llm_proxy"),
provider: AuthProvider::Password,
roles: vec!["user".to_string()],
token_version: 2,
phone_verified: true,
binding_status: BindingStatus::Active,
display_name: Some("LLM 代理用户".to_string()),
},
state.auth_jwt_config(),
OffsetDateTime::now_utc(),
)
.expect("claims should build");
sign_access_token(&claims, state.auth_jwt_config()).expect("token should sign")
}
#[test]
fn llm_router_stream_preserves_success_on_any_settlement_failure() {
let original = "event: response.completed\ndata: {\"type\":\"response.completed\"}\n\n";
for billing in [
Ok(()),
Err(insufficient_mud_points_error()),
Err(map_llm_router_billing_error(
SpacetimeClientError::ConnectDropped,
)),
] {
let mut terminal = Some(original.to_string());
let event = finalize_llm_router_terminal_event(
&mut terminal,
Some(&billing),
"request-1",
"user-1",
)
.expect("completed output must remain available");
assert_eq!(event.as_ref(), original.as_bytes());
assert!(terminal.is_none());
}
}
#[test]
fn responses_sse_done_marker_is_found_only_at_event_line_start() {
let bytes = b"data: {\"text\":\"data: [DONE]\"}\n\ndata: [DONE]\n\n";
let done_start = find_sse_done_marker(bytes).expect("done event should be found");
assert_eq!(
&bytes[done_start..done_start + SSE_DONE_MARKER.len()],
SSE_DONE_MARKER
);
assert_eq!(find_sse_event_end(bytes, done_start), Some(bytes.len()));
assert_eq!(sse_done_marker_suffix_len(b"data: [DON"), 10);
}
#[test]
fn append_utf8_chunk_preserves_multibyte_characters_split_across_chunks() {
let mut text = String::new();
let mut carry = Vec::new();
append_utf8_chunk(&mut text, &mut carry, &[0xe4, 0xb8]);
assert_eq!(text, "");
append_utf8_chunk(&mut text, &mut carry, &[0xad]);
assert_eq!(text, "中");
assert!(carry.is_empty());
}
fn spawn_mock_server(responses: Vec<MockResponse>) -> String {
let listener = TcpListener::bind("127.0.0.1:0").expect("listener should bind");
let address = listener.local_addr().expect("listener should have addr");
thread::spawn(move || {
for response in responses {
let (mut stream, _) = listener.accept().expect("request should connect");
let _ = read_request(&mut stream);
write_response(&mut stream, response);
}
});
format!("http://{address}")
}
fn spawn_capturing_mock_server(response: MockResponse) -> (String, Arc<Mutex<Option<String>>>) {
let listener = TcpListener::bind("127.0.0.1:0").expect("listener should bind");
let address = listener.local_addr().expect("listener should have addr");
let captured = Arc::new(Mutex::new(None));
let captured_for_thread = Arc::clone(&captured);
thread::spawn(move || {
let (mut stream, _) = listener.accept().expect("request should connect");
let request = read_request(&mut stream);
*captured_for_thread.lock().expect("captured request lock") = Some(request);
write_response(&mut stream, response);
});
(format!("http://{address}"), captured)
}
fn read_request(stream: &mut std::net::TcpStream) -> String {
stream
.set_read_timeout(Some(StdDuration::from_secs(1)))
.expect("read timeout should be set");
let mut buffer = Vec::new();
let mut chunk = [0_u8; 1024];
let mut expected_total = None;
loop {
match stream.read(&mut chunk) {
Ok(0) => break,
Ok(bytes_read) => {
buffer.extend_from_slice(&chunk[..bytes_read]);
if expected_total.is_none()
&& let Some(header_end) = find_header_end(&buffer)
{
let content_length =
read_content_length(&buffer[..header_end]).unwrap_or(0);
expected_total = Some(header_end + content_length);
}
if let Some(total_bytes) = expected_total
&& buffer.len() >= total_bytes
{
break;
}
}
Err(error)
if error.kind() == std::io::ErrorKind::WouldBlock
|| error.kind() == std::io::ErrorKind::TimedOut =>
{
break;
}
Err(error) => panic!("mock server failed to read request: {error}"),
}
}
String::from_utf8_lossy(&buffer).into_owned()
}
fn write_response(stream: &mut std::net::TcpStream, response: MockResponse) {
let body = response.body;
let mut raw_response = format!(
"HTTP/1.1 {}\r\nContent-Type: {}\r\nContent-Length: {}\r\nConnection: close\r\n",
response.status_line,
response.content_type,
body.len()
);
for (name, value) in response.extra_headers {
raw_response.push_str(format!("{name}: {value}\r\n").as_str());
}
raw_response.push_str("\r\n");
raw_response.push_str(body.as_str());
stream
.write_all(raw_response.as_bytes())
.expect("mock response should be written");
stream.flush().expect("mock response should flush");
}
fn find_header_end(buffer: &[u8]) -> Option<usize> {
buffer
.windows(4)
.position(|window| window == b"\r\n\r\n")
.map(|index| index + 4)
}
fn read_content_length(headers: &[u8]) -> Option<usize> {
let text = String::from_utf8_lossy(headers);
text.lines().find_map(|line| {
let (name, value) = line.split_once(':')?;
if name.eq_ignore_ascii_case("content-length") {
return value.trim().parse::<usize>().ok();
}
None
})
}
}