fix(usage): pricing routing, SSE lifecycle, and validation hardening

* model pricing routing: extend prefix-match families (gpt-/o1-o5/
  gemini-/deepseek-/qwen-/glm-/kimi-/minimax-) with per-family dash
  thresholds so short base IDs like gpt-5 no longer mis-match
  gpt-5-mini; strip ISO and 8-digit date suffixes via UTF-8-safe
  byte matching so claude-haiku-4-5-20251001 falls back to
  claude-haiku-4-5 pricing
* SSE collector: SseUsageFinishGuard (RAII) guarantees finish() on
  early return or panic; AtomicBool fast path lets push() skip the
  Mutex once first-event time is recorded
* validation: shared validate_cost_multiplier / validate_pricing_source
  helpers across DAO and service layers; PRICING_SOURCE_RESPONSE /
  PRICING_SOURCE_REQUEST constants replace string literals; price
  fields in update_model_pricing now reject empty / non-decimal /
  negative input before INSERT
* backfill: add backfill_missing_usage_costs_for_model so a single
  price edit only scans matching rows instead of the full log table;
  startup backfill remains full-scan
* session_usage{,_codex,_gemini}: share find_model_pricing helper from
  usage_stats; metadata_modified_nanos centralizes mtime precision
* frontend: NON_NEGATIVE_DECIMAL_REGEX + isNonNegativeDecimalString
  replace three copies of the same multiplier regex; isUnpricedUsage
  surfaces zero-cost rows that have usage tokens (cached per row to
  avoid double evaluation); invalidate usageKeys.all on pricing mutate
  so backfilled rows refresh
This commit is contained in:
Jason
2026-05-14 11:53:51 +08:00
parent 206125b4e3
commit 402570ce31
19 changed files with 590 additions and 375 deletions
+51 -8
View File
@@ -12,6 +12,7 @@ use super::{
usage::parser::TokenUsage,
ProxyError,
};
use crate::database::PRICING_SOURCE_REQUEST;
use axum::http::{header::HeaderMap, HeaderName};
use axum::response::{IntoResponse, Response};
use bytes::Bytes;
@@ -372,6 +373,7 @@ pub struct SseUsageCollector {
struct SseUsageCollectorInner {
events: Mutex<Vec<Value>>,
first_event_time: Mutex<Option<std::time::Instant>>,
first_event_set: AtomicBool,
start_time: std::time::Instant,
on_complete: UsageCallbackWithTiming,
should_collect: Option<StreamUsageEventFilter>,
@@ -390,6 +392,7 @@ impl SseUsageCollector {
inner: Arc::new(SseUsageCollectorInner {
events: Mutex::new(Vec::new()),
first_event_time: Mutex::new(None),
first_event_set: AtomicBool::new(false),
start_time,
on_complete,
should_collect,
@@ -405,15 +408,21 @@ impl SseUsageCollector {
.unwrap_or(true)
}
/// 标记首个被收集的 SSE 事件时间,沿用 `first_token_ms` 的既有近似语义。
async fn mark_first_collected_event_time(&self) {
if self.inner.first_event_set.load(Ordering::Acquire) {
return;
}
let mut first_time = self.inner.first_event_time.lock().await;
if first_time.is_none() {
*first_time = Some(std::time::Instant::now());
self.inner.first_event_set.store(true, Ordering::Release);
}
}
/// 推送 SSE 事件
pub async fn push(&self, event: Value) {
// 记录首个事件时间
{
let mut first_time = self.inner.first_event_time.lock().await;
if first_time.is_none() {
*first_time = Some(std::time::Instant::now());
}
}
self.mark_first_collected_event_time().await;
let mut events = self.inner.events.lock().await;
events.push(event);
}
@@ -438,6 +447,36 @@ impl SseUsageCollector {
}
}
struct SseUsageFinishGuard {
collector: Option<SseUsageCollector>,
}
impl SseUsageFinishGuard {
fn new(collector: SseUsageCollector) -> Self {
Self {
collector: Some(collector),
}
}
fn disarm(&mut self) {
self.collector = None;
}
}
impl Drop for SseUsageFinishGuard {
fn drop(&mut self) {
if let Some(collector) = self.collector.take() {
if let Ok(handle) = tokio::runtime::Handle::try_current() {
handle.spawn(async move {
collector.finish().await;
});
} else {
log::warn!("SSE 用量收尾保护触发时 Tokio runtime 不可用,跳过异步 finish");
}
}
}
}
// ============================================================================
// 内部辅助函数
// ============================================================================
@@ -598,7 +637,7 @@ async fn log_usage_internal(
let logger = UsageLogger::new(&state.db);
let (multiplier, pricing_model_source) =
logger.resolve_pricing_config(provider_id, app_type).await;
let pricing_model = if pricing_model_source == "request" {
let pricing_model = if pricing_model_source == PRICING_SOURCE_REQUEST {
request_model
} else {
model
@@ -648,6 +687,7 @@ pub fn create_logged_passthrough_stream(
let mut buffer = String::new();
let mut utf8_remainder: Vec<u8> = Vec::new();
let mut collector = usage_collector;
let mut finish_guard = collector.clone().map(SseUsageFinishGuard::new);
let inspect_sse_events =
collector.is_some() || log::log_enabled!(log::Level::Debug);
let mut is_first_chunk = true;
@@ -753,6 +793,9 @@ pub fn create_logged_passthrough_stream(
if let Some(c) = collector.take() {
c.finish().await;
}
if let Some(guard) = &mut finish_guard {
guard.disarm();
}
}
}