mirror of
https://github.com/farion1231/cc-switch.git
synced 2026-07-24 21:30:17 +08:00
fix(proxy): bill route-takeover traffic by the real upstream model
The model mapped for takeover (env mapping, Claude Desktop routes, Copilot normalization, Codex chat override) was discarded inside the forwarder, so usage attribution depended entirely on the upstream echoing it back. When the upstream omitted the model or mirrored the client alias, kimi/glm tokens were recorded and priced as claude-* (roughly 5-25x overstatement). - capture the final outbound model in forward(), return it via ForwardResult, and store it on the request context - attribution fallback order is now: upstream echo (empty string treated as missing) -> outbound model -> client-requested model - 'request' pricing mode anchors to the outbound model instead of the pre-mapping client alias; unchanged when no mapping applies - persist the resolved pricing_model on every usage row - Claude Desktop rows now log app_type "claude-desktop" on streaming and transform paths too (was hardcoded "claude", silently dropping desktop provider pricing overrides and splitting the cost basis by the stream flag); its global pricing defaults inherit the claude config since proxy_config only allows claude/codex/gemini rows
This commit is contained in:
@@ -38,6 +38,11 @@ pub struct ForwardResult {
|
||||
pub response: ProxyResponse,
|
||||
pub provider: Provider,
|
||||
pub claude_api_format: Option<String>,
|
||||
/// 实际发往上游的模型名(路由接管/模型映射后的真值)。
|
||||
///
|
||||
/// usage 归因不能依赖 ctx.request_model(映射前的客户端别名):上游响应
|
||||
/// 缺失 model 或回显别名时,接管流量会被记成 claude-* 并按其定价计费。
|
||||
pub outbound_model: Option<String>,
|
||||
/// 活跃连接 RAII guard:随响应一起流转到 response_processor / handle_claude_transform,
|
||||
/// 最终被 move 进流式 body future(或非流式响应作用域),覆盖整个响应生命周期。
|
||||
pub(crate) connection_guard: Option<ActiveConnectionGuard>,
|
||||
@@ -463,7 +468,7 @@ impl RequestForwarder {
|
||||
)
|
||||
.await
|
||||
{
|
||||
Ok((response, claude_api_format)) => {
|
||||
Ok((response, claude_api_format, outbound_model)) => {
|
||||
// 成功:普通闭合熔断状态异步记录,避免阻塞流式首包返回;
|
||||
// HalfOpen 探测仍同步等待,保证 permit 与熔断状态及时释放。
|
||||
self.record_success_result(&provider.id, app_type_str, used_half_open_permit)
|
||||
@@ -511,6 +516,7 @@ impl RequestForwarder {
|
||||
response,
|
||||
provider: provider.clone(),
|
||||
claude_api_format,
|
||||
outbound_model,
|
||||
connection_guard: None,
|
||||
});
|
||||
}
|
||||
@@ -561,7 +567,7 @@ impl RequestForwarder {
|
||||
)
|
||||
.await
|
||||
{
|
||||
Ok((response, claude_api_format)) => {
|
||||
Ok((response, claude_api_format, outbound_model)) => {
|
||||
log::info!(
|
||||
"[{app_type_str}] [Media] Unsupported-image retry succeeded"
|
||||
);
|
||||
@@ -613,6 +619,7 @@ impl RequestForwarder {
|
||||
response,
|
||||
provider: provider.clone(),
|
||||
claude_api_format,
|
||||
outbound_model,
|
||||
connection_guard: None,
|
||||
});
|
||||
}
|
||||
@@ -706,7 +713,7 @@ impl RequestForwarder {
|
||||
)
|
||||
.await
|
||||
{
|
||||
Ok((response, claude_api_format)) => {
|
||||
Ok((response, claude_api_format, outbound_model)) => {
|
||||
log::info!("[{app_type_str}] [RECT-002] 整流重试成功");
|
||||
self.record_success_result(
|
||||
&provider.id,
|
||||
@@ -761,6 +768,7 @@ impl RequestForwarder {
|
||||
response,
|
||||
provider: provider.clone(),
|
||||
claude_api_format,
|
||||
outbound_model,
|
||||
connection_guard: None,
|
||||
});
|
||||
}
|
||||
@@ -871,7 +879,7 @@ impl RequestForwarder {
|
||||
)
|
||||
.await
|
||||
{
|
||||
Ok((response, claude_api_format)) => {
|
||||
Ok((response, claude_api_format, outbound_model)) => {
|
||||
log::info!("[{app_type_str}] [RECT-011] budget 整流重试成功");
|
||||
self.record_success_result(
|
||||
&provider.id,
|
||||
@@ -920,6 +928,7 @@ impl RequestForwarder {
|
||||
response,
|
||||
provider: provider.clone(),
|
||||
claude_api_format,
|
||||
outbound_model,
|
||||
connection_guard: None,
|
||||
});
|
||||
}
|
||||
@@ -1077,6 +1086,9 @@ impl RequestForwarder {
|
||||
}
|
||||
|
||||
/// 转发单个请求(使用适配器)
|
||||
///
|
||||
/// 成功时返回 `(response, claude_api_format, outbound_model)`,其中
|
||||
/// `outbound_model` 是最终发往上游的模型名(所有映射/改写之后)。
|
||||
#[allow(clippy::too_many_arguments)]
|
||||
async fn forward(
|
||||
&self,
|
||||
@@ -1088,7 +1100,7 @@ impl RequestForwarder {
|
||||
headers: &axum::http::HeaderMap,
|
||||
extensions: &Extensions,
|
||||
adapter: &dyn ProviderAdapter,
|
||||
) -> Result<(ProxyResponse, Option<String>), ProxyError> {
|
||||
) -> Result<(ProxyResponse, Option<String>, Option<String>), ProxyError> {
|
||||
// 使用适配器提取 base_url
|
||||
let mut base_url = adapter.extract_base_url(provider)?;
|
||||
|
||||
@@ -1320,6 +1332,15 @@ impl RequestForwarder {
|
||||
adapter.build_url(&base_url, &effective_endpoint)
|
||||
};
|
||||
|
||||
// 记录映射后的出站模型名(此时 mapped_body 已完成接管映射 / [1m] 剥离 /
|
||||
// Copilot 归一化)。格式转换后若 body 仍带 model 字段会在下方刷新覆盖;
|
||||
// gemini_native 等模型在 URL 中的格式则保留此处的转换前真值。
|
||||
let mut outbound_model = mapped_body
|
||||
.get("model")
|
||||
.and_then(|m| m.as_str())
|
||||
.filter(|m| !m.is_empty())
|
||||
.map(str::to_string);
|
||||
|
||||
// 转换请求体(如果需要)
|
||||
let mut request_body = if codex_responses_to_chat {
|
||||
let mut mapped_body = mapped_body;
|
||||
@@ -1366,6 +1387,14 @@ impl RequestForwarder {
|
||||
// 过滤私有参数(以 `_` 开头的字段),防止内部信息泄露到上游
|
||||
// 默认使用空白名单,过滤所有 _ 前缀字段
|
||||
let filtered_body = prepare_upstream_request_body(request_body);
|
||||
// 出站 body 定稿后刷新真值(覆盖 Codex chat 上游模型覆写、转换层模型改写)
|
||||
if let Some(m) = filtered_body
|
||||
.get("model")
|
||||
.and_then(|m| m.as_str())
|
||||
.filter(|m| !m.is_empty())
|
||||
{
|
||||
outbound_model = Some(m.to_string());
|
||||
}
|
||||
log_prompt_cache_trace(
|
||||
app_type,
|
||||
provider,
|
||||
@@ -1878,7 +1907,7 @@ impl RequestForwarder {
|
||||
let response = self
|
||||
.prepare_success_response_for_failover(response, request_is_streaming)
|
||||
.await?;
|
||||
Ok((response, resolved_claude_api_format))
|
||||
Ok((response, resolved_claude_api_format, outbound_model))
|
||||
} else {
|
||||
let status_code = status.as_u16();
|
||||
let body_text = String::from_utf8(response.bytes().await?.to_vec()).ok();
|
||||
|
||||
@@ -60,37 +60,40 @@ fn gemini_stream_usage_event_filter(data: &str) -> bool {
|
||||
// ============================================================================
|
||||
|
||||
/// Claude 流式响应模型提取(优先使用 usage.model)
|
||||
fn claude_model_extractor(events: &[Value], request_model: &str) -> String {
|
||||
///
|
||||
/// 空字符串模型名视为缺失(转换层对无回显上游会合成 model:""),
|
||||
/// 落到 fallback_model(映射后的出站模型或客户端请求模型)。
|
||||
fn claude_model_extractor(events: &[Value], fallback_model: &str) -> String {
|
||||
// 首先尝试从解析的 usage 中获取模型
|
||||
if let Some(usage) = TokenUsage::from_claude_stream_events(events) {
|
||||
if let Some(model) = usage.model {
|
||||
if let Some(model) = usage.model.filter(|m| !m.is_empty()) {
|
||||
return model;
|
||||
}
|
||||
}
|
||||
request_model.to_string()
|
||||
fallback_model.to_string()
|
||||
}
|
||||
|
||||
/// OpenAI Chat Completions 流式响应模型提取(优先使用 usage.model)
|
||||
fn openai_model_extractor(events: &[Value], request_model: &str) -> String {
|
||||
fn openai_model_extractor(events: &[Value], fallback_model: &str) -> String {
|
||||
// 首先尝试从解析的 usage 中获取模型
|
||||
if let Some(usage) = TokenUsage::from_openai_stream_events(events) {
|
||||
if let Some(model) = usage.model {
|
||||
if let Some(model) = usage.model.filter(|m| !m.is_empty()) {
|
||||
return model;
|
||||
}
|
||||
}
|
||||
// 回退:从事件中直接提取
|
||||
events
|
||||
.iter()
|
||||
.find_map(|e| e.get("model")?.as_str())
|
||||
.unwrap_or(request_model)
|
||||
.find_map(|e| e.get("model")?.as_str().filter(|m| !m.is_empty()))
|
||||
.unwrap_or(fallback_model)
|
||||
.to_string()
|
||||
}
|
||||
|
||||
/// Codex 智能流式响应模型提取(自动检测格式)
|
||||
fn codex_auto_model_extractor(events: &[Value], request_model: &str) -> String {
|
||||
fn codex_auto_model_extractor(events: &[Value], fallback_model: &str) -> String {
|
||||
// 首先尝试从解析的 usage 中获取模型
|
||||
if let Some(usage) = TokenUsage::from_codex_stream_events_auto(events) {
|
||||
if let Some(model) = usage.model {
|
||||
if let Some(model) = usage.model.filter(|m| !m.is_empty()) {
|
||||
return model;
|
||||
}
|
||||
}
|
||||
@@ -99,28 +102,33 @@ fn codex_auto_model_extractor(events: &[Value], request_model: &str) -> String {
|
||||
.iter()
|
||||
.find_map(|e| {
|
||||
if e.get("type")?.as_str()? == "response.completed" {
|
||||
e.get("response")?.get("model")?.as_str()
|
||||
e.get("response")?
|
||||
.get("model")?
|
||||
.as_str()
|
||||
.filter(|m| !m.is_empty())
|
||||
} else {
|
||||
None
|
||||
}
|
||||
})
|
||||
.or_else(|| {
|
||||
// 再回退:从 OpenAI 格式事件中提取
|
||||
events.iter().find_map(|e| e.get("model")?.as_str())
|
||||
events
|
||||
.iter()
|
||||
.find_map(|e| e.get("model")?.as_str().filter(|m| !m.is_empty()))
|
||||
})
|
||||
.unwrap_or(request_model)
|
||||
.unwrap_or(fallback_model)
|
||||
.to_string()
|
||||
}
|
||||
|
||||
/// Gemini 流式响应模型提取(优先使用 usage.model)
|
||||
fn gemini_model_extractor(events: &[Value], request_model: &str) -> String {
|
||||
fn gemini_model_extractor(events: &[Value], fallback_model: &str) -> String {
|
||||
// 首先尝试从解析的 usage 中获取模型
|
||||
if let Some(usage) = TokenUsage::from_gemini_stream_chunks(events) {
|
||||
if let Some(model) = usage.model {
|
||||
if let Some(model) = usage.model.filter(|m| !m.is_empty()) {
|
||||
return model;
|
||||
}
|
||||
}
|
||||
request_model.to_string()
|
||||
fallback_model.to_string()
|
||||
}
|
||||
|
||||
// ============================================================================
|
||||
|
||||
@@ -48,6 +48,11 @@ pub struct RequestContext {
|
||||
pub current_provider_id: String,
|
||||
/// 请求中的模型名称
|
||||
pub request_model: String,
|
||||
/// 实际发往上游的模型名(路由接管/模型映射后的真值,forward 成功后回填)。
|
||||
///
|
||||
/// usage 归因的兜底顺序:上游响应回显 → outbound_model → request_model。
|
||||
/// 不能直接用 request_model 兜底:接管场景下它是映射前的客户端别名。
|
||||
pub outbound_model: Option<String>,
|
||||
/// 日志标签(如 "Claude"、"Codex"、"Gemini")
|
||||
pub tag: &'static str,
|
||||
/// 应用类型字符串(如 "claude"、"codex"、"gemini")
|
||||
@@ -159,6 +164,7 @@ impl RequestContext {
|
||||
providers,
|
||||
current_provider_id,
|
||||
request_model,
|
||||
outbound_model: None,
|
||||
tag,
|
||||
app_type_str,
|
||||
app_type,
|
||||
|
||||
@@ -207,6 +207,7 @@ async fn handle_messages_for_app(
|
||||
};
|
||||
|
||||
let connection_guard = result.connection_guard.take();
|
||||
ctx.outbound_model = result.outbound_model.take();
|
||||
ctx.provider = result.provider;
|
||||
let api_format = result
|
||||
.claude_api_format
|
||||
@@ -334,29 +335,44 @@ async fn handle_claude_transform(
|
||||
let state = state.clone();
|
||||
let provider_id = ctx.provider.id.clone();
|
||||
let request_model = ctx.request_model.clone();
|
||||
// 上游/转换层未回显模型时,优先用映射后的出站模型兜底(路由接管真值),
|
||||
// 其次才是客户端请求别名。空字符串视为缺失(转换器对无回显上游会合成 "")。
|
||||
let fallback_model = ctx
|
||||
.outbound_model
|
||||
.clone()
|
||||
.unwrap_or_else(|| ctx.request_model.clone());
|
||||
let status_code = status.as_u16();
|
||||
let start_time = ctx.start_time;
|
||||
let session_id = ctx.session_id.clone();
|
||||
// 用 ctx 的 app_type:Claude Desktop 网关也走此转换路径,硬编码
|
||||
// "claude" 会把 claude-desktop 的行错记到 claude 名下
|
||||
let app_type_str = ctx.app_type_str;
|
||||
|
||||
Some(SseUsageCollector::new(
|
||||
start_time,
|
||||
Some(claude_stream_usage_event_filter),
|
||||
move |events, first_token_ms| {
|
||||
if let Some(usage) = TokenUsage::from_claude_stream_events(&events) {
|
||||
let model = usage.model.clone().unwrap_or(request_model.clone());
|
||||
let model = usage
|
||||
.model
|
||||
.clone()
|
||||
.filter(|m| !m.is_empty())
|
||||
.unwrap_or_else(|| fallback_model.clone());
|
||||
let latency_ms = start_time.elapsed().as_millis() as u64;
|
||||
let state = state.clone();
|
||||
let provider_id = provider_id.clone();
|
||||
let session_id = session_id.clone();
|
||||
let request_model = request_model.clone();
|
||||
let outbound_model = fallback_model.clone();
|
||||
|
||||
tokio::spawn(async move {
|
||||
log_usage(
|
||||
&state,
|
||||
&provider_id,
|
||||
"claude",
|
||||
app_type_str,
|
||||
&model,
|
||||
&request_model,
|
||||
&outbound_model,
|
||||
usage,
|
||||
latency_ms,
|
||||
first_token_ms,
|
||||
@@ -442,25 +458,35 @@ async fn handle_claude_transform(
|
||||
|
||||
// 记录使用量
|
||||
if let Some(usage) = TokenUsage::from_claude_response(&anthropic_response) {
|
||||
// 转换后的响应缺失/合成空 model 时,回退到映射后的出站模型(接管真值),
|
||||
// 再回退到客户端请求别名
|
||||
let model = anthropic_response
|
||||
.get("model")
|
||||
.and_then(|m| m.as_str())
|
||||
.unwrap_or("unknown");
|
||||
.filter(|m| !m.is_empty())
|
||||
.map(str::to_string)
|
||||
.or_else(|| ctx.outbound_model.clone())
|
||||
.unwrap_or_else(|| ctx.request_model.clone());
|
||||
let latency_ms = ctx.latency_ms();
|
||||
|
||||
let request_model = ctx.request_model.clone();
|
||||
let outbound_model = ctx
|
||||
.outbound_model
|
||||
.clone()
|
||||
.unwrap_or_else(|| ctx.request_model.clone());
|
||||
let app_type_str = ctx.app_type_str;
|
||||
tokio::spawn({
|
||||
let state = state.clone();
|
||||
let provider_id = ctx.provider.id.clone();
|
||||
let model = model.to_string();
|
||||
let session_id = ctx.session_id.clone();
|
||||
async move {
|
||||
log_usage(
|
||||
&state,
|
||||
&provider_id,
|
||||
"claude",
|
||||
app_type_str,
|
||||
&model,
|
||||
&request_model,
|
||||
&outbound_model,
|
||||
usage,
|
||||
latency_ms,
|
||||
None,
|
||||
@@ -563,6 +589,7 @@ pub async fn handle_chat_completions(
|
||||
};
|
||||
|
||||
let connection_guard = result.connection_guard.take();
|
||||
ctx.outbound_model = result.outbound_model.take();
|
||||
ctx.provider = result.provider;
|
||||
let response = result.response;
|
||||
|
||||
@@ -628,6 +655,7 @@ pub async fn handle_responses(
|
||||
};
|
||||
|
||||
let connection_guard = result.connection_guard.take();
|
||||
ctx.outbound_model = result.outbound_model.take();
|
||||
ctx.provider = result.provider;
|
||||
let response = result.response;
|
||||
|
||||
@@ -705,6 +733,7 @@ pub async fn handle_responses_compact(
|
||||
};
|
||||
|
||||
let connection_guard = result.connection_guard.take();
|
||||
ctx.outbound_model = result.outbound_model.take();
|
||||
ctx.provider = result.provider;
|
||||
let response = result.response;
|
||||
|
||||
@@ -756,6 +785,12 @@ async fn handle_codex_chat_to_responses_transform(
|
||||
let state = state.clone();
|
||||
let provider_id = ctx.provider.id.clone();
|
||||
let request_model = ctx.request_model.clone();
|
||||
// 接管/模型覆写场景的归因兜底:出站真值优先于客户端请求别名
|
||||
let fallback_model = ctx
|
||||
.outbound_model
|
||||
.clone()
|
||||
.unwrap_or_else(|| ctx.request_model.clone());
|
||||
let app_type_str = ctx.app_type_str;
|
||||
let start_time = ctx.start_time;
|
||||
let session_id = ctx.session_id.clone();
|
||||
|
||||
@@ -774,21 +809,27 @@ async fn handle_codex_chat_to_responses_transform(
|
||||
log::debug!("[Codex] 流式响应 usage 全 0 或缺失,跳过消费记录");
|
||||
return;
|
||||
}
|
||||
let model = usage.model.clone().unwrap_or_else(|| request_model.clone());
|
||||
let model = usage
|
||||
.model
|
||||
.clone()
|
||||
.filter(|m| !m.is_empty())
|
||||
.unwrap_or_else(|| fallback_model.clone());
|
||||
let latency_ms = start_time.elapsed().as_millis() as u64;
|
||||
|
||||
let state = state.clone();
|
||||
let provider_id = provider_id.clone();
|
||||
let request_model = request_model.clone();
|
||||
let outbound_model = fallback_model.clone();
|
||||
let session_id = session_id.clone();
|
||||
|
||||
tokio::spawn(async move {
|
||||
log_usage(
|
||||
&state,
|
||||
&provider_id,
|
||||
"codex",
|
||||
app_type_str,
|
||||
&model,
|
||||
&request_model,
|
||||
&outbound_model,
|
||||
usage,
|
||||
latency_ms,
|
||||
first_token_ms,
|
||||
@@ -863,21 +904,29 @@ async fn handle_codex_chat_to_responses_transform(
|
||||
let model = responses_response
|
||||
.get("model")
|
||||
.and_then(|m| m.as_str())
|
||||
.unwrap_or(&ctx.request_model);
|
||||
.filter(|m| !m.is_empty())
|
||||
.map(str::to_string)
|
||||
.or_else(|| ctx.outbound_model.clone())
|
||||
.unwrap_or_else(|| ctx.request_model.clone());
|
||||
let request_model = ctx.request_model.clone();
|
||||
let outbound_model = ctx
|
||||
.outbound_model
|
||||
.clone()
|
||||
.unwrap_or_else(|| ctx.request_model.clone());
|
||||
let app_type_str = ctx.app_type_str;
|
||||
tokio::spawn({
|
||||
let state = state.clone();
|
||||
let provider_id = ctx.provider.id.clone();
|
||||
let model = model.to_string();
|
||||
let session_id = ctx.session_id.clone();
|
||||
let latency_ms = ctx.latency_ms();
|
||||
async move {
|
||||
log_usage(
|
||||
&state,
|
||||
&provider_id,
|
||||
"codex",
|
||||
app_type_str,
|
||||
&model,
|
||||
&request_model,
|
||||
&outbound_model,
|
||||
usage,
|
||||
latency_ms,
|
||||
None,
|
||||
@@ -1242,6 +1291,7 @@ pub async fn handle_gemini(
|
||||
};
|
||||
|
||||
let connection_guard = result.connection_guard.take();
|
||||
ctx.outbound_model = result.outbound_model.take();
|
||||
ctx.provider = result.provider;
|
||||
let response = result.response;
|
||||
|
||||
@@ -1370,6 +1420,9 @@ fn log_forward_error(
|
||||
}
|
||||
|
||||
/// 记录请求使用量
|
||||
///
|
||||
/// `outbound_model` 是「按请求计价」模式的锚点:实际发往上游的模型
|
||||
/// (路由接管映射后的真值,无映射时等于 request_model)。
|
||||
#[allow(clippy::too_many_arguments)]
|
||||
async fn log_usage(
|
||||
state: &ProxyState,
|
||||
@@ -1377,6 +1430,7 @@ async fn log_usage(
|
||||
app_type: &str,
|
||||
model: &str,
|
||||
request_model: &str,
|
||||
outbound_model: &str,
|
||||
usage: TokenUsage,
|
||||
latency_ms: u64,
|
||||
first_token_ms: Option<u64>,
|
||||
@@ -1395,7 +1449,7 @@ async fn log_usage(
|
||||
let (multiplier, pricing_model_source) =
|
||||
logger.resolve_pricing_config(provider_id, app_type).await;
|
||||
let pricing_model = if pricing_model_source == PRICING_SOURCE_REQUEST {
|
||||
request_model
|
||||
outbound_model
|
||||
} else {
|
||||
model
|
||||
};
|
||||
|
||||
@@ -270,14 +270,21 @@ pub async fn handle_non_streaming(
|
||||
if let Ok(json_value) = serde_json::from_slice::<Value>(&body_bytes) {
|
||||
// 解析使用量
|
||||
if let Some(usage) = (parser_config.response_parser)(&json_value) {
|
||||
// 优先使用 usage 中解析出的模型名称,其次使用响应中的 model 字段,最后回退到请求模型
|
||||
let model = if let Some(ref m) = usage.model {
|
||||
m.clone()
|
||||
} else if let Some(m) = json_value.get("model").and_then(|m| m.as_str()) {
|
||||
m.to_string()
|
||||
} else {
|
||||
ctx.request_model.clone()
|
||||
};
|
||||
// 归因优先级:usage 解析出的模型 → 响应 model 字段 → 映射后的出站
|
||||
// 模型(路由接管真值)→ 客户端请求模型。空字符串视为缺失。
|
||||
let model = usage
|
||||
.model
|
||||
.clone()
|
||||
.filter(|m| !m.is_empty())
|
||||
.or_else(|| {
|
||||
json_value
|
||||
.get("model")
|
||||
.and_then(|m| m.as_str())
|
||||
.filter(|m| !m.is_empty())
|
||||
.map(str::to_string)
|
||||
})
|
||||
.or_else(|| ctx.outbound_model.clone())
|
||||
.unwrap_or_else(|| ctx.request_model.clone());
|
||||
|
||||
spawn_log_usage(
|
||||
state,
|
||||
@@ -292,8 +299,10 @@ pub async fn handle_non_streaming(
|
||||
let model = json_value
|
||||
.get("model")
|
||||
.and_then(|m| m.as_str())
|
||||
.unwrap_or(&ctx.request_model)
|
||||
.to_string();
|
||||
.filter(|m| !m.is_empty())
|
||||
.map(str::to_string)
|
||||
.or_else(|| ctx.outbound_model.clone())
|
||||
.unwrap_or_else(|| ctx.request_model.clone());
|
||||
spawn_log_usage(
|
||||
state,
|
||||
ctx,
|
||||
@@ -318,7 +327,7 @@ pub async fn handle_non_streaming(
|
||||
state,
|
||||
ctx,
|
||||
TokenUsage::default(),
|
||||
&ctx.request_model,
|
||||
ctx.outbound_model.as_deref().unwrap_or(&ctx.request_model),
|
||||
&ctx.request_model,
|
||||
status.as_u16(),
|
||||
false,
|
||||
@@ -500,7 +509,16 @@ fn create_usage_collector(
|
||||
let state = state.clone();
|
||||
let provider_id = ctx.provider.id.clone();
|
||||
let request_model = ctx.request_model.clone();
|
||||
let app_type_str = parser_config.app_type_str;
|
||||
// 流式事件缺失模型名时的归因兜底:映射后的出站模型(路由接管真值)优先,
|
||||
// 其次才是客户端请求别名
|
||||
let fallback_model = ctx
|
||||
.outbound_model
|
||||
.clone()
|
||||
.unwrap_or_else(|| ctx.request_model.clone());
|
||||
// 用 ctx 的 app_type 而不是 parser_config 的:Claude Desktop 流式透传复用
|
||||
// CLAUDE_PARSER_CONFIG(app_type_str="claude"),按 parser_config 记账会把
|
||||
// claude-desktop 的行错记到 claude 名下,导致供应商计价覆盖解析不到。
|
||||
let app_type_str = ctx.app_type_str;
|
||||
let tag = ctx.tag;
|
||||
let start_time = ctx.start_time;
|
||||
let stream_parser = parser_config.stream_parser;
|
||||
@@ -512,13 +530,14 @@ fn create_usage_collector(
|
||||
parser_config.stream_event_filter,
|
||||
move |events, first_token_ms| {
|
||||
if let Some(usage) = stream_parser(&events) {
|
||||
let model = model_extractor(&events, &request_model);
|
||||
let model = model_extractor(&events, &fallback_model);
|
||||
let latency_ms = start_time.elapsed().as_millis() as u64;
|
||||
|
||||
let state = state.clone();
|
||||
let provider_id = provider_id.clone();
|
||||
let session_id = session_id.clone();
|
||||
let request_model = request_model.clone();
|
||||
let outbound_model = fallback_model.clone();
|
||||
|
||||
tokio::spawn(async move {
|
||||
log_usage_internal(
|
||||
@@ -527,6 +546,7 @@ fn create_usage_collector(
|
||||
app_type_str,
|
||||
&model,
|
||||
&request_model,
|
||||
&outbound_model,
|
||||
usage,
|
||||
latency_ms,
|
||||
first_token_ms,
|
||||
@@ -537,12 +557,13 @@ fn create_usage_collector(
|
||||
.await;
|
||||
});
|
||||
} else {
|
||||
let model = model_extractor(&events, &request_model);
|
||||
let model = model_extractor(&events, &fallback_model);
|
||||
let latency_ms = start_time.elapsed().as_millis() as u64;
|
||||
let state = state.clone();
|
||||
let provider_id = provider_id.clone();
|
||||
let session_id = session_id.clone();
|
||||
let request_model = request_model.clone();
|
||||
let outbound_model = fallback_model.clone();
|
||||
|
||||
tokio::spawn(async move {
|
||||
log_usage_internal(
|
||||
@@ -551,6 +572,7 @@ fn create_usage_collector(
|
||||
app_type_str,
|
||||
&model,
|
||||
&request_model,
|
||||
&outbound_model,
|
||||
TokenUsage::default(),
|
||||
latency_ms,
|
||||
first_token_ms,
|
||||
@@ -588,6 +610,11 @@ fn spawn_log_usage(
|
||||
let app_type_str = ctx.app_type_str.to_string();
|
||||
let model = model.to_string();
|
||||
let request_model = request_model.to_string();
|
||||
// 「按请求计价」模式的锚点:映射后的出站模型,无映射时等于 request_model
|
||||
let outbound_model = ctx
|
||||
.outbound_model
|
||||
.clone()
|
||||
.unwrap_or_else(|| ctx.request_model.clone());
|
||||
let latency_ms = ctx.latency_ms();
|
||||
let session_id = ctx.session_id.clone();
|
||||
|
||||
@@ -598,6 +625,7 @@ fn spawn_log_usage(
|
||||
&app_type_str,
|
||||
&model,
|
||||
&request_model,
|
||||
&outbound_model,
|
||||
usage,
|
||||
latency_ms,
|
||||
None,
|
||||
@@ -618,6 +646,11 @@ pub(crate) fn usage_logging_enabled(state: &ProxyState) -> bool {
|
||||
}
|
||||
|
||||
/// 内部使用量记录函数
|
||||
///
|
||||
/// `outbound_model` 是「按请求计价」模式的锚点:实际发往上游的模型
|
||||
/// (路由接管映射后的真值,无映射时等于 request_model)。该模式的语义是
|
||||
/// 「按代理发出的请求计价、不信任上游回显」,接管场景下发出的请求模型是
|
||||
/// 映射后的 Y 而非客户端别名 X,按 X 计价会用错定价表行。
|
||||
#[allow(clippy::too_many_arguments)]
|
||||
async fn log_usage_internal(
|
||||
state: &ProxyState,
|
||||
@@ -625,6 +658,7 @@ async fn log_usage_internal(
|
||||
app_type: &str,
|
||||
model: &str,
|
||||
request_model: &str,
|
||||
outbound_model: &str,
|
||||
usage: TokenUsage,
|
||||
latency_ms: u64,
|
||||
first_token_ms: Option<u64>,
|
||||
@@ -638,7 +672,7 @@ async fn log_usage_internal(
|
||||
let (multiplier, pricing_model_source) =
|
||||
logger.resolve_pricing_config(provider_id, app_type).await;
|
||||
let pricing_model = if pricing_model_source == PRICING_SOURCE_REQUEST {
|
||||
request_model
|
||||
outbound_model
|
||||
} else {
|
||||
model
|
||||
};
|
||||
@@ -1015,6 +1049,7 @@ mod tests {
|
||||
app_type,
|
||||
"resp-model",
|
||||
"req-model",
|
||||
"req-model",
|
||||
usage,
|
||||
10,
|
||||
None,
|
||||
@@ -1047,6 +1082,95 @@ mod tests {
|
||||
Ok(())
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_request_pricing_mode_anchors_to_outbound_model() -> Result<(), AppError> {
|
||||
let db = Arc::new(Database::memory()?);
|
||||
let app_type = "claude";
|
||||
|
||||
db.set_pricing_model_source(app_type, "request").await?;
|
||||
seed_pricing(&db)?;
|
||||
{
|
||||
let conn = crate::database::lock_conn!(db.conn);
|
||||
conn.execute(
|
||||
"INSERT OR REPLACE INTO model_pricing (model_id, display_name, input_cost_per_million, output_cost_per_million)
|
||||
VALUES ('outbound-model', 'Outbound Model', '4.0', '0')",
|
||||
[],
|
||||
)
|
||||
.map_err(|e| AppError::Database(e.to_string()))?;
|
||||
}
|
||||
|
||||
insert_provider(&db, "provider-3", app_type, ProviderMeta::default())?;
|
||||
|
||||
let state = build_state(db.clone());
|
||||
let usage = TokenUsage {
|
||||
input_tokens: 1_000_000,
|
||||
output_tokens: 0,
|
||||
cache_read_tokens: 0,
|
||||
cache_creation_tokens: 0,
|
||||
model: None,
|
||||
message_id: None,
|
||||
};
|
||||
|
||||
// 路由接管场景:客户端请求 req-model($2/M),代理实际发出 outbound-model
|
||||
// ($4/M),上游回显 resp-model。「按请求计价」必须锚定实际发出的模型。
|
||||
log_usage_internal(
|
||||
&state,
|
||||
"provider-3",
|
||||
app_type,
|
||||
"resp-model",
|
||||
"req-model",
|
||||
"outbound-model",
|
||||
usage,
|
||||
10,
|
||||
None,
|
||||
false,
|
||||
200,
|
||||
None,
|
||||
)
|
||||
.await;
|
||||
|
||||
let conn = crate::database::lock_conn!(db.conn);
|
||||
let (model, request_model, total_cost): (String, String, String) = conn
|
||||
.query_row(
|
||||
"SELECT model, request_model, total_cost_usd
|
||||
FROM proxy_request_logs WHERE provider_id = ?1",
|
||||
["provider-3"],
|
||||
|row| Ok((row.get(0)?, row.get(1)?, row.get(2)?)),
|
||||
)
|
||||
.map_err(|e| AppError::Database(e.to_string()))?;
|
||||
|
||||
// model / request_model 列不受计价锚点影响
|
||||
assert_eq!(model, "resp-model");
|
||||
assert_eq!(request_model, "req-model");
|
||||
// 按 outbound-model($4/M)计价,而不是 req-model($2/M)或 resp-model($1/M)
|
||||
assert_eq!(
|
||||
Decimal::from_str(&total_cost).unwrap(),
|
||||
Decimal::from_str("4").unwrap()
|
||||
);
|
||||
Ok(())
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_claude_desktop_inherits_claude_global_defaults() -> Result<(), AppError> {
|
||||
use crate::proxy::usage::logger::UsageLogger;
|
||||
|
||||
let db = Arc::new(Database::memory()?);
|
||||
|
||||
// 全局计费配置只有 claude/codex/gemini 三行;claude-desktop 的
|
||||
// 全局默认必须继承 claude,而不是静默落回工厂默认(1 / response)
|
||||
db.set_default_cost_multiplier("claude", "1.5").await?;
|
||||
db.set_pricing_model_source("claude", "request").await?;
|
||||
|
||||
let logger = UsageLogger::new(&db);
|
||||
let (multiplier, source) = logger
|
||||
.resolve_pricing_config("nonexistent-provider", "claude-desktop")
|
||||
.await;
|
||||
|
||||
assert_eq!(multiplier, Decimal::from_str("1.5").unwrap());
|
||||
assert_eq!(source, "request");
|
||||
Ok(())
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_log_usage_falls_back_to_global_defaults() -> Result<(), AppError> {
|
||||
let db = Arc::new(Database::memory()?);
|
||||
@@ -1075,6 +1199,7 @@ mod tests {
|
||||
app_type,
|
||||
"resp-model",
|
||||
"req-model",
|
||||
"req-model",
|
||||
usage,
|
||||
10,
|
||||
None,
|
||||
|
||||
@@ -16,6 +16,11 @@ pub struct RequestLog {
|
||||
pub app_type: String,
|
||||
pub model: String,
|
||||
pub request_model: String,
|
||||
/// 写入时实际用于计价的模型名(pricing_model_source 解析后的结果)。
|
||||
/// 落库供回填使用:缺价行补价后必须按写入时的基准重算,而不是
|
||||
/// 用 model/request_model 猜——路由接管下三者可能各不相同。
|
||||
/// 错误行(未计价)为空字符串。
|
||||
pub pricing_model: String,
|
||||
pub usage: TokenUsage,
|
||||
pub cost: Option<CostBreakdown>,
|
||||
pub latency_ms: u64,
|
||||
@@ -68,18 +73,19 @@ impl<'a> UsageLogger<'a> {
|
||||
|
||||
conn.execute(
|
||||
"INSERT OR REPLACE INTO proxy_request_logs (
|
||||
request_id, provider_id, app_type, model, request_model,
|
||||
request_id, provider_id, app_type, model, request_model, pricing_model,
|
||||
input_tokens, output_tokens, cache_read_tokens, cache_creation_tokens,
|
||||
input_cost_usd, output_cost_usd, cache_read_cost_usd, cache_creation_cost_usd, total_cost_usd,
|
||||
latency_ms, first_token_ms, status_code, error_message, session_id,
|
||||
provider_type, is_streaming, cost_multiplier, created_at
|
||||
) VALUES (?1, ?2, ?3, ?4, ?5, ?6, ?7, ?8, ?9, ?10, ?11, ?12, ?13, ?14, ?15, ?16, ?17, ?18, ?19, ?20, ?21, ?22, ?23)",
|
||||
) VALUES (?1, ?2, ?3, ?4, ?5, ?6, ?7, ?8, ?9, ?10, ?11, ?12, ?13, ?14, ?15, ?16, ?17, ?18, ?19, ?20, ?21, ?22, ?23, ?24)",
|
||||
rusqlite::params![
|
||||
log.request_id,
|
||||
log.provider_id,
|
||||
log.app_type,
|
||||
log.model,
|
||||
log.request_model,
|
||||
log.pricing_model,
|
||||
log.usage.input_tokens,
|
||||
log.usage.output_tokens,
|
||||
log.usage.cache_read_tokens,
|
||||
@@ -129,6 +135,8 @@ impl<'a> UsageLogger<'a> {
|
||||
app_type,
|
||||
model,
|
||||
request_model,
|
||||
// 错误行未经过计价,留空(回填的 has_usage 闸门也不会碰全 0 行)
|
||||
pricing_model: String::new(),
|
||||
usage: TokenUsage::default(),
|
||||
cost: None,
|
||||
latency_ms,
|
||||
@@ -168,6 +176,8 @@ impl<'a> UsageLogger<'a> {
|
||||
app_type,
|
||||
model,
|
||||
request_model,
|
||||
// 错误行未经过计价,留空(回填的 has_usage 闸门也不会碰全 0 行)
|
||||
pricing_model: String::new(),
|
||||
usage: TokenUsage::default(),
|
||||
cost: None,
|
||||
latency_ms,
|
||||
@@ -203,13 +213,22 @@ impl<'a> UsageLogger<'a> {
|
||||
provider_id: &str,
|
||||
app_type: &str,
|
||||
) -> (Decimal, String) {
|
||||
let default_multiplier_raw = match self.db.get_default_cost_multiplier(app_type).await {
|
||||
Ok(value) => value,
|
||||
Err(e) => {
|
||||
log::warn!("[USG-003] 获取默认倍率失败 (app_type={app_type}): {e}");
|
||||
"1".to_string()
|
||||
}
|
||||
// Claude Desktop 网关没有独立的全局计费配置(proxy_config 的 CHECK 仅
|
||||
// 允许 claude/codex/gemini,前端也只暴露三项),全局默认继承 claude;
|
||||
// 供应商级 meta 覆盖仍按 claude-desktop 查找(providers 表按该 app_type 存)。
|
||||
let default_app_type = if app_type == "claude-desktop" {
|
||||
"claude"
|
||||
} else {
|
||||
app_type
|
||||
};
|
||||
let default_multiplier_raw =
|
||||
match self.db.get_default_cost_multiplier(default_app_type).await {
|
||||
Ok(value) => value,
|
||||
Err(e) => {
|
||||
log::warn!("[USG-003] 获取默认倍率失败 (app_type={app_type}): {e}");
|
||||
"1".to_string()
|
||||
}
|
||||
};
|
||||
let default_multiplier = match Decimal::from_str(&default_multiplier_raw) {
|
||||
Ok(value) => value,
|
||||
Err(e) => {
|
||||
@@ -220,13 +239,14 @@ impl<'a> UsageLogger<'a> {
|
||||
}
|
||||
};
|
||||
|
||||
let default_pricing_source_raw = match self.db.get_pricing_model_source(app_type).await {
|
||||
Ok(value) => value,
|
||||
Err(e) => {
|
||||
log::warn!("[USG-003] 获取默认计费模式失败 (app_type={app_type}): {e}");
|
||||
PRICING_SOURCE_RESPONSE.to_string()
|
||||
}
|
||||
};
|
||||
let default_pricing_source_raw =
|
||||
match self.db.get_pricing_model_source(default_app_type).await {
|
||||
Ok(value) => value,
|
||||
Err(e) => {
|
||||
log::warn!("[USG-003] 获取默认计费模式失败 (app_type={app_type}): {e}");
|
||||
PRICING_SOURCE_RESPONSE.to_string()
|
||||
}
|
||||
};
|
||||
let default_pricing_source = if default_pricing_source_raw == PRICING_SOURCE_RESPONSE
|
||||
|| default_pricing_source_raw == PRICING_SOURCE_REQUEST
|
||||
{
|
||||
@@ -325,6 +345,7 @@ impl<'a> UsageLogger<'a> {
|
||||
app_type,
|
||||
model,
|
||||
request_model,
|
||||
pricing_model,
|
||||
usage,
|
||||
cost,
|
||||
latency_ms,
|
||||
|
||||
Reference in New Issue
Block a user