Files
Genarrative/server-rs/crates/api-server/src/llm/mod.rs
T
suzmii 4f3f0f24ff 接入AGC后台模型目录与对话模型选择
新增后台 AGC 模型目录、别名、启停和默认项管理

客户端设置页恢复原状,对话框右下角按别名选择模型

服务端按稳定模型标识映射并校验实际模型白名单

修复 AGC 配套后端端口漂移、启动等待和 SpacetimeDB 版本检查

补充迁移、文档、启动与模型选择测试
2026-09-05 19:10:49 +08:00

1962 lines
70 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, LlmTokenUsage};
use serde_json::{Value, json};
use sha2::{Digest, Sha256};
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) 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.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!(!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>,
headers: HeaderMap,
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),
));
}
};
let api_kind = LlmApiKind::OpenAiResponses;
let request_fingerprint = serde_json::to_vec(&payload).unwrap_or_default();
let request = LlmRunRequest {
model: None,
api_kind,
messages: payload
.messages
.into_iter()
.map(map_chat_message)
.collect::<Vec<_>>(),
max_output_tokens: None,
enable_web_search: false,
request_timeout_ms: None,
response_reasoning_effort: None,
response_text_verbosity: None,
function_tools: Vec::new(),
tool_choice: None,
};
if payload.stream {
let billing_key = request_billing_key(
&headers,
request_context.request_id(),
"chat-completions",
request_fingerprint.as_slice(),
);
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(),
billing_key,
)
.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(),
request_context.request_id(),
request_billing_key(
&headers,
request_context.request_id(),
"chat-completions",
request_fingerprint.as_slice(),
),
response.usage.as_ref(),
)
.await
{
if is_mud_points_insufficient_app_error(&error) {
return Err(llm_error_response(&request_context, error));
}
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(),
}
}
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> {
const MAX_REQUEST_BYTES: usize = 32 * 1024 * 1024;
if body.len() > MAX_REQUEST_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),
)
})?;
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 request_fingerprint = serde_json::to_vec(&payload).unwrap_or_default();
let billing_key = request_billing_key(
&headers,
request_context.request_id(),
"responses",
request_fingerprint.as_slice(),
);
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}")),
)
})?;
let usage = extract_llm_usage_from_json_bytes(&body);
if status.is_success() {
if let Err(error) = settle_llm_router_usage(
&state,
authenticated.claims().user_id(),
request_context.request_id(),
billing_key,
usage.as_ref(),
)
.await
{
if is_mud_points_insufficient_app_error(&error) {
return Err(llm_error_response(&request_context, error));
}
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(),
billing_key,
);
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 request_billing_key(
headers: &HeaderMap,
fallback: &str,
endpoint: &str,
request_bytes: &[u8],
) -> String {
let client_key = headers
.get("idempotency-key")
.or_else(|| headers.get("x-idempotency-key"))
.and_then(|value| value.to_str().ok())
.map(str::trim)
.filter(|value| {
!value.is_empty()
&& value.len() <= 256
&& value.bytes().all(|byte| (0x21..=0x7e).contains(&byte))
});
let Some(client_key) = client_key else {
return fallback.to_string();
};
let mut hasher = Sha256::new();
hasher.update(b"llm-router-billing-key:v2\n");
hasher.update(endpoint.as_bytes());
hasher.update(b"\n");
hasher.update(client_key.as_bytes());
hasher.update(b"\n");
hasher.update(request_bytes);
hex::encode(hasher.finalize())
}
fn llm_router_ledger_id(owner_user_id: &str, idempotency_key: &str) -> String {
let seed = format!(
"llm-router\n{}\n{}",
owner_user_id.trim(),
idempotency_key.trim()
);
let digest = Sha256::digest(seed.as_bytes());
// Reuse the existing wallet consume procedure contract. It accepts the
// asset-operation consume namespace and records the LLM Router operation in
// metadata, so no parallel billing reducer/table is introduced.
format!("asset_operation_consume:llm-router-{}", hex::encode(digest))
}
fn llm_router_points_for_usage(usage: Option<&LlmTokenUsage>) -> u64 {
let total_tokens = usage.map(|value| value.total_tokens).unwrap_or(0);
// Temporary product default until Router pricing is wired to the account
// service: one mud point per started 10,000 tokens, with a one-point minimum
// for a successful response whose gateway omitted usage. Keep the unit in
// one constant so the product can tune it without changing the ledger
// semantics or idempotency contract.
const LLM_ROUTER_BILLING_TOKEN_UNIT: u64 = 10_000;
total_tokens
.saturating_add(LLM_ROUTER_BILLING_TOKEN_UNIT - 1)
.checked_div(LLM_ROUTER_BILLING_TOKEN_UNIT)
.unwrap_or(1)
.max(1)
}
fn is_insufficient_mud_points_error(error: &SpacetimeClientError) -> bool {
match error {
SpacetimeClientError::Procedure(message)
| SpacetimeClientError::Runtime(message)
| SpacetimeClientError::Build(message) => {
message.contains("泥点余额不足") || message.contains("可消费泥点不足:")
}
SpacetimeClientError::ConnectDropped | SpacetimeClientError::Timeout(_) => false,
}
}
fn map_llm_router_billing_error(error: SpacetimeClientError) -> AppError {
if is_insufficient_mud_points_error(&error) {
return AppError::from_status(StatusCode::CONFLICT)
.with_code("MUD_POINTS_INSUFFICIENT")
.with_message("泥点余额不足")
.with_details(json!({
"reason": "insufficient-mud-points",
}));
}
AppError::from_status(StatusCode::SERVICE_UNAVAILABLE)
.with_code("LLM_BILLING_FAILED")
.with_message("LLM 已返回,但泥点扣费未完成")
}
fn is_mud_points_insufficient_app_error(error: &AppError) -> bool {
error.code() == "MUD_POINTS_INSUFFICIENT"
}
async fn settle_llm_router_usage(
state: &AppState,
owner_user_id: &str,
request_id: &str,
idempotency_key: String,
usage: Option<&LlmTokenUsage>,
) -> Result<(), AppError> {
#[cfg(test)]
if test_provisioned_router_credentials()
.lock()
.expect("test Router credential fixture lock should not poison")
.contains_key(owner_user_id)
{
// Existing HTTP proxy tests intentionally use an in-process Router
// fixture without a wallet procedure. Keep those transport tests
// focused; billing arithmetic and ledger-id behavior are covered by
// the pure unit tests below.
return Ok(());
}
let points = llm_router_points_for_usage(usage);
let ledger_id = llm_router_ledger_id(owner_user_id, idempotency_key.as_str());
let usage_json = usage.map(|value| {
json!({
"inputTokens": value.prompt_tokens,
"outputTokens": value.completion_tokens,
"totalTokens": value.total_tokens,
})
});
let metadata = json!({
"operation": "llm-router",
"billingMode": "llm-router-best-effort",
"requestedPoints": points,
"requestId": request_id,
"idempotencyKey": idempotency_key,
"usage": usage_json,
"billingRule": "temporary-1-point-per-started-10000-tokens",
});
state
.spacetime_client()
.consume_profile_wallet_points_with_metadata(
owner_user_id.to_string(),
points,
ledger_id,
crate::editor_project::current_utc_micros(),
metadata.to_string(),
)
.await
.map(|_| ())
.map_err(map_llm_router_billing_error)
}
fn extract_llm_usage_from_json_bytes(body: &[u8]) -> Option<LlmTokenUsage> {
let value = serde_json::from_slice::<Value>(body).ok()?;
extract_llm_usage_from_value(&value)
}
fn extract_llm_usage_from_value(value: &Value) -> Option<LlmTokenUsage> {
let usage = value.get("usage").or_else(|| {
value
.get("response")
.and_then(|response| response.get("usage"))
})?;
let prompt_tokens = usage
.get("input_tokens")
.or_else(|| usage.get("prompt_tokens"))
.and_then(Value::as_u64)
.unwrap_or(0);
let completion_tokens = usage
.get("output_tokens")
.or_else(|| usage.get("completion_tokens"))
.and_then(Value::as_u64)
.unwrap_or(0);
let total_tokens = usage
.get("total_tokens")
.and_then(Value::as_u64)
.unwrap_or_else(|| prompt_tokens.saturating_add(completion_tokens));
Some(LlmTokenUsage {
prompt_tokens,
completion_tokens,
total_tokens,
})
}
fn llm_router_insufficient_mud_points_sse_event() -> Bytes {
let error = json!({
"code": "insufficient_mud_points",
"message": "泥点余额不足",
});
let payload = json!({
"type": "response.failed",
"error": error,
"response": {
"status": "failed",
"error": {
"code": "insufficient_mud_points",
"message": "泥点余额不足",
},
},
});
Bytes::from(format!("event: response.failed\ndata: {}\n\n", payload))
}
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) if is_mud_points_insufficient_app_error(error) => {
terminal_event.take();
Some(llm_router_insufficient_mud_points_sse_event())
}
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,
idempotency_key: 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 usage = None;
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 let Some(event_usage) = extract_llm_usage_from_sse_event(&event) {
usage = Some(event_usage);
}
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(),
request_id.as_str(),
idempotency_key.clone(),
usage.as_ref(),
).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(),
request_id.as_str(),
idempotency_key.clone(),
usage.as_ref(),
).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(),
request_id.as_str(),
idempotency_key.clone(),
usage.as_ref(),
).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 let Some(event_usage) = extract_llm_usage_from_sse_event(&event) {
usage = Some(event_usage);
}
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(),
request_id.as_str(),
idempotency_key,
usage.as_ref(),
).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)
}
fn extract_llm_usage_from_sse_event(event: &str) -> Option<LlmTokenUsage> {
let mut data_lines = Vec::new();
for line in event.lines() {
if let Some(data) = line.strip_prefix("data:") {
data_lines.push(data.trim_start());
}
}
if data_lines.is_empty() {
return None;
}
let data = data_lines.join("\n");
if data.trim() == "[DONE]" {
return None;
}
serde_json::from_str::<Value>(data.as_str())
.ok()
.and_then(|value| extract_llm_usage_from_value(&value))
}
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,
idempotency_key: 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(),
request_id.as_str(),
idempotency_key,
response.usage.as_ref(),
)
.await
{
if is_mud_points_insufficient_app_error(&error) {
yield Ok::<Event, Infallible>(llm_sse_json_event_or_error(
"error",
json!({
"code": "MUD_POINTS_INSUFFICIENT",
"message": "泥点余额不足",
}),
));
} else {
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,
}),
));
}
} else {
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_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_billing_rounds_started_ten_thousand_tokens_and_has_minimum() {
assert_eq!(llm_router_points_for_usage(None), 1);
assert_eq!(
llm_router_points_for_usage(Some(&LlmTokenUsage {
prompt_tokens: 0,
completion_tokens: 0,
total_tokens: 1,
})),
1
);
assert_eq!(
llm_router_points_for_usage(Some(&LlmTokenUsage {
prompt_tokens: 900,
completion_tokens: 100,
total_tokens: 10_000,
})),
1
);
assert_eq!(
llm_router_points_for_usage(Some(&LlmTokenUsage {
prompt_tokens: 10_001,
completion_tokens: 0,
total_tokens: 10_001,
})),
2
);
assert_eq!(
llm_router_points_for_usage(Some(&LlmTokenUsage {
prompt_tokens: 30_001,
completion_tokens: 0,
total_tokens: 30_001,
})),
4
);
}
#[test]
fn llm_router_billing_insufficient_balance_has_stable_public_error() {
let error = map_llm_router_billing_error(SpacetimeClientError::Procedure(
"可消费泥点不足:需要 31,扣除退款占用后可用 11".to_string(),
));
assert_eq!(error.status_code(), StatusCode::CONFLICT);
assert_eq!(error.code(), "MUD_POINTS_INSUFFICIENT");
assert_eq!(error.message(), "泥点余额不足");
assert_eq!(
error
.details()
.and_then(|details| details.get("reason"))
.and_then(Value::as_str),
Some("insufficient-mud-points")
);
}
#[test]
fn llm_router_stream_terminal_is_replaced_only_for_insufficient_balance() {
let mut terminal = Some(
"event: response.completed\ndata: {\"type\":\"response.completed\"}\n\n".to_string(),
);
let billing = Err(map_llm_router_billing_error(
SpacetimeClientError::Procedure("泥点余额不足".to_string()),
));
let event = finalize_llm_router_terminal_event(
&mut terminal,
Some(&billing),
"request-1",
"user-1",
)
.expect("insufficient balance should emit a failure event");
let event_text = String::from_utf8(event.to_vec()).expect("failure event should be utf8");
assert!(event_text.contains("response.failed"));
assert!(event_text.contains("insufficient_mud_points"));
assert!(!event_text.contains("response.completed"));
assert!(terminal.is_none());
}
#[test]
fn responses_usage_parser_accepts_input_output_token_names() {
let usage = extract_llm_usage_from_value(&json!({
"usage": {"input_tokens": 12, "output_tokens": 8}
}))
.expect("usage should parse");
assert_eq!(usage.prompt_tokens, 12);
assert_eq!(usage.completion_tokens, 8);
assert_eq!(usage.total_tokens, 20);
let nested = extract_llm_usage_from_value(&json!({
"response": {"usage": {"prompt_tokens": 3, "completion_tokens": 4, "total_tokens": 9}}
}))
.expect("nested usage should parse");
assert_eq!(nested.total_tokens, 9);
}
#[test]
fn responses_sse_usage_parser_handles_crlf_and_multiline_data() {
let usage = extract_llm_usage_from_sse_event(
"event: response.completed\r\ndata: {\"response\":\r\ndata: {\"usage\":{\"input_tokens\":5,\"output_tokens\":7}}}\r\n",
)
.expect("SSE usage should parse");
assert_eq!(usage.prompt_tokens, 5);
assert_eq!(usage.completion_tokens, 7);
assert_eq!(usage.total_tokens, 12);
assert!(extract_llm_usage_from_sse_event("data: [DONE]").is_none());
assert!(extract_llm_usage_from_sse_event("event: response.output_text.delta").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 llm_router_ledger_id_is_stable_and_does_not_expose_raw_key() {
let first = llm_router_ledger_id("user-1", "request-1");
assert_eq!(first, llm_router_ledger_id("user-1", "request-1"));
assert_ne!(first, llm_router_ledger_id("user-1", "request-2"));
assert_ne!(first, llm_router_ledger_id("user-2", "request-1"));
assert!(first.starts_with("asset_operation_consume:llm-router-"));
assert!(!first.contains("request-1"));
}
#[test]
fn request_billing_key_binds_client_key_to_endpoint_and_payload() {
let mut headers = HeaderMap::new();
headers.insert("idempotency-key", HeaderValue::from_static("fixed"));
let first = request_billing_key(&headers, "request-1", "responses", br#"{"input":"a"}"#);
let replay = request_billing_key(&headers, "request-2", "responses", br#"{"input":"a"}"#);
let different_payload =
request_billing_key(&headers, "request-3", "responses", br#"{"input":"b"}"#);
let different_endpoint = request_billing_key(
&headers,
"request-4",
"chat-completions",
br#"{"input":"a"}"#,
);
assert_eq!(first, replay);
assert_ne!(first, different_payload);
assert_ne!(first, different_endpoint);
}
#[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
})
}
}