Merge remote-tracking branch 'upstream/main'

# Conflicts:
#	crates/aether-data/src/repository/usage/mysql.rs
This commit is contained in:
zhefox
2026-05-25 17:29:12 +08:00
17 changed files with 1146 additions and 210 deletions
@@ -29,7 +29,7 @@ use crate::ai_serving::api::StreamingStandardTerminalObserver;
use crate::clock::current_unix_secs; use crate::clock::current_unix_secs;
use crate::execution_runtime::ndjson::encode_stream_frame_ndjson; use crate::execution_runtime::ndjson::encode_stream_frame_ndjson;
use crate::execution_runtime::transport::{ use crate::execution_runtime::transport::{
DirectSyncExecutionRuntime, ExecutionRuntimeTransportError, with_non_stream_total_timeout, DirectSyncExecutionRuntime, ExecutionRuntimeTransportError,
}; };
use crate::handlers::shared::{ use crate::handlers::shared::{
sync_provider_key_oauth_status_snapshot, sync_provider_key_quota_status_snapshot, sync_provider_key_oauth_status_snapshot, sync_provider_key_quota_status_snapshot,
@@ -108,12 +108,16 @@ pub(crate) async fn maybe_execute_chatgpt_web_image_sync(
if !is_chatgpt_web_image_plan(plan, report_context) { if !is_chatgpt_web_image_plan(plan, report_context) {
return Ok(None); return Ok(None);
} }
let started_at = Instant::now(); with_non_stream_total_timeout(plan, async move {
let result = match execute_chatgpt_web_image(state, plan, report_context, started_at).await { let started_at = Instant::now();
Ok(result) => result, let result = match execute_chatgpt_web_image(state, plan, report_context, started_at).await
Err(err) => chatgpt_web_transport_error_execution_result(plan, started_at, &err), {
}; Ok(result) => result,
Ok(Some(result)) Err(err) => chatgpt_web_transport_error_execution_result(plan, started_at, &err),
};
Ok(Some(result))
})
.await
} }
pub(crate) async fn maybe_execute_chatgpt_web_image_stream( pub(crate) async fn maybe_execute_chatgpt_web_image_stream(
@@ -1,4 +1,5 @@
use std::collections::BTreeMap; use std::collections::BTreeMap;
use std::future::Future;
use std::io::Error as IoError; use std::io::Error as IoError;
use std::net::IpAddr; use std::net::IpAddr;
use std::sync::OnceLock; use std::sync::OnceLock;
@@ -29,7 +30,9 @@ use crate::execution_runtime::ndjson::encode_stream_frame_ndjson;
use crate::execution_runtime::transport::{ use crate::execution_runtime::transport::{
build_browser_wreq_client, build_request_body, build_request_headers, build_browser_wreq_client, build_request_body, build_request_headers,
decode_response_body_bytes, format_upstream_request_error, format_wreq_upstream_request_error, decode_response_body_bytes, format_upstream_request_error, format_wreq_upstream_request_error,
send_request, DirectHttpResponse, ExecutionRuntimeTransportError, ExecutionTransportControls, resolve_stream_first_byte_timeout, send_request, stream_first_byte_timeout_message,
with_non_stream_total_timeout, DirectHttpResponse, ExecutionRuntimeTransportError,
ExecutionTransportControls,
}; };
const GROK_INTERNAL_HEADER: &str = "x-aether-grok-runtime"; const GROK_INTERNAL_HEADER: &str = "x-aether-grok-runtime";
@@ -148,9 +151,12 @@ pub(crate) async fn maybe_execute_grok_sync(
if !is_grok_plan(plan, report_context) { if !is_grok_plan(plan, report_context) {
return Ok(None); return Ok(None);
} }
let mut collected = execute_grok_app_chat(plan, report_context).await?; with_non_stream_total_timeout(plan, async move {
materialize_grok_image_assets(plan, &mut collected).await; let mut collected = execute_grok_app_chat(plan, report_context).await?;
Ok(Some(grok_execution_result(plan, collected, report_context))) materialize_grok_image_assets(plan, &mut collected).await;
Ok(Some(grok_execution_result(plan, collected, report_context)))
})
.await
} }
pub(crate) async fn maybe_execute_grok_stream( pub(crate) async fn maybe_execute_grok_stream(
@@ -387,6 +393,7 @@ async fn grok_imagine_websocket_images(
plan.proxy.as_ref(), plan.proxy.as_ref(),
profile, profile,
ExecutionTransportControls::default(), ExecutionTransportControls::default(),
true,
)?; )?;
let response = client let response = client
.websocket(GROK_IMAGINE_WS_URL) .websocket(GROK_IMAGINE_WS_URL)
@@ -561,6 +568,7 @@ fn grok_success_frame_stream(
started_at: Instant, started_at: Instant,
mut body_stream: GrokUpstreamBodyStream, mut body_stream: GrokUpstreamBodyStream,
) -> BoxStream<'static, Result<Bytes, IoError>> { ) -> BoxStream<'static, Result<Bytes, IoError>> {
let stream_first_byte_timeout = resolve_stream_first_byte_timeout(&plan);
async_stream::stream! { async_stream::stream! {
match encode_grok_headers_frame( match encode_grok_headers_frame(
status_code, status_code,
@@ -581,8 +589,36 @@ fn grok_success_frame_stream(
let mut text_len = 0usize; let mut text_len = 0usize;
let mut thinking_len = 0usize; let mut thinking_len = 0usize;
let mut image_len = 0usize; let mut image_len = 0usize;
let mut terminal_error_emitted = false;
while let Some(item) = body_stream.next().await { loop {
let item = if ttfb_ms.is_none() {
match await_grok_stream_first_byte(
body_stream.next(),
started_at,
stream_first_byte_timeout,
)
.await
{
Ok(item) => item,
Err(timeout) => {
match encode_grok_first_byte_timeout_frame(timeout) {
Ok(frame) => yield Ok(frame),
Err(err) => {
yield Err(err);
return;
}
}
terminal_error_emitted = true;
break;
}
}
} else {
body_stream.next().await
};
let Some(item) = item else {
break;
};
let chunk = match item { let chunk = match item {
Ok(chunk) => chunk, Ok(chunk) => chunk,
Err(message) => { Err(message) => {
@@ -593,6 +629,7 @@ fn grok_success_frame_stream(
return; return;
} }
} }
terminal_error_emitted = true;
break; break;
} }
}; };
@@ -630,6 +667,25 @@ fn grok_success_frame_stream(
} }
} }
if terminal_error_emitted {
match encode_grok_telemetry_frame(
ttfb_ms,
Some(started_at.elapsed().as_millis() as u64),
upstream_bytes,
) {
Ok(frame) => yield Ok(frame),
Err(err) => {
yield Err(err);
return;
}
}
match encode_stream_frame_ndjson(&StreamFrame::eof_with_summary(None)) {
Ok(frame) => yield Ok(frame),
Err(err) => yield Err(err),
}
return;
}
adapter.finish(); adapter.finish();
match emit_grok_adapter_deltas( match emit_grok_adapter_deltas(
&mut client_emitter, &mut client_emitter,
@@ -691,6 +747,28 @@ fn grok_success_frame_stream(
.boxed() .boxed()
} }
async fn await_grok_stream_first_byte<T, F>(
future: F,
started_at: Instant,
timeout: Option<Duration>,
) -> Result<T, Duration>
where
F: Future<Output = T>,
{
let Some(timeout) = timeout else {
return Ok(future.await);
};
let Some(remaining) = timeout.checked_sub(started_at.elapsed()) else {
return Err(timeout);
};
if remaining.is_zero() {
return Err(timeout);
}
tokio::time::timeout(remaining, future)
.await
.map_err(|_| timeout)
}
fn emit_grok_adapter_deltas( fn emit_grok_adapter_deltas(
client_emitter: &mut GrokClientStreamEmitter, client_emitter: &mut GrokClientStreamEmitter,
adapter: &GrokStreamAdapter, adapter: &GrokStreamAdapter,
@@ -790,6 +868,22 @@ fn encode_grok_error_frame(status_code: u16, message: String) -> Result<Bytes, I
}) })
} }
fn encode_grok_first_byte_timeout_frame(timeout: Duration) -> Result<Bytes, IoError> {
encode_stream_frame_ndjson(&StreamFrame {
frame_type: StreamFrameType::Error,
payload: StreamFramePayload::Error {
error: aether_contracts::ExecutionError {
kind: aether_contracts::ExecutionErrorKind::FirstByteTimeout,
phase: aether_contracts::ExecutionPhase::FirstByte,
message: stream_first_byte_timeout_message(timeout),
upstream_status: Some(504),
retryable: true,
failover_recommended: true,
},
},
})
}
enum GrokClientStreamEmitter { enum GrokClientStreamEmitter {
OpenAiChat { OpenAiChat {
id: String, id: String,
@@ -1,6 +1,7 @@
use std::collections::BTreeMap; use std::collections::BTreeMap;
use std::future::Future;
use std::io::Error as IoError; use std::io::Error as IoError;
use std::time::Instant; use std::time::{Duration, Instant};
use aether_contracts::{ use aether_contracts::{
ExecutionError, ExecutionErrorKind, ExecutionPhase, ExecutionStreamTerminalSummary, ExecutionError, ExecutionErrorKind, ExecutionPhase, ExecutionStreamTerminalSummary,
@@ -19,7 +20,7 @@ use crate::ai_serving::api::{
}; };
use crate::execution_runtime::ndjson::encode_stream_frame_ndjson; use crate::execution_runtime::ndjson::encode_stream_frame_ndjson;
use crate::execution_runtime::transport::{ use crate::execution_runtime::transport::{
format_wreq_upstream_request_error, DirectUpstreamResponse, format_wreq_upstream_request_error, stream_first_byte_timeout_message, DirectUpstreamResponse,
}; };
use crate::execution_runtime::DirectUpstreamStreamExecution; use crate::execution_runtime::DirectUpstreamStreamExecution;
use crate::GatewayError; use crate::GatewayError;
@@ -37,6 +38,7 @@ pub(crate) fn build_direct_execution_frame_stream(
stream_summary_report_context, stream_summary_report_context,
response, response,
started_at, started_at,
stream_first_byte_timeout,
} = execution; } = execution;
let mut observer_context = stream_summary_report_context; let mut observer_context = stream_summary_report_context;
@@ -64,7 +66,7 @@ pub(crate) fn build_direct_execution_frame_stream(
if should_buffer_non_stream_response(&headers, &observer_context) { if should_buffer_non_stream_response(&headers, &observer_context) {
let original_headers = headers.clone(); let original_headers = headers.clone();
match buffer_non_sse_upstream_body(response, started_at).await { match buffer_non_sse_upstream_body(response, started_at, stream_first_byte_timeout).await {
Ok(buffered) => { Ok(buffered) => {
let mut response_headers = original_headers; let mut response_headers = original_headers;
let mut response_body = Bytes::from(buffered.body_bytes); let mut response_body = Bytes::from(buffered.body_bytes);
@@ -134,6 +136,7 @@ pub(crate) fn build_direct_execution_frame_stream(
message, message,
ttfb_ms, ttfb_ms,
upstream_bytes, upstream_bytes,
first_byte_timeout,
}) => { }) => {
match encode_headers_frame(status_code, original_headers) { match encode_headers_frame(status_code, original_headers) {
Ok(frame) => yield Ok(frame), Ok(frame) => yield Ok(frame),
@@ -142,7 +145,12 @@ pub(crate) fn build_direct_execution_frame_stream(
return; return;
} }
} }
match encode_error_frame(status_code, message) { let error_frame = if let Some(timeout) = first_byte_timeout {
encode_first_byte_timeout_frame(timeout)
} else {
encode_error_frame(status_code, message)
};
match error_frame {
Ok(frame) => yield Ok(frame), Ok(frame) => yield Ok(frame),
Err(err) => { Err(err) => {
yield Err(err); yield Err(err);
@@ -183,7 +191,33 @@ pub(crate) fn build_direct_execution_frame_stream(
match response { match response {
DirectUpstreamResponse::Reqwest(response) => { DirectUpstreamResponse::Reqwest(response) => {
let mut bytes_stream = response.bytes_stream(); let mut bytes_stream = response.bytes_stream();
while let Some(item) = bytes_stream.next().await { loop {
let item = if ttfb_ms.is_none() {
match await_stream_first_byte(
bytes_stream.next(),
started_at,
stream_first_byte_timeout,
)
.await
{
Ok(item) => item,
Err(timeout) => {
match encode_first_byte_timeout_frame(timeout) {
Ok(frame) => yield Ok(frame),
Err(err) => {
yield Err(err);
return;
}
}
break;
}
}
} else {
bytes_stream.next().await
};
let Some(item) = item else {
break;
};
match item { match item {
Ok(chunk) => { Ok(chunk) => {
if ttfb_ms.is_none() { if ttfb_ms.is_none() {
@@ -239,7 +273,33 @@ pub(crate) fn build_direct_execution_frame_stream(
} }
DirectUpstreamResponse::BrowserWreq(response) => { DirectUpstreamResponse::BrowserWreq(response) => {
let mut bytes_stream = response.bytes_stream(); let mut bytes_stream = response.bytes_stream();
while let Some(item) = bytes_stream.next().await { loop {
let item = if ttfb_ms.is_none() {
match await_stream_first_byte(
bytes_stream.next(),
started_at,
stream_first_byte_timeout,
)
.await
{
Ok(item) => item,
Err(timeout) => {
match encode_first_byte_timeout_frame(timeout) {
Ok(frame) => yield Ok(frame),
Err(err) => {
yield Err(err);
return;
}
}
break;
}
}
} else {
bytes_stream.next().await
};
let Some(item) = item else {
break;
};
match item { match item {
Ok(chunk) => { Ok(chunk) => {
if ttfb_ms.is_none() { if ttfb_ms.is_none() {
@@ -294,7 +354,30 @@ pub(crate) fn build_direct_execution_frame_stream(
} }
} }
DirectUpstreamResponse::LocalTunnel(mut response) => loop { DirectUpstreamResponse::LocalTunnel(mut response) => loop {
match response.next_chunk().await { let item = if ttfb_ms.is_none() {
match await_stream_first_byte(
response.next_chunk(),
started_at,
stream_first_byte_timeout,
)
.await
{
Ok(item) => item,
Err(timeout) => {
match encode_first_byte_timeout_frame(timeout) {
Ok(frame) => yield Ok(frame),
Err(err) => {
yield Err(err);
return;
}
}
break;
}
}
} else {
response.next_chunk().await
};
match item {
Ok(Some(chunk)) => { Ok(Some(chunk)) => {
if ttfb_ms.is_none() { if ttfb_ms.is_none() {
ttfb_ms = Some(started_at.elapsed().as_millis() as u64); ttfb_ms = Some(started_at.elapsed().as_millis() as u64);
@@ -428,6 +511,44 @@ fn encode_error_frame(status_code: u16, message: String) -> Result<Bytes, IoErro
}) })
} }
fn encode_first_byte_timeout_frame(timeout: Duration) -> Result<Bytes, IoError> {
encode_stream_frame_ndjson(&StreamFrame {
frame_type: StreamFrameType::Error,
payload: StreamFramePayload::Error {
error: ExecutionError {
kind: ExecutionErrorKind::FirstByteTimeout,
phase: ExecutionPhase::FirstByte,
message: stream_first_byte_timeout_message(timeout),
upstream_status: Some(504),
retryable: true,
failover_recommended: true,
},
},
})
}
async fn await_stream_first_byte<T, F>(
future: F,
started_at: Instant,
timeout: Option<Duration>,
) -> Result<T, Duration>
where
F: Future<Output = T>,
{
let Some(timeout) = timeout else {
return Ok(future.await);
};
let Some(remaining) = timeout.checked_sub(started_at.elapsed()) else {
return Err(timeout);
};
if remaining.is_zero() {
return Err(timeout);
}
tokio::time::timeout(remaining, future)
.await
.map_err(|_| timeout)
}
struct BufferedUpstreamBody { struct BufferedUpstreamBody {
body_bytes: Vec<u8>, body_bytes: Vec<u8>,
ttfb_ms: Option<u64>, ttfb_ms: Option<u64>,
@@ -438,6 +559,7 @@ struct BufferedUpstreamBodyError {
message: String, message: String,
ttfb_ms: Option<u64>, ttfb_ms: Option<u64>,
upstream_bytes: u64, upstream_bytes: u64,
first_byte_timeout: Option<Duration>,
} }
fn response_headers_indicate_sse(headers: &BTreeMap<String, String>) -> bool { fn response_headers_indicate_sse(headers: &BTreeMap<String, String>) -> bool {
@@ -477,6 +599,7 @@ fn should_buffer_non_stream_response(
async fn buffer_non_sse_upstream_body( async fn buffer_non_sse_upstream_body(
response: DirectUpstreamResponse, response: DirectUpstreamResponse,
started_at: Instant, started_at: Instant,
stream_first_byte_timeout: Option<Duration>,
) -> Result<BufferedUpstreamBody, BufferedUpstreamBodyError> { ) -> Result<BufferedUpstreamBody, BufferedUpstreamBodyError> {
let mut body_bytes = Vec::new(); let mut body_bytes = Vec::new();
let mut upstream_bytes = 0u64; let mut upstream_bytes = 0u64;
@@ -485,7 +608,31 @@ async fn buffer_non_sse_upstream_body(
match response { match response {
DirectUpstreamResponse::Reqwest(response) => { DirectUpstreamResponse::Reqwest(response) => {
let mut bytes_stream = response.bytes_stream(); let mut bytes_stream = response.bytes_stream();
while let Some(item) = bytes_stream.next().await { loop {
let item = if ttfb_ms.is_none() {
match await_stream_first_byte(
bytes_stream.next(),
started_at,
stream_first_byte_timeout,
)
.await
{
Ok(item) => item,
Err(timeout) => {
return Err(BufferedUpstreamBodyError {
message: stream_first_byte_timeout_message(timeout),
ttfb_ms,
upstream_bytes,
first_byte_timeout: Some(timeout),
});
}
}
} else {
bytes_stream.next().await
};
let Some(item) = item else {
break;
};
match item { match item {
Ok(chunk) => { Ok(chunk) => {
if ttfb_ms.is_none() { if ttfb_ms.is_none() {
@@ -507,6 +654,7 @@ async fn buffer_non_sse_upstream_body(
message, message,
ttfb_ms, ttfb_ms,
upstream_bytes, upstream_bytes,
first_byte_timeout: None,
}); });
} }
} }
@@ -514,7 +662,31 @@ async fn buffer_non_sse_upstream_body(
} }
DirectUpstreamResponse::BrowserWreq(response) => { DirectUpstreamResponse::BrowserWreq(response) => {
let mut bytes_stream = response.bytes_stream(); let mut bytes_stream = response.bytes_stream();
while let Some(item) = bytes_stream.next().await { loop {
let item = if ttfb_ms.is_none() {
match await_stream_first_byte(
bytes_stream.next(),
started_at,
stream_first_byte_timeout,
)
.await
{
Ok(item) => item,
Err(timeout) => {
return Err(BufferedUpstreamBodyError {
message: stream_first_byte_timeout_message(timeout),
ttfb_ms,
upstream_bytes,
first_byte_timeout: Some(timeout),
});
}
}
} else {
bytes_stream.next().await
};
let Some(item) = item else {
break;
};
match item { match item {
Ok(chunk) => { Ok(chunk) => {
if ttfb_ms.is_none() { if ttfb_ms.is_none() {
@@ -536,13 +708,35 @@ async fn buffer_non_sse_upstream_body(
message, message,
ttfb_ms, ttfb_ms,
upstream_bytes, upstream_bytes,
first_byte_timeout: None,
}); });
} }
} }
} }
} }
DirectUpstreamResponse::LocalTunnel(mut response) => loop { DirectUpstreamResponse::LocalTunnel(mut response) => loop {
match response.next_chunk().await { let item = if ttfb_ms.is_none() {
match await_stream_first_byte(
response.next_chunk(),
started_at,
stream_first_byte_timeout,
)
.await
{
Ok(item) => item,
Err(timeout) => {
return Err(BufferedUpstreamBodyError {
message: stream_first_byte_timeout_message(timeout),
ttfb_ms,
upstream_bytes,
first_byte_timeout: Some(timeout),
});
}
}
} else {
response.next_chunk().await
};
match item {
Ok(Some(chunk)) => { Ok(Some(chunk)) => {
if ttfb_ms.is_none() { if ttfb_ms.is_none() {
ttfb_ms = Some(started_at.elapsed().as_millis() as u64); ttfb_ms = Some(started_at.elapsed().as_millis() as u64);
@@ -563,6 +757,7 @@ async fn buffer_non_sse_upstream_body(
message, message,
ttfb_ms, ttfb_ms,
upstream_bytes, upstream_bytes,
first_byte_timeout: None,
}); });
} }
} }
@@ -761,6 +956,7 @@ mod tests {
use base64::Engine as _; use base64::Engine as _;
use futures_util::StreamExt; use futures_util::StreamExt;
use serde_json::Value; use serde_json::Value;
use tokio::io::{AsyncReadExt, AsyncWriteExt};
use tokio::sync::watch; use tokio::sync::watch;
use super::{ use super::{
@@ -903,6 +1099,93 @@ mod tests {
); );
} }
#[tokio::test]
async fn direct_execution_frame_stream_applies_first_byte_timeout_after_headers() {
let listener = crate::test_support::bind_loopback_listener()
.await
.expect("listener should bind");
let addr = listener.local_addr().expect("local addr should resolve");
let server = tokio::spawn(async move {
let (mut socket, _) = listener.accept().await.expect("client should connect");
let mut request = [0_u8; 1024];
let _ = socket
.read(&mut request)
.await
.expect("request should read");
socket
.write_all(
b"HTTP/1.1 200 OK\r\ncontent-type: text/event-stream\r\ntransfer-encoding: chunked\r\n\r\n",
)
.await
.expect("headers should write");
socket.flush().await.expect("headers should flush");
tokio::time::sleep(Duration::from_millis(200)).await;
let _ = socket.write_all(b"d\r\ndata: hello\n\n\r\n0\r\n\r\n").await;
});
let execution = DirectSyncExecutionRuntime::new()
.execute_stream(&ExecutionPlan {
request_id: "req-stream-first-byte-timeout".into(),
candidate_id: Some("cand-stream-first-byte-timeout".into()),
provider_name: Some("openai".into()),
provider_id: "prov-1".into(),
endpoint_id: "ep-1".into(),
key_id: "key-1".into(),
method: "POST".into(),
url: format!("http://{addr}/chat"),
headers: BTreeMap::from([("content-type".into(), "application/json".into())]),
content_type: Some("application/json".into()),
content_encoding: None,
body: RequestBody::from_json(serde_json::json!({"stream": true})),
stream: true,
client_api_format: "openai:chat".into(),
provider_api_format: "openai:chat".into(),
model_name: Some("gpt-5".into()),
proxy: None,
transport_profile: None,
timeouts: Some(ExecutionTimeouts {
first_byte_ms: Some(50),
total_ms: Some(5_000),
..ExecutionTimeouts::default()
}),
})
.await
.expect("stream execution should receive response headers");
let frames = build_direct_execution_frame_stream(execution)
.map(|item| item.expect("frame should encode"))
.collect::<Vec<_>>()
.await
.into_iter()
.map(|bytes| String::from_utf8(bytes.to_vec()).expect("frame should be utf8"))
.collect::<Vec<_>>();
server.abort();
let error_frame = frames
.iter()
.map(|line| serde_json::from_str::<Value>(line).expect("frame should parse"))
.find(|frame| frame.get("type").and_then(Value::as_str) == Some("error"))
.expect("timeout should emit an error frame");
assert_eq!(
error_frame
.get("payload")
.and_then(|payload| payload.get("error"))
.and_then(|error| error.get("kind"))
.and_then(Value::as_str),
Some("first_byte_timeout")
);
assert!(error_frame
.get("payload")
.and_then(|payload| payload.get("error"))
.and_then(|error| error.get("message"))
.and_then(Value::as_str)
.is_some_and(
|message| message.contains("provider stream first byte timeout after 50 ms")
));
}
#[tokio::test] #[tokio::test]
async fn direct_execution_frame_stream_emits_telemetry_before_first_data_frame() { async fn direct_execution_frame_stream_emits_telemetry_before_first_data_frame() {
let listener = crate::test_support::bind_loopback_listener() let listener = crate::test_support::bind_loopback_listener()
@@ -1,5 +1,6 @@
use std::collections::BTreeMap; use std::collections::BTreeMap;
use std::error::Error as _; use std::error::Error as _;
use std::future::Future;
use std::io::Read; use std::io::Read;
use std::io::Write; use std::io::Write;
use std::time::{Duration, Instant}; use std::time::{Duration, Instant};
@@ -226,6 +227,7 @@ pub(crate) struct DirectUpstreamStreamExecution {
pub(crate) stream_summary_report_context: Value, pub(crate) stream_summary_report_context: Value,
pub(crate) response: DirectUpstreamResponse, pub(crate) response: DirectUpstreamResponse,
pub(crate) started_at: Instant, pub(crate) started_at: Instant,
pub(crate) stream_first_byte_timeout: Option<Duration>,
} }
impl DirectSyncExecutionRuntime { impl DirectSyncExecutionRuntime {
@@ -240,32 +242,39 @@ impl DirectSyncExecutionRuntime {
let body_bytes = build_request_body(plan)?; let body_bytes = build_request_body(plan)?;
let started_at = Instant::now(); let started_at = Instant::now();
let response = send_request(plan, body_bytes).await?; with_non_stream_total_timeout(plan, async move {
let ttfb_ms = started_at.elapsed().as_millis() as u64; let response = send_request_inner(plan, body_bytes, false).await?;
let status_code = response.status_code(); let ttfb_ms = started_at.elapsed().as_millis() as u64;
let headers = response.headers(); let status_code = response.status_code();
let body_bytes = response.bytes().await?; let headers = response.headers();
let decoded_body_bytes = decode_response_body_bytes(&headers, &body_bytes) let body_bytes = response.bytes().await?;
.unwrap_or_else(|| body_bytes.to_vec()); let decoded_body_bytes = decode_response_body_bytes(&headers, &body_bytes)
let elapsed_ms = started_at.elapsed().as_millis() as u64; .unwrap_or_else(|| body_bytes.to_vec());
let upstream_bytes = body_bytes.len() as u64; let elapsed_ms = started_at.elapsed().as_millis() as u64;
let upstream_bytes = body_bytes.len() as u64;
let body = let body = build_execution_response_body(
build_execution_response_body(&headers, &body_bytes, &decoded_body_bytes, plan.stream)?; &headers,
&body_bytes,
&decoded_body_bytes,
plan.stream,
)?;
Ok(ExecutionResult { Ok(ExecutionResult {
request_id: plan.request_id.clone(), request_id: plan.request_id.clone(),
candidate_id: plan.candidate_id.clone(), candidate_id: plan.candidate_id.clone(),
status_code, status_code,
headers, headers,
body, body,
telemetry: Some(ExecutionTelemetry { telemetry: Some(ExecutionTelemetry {
ttfb_ms: Some(ttfb_ms), ttfb_ms: Some(ttfb_ms),
elapsed_ms: Some(elapsed_ms), elapsed_ms: Some(elapsed_ms),
upstream_bytes: Some(upstream_bytes), upstream_bytes: Some(upstream_bytes),
}), }),
error: None, error: None,
})
}) })
.await
} }
pub(crate) async fn execute_stream( pub(crate) async fn execute_stream(
@@ -294,6 +303,7 @@ impl DirectSyncExecutionRuntime {
stream_summary_report_context, stream_summary_report_context,
response: response.into_direct_upstream_response(), response: response.into_direct_upstream_response(),
started_at, started_at,
stream_first_byte_timeout: resolve_stream_first_byte_timeout(plan),
}) })
} }
} }
@@ -405,6 +415,7 @@ pub(crate) async fn execute_stream_plan_via_local_tunnel(
stream_summary_report_context: build_stream_summary_report_context(plan), stream_summary_report_context: build_stream_summary_report_context(plan),
response: DirectUpstreamResponse::LocalTunnel(response), response: DirectUpstreamResponse::LocalTunnel(response),
started_at, started_at,
stream_first_byte_timeout: resolve_stream_first_byte_timeout(plan),
})) }))
} }
@@ -481,6 +492,13 @@ fn manual_proxy_node_id(proxy: Option<&ProxySnapshot>) -> Option<String> {
async fn execute_sync_plan_via_local_tunnel( async fn execute_sync_plan_via_local_tunnel(
state: &AppState, state: &AppState,
plan: &ExecutionPlan, plan: &ExecutionPlan,
) -> Result<ExecutionResult, ExecutionRuntimeTransportError> {
with_non_stream_total_timeout(plan, execute_sync_plan_via_local_tunnel_inner(state, plan)).await
}
async fn execute_sync_plan_via_local_tunnel_inner(
state: &AppState,
plan: &ExecutionPlan,
) -> Result<ExecutionResult, ExecutionRuntimeTransportError> { ) -> Result<ExecutionResult, ExecutionRuntimeTransportError> {
let node_id = resolve_local_tunnel_node_id(state, plan.proxy.as_ref()).ok_or_else(|| { let node_id = resolve_local_tunnel_node_id(state, plan.proxy.as_ref()).ok_or_else(|| {
ExecutionRuntimeTransportError::RelayError("local tunnel node unavailable".to_string()) ExecutionRuntimeTransportError::RelayError("local tunnel node unavailable".to_string())
@@ -616,6 +634,14 @@ fn build_direct_tunnel_request_meta(
pub(crate) async fn send_request( pub(crate) async fn send_request(
plan: &ExecutionPlan, plan: &ExecutionPlan,
body_bytes: Vec<u8>, body_bytes: Vec<u8>,
) -> Result<DirectHttpResponse, ExecutionRuntimeTransportError> {
send_request_inner(plan, body_bytes, true).await
}
async fn send_request_inner(
plan: &ExecutionPlan,
body_bytes: Vec<u8>,
apply_request_total_timeout: bool,
) -> Result<DirectHttpResponse, ExecutionRuntimeTransportError> { ) -> Result<DirectHttpResponse, ExecutionRuntimeTransportError> {
if let Some(detail) = gateway_frontdoor_self_loop_guard_error(plan.url.as_str()) { if let Some(detail) = gateway_frontdoor_self_loop_guard_error(plan.url.as_str()) {
return Err(ExecutionRuntimeTransportError::UpstreamRequest(detail)); return Err(ExecutionRuntimeTransportError::UpstreamRequest(detail));
@@ -628,11 +654,12 @@ pub(crate) async fn send_request(
plan.content_encoding.as_deref(), plan.content_encoding.as_deref(),
plan.body.body_bytes_b64.is_some(), plan.body.body_bytes_b64.is_some(),
)?; )?;
let total_timeout = plan let total_timeout = if apply_request_total_timeout {
.timeouts resolve_non_stream_total_timeout(plan)
.as_ref() } else {
.and_then(|timeouts| timeouts.total_ms) None
.map(Duration::from_millis); };
let stream_first_byte_timeout = resolve_stream_first_byte_timeout(plan);
if transport_profile_uses_browser_wreq(plan.transport_profile.as_ref()) { if transport_profile_uses_browser_wreq(plan.transport_profile.as_ref()) {
return send_via_browser_wreq_transport( return send_via_browser_wreq_transport(
@@ -641,7 +668,9 @@ pub(crate) async fn send_request(
headers, headers,
body_bytes, body_bytes,
total_timeout, total_timeout,
stream_first_byte_timeout,
transport_controls, transport_controls,
apply_request_total_timeout,
) )
.await; .await;
} }
@@ -654,6 +683,7 @@ pub(crate) async fn send_request(
body_bytes, body_bytes,
&node_id, &node_id,
total_timeout, total_timeout,
stream_first_byte_timeout,
transport_controls, transport_controls,
) )
.await .await
@@ -671,13 +701,9 @@ pub(crate) async fn send_request(
if let Some(timeout) = total_timeout { if let Some(timeout) = total_timeout {
request = request.timeout(timeout); request = request.timeout(timeout);
} }
request send_reqwest_request(request, stream_first_byte_timeout)
.send()
.await .await
.map(DirectHttpResponse::Reqwest) .map(DirectHttpResponse::Reqwest)
.map_err(|err| {
ExecutionRuntimeTransportError::UpstreamRequest(format_upstream_request_error(&err))
})
} }
pub(crate) enum DirectHttpResponse { pub(crate) enum DirectHttpResponse {
@@ -731,7 +757,9 @@ async fn send_via_browser_wreq_transport(
headers: HeaderMap, headers: HeaderMap,
body_bytes: Vec<u8>, body_bytes: Vec<u8>,
total_timeout: Option<Duration>, total_timeout: Option<Duration>,
stream_first_byte_timeout: Option<Duration>,
transport_controls: ExecutionTransportControls, transport_controls: ExecutionTransportControls,
apply_request_total_timeout: bool,
) -> Result<DirectHttpResponse, ExecutionRuntimeTransportError> { ) -> Result<DirectHttpResponse, ExecutionRuntimeTransportError> {
let profile = plan.transport_profile.as_ref().ok_or_else(|| { let profile = plan.transport_profile.as_ref().ok_or_else(|| {
ExecutionRuntimeTransportError::UnsupportedTransportProfile(String::new()) ExecutionRuntimeTransportError::UnsupportedTransportProfile(String::new())
@@ -741,6 +769,7 @@ async fn send_via_browser_wreq_transport(
plan.proxy.as_ref(), plan.proxy.as_ref(),
profile, profile,
transport_controls, transport_controls,
apply_request_total_timeout && !plan.stream,
)?; )?;
let method = wreq::Method::from_bytes(method.as_str().as_bytes()) let method = wreq::Method::from_bytes(method.as_str().as_bytes())
.map_err(ExecutionRuntimeTransportError::InvalidMethod)?; .map_err(ExecutionRuntimeTransportError::InvalidMethod)?;
@@ -751,15 +780,9 @@ async fn send_via_browser_wreq_transport(
if let Some(timeout) = total_timeout { if let Some(timeout) = total_timeout {
request = request.timeout(timeout); request = request.timeout(timeout);
} }
request send_wreq_request(request, stream_first_byte_timeout)
.send()
.await .await
.map(DirectHttpResponse::BrowserWreq) .map(DirectHttpResponse::BrowserWreq)
.map_err(|err| {
ExecutionRuntimeTransportError::UpstreamRequest(format_wreq_upstream_request_error(
&err,
))
})
} }
async fn send_via_tunnel_relay( async fn send_via_tunnel_relay(
@@ -769,6 +792,7 @@ async fn send_via_tunnel_relay(
body_bytes: Vec<u8>, body_bytes: Vec<u8>,
node_id: &str, node_id: &str,
total_timeout: Option<Duration>, total_timeout: Option<Duration>,
stream_first_byte_timeout: Option<Duration>,
transport_controls: ExecutionTransportControls, transport_controls: ExecutionTransportControls,
) -> Result<reqwest::Response, ExecutionRuntimeTransportError> { ) -> Result<reqwest::Response, ExecutionRuntimeTransportError> {
let client = build_relay_client(plan.timeouts.as_ref())?; let client = build_relay_client(plan.timeouts.as_ref())?;
@@ -822,7 +846,7 @@ async fn send_via_tunnel_relay(
} }
let first_byte_timeout = if plan.stream { let first_byte_timeout = if plan.stream {
resolve_tunnel_first_byte_timeout(plan) stream_first_byte_timeout.or_else(|| resolve_tunnel_first_byte_timeout(plan))
} else { } else {
None None
}; };
@@ -1022,6 +1046,101 @@ fn resolve_tunnel_first_byte_timeout(plan: &ExecutionPlan) -> Option<Duration> {
}) })
} }
fn resolve_non_stream_total_timeout(plan: &ExecutionPlan) -> Option<Duration> {
if plan.stream {
return None;
}
plan.timeouts
.as_ref()
.and_then(|timeouts| timeouts.total_ms)
.map(|value| Duration::from_millis(value.max(1)))
}
pub(crate) fn resolve_stream_first_byte_timeout(plan: &ExecutionPlan) -> Option<Duration> {
if !plan.stream {
return None;
}
plan.timeouts
.as_ref()
.and_then(|timeouts| timeouts.first_byte_ms.or(timeouts.total_ms))
.map(|value| Duration::from_millis(value.max(1)))
}
pub(crate) async fn with_non_stream_total_timeout<T, F>(
plan: &ExecutionPlan,
future: F,
) -> Result<T, ExecutionRuntimeTransportError>
where
F: Future<Output = Result<T, ExecutionRuntimeTransportError>>,
{
let Some(timeout) = resolve_non_stream_total_timeout(plan) else {
return future.await;
};
match tokio::time::timeout(timeout, future).await {
Ok(result) => result,
Err(_) => Err(ExecutionRuntimeTransportError::UpstreamRequest(
non_stream_total_timeout_message(timeout),
)),
}
}
async fn send_reqwest_request(
request: reqwest::RequestBuilder,
stream_first_byte_timeout: Option<Duration>,
) -> Result<reqwest::Response, ExecutionRuntimeTransportError> {
if let Some(timeout) = stream_first_byte_timeout {
return match tokio::time::timeout(timeout, request.send()).await {
Ok(Ok(response)) => Ok(response),
Ok(Err(error)) => Err(ExecutionRuntimeTransportError::UpstreamRequest(
format_upstream_request_error(&error),
)),
Err(_) => Err(ExecutionRuntimeTransportError::UpstreamRequest(
stream_first_byte_timeout_message(timeout),
)),
};
}
request.send().await.map_err(|err| {
ExecutionRuntimeTransportError::UpstreamRequest(format_upstream_request_error(&err))
})
}
async fn send_wreq_request(
request: wreq::RequestBuilder,
stream_first_byte_timeout: Option<Duration>,
) -> Result<wreq::Response, ExecutionRuntimeTransportError> {
if let Some(timeout) = stream_first_byte_timeout {
return match tokio::time::timeout(timeout, request.send()).await {
Ok(Ok(response)) => Ok(response),
Ok(Err(error)) => Err(ExecutionRuntimeTransportError::UpstreamRequest(
format_wreq_upstream_request_error(&error),
)),
Err(_) => Err(ExecutionRuntimeTransportError::UpstreamRequest(
stream_first_byte_timeout_message(timeout),
)),
};
}
request.send().await.map_err(|err| {
ExecutionRuntimeTransportError::UpstreamRequest(format_wreq_upstream_request_error(&err))
})
}
fn non_stream_total_timeout_message(timeout: Duration) -> String {
format!(
"provider non-stream request total timeout after {} ms",
timeout.as_millis()
)
}
pub(crate) fn stream_first_byte_timeout_message(timeout: Duration) -> String {
format!(
"provider stream first byte timeout after {} ms",
timeout.as_millis()
)
}
fn resolve_tunnel_timeout_metadata(plan: &ExecutionPlan) -> TunnelTimeoutMetadata { fn resolve_tunnel_timeout_metadata(plan: &ExecutionPlan) -> TunnelTimeoutMetadata {
TunnelTimeoutMetadata { TunnelTimeoutMetadata {
request_timeout_ms: plan request_timeout_ms: plan
@@ -1128,6 +1247,7 @@ pub(crate) fn build_browser_wreq_client(
proxy: Option<&ProxySnapshot>, proxy: Option<&ProxySnapshot>,
transport_profile: &ResolvedTransportProfile, transport_profile: &ResolvedTransportProfile,
transport_controls: ExecutionTransportControls, transport_controls: ExecutionTransportControls,
apply_total_timeout: bool,
) -> Result<wreq::Client, ExecutionRuntimeTransportError> { ) -> Result<wreq::Client, ExecutionRuntimeTransportError> {
let emulation = browser_wreq_emulation_from_profile(transport_profile)?; let emulation = browser_wreq_emulation_from_profile(transport_profile)?;
let mut builder = wreq::Client::builder().emulation(emulation); let mut builder = wreq::Client::builder().emulation(emulation);
@@ -1143,8 +1263,10 @@ pub(crate) fn build_browser_wreq_client(
if let Some(connect_ms) = timeouts.and_then(|timeouts| timeouts.connect_ms) { if let Some(connect_ms) = timeouts.and_then(|timeouts| timeouts.connect_ms) {
builder = builder.connect_timeout(Duration::from_millis(connect_ms)); builder = builder.connect_timeout(Duration::from_millis(connect_ms));
} }
if let Some(total_ms) = timeouts.and_then(|timeouts| timeouts.total_ms) { if apply_total_timeout {
builder = builder.timeout(Duration::from_millis(total_ms)); if let Some(total_ms) = timeouts.and_then(|timeouts| timeouts.total_ms) {
builder = builder.timeout(Duration::from_millis(total_ms));
}
} }
if let Some(read_ms) = timeouts.and_then(|timeouts| timeouts.read_ms) { if let Some(read_ms) = timeouts.and_then(|timeouts| timeouts.read_ms) {
builder = builder.read_timeout(Duration::from_millis(read_ms)); builder = builder.read_timeout(Duration::from_millis(read_ms));
@@ -1764,6 +1886,34 @@ mod tests {
} }
} }
fn direct_timeout_plan(
url: String,
stream: bool,
timeouts: ExecutionTimeouts,
) -> ExecutionPlan {
ExecutionPlan {
request_id: "req-direct-timeout".into(),
candidate_id: None,
provider_name: Some("provider".into()),
provider_id: "prov-direct-timeout".into(),
endpoint_id: "ep-direct-timeout".into(),
key_id: "key-direct-timeout".into(),
method: "POST".into(),
url,
headers: BTreeMap::from([("content-type".into(), "application/json".into())]),
content_type: Some("application/json".into()),
content_encoding: None,
body: RequestBody::from_json(json!({"model": "gpt-4.1"})),
stream,
client_api_format: "openai:chat".into(),
provider_api_format: "openai:chat".into(),
model_name: Some("gpt-4.1".into()),
proxy: None,
transport_profile: None,
timeouts: Some(timeouts),
}
}
fn tunnel_proxy_snapshot(base_url: String) -> ProxySnapshot { fn tunnel_proxy_snapshot(base_url: String) -> ProxySnapshot {
ProxySnapshot { ProxySnapshot {
enabled: Some(true), enabled: Some(true),
@@ -1894,6 +2044,147 @@ mod tests {
); );
} }
#[tokio::test]
async fn direct_sync_execution_runtime_applies_non_stream_total_timeout_to_body() {
let listener = crate::test_support::bind_loopback_listener()
.await
.expect("listener should bind");
let addr = listener.local_addr().expect("local addr should resolve");
let app = Router::new().route(
"/chat",
post(|| async {
let body = Body::from_stream(async_stream::stream! {
tokio::time::sleep(std::time::Duration::from_millis(200)).await;
yield Ok::<Bytes, std::io::Error>(Bytes::from_static(br#"{"ok":true}"#));
});
axum::response::Response::builder()
.status(http::StatusCode::OK)
.header("content-type", "application/json")
.body(body)
.expect("response should build")
}),
);
let server = tokio::spawn(async move {
axum::serve(listener, app)
.await
.expect("test server should run");
});
let result = DirectSyncExecutionRuntime::new()
.execute_sync(&direct_timeout_plan(
format!("http://{addr}/chat"),
false,
ExecutionTimeouts {
total_ms: Some(50),
..ExecutionTimeouts::default()
},
))
.await;
server.abort();
let error = match result {
Ok(_) => panic!("non-stream body should hit total timeout"),
Err(error) => error,
};
assert!(
error
.to_string()
.contains("provider non-stream request total timeout after 50 ms"),
"unexpected error: {error}"
);
}
#[tokio::test]
async fn direct_stream_execution_runtime_applies_first_byte_timeout() {
let listener = crate::test_support::bind_loopback_listener()
.await
.expect("listener should bind");
let addr = listener.local_addr().expect("local addr should resolve");
let app = Router::new().route(
"/chat",
post(|| async {
tokio::time::sleep(std::time::Duration::from_millis(200)).await;
axum::response::Response::builder()
.status(http::StatusCode::OK)
.header("content-type", "text/event-stream")
.body(Body::from(Bytes::from_static(b"data: {}\n\n")))
.expect("response should build")
}),
);
let server = tokio::spawn(async move {
axum::serve(listener, app)
.await
.expect("test server should run");
});
let result = DirectSyncExecutionRuntime::new()
.execute_stream(&direct_timeout_plan(
format!("http://{addr}/chat"),
true,
ExecutionTimeouts {
first_byte_ms: Some(50),
total_ms: Some(5_000),
..ExecutionTimeouts::default()
},
))
.await;
server.abort();
let error = match result {
Ok(_) => panic!("stream should hit first-byte timeout"),
Err(error) => error,
};
assert!(
error
.to_string()
.contains("provider stream first byte timeout after 50 ms"),
"unexpected error: {error}"
);
}
#[tokio::test]
async fn direct_stream_execution_runtime_prefers_first_byte_timeout_over_total_timeout() {
let listener = crate::test_support::bind_loopback_listener()
.await
.expect("listener should bind");
let addr = listener.local_addr().expect("local addr should resolve");
let app = Router::new().route(
"/chat",
post(|| async {
tokio::time::sleep(std::time::Duration::from_millis(80)).await;
axum::response::Response::builder()
.status(http::StatusCode::OK)
.header("content-type", "text/event-stream")
.body(Body::from(Bytes::from_static(b"data: {}\n\n")))
.expect("response should build")
}),
);
let server = tokio::spawn(async move {
axum::serve(listener, app)
.await
.expect("test server should run");
});
let execution = DirectSyncExecutionRuntime::new()
.execute_stream(&direct_timeout_plan(
format!("http://{addr}/chat"),
true,
ExecutionTimeouts {
first_byte_ms: Some(250),
total_ms: Some(25),
..ExecutionTimeouts::default()
},
))
.await
.expect("stream should use first-byte timeout instead of total timeout");
server.abort();
assert_eq!(execution.status_code, http::StatusCode::OK.as_u16());
}
#[tokio::test] #[tokio::test]
async fn direct_sync_execution_runtime_routes_browser_wreq_transport_in_process() { async fn direct_sync_execution_runtime_routes_browser_wreq_transport_in_process() {
async fn browser_upstream(headers: AxumHeaderMap, body: Bytes) -> axum::response::Response { async fn browser_upstream(headers: AxumHeaderMap, body: Bytes) -> axum::response::Response {
@@ -2007,6 +2298,7 @@ mod tests {
None, None,
&profile, &profile,
ExecutionTransportControls::default(), ExecutionTransportControls::default(),
true,
) { ) {
Ok(_) => panic!("unknown browser profile should fail loudly"), Ok(_) => panic!("unknown browser profile should fail loudly"),
Err(error) => error, Err(error) => error,
@@ -36,7 +36,7 @@ use tracing::{debug, error, info, warn};
use uuid::Uuid; use uuid::Uuid;
use super::ndjson::encode_stream_frame_ndjson; use super::ndjson::encode_stream_frame_ndjson;
use super::transport::ExecutionRuntimeTransportError; use super::transport::{with_non_stream_total_timeout, ExecutionRuntimeTransportError};
use crate::AppState; use crate::AppState;
const LS_SERVICE: &str = "/exa.language_server_pb.LanguageServerService"; const LS_SERVICE: &str = "/exa.language_server_pb.LanguageServerService";
@@ -215,66 +215,69 @@ pub(crate) async fn maybe_execute_windsurf_sync(
let Some(input) = detect_windsurf_request(plan, report_context) else { let Some(input) = detect_windsurf_request(plan, report_context) else {
return Ok(None); return Ok(None);
}; };
let key_upstream_metadata = read_windsurf_key_upstream_metadata(state, plan).await; with_non_stream_total_timeout(plan, async move {
let prepared = prepare_windsurf_cascade(plan, input, key_upstream_metadata).await?; let key_upstream_metadata = read_windsurf_key_upstream_metadata(state, plan).await;
let started_at = Instant::now(); let prepared = prepare_windsurf_cascade(plan, input, key_upstream_metadata).await?;
let mut deltas = Vec::new(); let started_at = Instant::now();
let poll_result = poll_windsurf_cascade_with_transport_recovery(&prepared, |event| { let mut deltas = Vec::new();
if let WindsurfPollEvent::TextDelta(delta) = event { let poll_result = poll_windsurf_cascade_with_transport_recovery(&prepared, |event| {
deltas.push(sanitize_windsurf_text(&delta)); if let WindsurfPollEvent::TextDelta(delta) = event {
deltas.push(sanitize_windsurf_text(&delta));
}
Ok(())
})
.await?;
let elapsed_ms = started_at.elapsed().as_millis() as u64;
let content = deltas.concat();
let parsed_tool_calls = parse_and_filter_windsurf_tool_calls(&content, &prepared.input);
let mut tool_calls = poll_result.native_tool_calls;
tool_calls.extend(parsed_tool_calls.tool_calls);
let has_tool_calls = !tool_calls.is_empty();
let message = if has_tool_calls {
json!({
"role": "assistant",
"content": Value::Null,
"tool_calls": openai_tool_call_values(&tool_calls),
})
} else {
json!({
"role": "assistant",
"content": content,
})
};
let mut body_json = json!({
"id": format!("chatcmpl-{}", prepared.request_id),
"object": "chat.completion",
"created": current_unix_secs(),
"model": prepared.model,
"choices": [{
"index": 0,
"message": message,
"finish_reason": if has_tool_calls { "tool_calls" } else { "stop" },
}],
});
if let Some(usage) = poll_result.usage {
body_json["usage"] = windsurf_openai_usage_json(&usage);
} }
Ok(())
})
.await?;
let elapsed_ms = started_at.elapsed().as_millis() as u64;
let content = deltas.concat();
let parsed_tool_calls = parse_and_filter_windsurf_tool_calls(&content, &prepared.input);
let mut tool_calls = poll_result.native_tool_calls;
tool_calls.extend(parsed_tool_calls.tool_calls);
let has_tool_calls = !tool_calls.is_empty();
let message = if has_tool_calls {
json!({
"role": "assistant",
"content": Value::Null,
"tool_calls": openai_tool_call_values(&tool_calls),
})
} else {
json!({
"role": "assistant",
"content": content,
})
};
let mut body_json = json!({
"id": format!("chatcmpl-{}", prepared.request_id),
"object": "chat.completion",
"created": current_unix_secs(),
"model": prepared.model,
"choices": [{
"index": 0,
"message": message,
"finish_reason": if has_tool_calls { "tool_calls" } else { "stop" },
}],
});
if let Some(usage) = poll_result.usage {
body_json["usage"] = windsurf_openai_usage_json(&usage);
}
Ok(Some(ExecutionResult { Ok(Some(ExecutionResult {
request_id: prepared.request_id, request_id: prepared.request_id,
candidate_id: prepared.candidate_id, candidate_id: prepared.candidate_id,
status_code: 200, status_code: 200,
headers: BTreeMap::from([("content-type".to_string(), "application/json".to_string())]), headers: BTreeMap::from([("content-type".to_string(), "application/json".to_string())]),
body: Some(ResponseBody { body: Some(ResponseBody {
json_body: Some(body_json), json_body: Some(body_json),
body_bytes_b64: None, body_bytes_b64: None,
}), }),
telemetry: Some(ExecutionTelemetry { telemetry: Some(ExecutionTelemetry {
ttfb_ms: None, ttfb_ms: None,
elapsed_ms: Some(elapsed_ms), elapsed_ms: Some(elapsed_ms),
upstream_bytes: None, upstream_bytes: None,
}), }),
error: None, error: None,
})) }))
})
.await
} }
async fn prepare_windsurf_cascade( async fn prepare_windsurf_cascade(
@@ -40,6 +40,7 @@ impl<'a> AdminAppState<'a> {
.map(|model| AdminSystemConfigGlobalModel { .map(|model| AdminSystemConfigGlobalModel {
name: model.name.clone(), name: model.name.clone(),
display_name: model.display_name.clone(), display_name: model.display_name.clone(),
usage_count: Some(model.usage_count),
default_price_per_request: model.default_price_per_request, default_price_per_request: model.default_price_per_request,
default_tiered_pricing: model.default_tiered_pricing.clone(), default_tiered_pricing: model.default_tiered_pricing.clone(),
supported_capabilities: model.supported_capabilities.as_ref().and_then(|value| { supported_capabilities: model.supported_capabilities.as_ref().and_then(|value| {
@@ -169,6 +170,13 @@ impl<'a> AdminAppState<'a> {
) -> Result<serde_json::Value, GatewayError> { ) -> Result<serde_json::Value, GatewayError> {
let users = self.list_non_admin_export_users().await?; let users = self.list_non_admin_export_users().await?;
let user_ids = users.iter().map(|user| user.id.clone()).collect::<Vec<_>>(); let user_ids = users.iter().map(|user| user.id.clone()).collect::<Vec<_>>();
let user_usage_totals = self
.app
.summarize_usage_totals_by_user_ids(&user_ids)
.await?
.into_iter()
.map(|totals| (totals.user_id.clone(), totals))
.collect::<BTreeMap<_, _>>();
let user_wallets = self.list_wallet_snapshots_by_user_ids(&user_ids).await?; let user_wallets = self.list_wallet_snapshots_by_user_ids(&user_ids).await?;
let user_api_keys = self let user_api_keys = self
.list_auth_api_key_export_records_by_user_ids(&user_ids) .list_auth_api_key_export_records_by_user_ids(&user_ids)
@@ -261,6 +269,7 @@ impl<'a> AdminAppState<'a> {
self.build_admin_system_users_export_api_key_payload(key, None, true) self.build_admin_system_users_export_api_key_payload(key, None, true)
}) })
.collect::<Vec<_>>(); .collect::<Vec<_>>();
let usage_totals = user_usage_totals.get(&user.id);
json!({ json!({
"id": user.id.clone(), "id": user.id.clone(),
@@ -286,6 +295,12 @@ impl<'a> AdminAppState<'a> {
.unwrap_or(false), .unwrap_or(false),
"wallet": wallet_payload, "wallet": wallet_payload,
"is_active": user.is_active, "is_active": user.is_active,
"request_count": usage_totals
.map(|totals| totals.request_count)
.unwrap_or(0),
"total_tokens": usage_totals
.map(|totals| totals.total_tokens)
.unwrap_or(0),
"api_keys": api_keys_payload, "api_keys": api_keys_payload,
}) })
}) })
@@ -39,8 +39,8 @@ use aether_data::repository::oauth_providers::{
EncryptedSecretUpdate, UpsertOAuthProviderConfigRecord, EncryptedSecretUpdate, UpsertOAuthProviderConfigRecord,
}; };
use aether_data::repository::system::{ use aether_data::repository::system::{
AdminSystemUsageAggregateImportMode, AdminSystemUsageAggregateImportSummary, AdminSystemStatsUserDailyAggregate, AdminSystemUsageAggregateImportMode,
AdminSystemUsageAggregateSnapshot, AdminSystemUsageAggregateImportSummary, AdminSystemUsageAggregateSnapshot,
}; };
use aether_data::repository::wallet::WalletLookupKey; use aether_data::repository::wallet::WalletLookupKey;
use aether_data_contracts::repository::global_models::{ use aether_data_contracts::repository::global_models::{
@@ -694,6 +694,58 @@ fn imported_user_rate_limit_policy_mode(
}) })
} }
fn build_imported_user_usage_total_aggregates(
users: &[Value],
exported_at: Option<&Value>,
) -> Result<Vec<AdminSystemStatsUserDailyAggregate>, String> {
let date_unix_secs = imported_export_day_unix_secs(exported_at);
let mut rows = Vec::new();
for (index, raw_user) in users.iter().enumerate() {
let user = imported_object_field(raw_user, &format!("users[{index}]"))?;
let Some(user_id) = imported_optional_string(user.get("id"))? else {
continue;
};
let request_count = imported_optional_u64(user.get("request_count"), "request_count")?;
let total_tokens = imported_optional_u64(user.get("total_tokens"), "total_tokens")?;
if request_count.is_none() && total_tokens.is_none() {
continue;
}
let total_requests = request_count.unwrap_or(0);
let input_tokens = total_tokens.unwrap_or(0);
if total_requests == 0 && input_tokens == 0 {
continue;
}
rows.push(AdminSystemStatsUserDailyAggregate {
user_id,
username: imported_optional_string(user.get("username"))?,
date_unix_secs,
total_requests,
success_requests: total_requests,
error_requests: 0,
input_tokens,
output_tokens: 0,
cache_creation_tokens: 0,
cache_read_tokens: 0,
total_cost: 0.0,
});
}
Ok(rows)
}
fn imported_export_day_unix_secs(exported_at: Option<&Value>) -> u64 {
imported_optional_string(exported_at)
.ok()
.flatten()
.and_then(|value| chrono::DateTime::parse_from_rfc3339(&value).ok())
.map(|value| unix_day_start_secs(value.timestamp()))
.unwrap_or_else(|| unix_day_start_secs(chrono::Utc::now().timestamp()))
}
fn unix_day_start_secs(timestamp: i64) -> u64 {
let timestamp = timestamp.max(0) as u64;
timestamp - (timestamp % 86_400)
}
fn imported_rfc3339_to_unix_secs( fn imported_rfc3339_to_unix_secs(
value: Option<&Value>, value: Option<&Value>,
field_name: &str, field_name: &str,
@@ -1194,7 +1246,7 @@ impl<'a> AdminAppState<'a> {
return Ok(Err(invalid_request(format!("GlobalModel '{name}' 已存在")))); return Ok(Err(invalid_request(format!("GlobalModel '{name}' 已存在"))));
} }
AdminImportMergeMode::Overwrite => { AdminImportMergeMode::Overwrite => {
let record = invalid!(UpdateAdminGlobalModelRecord::new( let mut record = invalid!(UpdateAdminGlobalModelRecord::new(
existing.id.clone(), existing.id.clone(),
display_name, display_name,
model.is_active, model.is_active,
@@ -1204,6 +1256,7 @@ impl<'a> AdminAppState<'a> {
config, config,
) )
.map_err(|err| err.to_string())); .map_err(|err| err.to_string()));
record.usage_count = model.usage_count;
let Some(updated) = self.update_admin_global_model(&record).await? else { let Some(updated) = self.update_admin_global_model(&record).await? else {
return Ok(Err(invalid_request(format!( return Ok(Err(invalid_request(format!(
"更新 GlobalModel '{name}' 失败" "更新 GlobalModel '{name}' 失败"
@@ -1216,7 +1269,7 @@ impl<'a> AdminAppState<'a> {
continue; continue;
} }
let record = invalid!(CreateAdminGlobalModelRecord::new( let mut record = invalid!(CreateAdminGlobalModelRecord::new(
Uuid::new_v4().to_string(), Uuid::new_v4().to_string(),
name.clone(), name.clone(),
display_name, display_name,
@@ -1227,6 +1280,7 @@ impl<'a> AdminAppState<'a> {
config, config,
) )
.map_err(|err| err.to_string())); .map_err(|err| err.to_string()));
record.usage_count = model.usage_count;
let Some(created) = self.create_admin_global_model(&record).await? else { let Some(created) = self.create_admin_global_model(&record).await? else {
return Ok(Err(invalid_request(format!( return Ok(Err(invalid_request(format!(
"创建 GlobalModel '{name}' 失败" "创建 GlobalModel '{name}' 失败"
@@ -2081,6 +2135,9 @@ impl<'a> AdminAppState<'a> {
root.get("version") root.get("version")
)); ));
let supplemental_user_usage_aggregates = invalid_value!(
build_imported_user_usage_total_aggregates(users, root.get("exported_at"))
);
let mut stats = AdminSystemUsersImportStats::default(); let mut stats = AdminSystemUsersImportStats::default();
let mut imported_user_id_map = BTreeMap::<String, String>::new(); let mut imported_user_id_map = BTreeMap::<String, String>::new();
let mut imported_api_key_id_map = BTreeMap::<String, String>::new(); let mut imported_api_key_id_map = BTreeMap::<String, String>::new();
@@ -2758,6 +2815,7 @@ impl<'a> AdminAppState<'a> {
if let Some(summary) = self if let Some(summary) = self
.import_admin_system_user_usage_aggregates( .import_admin_system_user_usage_aggregates(
root.get("usage_aggregates"), root.get("usage_aggregates"),
&supplemental_user_usage_aggregates,
&imported_user_id_map, &imported_user_id_map,
&imported_api_key_id_map, &imported_api_key_id_map,
merge_mode, merge_mode,
@@ -3012,6 +3070,7 @@ impl<'a> AdminAppState<'a> {
if let Some(summary) = self if let Some(summary) = self
.import_admin_system_user_usage_aggregates( .import_admin_system_user_usage_aggregates(
root.get("usage_aggregates"), root.get("usage_aggregates"),
&supplemental_user_usage_aggregates,
&imported_user_id_map, &imported_user_id_map,
&imported_api_key_id_map, &imported_api_key_id_map,
merge_mode, merge_mode,
@@ -3030,21 +3089,63 @@ impl<'a> AdminAppState<'a> {
async fn import_admin_system_user_usage_aggregates( async fn import_admin_system_user_usage_aggregates(
&self, &self,
value: Option<&Value>, value: Option<&Value>,
supplemental_user_daily: &[AdminSystemStatsUserDailyAggregate],
user_id_map: &BTreeMap<String, String>, user_id_map: &BTreeMap<String, String>,
api_key_id_map: &BTreeMap<String, String>, api_key_id_map: &BTreeMap<String, String>,
merge_mode: AdminImportMergeMode, merge_mode: AdminImportMergeMode,
) -> Result<Option<AdminSystemUsageAggregateImportSummary>, GatewayError> { ) -> Result<Option<AdminSystemUsageAggregateImportSummary>, GatewayError> {
let Some(value) = value else { let mut snapshot = match value {
return Ok(None); Some(value) if !value.is_null() => serde_json::from_value::<
}; AdminSystemUsageAggregateSnapshot,
if value.is_null() { >(value.clone())
return Ok(None);
}
let snapshot = serde_json::from_value::<AdminSystemUsageAggregateSnapshot>(value.clone())
.map_err(|err| GatewayError::Client { .map_err(|err| GatewayError::Client {
status: http::StatusCode::BAD_REQUEST, status: http::StatusCode::BAD_REQUEST,
message: format!("usage_aggregates 格式无效: {err}"), message: format!("usage_aggregates 格式无效: {err}"),
})?; })?,
_ => AdminSystemUsageAggregateSnapshot::default(),
};
let mut existing_user_totals = BTreeMap::<String, (u64, u64)>::new();
for row in &snapshot.stats_user_daily {
let total_tokens = row
.input_tokens
.saturating_add(row.output_tokens)
.saturating_add(row.cache_creation_tokens)
.saturating_add(row.cache_read_tokens);
let entry = existing_user_totals
.entry(row.user_id.clone())
.or_insert((0, 0));
entry.0 = entry.0.saturating_add(row.total_requests);
entry.1 = entry.1.saturating_add(total_tokens);
}
for row in supplemental_user_daily {
let existing = existing_user_totals
.get(&row.user_id)
.copied()
.unwrap_or_default();
let request_delta = row.total_requests.saturating_sub(existing.0);
let token_delta = row.input_tokens.saturating_sub(existing.1);
if request_delta == 0 && token_delta == 0 {
continue;
}
if let Some(existing_row) = snapshot
.stats_user_daily
.iter_mut()
.rev()
.find(|existing_row| existing_row.user_id == row.user_id)
{
existing_row.total_requests =
existing_row.total_requests.saturating_add(request_delta);
existing_row.success_requests =
existing_row.success_requests.saturating_add(request_delta);
existing_row.input_tokens = existing_row.input_tokens.saturating_add(token_delta);
} else {
let mut row = row.clone();
row.total_requests = request_delta;
row.success_requests = request_delta;
row.input_tokens = token_delta;
snapshot.stats_user_daily.push(row);
}
}
if snapshot.stats_daily.is_empty() if snapshot.stats_daily.is_empty()
&& snapshot.stats_user_daily.is_empty() && snapshot.stats_user_daily.is_empty()
&& snapshot.stats_daily_api_key.is_empty() && snapshot.stats_daily_api_key.is_empty()
@@ -3185,20 +3286,22 @@ mod tests {
use serde_json::json; use serde_json::json;
use super::{ use super::{
imported_optional_bool, imported_optional_f64, imported_optional_i32, build_imported_user_usage_total_aggregates, imported_optional_bool, imported_optional_f64,
imported_optional_u64, imported_rfc3339_to_unix_secs, imported_string_list_from_value, imported_optional_i32, imported_optional_u64, imported_rfc3339_to_unix_secs,
normalize_import_endpoint_format, normalize_import_key_formats, imported_string_list_from_value, normalize_import_endpoint_format,
normalize_import_key_raw_payload, normalize_imported_wallet_target, normalize_import_key_formats, normalize_import_key_raw_payload,
validate_imported_system_users_export_version, ImportedProviderKey, normalize_imported_wallet_target, validate_imported_system_users_export_version,
ImportedProviderKey,
}; };
#[test] #[test]
fn users_import_requires_supported_export_version() { fn users_import_requires_supported_export_version() {
assert!(validate_imported_system_users_export_version(Some(&json!("1.3"))).is_ok()); assert!(validate_imported_system_users_export_version(Some(&json!("1.3"))).is_ok());
assert!(validate_imported_system_users_export_version(Some(&json!("1.4"))).is_ok()); assert!(validate_imported_system_users_export_version(Some(&json!("1.4"))).is_ok());
assert!(validate_imported_system_users_export_version(Some(&json!("1.5"))).is_ok());
assert_eq!( assert_eq!(
validate_imported_system_users_export_version(Some(&json!("2.2"))).unwrap_err(), validate_imported_system_users_export_version(Some(&json!("2.2"))).unwrap_err(),
"不支持的用户数据版本: 2.2,支持的版本: 1.3, 1.4" "不支持的用户数据版本: 2.2,支持的版本: 1.3, 1.4, 1.5"
); );
assert_eq!( assert_eq!(
validate_imported_system_users_export_version(Some(&json!(null))).unwrap_err(), validate_imported_system_users_export_version(Some(&json!(null))).unwrap_err(),
@@ -3206,6 +3309,43 @@ mod tests {
); );
} }
#[test]
fn users_import_builds_supplemental_usage_aggregates_from_summary_fields() {
let users = vec![
json!({
"id": "source-user-1",
"username": "alice",
"request_count": 12,
"total_tokens": 3456
}),
json!({
"id": "source-user-zero",
"username": "zero",
"request_count": 0,
"total_tokens": 0
}),
json!({
"username": "no-source-id",
"request_count": 5,
"total_tokens": 6
}),
];
let rows = build_imported_user_usage_total_aggregates(
&users,
Some(&json!("2026-05-25T12:34:56Z")),
)
.expect("supplemental usage aggregates should build");
assert_eq!(rows.len(), 1);
assert_eq!(rows[0].user_id, "source-user-1");
assert_eq!(rows[0].username.as_deref(), Some("alice"));
assert_eq!(rows[0].total_requests, 12);
assert_eq!(rows[0].success_requests, 12);
assert_eq!(rows[0].input_tokens, 3456);
assert_eq!(rows[0].date_unix_secs % 86_400, 0);
}
#[test] #[test]
fn config_import_normalizes_python_cli_api_format_aliases() { fn config_import_normalizes_python_cli_api_format_aliases() {
for (raw, expected) in [ for (raw, expected) in [
@@ -735,11 +735,11 @@ async fn gateway_handles_admin_system_config_export_locally_with_trusted_admin_p
)); ));
let global_model_repository = Arc::new( let global_model_repository = Arc::new(
InMemoryGlobalModelReadRepository::seed(Vec::<StoredPublicGlobalModel>::new()) InMemoryGlobalModelReadRepository::seed(Vec::<StoredPublicGlobalModel>::new())
.with_admin_global_models(vec![sample_admin_global_model( .with_admin_global_models(vec![{
"global-gpt-5", let mut model = sample_admin_global_model("global-gpt-5", "gpt-5", "GPT 5");
"gpt-5", model.usage_count = 7;
"GPT 5", model
)]) }])
.with_admin_provider_models(vec![sample_admin_provider_model( .with_admin_provider_models(vec![sample_admin_provider_model(
"model-gpt-5", "model-gpt-5",
&provider_id, &provider_id,
@@ -802,9 +802,10 @@ async fn gateway_handles_admin_system_config_export_locally_with_trusted_admin_p
assert_eq!(response.status(), StatusCode::OK); assert_eq!(response.status(), StatusCode::OK);
let payload: serde_json::Value = response.json().await.expect("json body should parse"); let payload: serde_json::Value = response.json().await.expect("json body should parse");
assert_eq!(payload["version"], "2.2"); assert_eq!(payload["version"], "2.3");
assert!(payload["exported_at"].as_str().is_some()); assert!(payload["exported_at"].as_str().is_some());
assert_eq!(payload["global_models"][0]["name"], "gpt-5"); assert_eq!(payload["global_models"][0]["name"], "gpt-5");
assert_eq!(payload["global_models"][0]["usage_count"], json!(7));
assert_eq!(payload["providers"][0]["name"], "openai"); assert_eq!(payload["providers"][0]["name"], "openai");
assert_eq!( assert_eq!(
payload["providers"][0]["config"]["provider_ops"]["connector"]["credentials"] payload["providers"][0]["config"]["provider_ops"]["connector"]["credentials"]
@@ -1025,7 +1026,7 @@ async fn gateway_handles_admin_system_users_export_locally_with_trusted_admin_pr
assert_eq!(response.status(), StatusCode::OK); assert_eq!(response.status(), StatusCode::OK);
let payload: serde_json::Value = response.json().await.expect("json body should parse"); let payload: serde_json::Value = response.json().await.expect("json body should parse");
assert_eq!(payload["version"], "1.4"); assert_eq!(payload["version"], "1.5");
assert!(payload["exported_at"].as_str().is_some()); assert!(payload["exported_at"].as_str().is_some());
assert_eq!(payload["user_groups"][0]["name"], "Restricted GPT"); assert_eq!(payload["user_groups"][0]["name"], "Restricted GPT");
assert!(payload["user_groups"][0].get("priority").is_none()); assert!(payload["user_groups"][0].get("priority").is_none());
@@ -1044,6 +1045,8 @@ async fn gateway_handles_admin_system_users_export_locally_with_trusted_admin_pr
json!(["Restricted GPT"]) json!(["Restricted GPT"])
); );
assert_eq!(payload["users"][0]["id"], json!("user-1")); assert_eq!(payload["users"][0]["id"], json!("user-1"));
assert_eq!(payload["users"][0]["request_count"], json!(0));
assert_eq!(payload["users"][0]["total_tokens"], json!(0));
assert_eq!(payload["users"][0]["wallet"]["balance"], json!(12.5)); assert_eq!(payload["users"][0]["wallet"]["balance"], json!(12.5));
assert_eq!( assert_eq!(
payload["users"][0]["wallet"]["recharge_balance"], payload["users"][0]["wallet"]["recharge_balance"],
@@ -81,6 +81,7 @@ fn sample_system_import_payload() -> Value {
"global_models": [{ "global_models": [{
"name": "gpt-5", "name": "gpt-5",
"display_name": "GPT 5", "display_name": "GPT 5",
"usage_count": 123,
"default_price_per_request": 0.03, "default_price_per_request": 0.03,
"default_tiered_pricing": { "default_tiered_pricing": {
"tiers": [{ "tiers": [{
@@ -345,6 +346,7 @@ async fn gateway_imports_admin_system_config_locally_and_persists_data() {
.expect("global models should load"); .expect("global models should load");
assert_eq!(global_models.items.len(), 1); assert_eq!(global_models.items.len(), 1);
assert_eq!(global_models.items[0].name, "gpt-5"); assert_eq!(global_models.items[0].name, "gpt-5");
assert_eq!(global_models.items[0].usage_count, 123);
let providers = provider_catalog_repository let providers = provider_catalog_repository
.list_providers(false) .list_providers(false)
@@ -836,7 +838,7 @@ async fn gateway_rejects_unknown_admin_system_config_import_versions() {
let (gateway_url, gateway_handle) = start_server(gateway).await; let (gateway_url, gateway_handle) = start_server(gateway).await;
let client = reqwest::Client::new(); let client = reqwest::Client::new();
for version in ["1.9", "2.3"] { for version in ["1.9", "2.4"] {
let response = client let response = client
.post(format!("{gateway_url}/api/admin/system/config/import")) .post(format!("{gateway_url}/api/admin/system/config/import"))
.header(GATEWAY_HEADER, "rust-phase3b") .header(GATEWAY_HEADER, "rust-phase3b")
@@ -859,7 +861,7 @@ async fn gateway_rejects_unknown_admin_system_config_import_versions() {
.as_str() .as_str()
.expect("detail should be a string"); .expect("detail should be a string");
assert!(detail.contains(&format!("不支持的配置版本: {version}"))); assert!(detail.contains(&format!("不支持的配置版本: {version}")));
assert!(detail.contains("支持的版本: 2.0, 2.1, 2.2")); assert!(detail.contains("支持的版本: 2.0, 2.1, 2.2, 2.3"));
} }
gateway_handle.abort(); gateway_handle.abort();
+31 -5
View File
@@ -43,12 +43,12 @@ pub struct AdminEmailTemplateUpdate {
pub html: Option<String>, pub html: Option<String>,
} }
pub const ADMIN_SYSTEM_CONFIG_EXPORT_VERSION: &str = "2.2"; pub const ADMIN_SYSTEM_CONFIG_EXPORT_VERSION: &str = "2.3";
pub const ADMIN_SYSTEM_CONFIG_SUPPORTED_VERSIONS: &[&str] = pub const ADMIN_SYSTEM_CONFIG_SUPPORTED_VERSIONS: &[&str] =
&["2.0", "2.1", ADMIN_SYSTEM_CONFIG_EXPORT_VERSION]; &["2.0", "2.1", "2.2", ADMIN_SYSTEM_CONFIG_EXPORT_VERSION];
pub const ADMIN_SYSTEM_USERS_EXPORT_VERSION: &str = "1.4"; pub const ADMIN_SYSTEM_USERS_EXPORT_VERSION: &str = "1.5";
pub const ADMIN_SYSTEM_USERS_SUPPORTED_VERSIONS: &[&str] = pub const ADMIN_SYSTEM_USERS_SUPPORTED_VERSIONS: &[&str] =
&["1.3", ADMIN_SYSTEM_USERS_EXPORT_VERSION]; &["1.3", "1.4", ADMIN_SYSTEM_USERS_EXPORT_VERSION];
pub const ADMIN_SYSTEM_PROVIDER_OPS_SENSITIVE_CREDENTIAL_FIELDS: &[&str] = &[ pub const ADMIN_SYSTEM_PROVIDER_OPS_SENSITIVE_CREDENTIAL_FIELDS: &[&str] = &[
"api_key", "api_key",
"password", "password",
@@ -271,6 +271,28 @@ where
} }
} }
fn deserialize_optional_u64_from_number<'de, D>(deserializer: D) -> Result<Option<u64>, D::Error>
where
D: serde::Deserializer<'de>,
{
let value = Option::<Value>::deserialize(deserializer)?;
match value {
None | Some(Value::Null) => Ok(None),
Some(Value::Number(number)) => number
.as_u64()
.map(Some)
.ok_or_else(|| de::Error::custom("expected a non-negative integer or numeric string")),
Some(Value::String(raw)) if !raw.trim().is_empty() => raw
.trim()
.parse::<u64>()
.map(Some)
.map_err(|_| de::Error::custom("expected a non-negative integer or numeric string")),
Some(_) => Err(de::Error::custom(
"expected a non-negative integer or numeric string",
)),
}
}
#[derive(Debug, Clone, Copy, Default, PartialEq, Eq, Serialize)] #[derive(Debug, Clone, Copy, Default, PartialEq, Eq, Serialize)]
#[serde(rename_all = "snake_case")] #[serde(rename_all = "snake_case")]
pub enum AdminImportMergeMode { pub enum AdminImportMergeMode {
@@ -347,6 +369,8 @@ pub struct AdminSystemConfigImportStats {
pub struct AdminSystemConfigGlobalModel { pub struct AdminSystemConfigGlobalModel {
pub name: String, pub name: String,
pub display_name: String, pub display_name: String,
#[serde(default, deserialize_with = "deserialize_optional_u64_from_number")]
pub usage_count: Option<u64>,
#[serde(default, deserialize_with = "deserialize_optional_f64_from_number")] #[serde(default, deserialize_with = "deserialize_optional_f64_from_number")]
pub default_price_per_request: Option<f64>, pub default_price_per_request: Option<f64>,
#[serde(default)] #[serde(default)]
@@ -3147,7 +3171,7 @@ mod tests {
#[test] #[test]
fn parse_admin_system_config_import_request_rejects_unknown_versions() { fn parse_admin_system_config_import_request_rejects_unknown_versions() {
for version in ["1.9", "2.3"] { for version in ["1.9", "2.4"] {
let err = parse_admin_system_config_import_request( let err = parse_admin_system_config_import_request(
json!({ json!({
"version": version, "version": version,
@@ -3223,6 +3247,7 @@ mod tests {
"global_models": [{ "global_models": [{
"name": "veo3.1", "name": "veo3.1",
"display_name": "Veo 3.1", "display_name": "Veo 3.1",
"usage_count": "42",
"default_price_per_request": "1.80000000", "default_price_per_request": "1.80000000",
}], }],
"providers": [{ "providers": [{
@@ -3243,6 +3268,7 @@ mod tests {
.expect("numeric string fields from Python exports should parse"); .expect("numeric string fields from Python exports should parse");
let global_model = &parsed.request.document.global_models[0]; let global_model = &parsed.request.document.global_models[0];
assert_eq!(global_model.usage_count, Some(42));
assert_eq!(global_model.default_price_per_request, Some(1.8)); assert_eq!(global_model.default_price_per_request, Some(1.8));
let provider = &parsed.request.document.providers[0]; let provider = &parsed.request.document.providers[0];
@@ -600,6 +600,8 @@ pub struct CreateAdminGlobalModelRecord {
pub default_tiered_pricing: Option<Value>, pub default_tiered_pricing: Option<Value>,
pub supported_capabilities: Option<Value>, pub supported_capabilities: Option<Value>,
pub config: Option<Value>, pub config: Option<Value>,
#[serde(default)]
pub usage_count: Option<u64>,
} }
impl CreateAdminGlobalModelRecord { impl CreateAdminGlobalModelRecord {
@@ -645,6 +647,7 @@ impl CreateAdminGlobalModelRecord {
default_tiered_pricing, default_tiered_pricing,
supported_capabilities, supported_capabilities,
config, config,
usage_count: None,
}) })
} }
} }
@@ -658,6 +661,8 @@ pub struct UpdateAdminGlobalModelRecord {
pub default_tiered_pricing: Option<Value>, pub default_tiered_pricing: Option<Value>,
pub supported_capabilities: Option<Value>, pub supported_capabilities: Option<Value>,
pub config: Option<Value>, pub config: Option<Value>,
#[serde(default)]
pub usage_count: Option<u64>,
} }
impl UpdateAdminGlobalModelRecord { impl UpdateAdminGlobalModelRecord {
@@ -696,6 +701,7 @@ impl UpdateAdminGlobalModelRecord {
default_tiered_pricing, default_tiered_pricing,
supported_capabilities, supported_capabilities,
config, config,
usage_count: None,
}) })
} }
} }
@@ -577,7 +577,7 @@ impl GlobalModelWriteRepository for InMemoryGlobalModelReadRepository {
record.config.clone(), record.config.clone(),
0, 0,
0, 0,
0, record.usage_count.unwrap_or(0),
Some(1_711_000_000), Some(1_711_000_000),
Some(1_711_000_000), Some(1_711_000_000),
)?; )?;
@@ -606,6 +606,9 @@ impl GlobalModelWriteRepository for InMemoryGlobalModelReadRepository {
existing.default_tiered_pricing = record.default_tiered_pricing.clone(); existing.default_tiered_pricing = record.default_tiered_pricing.clone();
existing.supported_capabilities = record.supported_capabilities.clone(); existing.supported_capabilities = record.supported_capabilities.clone();
existing.config = record.config.clone(); existing.config = record.config.clone();
if let Some(usage_count) = record.usage_count {
existing.usage_count = usage_count;
}
existing.updated_at_unix_secs = Some(1_711_000_100); existing.updated_at_unix_secs = Some(1_711_000_100);
} }
self.get_admin_global_model_by_id(&record.id).await self.get_admin_global_model_by_id(&record.id).await
@@ -336,6 +336,8 @@ WHERE provider_id = ?
record: &CreateAdminGlobalModelRecord, record: &CreateAdminGlobalModelRecord,
) -> Result<Option<StoredAdminGlobalModel>, DataLayerError> { ) -> Result<Option<StoredAdminGlobalModel>, DataLayerError> {
let now = current_unix_secs(); let now = current_unix_secs();
let usage_count =
optional_admin_global_model_usage_count_i64(record.usage_count)?.unwrap_or_default();
sqlx::query( sqlx::query(
r#" r#"
INSERT INTO global_models ( INSERT INTO global_models (
@@ -351,7 +353,7 @@ INSERT INTO global_models (
created_at, created_at,
updated_at updated_at
) )
VALUES (?, ?, ?, ?, ?, ?, ?, 0, ?, ?, ?) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
"#, "#,
) )
.bind(&record.id) .bind(&record.id)
@@ -367,6 +369,7 @@ VALUES (?, ?, ?, ?, ?, ?, ?, 0, ?, ?, ?)
&record.supported_capabilities, &record.supported_capabilities,
"global_models.supported_capabilities", "global_models.supported_capabilities",
)?) )?)
.bind(usage_count)
.bind(optional_json_to_string( .bind(optional_json_to_string(
&record.config, &record.config,
"global_models.config", "global_models.config",
@@ -385,6 +388,7 @@ VALUES (?, ?, ?, ?, ?, ?, ?, 0, ?, ?, ?)
record: &UpdateAdminGlobalModelRecord, record: &UpdateAdminGlobalModelRecord,
) -> Result<Option<StoredAdminGlobalModel>, DataLayerError> { ) -> Result<Option<StoredAdminGlobalModel>, DataLayerError> {
let now = current_unix_secs(); let now = current_unix_secs();
let usage_count = optional_admin_global_model_usage_count_i64(record.usage_count)?;
let updated = sqlx::query( let updated = sqlx::query(
r#" r#"
UPDATE global_models UPDATE global_models
@@ -395,6 +399,7 @@ SET
default_tiered_pricing = ?, default_tiered_pricing = ?,
supported_capabilities = ?, supported_capabilities = ?,
config = ?, config = ?,
usage_count = COALESCE(?, usage_count),
updated_at = ? updated_at = ?
WHERE id = ? WHERE id = ?
"#, "#,
@@ -414,6 +419,7 @@ WHERE id = ?
&record.config, &record.config,
"global_models.config", "global_models.config",
)?) )?)
.bind(usage_count)
.bind(now as i64) .bind(now as i64)
.bind(&record.id) .bind(&record.id)
.execute(&self.pool) .execute(&self.pool)
@@ -880,6 +886,20 @@ fn map_active_global_model_row(
) )
} }
fn optional_admin_global_model_usage_count_i64(
value: Option<u64>,
) -> Result<Option<i64>, DataLayerError> {
value
.map(|value| {
i64::try_from(value).map_err(|_| {
DataLayerError::InvalidInput(
"global_models.usage_count exceeds i64 range".to_string(),
)
})
})
.transpose()
}
#[cfg(test)] #[cfg(test)]
mod tests { mod tests {
use super::MysqlGlobalModelReadRepository; use super::MysqlGlobalModelReadRepository;
@@ -664,6 +664,8 @@ RETURNING id
&self, &self,
record: &CreateAdminGlobalModelRecord, record: &CreateAdminGlobalModelRecord,
) -> Result<Option<StoredAdminGlobalModel>, DataLayerError> { ) -> Result<Option<StoredAdminGlobalModel>, DataLayerError> {
let usage_count =
optional_admin_global_model_usage_count_i64(record.usage_count)?.unwrap_or_default();
let inserted = sqlx::query( let inserted = sqlx::query(
r#" r#"
INSERT INTO global_models ( INSERT INTO global_models (
@@ -690,7 +692,7 @@ RETURNING id
.bind(record.default_price_per_request) .bind(record.default_price_per_request)
.bind(record.default_tiered_pricing.clone()) .bind(record.default_tiered_pricing.clone())
.bind(record.supported_capabilities.clone()) .bind(record.supported_capabilities.clone())
.bind(0_i32) .bind(usage_count)
.bind(record.config.clone()) .bind(record.config.clone())
.fetch_optional(&self.pool) .fetch_optional(&self.pool)
.await .await
@@ -707,6 +709,7 @@ RETURNING id
&self, &self,
record: &UpdateAdminGlobalModelRecord, record: &UpdateAdminGlobalModelRecord,
) -> Result<Option<StoredAdminGlobalModel>, DataLayerError> { ) -> Result<Option<StoredAdminGlobalModel>, DataLayerError> {
let usage_count = optional_admin_global_model_usage_count_i64(record.usage_count)?;
let updated = sqlx::query( let updated = sqlx::query(
r#" r#"
UPDATE global_models UPDATE global_models
@@ -717,6 +720,7 @@ SET
default_tiered_pricing = $5, default_tiered_pricing = $5,
supported_capabilities = $6, supported_capabilities = $6,
config = $7, config = $7,
usage_count = COALESCE($8, usage_count),
updated_at = NOW() updated_at = NOW()
WHERE id = $1 WHERE id = $1
RETURNING id RETURNING id
@@ -729,6 +733,7 @@ RETURNING id
.bind(record.default_tiered_pricing.clone()) .bind(record.default_tiered_pricing.clone())
.bind(record.supported_capabilities.clone()) .bind(record.supported_capabilities.clone())
.bind(record.config.clone()) .bind(record.config.clone())
.bind(usage_count)
.fetch_optional(&self.pool) .fetch_optional(&self.pool)
.await .await
.map_postgres_err()?; .map_postgres_err()?;
@@ -1203,6 +1208,20 @@ fn map_provider_active_global_model_row(
) )
} }
fn optional_admin_global_model_usage_count_i64(
value: Option<u64>,
) -> Result<Option<i64>, DataLayerError> {
value
.map(|value| {
i64::try_from(value).map_err(|_| {
DataLayerError::InvalidInput(
"global_models.usage_count exceeds i64 range".to_string(),
)
})
})
.transpose()
}
#[cfg(test)] #[cfg(test)]
mod tests { mod tests {
use super::{ use super::{
@@ -321,6 +321,8 @@ WHERE provider_id = ?
record: &CreateAdminGlobalModelRecord, record: &CreateAdminGlobalModelRecord,
) -> Result<Option<StoredAdminGlobalModel>, DataLayerError> { ) -> Result<Option<StoredAdminGlobalModel>, DataLayerError> {
let now = current_unix_secs(); let now = current_unix_secs();
let usage_count =
optional_admin_global_model_usage_count_i64(record.usage_count)?.unwrap_or_default();
sqlx::query( sqlx::query(
r#" r#"
INSERT INTO global_models ( INSERT INTO global_models (
@@ -336,7 +338,7 @@ INSERT INTO global_models (
created_at, created_at,
updated_at updated_at
) )
VALUES (?, ?, ?, ?, ?, ?, ?, 0, ?, ?, ?) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
"#, "#,
) )
.bind(&record.id) .bind(&record.id)
@@ -352,6 +354,7 @@ VALUES (?, ?, ?, ?, ?, ?, ?, 0, ?, ?, ?)
&record.supported_capabilities, &record.supported_capabilities,
"global_models.supported_capabilities", "global_models.supported_capabilities",
)?) )?)
.bind(usage_count)
.bind(optional_json_to_string( .bind(optional_json_to_string(
&record.config, &record.config,
"global_models.config", "global_models.config",
@@ -370,6 +373,7 @@ VALUES (?, ?, ?, ?, ?, ?, ?, 0, ?, ?, ?)
record: &UpdateAdminGlobalModelRecord, record: &UpdateAdminGlobalModelRecord,
) -> Result<Option<StoredAdminGlobalModel>, DataLayerError> { ) -> Result<Option<StoredAdminGlobalModel>, DataLayerError> {
let now = current_unix_secs(); let now = current_unix_secs();
let usage_count = optional_admin_global_model_usage_count_i64(record.usage_count)?;
let updated = sqlx::query( let updated = sqlx::query(
r#" r#"
UPDATE global_models UPDATE global_models
@@ -380,6 +384,7 @@ SET
default_tiered_pricing = ?, default_tiered_pricing = ?,
supported_capabilities = ?, supported_capabilities = ?,
config = ?, config = ?,
usage_count = COALESCE(?, usage_count),
updated_at = ? updated_at = ?
WHERE id = ? WHERE id = ?
"#, "#,
@@ -399,6 +404,7 @@ WHERE id = ?
&record.config, &record.config,
"global_models.config", "global_models.config",
)?) )?)
.bind(usage_count)
.bind(now as i64) .bind(now as i64)
.bind(&record.id) .bind(&record.id)
.execute(&self.pool) .execute(&self.pool)
@@ -1124,6 +1130,20 @@ fn map_active_global_model_row(
) )
} }
fn optional_admin_global_model_usage_count_i64(
value: Option<u64>,
) -> Result<Option<i64>, DataLayerError> {
value
.map(|value| {
i64::try_from(value).map_err(|_| {
DataLayerError::InvalidInput(
"global_models.usage_count exceeds i64 range".to_string(),
)
})
})
.transpose()
}
#[cfg(test)] #[cfg(test)]
mod tests { mod tests {
use super::SqliteGlobalModelReadRepository; use super::SqliteGlobalModelReadRepository;
@@ -245,7 +245,7 @@ impl MysqlUsageReadRepository {
SELECT SELECT
DATE_FORMAT(FROM_UNIXTIME(created_at_unix_ms), '%Y-%m-%d') AS date, DATE_FORMAT(FROM_UNIXTIME(created_at_unix_ms), '%Y-%m-%d') AS date,
COUNT(*) AS requests, COUNT(*) AS requests,
CAST(COALESCE(SUM( CAST(COALESCE(SUM(
GREATEST(COALESCE(input_tokens, 0), 0) GREATEST(COALESCE(input_tokens, 0), 0)
+ GREATEST(COALESCE(output_tokens, 0), 0) + GREATEST(COALESCE(output_tokens, 0), 0)
+ CASE + CASE
@@ -255,9 +255,9 @@ SELECT
ELSE GREATEST(COALESCE(cache_creation_input_tokens, 0), 0) ELSE GREATEST(COALESCE(cache_creation_input_tokens, 0), 0)
END END
+ GREATEST(COALESCE(cache_read_input_tokens, 0), 0) + GREATEST(COALESCE(cache_read_input_tokens, 0), 0)
), 0) AS SIGNED) AS total_tokens, ), 0) AS SIGNED) AS total_tokens,
CAST(COALESCE(SUM(COALESCE(total_cost_usd, 0)), 0) AS DOUBLE) AS total_cost_usd, CAST(COALESCE(SUM(COALESCE(total_cost_usd, 0)), 0) AS DOUBLE) AS total_cost_usd,
CAST(COALESCE(SUM(COALESCE(actual_total_cost_usd, 0)), 0) AS DOUBLE) AS actual_total_cost_usd CAST(COALESCE(SUM(COALESCE(actual_total_cost_usd, 0)), 0) AS DOUBLE) AS actual_total_cost_usd
FROM `usage` FROM `usage`
WHERE created_at_unix_ms >= ? WHERE created_at_unix_ms >= ?
AND created_at_unix_ms < ? AND created_at_unix_ms < ?
@@ -375,21 +375,21 @@ ORDER BY `date` ASC
sqlx::query( sqlx::query(
r#" r#"
SELECT SELECT
CAST(COALESCE(SUM(total_requests), 0) AS SIGNED) AS total_requests, CAST(COALESCE(SUM(total_requests), 0) AS SIGNED) AS total_requests,
CAST(COALESCE(SUM(input_tokens), 0) AS SIGNED) AS input_tokens, CAST(COALESCE(SUM(input_tokens), 0) AS SIGNED) AS input_tokens,
CAST(COALESCE(SUM(input_tokens), 0) AS SIGNED) AS effective_input_tokens, CAST(COALESCE(SUM(input_tokens), 0) AS SIGNED) AS effective_input_tokens,
CAST(COALESCE(SUM(output_tokens), 0) AS SIGNED) AS output_tokens, CAST(COALESCE(SUM(output_tokens), 0) AS SIGNED) AS output_tokens,
CAST(COALESCE(SUM(input_tokens + output_tokens + cache_creation_tokens + cache_read_tokens), 0) AS SIGNED) AS total_tokens, CAST(COALESCE(SUM(input_tokens + output_tokens + cache_creation_tokens + cache_read_tokens), 0) AS SIGNED) AS total_tokens,
CAST(COALESCE(SUM(cache_creation_tokens), 0) AS SIGNED) AS cache_creation_tokens, CAST(COALESCE(SUM(cache_creation_tokens), 0) AS SIGNED) AS cache_creation_tokens,
CAST(COALESCE(SUM(cache_read_tokens), 0) AS SIGNED) AS cache_read_tokens, CAST(COALESCE(SUM(cache_read_tokens), 0) AS SIGNED) AS cache_read_tokens,
CAST(COALESCE(SUM(input_tokens + cache_creation_tokens + cache_read_tokens), 0) AS SIGNED) AS total_input_context, CAST(COALESCE(SUM(input_tokens + cache_creation_tokens + cache_read_tokens), 0) AS SIGNED) AS total_input_context,
CAST(0 AS DOUBLE) AS cache_creation_cost_usd, CAST(0.0 AS DOUBLE) AS cache_creation_cost_usd,
CAST(0 AS DOUBLE) AS cache_read_cost_usd, CAST(0.0 AS DOUBLE) AS cache_read_cost_usd,
CAST(COALESCE(SUM(COALESCE(total_cost, 0)), 0) AS DOUBLE) AS total_cost_usd, CAST(COALESCE(SUM(COALESCE(total_cost, 0)), 0) AS DOUBLE) AS total_cost_usd,
CAST(COALESCE(SUM(COALESCE(total_cost, 0)), 0) AS DOUBLE) AS actual_total_cost_usd, CAST(COALESCE(SUM(COALESCE(total_cost, 0)), 0) AS DOUBLE) AS actual_total_cost_usd,
CAST(COALESCE(SUM(error_requests), 0) AS SIGNED) AS error_requests, CAST(COALESCE(SUM(error_requests), 0) AS SIGNED) AS error_requests,
CAST(0 AS DOUBLE) AS response_time_sum_ms, CAST(0.0 AS DOUBLE) AS response_time_sum_ms,
CAST(0 AS SIGNED) AS response_time_samples CAST(0 AS SIGNED) AS response_time_samples
FROM stats_user_daily FROM stats_user_daily
WHERE user_id = ? WHERE user_id = ?
AND `date` >= ? AND `date` >= ?
@@ -412,21 +412,21 @@ WHERE user_id = ?
sqlx::query( sqlx::query(
r#" r#"
SELECT SELECT
CAST(COALESCE(SUM(total_requests), 0) AS SIGNED) AS total_requests, CAST(COALESCE(SUM(total_requests), 0) AS SIGNED) AS total_requests,
CAST(COALESCE(SUM(input_tokens), 0) AS SIGNED) AS input_tokens, CAST(COALESCE(SUM(input_tokens), 0) AS SIGNED) AS input_tokens,
CAST(COALESCE(SUM(input_tokens), 0) AS SIGNED) AS effective_input_tokens, CAST(COALESCE(SUM(input_tokens), 0) AS SIGNED) AS effective_input_tokens,
CAST(COALESCE(SUM(output_tokens), 0) AS SIGNED) AS output_tokens, CAST(COALESCE(SUM(output_tokens), 0) AS SIGNED) AS output_tokens,
CAST(COALESCE(SUM(input_tokens + output_tokens + cache_creation_tokens + cache_read_tokens), 0) AS SIGNED) AS total_tokens, CAST(COALESCE(SUM(input_tokens + output_tokens + cache_creation_tokens + cache_read_tokens), 0) AS SIGNED) AS total_tokens,
CAST(COALESCE(SUM(cache_creation_tokens), 0) AS SIGNED) AS cache_creation_tokens, CAST(COALESCE(SUM(cache_creation_tokens), 0) AS SIGNED) AS cache_creation_tokens,
CAST(COALESCE(SUM(cache_read_tokens), 0) AS SIGNED) AS cache_read_tokens, CAST(COALESCE(SUM(cache_read_tokens), 0) AS SIGNED) AS cache_read_tokens,
CAST(COALESCE(SUM(input_tokens + cache_creation_tokens + cache_read_tokens), 0) AS SIGNED) AS total_input_context, CAST(COALESCE(SUM(input_tokens + cache_creation_tokens + cache_read_tokens), 0) AS SIGNED) AS total_input_context,
CAST(COALESCE(SUM(COALESCE(cache_creation_cost, 0)), 0) AS DOUBLE) AS cache_creation_cost_usd, CAST(COALESCE(SUM(COALESCE(cache_creation_cost, 0)), 0) AS DOUBLE) AS cache_creation_cost_usd,
CAST(COALESCE(SUM(COALESCE(cache_read_cost, 0)), 0) AS DOUBLE) AS cache_read_cost_usd, CAST(COALESCE(SUM(COALESCE(cache_read_cost, 0)), 0) AS DOUBLE) AS cache_read_cost_usd,
CAST(COALESCE(SUM(COALESCE(total_cost, 0)), 0) AS DOUBLE) AS total_cost_usd, CAST(COALESCE(SUM(COALESCE(total_cost, 0)), 0) AS DOUBLE) AS total_cost_usd,
CAST(COALESCE(SUM(COALESCE(actual_total_cost, 0)), 0) AS DOUBLE) AS actual_total_cost_usd, CAST(COALESCE(SUM(COALESCE(actual_total_cost, 0)), 0) AS DOUBLE) AS actual_total_cost_usd,
CAST(COALESCE(SUM(error_requests), 0) AS SIGNED) AS error_requests, CAST(COALESCE(SUM(error_requests), 0) AS SIGNED) AS error_requests,
CAST(0 AS DOUBLE) AS response_time_sum_ms, CAST(0.0 AS DOUBLE) AS response_time_sum_ms,
CAST(0 AS SIGNED) AS response_time_samples CAST(0 AS SIGNED) AS response_time_samples
FROM stats_daily FROM stats_daily
WHERE `date` >= ? WHERE `date` >= ?
AND `date` < ? AND `date` < ?
@@ -474,11 +474,11 @@ SELECT
DATE_FORMAT(FROM_UNIXTIME(`date`), '%Y-%m-%d') AS date, DATE_FORMAT(FROM_UNIXTIME(`date`), '%Y-%m-%d') AS date,
'aggregate' AS model, 'aggregate' AS model,
'aggregate' AS provider, 'aggregate' AS provider,
CAST(COALESCE(SUM(total_requests), 0) AS SIGNED) AS requests, CAST(COALESCE(SUM(total_requests), 0) AS SIGNED) AS requests,
CAST(COALESCE(SUM(input_tokens + output_tokens + cache_creation_tokens + cache_read_tokens), 0) AS SIGNED) AS total_tokens, CAST(COALESCE(SUM(input_tokens + output_tokens + cache_creation_tokens + cache_read_tokens), 0) AS SIGNED) AS total_tokens,
CAST(COALESCE(SUM(COALESCE(total_cost, 0)), 0) AS DOUBLE) AS total_cost_usd, CAST(COALESCE(SUM(COALESCE(total_cost, 0)), 0) AS DOUBLE) AS total_cost_usd,
CAST(0 AS DOUBLE) AS response_time_sum_ms, CAST(0.0 AS DOUBLE) AS response_time_sum_ms,
CAST(0 AS SIGNED) AS response_time_samples CAST(0 AS SIGNED) AS response_time_samples
FROM stats_user_daily FROM stats_user_daily
WHERE user_id = ? WHERE user_id = ?
AND `date` >= ? AND `date` >= ?
@@ -507,11 +507,11 @@ SELECT
DATE_FORMAT(FROM_UNIXTIME(`date`), '%Y-%m-%d') AS date, DATE_FORMAT(FROM_UNIXTIME(`date`), '%Y-%m-%d') AS date,
'aggregate' AS model, 'aggregate' AS model,
'aggregate' AS provider, 'aggregate' AS provider,
CAST(COALESCE(SUM(total_requests), 0) AS SIGNED) AS requests, CAST(COALESCE(SUM(total_requests), 0) AS SIGNED) AS requests,
CAST(COALESCE(SUM(input_tokens + output_tokens + cache_creation_tokens + cache_read_tokens), 0) AS SIGNED) AS total_tokens, CAST(COALESCE(SUM(input_tokens + output_tokens + cache_creation_tokens + cache_read_tokens), 0) AS SIGNED) AS total_tokens,
CAST(COALESCE(SUM(COALESCE(total_cost, 0)), 0) AS DOUBLE) AS total_cost_usd, CAST(COALESCE(SUM(COALESCE(total_cost, 0)), 0) AS DOUBLE) AS total_cost_usd,
CAST(0 AS DOUBLE) AS response_time_sum_ms, CAST(0.0 AS DOUBLE) AS response_time_sum_ms,
CAST(0 AS SIGNED) AS response_time_samples CAST(0 AS SIGNED) AS response_time_samples
FROM stats_daily FROM stats_daily
WHERE `date` >= ? WHERE `date` >= ?
AND `date` < ? AND `date` < ?
@@ -600,11 +600,11 @@ ORDER BY `date` ASC
r#" r#"
SELECT SELECT
user_id, user_id,
CAST(COALESCE(SUM(total_requests), 0) AS SIGNED) AS request_count, CAST(COALESCE(SUM(total_requests), 0) AS SIGNED) AS request_count,
CAST(COALESCE( CAST(COALESCE(
SUM(input_tokens + output_tokens + cache_creation_tokens + cache_read_tokens), SUM(input_tokens + output_tokens + cache_creation_tokens + cache_read_tokens),
0 0
) AS SIGNED) AS total_tokens, ) AS SIGNED) AS total_tokens,
MAX(`date`) AS latest_date MAX(`date`) AS latest_date
FROM stats_user_daily FROM stats_user_daily
WHERE user_id IN ( WHERE user_id IN (
@@ -642,7 +642,7 @@ WHERE user_id IN (
SELECT SELECT
`usage`.user_id, `usage`.user_id,
COUNT(*) AS request_count, COUNT(*) AS request_count,
CAST(COALESCE(SUM(GREATEST(COALESCE(`usage`.total_tokens, 0), 0)), 0) AS SIGNED) AS total_tokens CAST(COALESCE(SUM(GREATEST(COALESCE(`usage`.total_tokens, 0), 0)), 0) AS SIGNED) AS total_tokens
FROM `usage` FROM `usage`
JOIN ( JOIN (
"#, "#,
@@ -1509,6 +1509,7 @@ mod tests {
assert!(source.contains("summarize_usage_daily_heatmap_from_daily_aggregates")); assert!(source.contains("summarize_usage_daily_heatmap_from_daily_aggregates"));
assert!(source.contains("FROM stats_daily")); assert!(source.contains("FROM stats_daily"));
assert!(source.contains("FROM stats_user_daily")); assert!(source.contains("FROM stats_user_daily"));
assert!(source.contains("AS SIGNED) AS total_tokens"));
assert!(source.contains("summaries.entry(item.date.clone()).or_insert(item)")); assert!(source.contains("summaries.entry(item.date.clone()).or_insert(item)"));
} }
@@ -1518,6 +1519,7 @@ mod tests {
assert!(source.contains("async fn summarize_usage_totals_by_user_ids")); assert!(source.contains("async fn summarize_usage_totals_by_user_ids"));
assert!(source.contains("FROM stats_user_daily")); assert!(source.contains("FROM stats_user_daily"));
assert!(source.contains("MAX(`date`) AS latest_date")); assert!(source.contains("MAX(`date`) AS latest_date"));
assert!(source.contains("AS SIGNED) AS request_count"));
assert!(source.contains("requested.cutoff_unix_secs")); assert!(source.contains("requested.cutoff_unix_secs"));
} }
@@ -1529,6 +1531,7 @@ mod tests {
assert!(source.contains("FROM stats_daily")); assert!(source.contains("FROM stats_daily"));
assert!(source.contains("FROM stats_user_daily")); assert!(source.contains("FROM stats_user_daily"));
assert!(source.contains("'aggregate' AS model")); assert!(source.contains("'aggregate' AS model"));
assert!(source.contains("AS SIGNED) AS total_requests"));
} }
#[tokio::test] #[tokio::test]
+3
View File
@@ -142,6 +142,8 @@ export interface UserExport {
unlimited?: boolean unlimited?: boolean
wallet?: BillingSummary | null wallet?: BillingSummary | null
is_active: boolean is_active: boolean
request_count?: number
total_tokens?: number
api_keys: UserApiKeyExport[] api_keys: UserApiKeyExport[]
} }
@@ -237,6 +239,7 @@ export interface UsageAggregateImportSummary {
export interface GlobalModelExport { export interface GlobalModelExport {
name: string name: string
display_name: string display_name: string
usage_count?: number | null
default_price_per_request?: number | null default_price_per_request?: number | null
default_tiered_pricing: Record<string, unknown> default_tiered_pricing: Record<string, unknown>
supported_capabilities?: string[] | null supported_capabilities?: string[] | null