Merge remote-tracking branch 'upstream/main'

This commit is contained in:
AAEE86
2026-06-05 08:28:43 +08:00
101 changed files with 6967 additions and 602 deletions
@@ -56,6 +56,7 @@ pub(crate) fn resolve_same_format_provider_transport_unsupported_reason_for_trac
"jina:embedding" => "jina:embedding",
"jina:rerank" => "jina:rerank",
"doubao:embedding" => "doubao:embedding",
"aliyun:multimodal_embedding" => "aliyun:multimodal_embedding",
_ => return Some("transport_api_format_unsupported"),
};
let behavior = policy::classify_same_format_provider_request_behavior(
+15
View File
@@ -0,0 +1,15 @@
pub(crate) fn normalized_signature(api_format: &str) -> Option<&'static str> {
match crate::ai_serving::normalize_api_format_alias(api_format).as_str() {
"aliyun:multimodal_embedding" => Some("aliyun:multimodal_embedding"),
_ => None,
}
}
pub(crate) fn local_path(api_format: &str) -> Option<&'static str> {
match crate::ai_serving::normalize_api_format_alias(api_format).as_str() {
"aliyun:multimodal_embedding" => {
Some("/api/v1/services/embeddings/multimodal-embedding/multimodal-embedding")
}
_ => None,
}
}
+1
View File
@@ -1,3 +1,4 @@
mod aliyun;
mod claude;
mod doubao;
mod gemini;
+9 -1
View File
@@ -1,7 +1,7 @@
use axum::routing::{any, post};
use axum::Router;
use super::{claude, doubao, gemini, jina, openai};
use super::{aliyun, claude, doubao, gemini, jina, openai};
use crate::{handlers::proxy::proxy_request, state::AppState};
// Router registration patterns live here so AI public ingress has a single mount registry.
@@ -56,6 +56,7 @@ pub(crate) fn public_api_format_local_path(api_format: &str) -> &'static str {
.or_else(|| gemini::local_path(&normalized))
.or_else(|| jina::local_path(&normalized))
.or_else(|| doubao::local_path(&normalized))
.or_else(|| aliyun::local_path(&normalized))
.unwrap_or("/")
}
@@ -66,6 +67,7 @@ pub(crate) fn normalize_admin_endpoint_signature(api_format: &str) -> Option<&'s
.or_else(|| gemini::normalized_signature(&normalized))
.or_else(|| jina::normalized_signature(&normalized))
.or_else(|| doubao::normalized_signature(&normalized))
.or_else(|| aliyun::normalized_signature(&normalized))
}
pub(crate) fn admin_endpoint_signature_parts(
@@ -101,6 +103,12 @@ mod tests {
),
("jina:embedding", "jina", "embedding", "/v1/embeddings"),
("doubao:embedding", "doubao", "embedding", "/v1/embeddings"),
(
"aliyun:multimodal_embedding",
"aliyun",
"multimodal_embedding",
"/api/v1/services/embeddings/multimodal-embedding/multimodal-embedding",
),
("openai:rerank", "openai", "rerank", "/v1/rerank"),
("jina:rerank", "jina", "rerank", "/v1/rerank"),
] {
@@ -715,6 +715,21 @@ impl GatewayDataState {
}
}
pub(crate) async fn find_pending_plan_purchase_order_by_user_id(
&self,
user_id: &str,
product_id: &str,
) -> Result<Option<StoredAdminPaymentOrder>, DataLayerError> {
match &self.wallet_reader {
Some(repository) => {
repository
.find_pending_plan_purchase_order_by_user_id(user_id, product_id)
.await
}
None => Ok(None),
}
}
pub(crate) async fn find_wallet_refund(
&self,
wallet_id: &str,
@@ -958,6 +958,64 @@ impl GatewayDataState {
}
}
#[cfg(test)]
pub(crate) fn with_user_billing_and_wallet_for_tests<T>(
user_repository: Arc<dyn UserReadRepository>,
billing_repository: Arc<dyn BillingReadRepository>,
wallet_repository: Arc<T>,
) -> Self
where
T: aether_data::repository::wallet::WalletRepository + 'static,
{
let wallet_reader: Arc<dyn WalletReadRepository> = wallet_repository.clone();
let wallet_writer: Arc<dyn WalletWriteRepository> = wallet_repository;
Self {
config: GatewayDataConfig::disabled(),
backends: None,
auth_api_key_reader: None,
auth_api_key_writer: None,
auth_module_reader: None,
auth_module_writer: None,
announcement_reader: None,
announcement_writer: None,
management_token_reader: None,
management_token_writer: None,
oauth_provider_reader: None,
oauth_provider_writer: None,
proxy_node_reader: None,
proxy_node_writer: None,
billing_reader: Some(billing_repository),
gemini_file_mapping_reader: None,
gemini_file_mapping_writer: None,
global_model_reader: None,
global_model_writer: None,
minimal_candidate_selection_reader: None,
request_candidate_reader: None,
request_candidate_writer: None,
provider_catalog_reader: None,
provider_catalog_writer: None,
pool_score_reader: None,
pool_score_writer: None,
provider_quota_reader: None,
provider_quota_writer: None,
routing_group_reader: None,
routing_group_writer: None,
usage_reader: None,
usage_writer: None,
user_reader: Some(user_repository),
user_preferences: None,
usage_worker_queue: None,
video_task_reader: None,
video_task_writer: None,
background_task_reader: None,
background_task_writer: None,
wallet_reader: Some(wallet_reader),
wallet_writer: Some(wallet_writer),
settlement_writer: None,
system_config_values: None,
}
}
#[cfg(test)]
pub(crate) fn with_user_wallet_and_usage_for_tests<TUsage, TWallet>(
user_repository: Arc<dyn UserReadRepository>,
@@ -1,14 +1,21 @@
use std::collections::{BTreeMap, HashMap};
use std::fmt::Write as _;
use std::sync::{Mutex, OnceLock};
use std::time::{Duration, Instant};
use aether_runtime_state::{DataLayerError, RuntimeState};
use serde::{Deserialize, Serialize};
use serde_json::Value;
use sha2::{Digest, Sha256};
use tracing::warn;
use crate::clock::current_unix_ms;
const DEFAULT_CACHE_TTL: Duration = Duration::from_secs(300);
const ONE_HOUR_CACHE_TTL: Duration = Duration::from_secs(3600);
const MAX_ENTRIES: usize = 2048;
const PREFIX_LOOKBACK_LIMIT: usize = 10;
const KIRO_PROMPT_CACHE_INDEX_KEY: &str = "kiro:prompt-cache:index";
const PREFIX_LOOKBACK_WINDOW: usize = 20;
const TOKENS_PER_TOOL: u64 = 150;
const TOKENS_PER_MESSAGE: u64 = 4;
const INLINE_IMAGE_DATA_TOKEN_PLACEHOLDER: &str = "[inline-image-data]";
@@ -21,6 +28,7 @@ pub(crate) struct KiroPromptCacheProfile {
total_input_tokens: u64,
min_cacheable_tokens: u64,
breakpoints: Vec<KiroPromptCacheBreakpoint>,
match_candidates: Vec<KiroPromptCacheCandidate>,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
@@ -37,6 +45,12 @@ struct KiroPromptCacheEntry {
expires_at: Instant,
}
#[derive(Debug, Clone, Copy, Deserialize, Serialize)]
struct KiroPromptCacheRuntimeEntry {
token_count: u64,
ttl_secs: u64,
}
#[derive(Debug, Clone, Copy, Default, PartialEq, Eq)]
pub(crate) struct KiroPromptCacheUsage {
pub(crate) cache_creation_input_tokens: u64,
@@ -53,11 +67,10 @@ struct PendingBlock {
value: Value,
tokens: u64,
breakpoint_ttl: Option<Duration>,
is_message_end: bool,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
struct PrefixCandidate {
struct KiroPromptCacheCandidate {
fingerprint: [u8; 32],
cumulative_tokens: u64,
}
@@ -66,6 +79,230 @@ pub(crate) fn kiro_prompt_cache_tracker() -> &'static KiroPromptCacheTracker {
KIRO_PROMPT_CACHE_TRACKER.get_or_init(KiroPromptCacheTracker::default)
}
pub(crate) async fn compute_kiro_prompt_cache_usage(
runtime_state: &RuntimeState,
credential_id: String,
profile: &KiroPromptCacheProfile,
) -> KiroPromptCacheUsage {
match compute_kiro_prompt_cache_usage_with_runtime_state(
runtime_state,
credential_id.as_str(),
profile,
)
.await
{
Ok(usage) => usage,
Err(err) => {
warn!(
event_name = "kiro_simulated_cache_runtime_state_failed",
log_type = "event",
error = ?err,
"failed to update Kiro simulated cache runtime state; falling back to process-local tracker"
);
kiro_prompt_cache_tracker().compute_and_update(credential_id, profile)
}
}
}
async fn compute_kiro_prompt_cache_usage_with_runtime_state(
runtime_state: &RuntimeState,
credential_id: &str,
profile: &KiroPromptCacheProfile,
) -> Result<KiroPromptCacheUsage, DataLayerError> {
let last_breakpoint = match profile.breakpoints.last().copied() {
Some(last_breakpoint) => last_breakpoint,
None => return Ok(KiroPromptCacheUsage::default()),
};
let reversed_candidates = profile
.match_candidates
.iter()
.rev()
.copied()
.collect::<Vec<_>>();
let candidate_keys = reversed_candidates
.iter()
.map(|candidate| kiro_prompt_cache_runtime_key(credential_id, &candidate.fingerprint))
.collect::<Vec<_>>();
let candidate_values = runtime_state.kv_get_many(&candidate_keys).await?;
let mut existing_entries = HashMap::<String, KiroPromptCacheRuntimeEntry>::new();
let mut matched_tokens = 0u64;
let mut matched_refresh: Option<(String, KiroPromptCacheRuntimeEntry)> = None;
for ((candidate, key), value) in reversed_candidates
.iter()
.zip(candidate_keys.iter())
.zip(candidate_values)
{
let Some(entry) = value
.as_deref()
.and_then(parse_kiro_prompt_cache_runtime_entry)
else {
continue;
};
existing_entries.insert(key.clone(), entry);
if matched_tokens == 0 {
matched_tokens = entry
.token_count
.min(candidate.cumulative_tokens)
.min(profile.total_input_tokens);
matched_refresh = Some((key.clone(), entry));
}
}
if let Some((key, entry)) = matched_refresh {
store_kiro_prompt_cache_runtime_entry(runtime_state, key.as_str(), entry).await?;
}
let creation_tokens = last_breakpoint
.cumulative_tokens
.min(profile.total_input_tokens)
.saturating_sub(matched_tokens);
for breakpoint in &profile.breakpoints {
let key = kiro_prompt_cache_runtime_key(credential_id, &breakpoint.fingerprint);
let ttl_secs = breakpoint.ttl.as_secs().max(1);
let entry = existing_entries
.get(&key)
.copied()
.map(|existing| KiroPromptCacheRuntimeEntry {
token_count: existing.token_count.max(breakpoint.cumulative_tokens),
ttl_secs: existing.ttl_secs.max(ttl_secs),
})
.unwrap_or(KiroPromptCacheRuntimeEntry {
token_count: breakpoint.cumulative_tokens,
ttl_secs,
});
store_kiro_prompt_cache_runtime_entry(runtime_state, key.as_str(), entry).await?;
}
trim_kiro_prompt_cache_runtime_state(runtime_state, MAX_ENTRIES).await;
Ok(KiroPromptCacheUsage {
cache_creation_input_tokens: creation_tokens,
cache_read_input_tokens: matched_tokens,
})
}
async fn store_kiro_prompt_cache_runtime_entry(
runtime_state: &RuntimeState,
key: &str,
entry: KiroPromptCacheRuntimeEntry,
) -> Result<(), DataLayerError> {
let ttl = Duration::from_secs(entry.ttl_secs.max(1));
runtime_state
.kv_set(
key,
encode_kiro_prompt_cache_runtime_entry(entry),
Some(ttl),
)
.await?;
let expires_at_ms = current_unix_ms().saturating_add(entry.ttl_secs.saturating_mul(1000));
if let Err(err) = runtime_state
.score_set(KIRO_PROMPT_CACHE_INDEX_KEY, key, expires_at_ms as f64)
.await
{
warn!(
event_name = "kiro_simulated_cache_index_update_failed",
log_type = "event",
cache_key = %key,
error = ?err,
"failed to update Kiro simulated cache index; cache entry was persisted but cleanup may lag"
);
}
Ok(())
}
async fn trim_kiro_prompt_cache_runtime_state(runtime_state: &RuntimeState, max_entries: usize) {
if let Err(err) = runtime_state
.score_remove_by_score(KIRO_PROMPT_CACHE_INDEX_KEY, current_unix_ms() as f64)
.await
{
warn!(
event_name = "kiro_simulated_cache_index_expiry_trim_failed",
log_type = "event",
error = ?err,
"failed to trim expired Kiro simulated cache index entries"
);
return;
}
let Ok(index_len) = runtime_state.score_len(KIRO_PROMPT_CACHE_INDEX_KEY).await else {
return;
};
if index_len <= max_entries {
return;
}
let Ok(all_members) = runtime_state
.score_range_by_min(KIRO_PROMPT_CACHE_INDEX_KEY, 0.0)
.await
else {
return;
};
let trim_count = index_len.saturating_sub(max_entries);
if trim_count == 0 {
return;
}
let trimmed_members = all_members.into_iter().take(trim_count).collect::<Vec<_>>();
if let Err(err) = runtime_state.kv_delete_many(&trimmed_members).await {
warn!(
event_name = "kiro_simulated_cache_kv_trim_failed",
log_type = "event",
error = ?err,
trim_count,
"failed to delete trimmed Kiro simulated cache KV entries"
);
}
if let Err(err) = runtime_state
.score_remove_by_rank(KIRO_PROMPT_CACHE_INDEX_KEY, 0, trim_count as i64 - 1)
.await
{
warn!(
event_name = "kiro_simulated_cache_index_trim_failed",
log_type = "event",
error = ?err,
trim_count,
"failed to delete trimmed Kiro simulated cache index entries"
);
}
}
fn parse_kiro_prompt_cache_runtime_entry(value: &str) -> Option<KiroPromptCacheRuntimeEntry> {
serde_json::from_str::<KiroPromptCacheRuntimeEntry>(value)
.ok()
.filter(|entry| entry.token_count > 0 && entry.ttl_secs > 0)
}
fn encode_kiro_prompt_cache_runtime_entry(entry: KiroPromptCacheRuntimeEntry) -> String {
serde_json::to_string(&entry).unwrap_or_else(|_| {
format!(
r#"{{"token_count":{},"ttl_secs":{}}}"#,
entry.token_count, entry.ttl_secs
)
})
}
fn kiro_prompt_cache_runtime_key(credential_id: &str, fingerprint: &[u8; 32]) -> String {
let credential_hash: [u8; 32] = Sha256::digest(credential_id.as_bytes()).into();
format!(
"kiro:prompt-cache:{}:{}",
hex_digest(&credential_hash),
hex_digest(fingerprint)
)
}
fn hex_digest(bytes: &[u8]) -> String {
let mut output = String::with_capacity(bytes.len() * 2);
for byte in bytes {
let _ = write!(&mut output, "{byte:02x}");
}
output
}
pub(crate) fn build_kiro_prompt_cache_profile(
request_body: &Value,
total_input_tokens: u64,
@@ -75,7 +312,8 @@ pub(crate) fn build_kiro_prompt_cache_profile(
.and_then(Value::as_str)
.unwrap_or_default();
let flattened = flatten_cacheable_blocks(request_body);
if flattened.iter().all(|block| block.breakpoint_ttl.is_none()) {
let automatic_ttl = extract_cache_ttl(request_body);
if automatic_ttl.is_none() && flattened.iter().all(|block| block.breakpoint_ttl.is_none()) {
return None;
}
@@ -89,12 +327,15 @@ pub(crate) fn build_kiro_prompt_cache_profile(
prefix_hasher.update(prelude_bytes);
let mut cumulative_tokens = 0u64;
let mut active_ttl: Option<Duration> = None;
let mut breakpoints = Vec::new();
let mut seen_fingerprints = std::collections::BTreeSet::<[u8; 32]>::new();
let mut lookback_candidates = Vec::new();
let mut match_candidates = Vec::new();
let last_block_index = flattened.len().saturating_sub(1);
for block in flattened {
for (block_index, mut block) in flattened.into_iter().enumerate() {
if block.breakpoint_ttl.is_none() && block_index == last_block_index {
block.breakpoint_ttl = automatic_ttl;
}
cumulative_tokens = cumulative_tokens.saturating_add(block.tokens);
let block_bytes = serde_json::to_vec(&block.value).unwrap_or_default();
let block_hash: [u8; 32] = Sha256::digest(block_bytes).into();
@@ -105,13 +346,6 @@ pub(crate) fn build_kiro_prompt_cache_profile(
prefix_hasher.update(fingerprint);
if let Some(ttl) = block.breakpoint_ttl {
push_lookback_breakpoints(
&mut breakpoints,
&mut seen_fingerprints,
&lookback_candidates,
ttl,
);
active_ttl = Some(ttl);
push_breakpoint(
&mut breakpoints,
&mut seen_fingerprints,
@@ -120,18 +354,7 @@ pub(crate) fn build_kiro_prompt_cache_profile(
ttl,
);
}
if block.is_message_end {
if let Some(ttl) = active_ttl {
push_breakpoint(
&mut breakpoints,
&mut seen_fingerprints,
fingerprint,
cumulative_tokens,
ttl,
);
}
}
push_prefix_candidate(&mut lookback_candidates, fingerprint, cumulative_tokens);
push_match_candidate(&mut match_candidates, fingerprint, cumulative_tokens);
}
let min_cacheable_tokens = minimum_cacheable_tokens_for_model(model);
@@ -139,13 +362,54 @@ pub(crate) fn build_kiro_prompt_cache_profile(
.into_iter()
.filter(|breakpoint| breakpoint.cumulative_tokens >= min_cacheable_tokens)
.collect::<Vec<_>>();
(!cacheable_breakpoints.is_empty()).then_some(KiroPromptCacheProfile {
if cacheable_breakpoints.is_empty() {
return None;
}
let match_candidates = build_lookback_match_candidates(
&match_candidates,
&cacheable_breakpoints,
min_cacheable_tokens,
);
Some(KiroPromptCacheProfile {
total_input_tokens,
min_cacheable_tokens,
breakpoints: cacheable_breakpoints,
match_candidates,
})
}
fn build_lookback_match_candidates(
candidates: &[KiroPromptCacheCandidate],
breakpoints: &[KiroPromptCacheBreakpoint],
min_cacheable_tokens: u64,
) -> Vec<KiroPromptCacheCandidate> {
let mut out = Vec::new();
let mut seen_fingerprints = std::collections::BTreeSet::<[u8; 32]>::new();
for breakpoint in breakpoints {
let Some(index) = candidates
.iter()
.position(|candidate| candidate.fingerprint == breakpoint.fingerprint)
else {
continue;
};
let start = index
.saturating_add(1)
.saturating_sub(PREFIX_LOOKBACK_WINDOW);
for candidate in &candidates[start..=index] {
if candidate.cumulative_tokens < min_cacheable_tokens
|| candidate.cumulative_tokens > breakpoint.cumulative_tokens
|| !seen_fingerprints.insert(candidate.fingerprint)
{
continue;
}
out.push(*candidate);
}
}
out
}
pub(crate) fn kiro_simulated_cache_enabled_from_provider_config(config: Option<&Value>) -> bool {
config
.and_then(Value::as_object)
@@ -269,35 +533,15 @@ fn push_breakpoint(
}
}
fn push_lookback_breakpoints(
breakpoints: &mut Vec<KiroPromptCacheBreakpoint>,
seen_fingerprints: &mut std::collections::BTreeSet<[u8; 32]>,
candidates: &[PrefixCandidate],
ttl: Duration,
) {
for candidate in candidates {
push_breakpoint(
breakpoints,
seen_fingerprints,
candidate.fingerprint,
candidate.cumulative_tokens,
ttl,
);
}
}
fn push_prefix_candidate(
candidates: &mut Vec<PrefixCandidate>,
fn push_match_candidate(
candidates: &mut Vec<KiroPromptCacheCandidate>,
fingerprint: [u8; 32],
cumulative_tokens: u64,
) {
candidates.push(PrefixCandidate {
candidates.push(KiroPromptCacheCandidate {
fingerprint,
cumulative_tokens,
});
if candidates.len() > PREFIX_LOOKBACK_LIMIT {
candidates.remove(0);
}
}
fn flatten_cacheable_blocks(request_body: &Value) -> Vec<PendingBlock> {
@@ -316,7 +560,6 @@ fn flatten_cacheable_blocks(request_body: &Value) -> Vec<PendingBlock> {
tokens: TOKENS_PER_TOOL,
value,
breakpoint_ttl,
is_message_end: false,
});
}
}
@@ -337,7 +580,6 @@ fn flatten_cacheable_blocks(request_body: &Value) -> Vec<PendingBlock> {
tokens: count_system_block_tokens(item),
value,
breakpoint_ttl,
is_message_end: false,
});
}
}
@@ -351,7 +593,6 @@ fn flatten_cacheable_blocks(request_body: &Value) -> Vec<PendingBlock> {
tokens: count_text_tokens(text),
value,
breakpoint_ttl: None,
is_message_end: false,
});
}
other => {
@@ -364,7 +605,6 @@ fn flatten_cacheable_blocks(request_body: &Value) -> Vec<PendingBlock> {
tokens: count_system_block_tokens(other),
value,
breakpoint_ttl: None,
is_message_end: false,
});
}
}
@@ -376,11 +616,17 @@ fn flatten_cacheable_blocks(request_body: &Value) -> Vec<PendingBlock> {
.get("role")
.and_then(Value::as_str)
.unwrap_or_default();
let message_breakpoint_ttl = extract_cache_ttl(message);
match message.get("content") {
Some(Value::Array(items)) => {
let last_block_index = items.len().saturating_sub(1);
for (block_index, item) in items.iter().enumerate() {
let breakpoint_ttl = extract_cache_ttl(item);
let breakpoint_ttl =
extract_cache_ttl(item).or(if block_index == last_block_index {
message_breakpoint_ttl
} else {
None
});
let mut normalized = item.clone();
strip_cache_control(&mut normalized);
let value = canonicalize_json(serde_json::json!({
@@ -394,7 +640,6 @@ fn flatten_cacheable_blocks(request_body: &Value) -> Vec<PendingBlock> {
tokens: count_message_content_tokens(item),
value,
breakpoint_ttl,
is_message_end: block_index == last_block_index,
});
}
}
@@ -409,8 +654,7 @@ fn flatten_cacheable_blocks(request_body: &Value) -> Vec<PendingBlock> {
blocks.push(PendingBlock {
tokens: count_text_tokens(text),
value,
breakpoint_ttl: None,
is_message_end: true,
breakpoint_ttl: message_breakpoint_ttl,
});
}
Some(other) => {
@@ -424,8 +668,7 @@ fn flatten_cacheable_blocks(request_body: &Value) -> Vec<PendingBlock> {
blocks.push(PendingBlock {
tokens: count_message_content_tokens(other),
value,
breakpoint_ttl: None,
is_message_end: true,
breakpoint_ttl: message_breakpoint_ttl,
});
}
None => {}
@@ -634,20 +877,16 @@ impl KiroPromptCacheTracker {
};
let mut matched_tokens = 0;
for breakpoint in profile
.breakpoints
.iter()
.rev()
.take(PREFIX_LOOKBACK_LIMIT.saturating_add(1))
{
let key = (credential_id.clone(), breakpoint.fingerprint);
let Some(entry) = entries.get(&key) else {
for candidate in profile.match_candidates.iter().rev() {
let key = (credential_id.clone(), candidate.fingerprint);
let Some(entry) = entries.get_mut(&key) else {
continue;
};
if entry.expires_at > now {
entry.expires_at = entry.expires_at.max(now + entry.ttl);
matched_tokens = entry
.token_count
.min(breakpoint.cumulative_tokens)
.min(candidate.cumulative_tokens)
.min(profile.total_input_tokens);
break;
}
@@ -664,6 +903,7 @@ impl KiroPromptCacheTracker {
Some(existing) => {
existing.token_count = existing.token_count.max(breakpoint.cumulative_tokens);
existing.ttl = existing.ttl.max(breakpoint.ttl);
existing.expires_at = existing.expires_at.max(now + existing.ttl);
}
None => {
self.evict_to_capacity(&mut entries);
@@ -702,6 +942,7 @@ impl KiroPromptCacheTracker {
#[cfg(test)]
mod tests {
use super::*;
use aether_runtime_state::MemoryRuntimeStateConfig;
fn long_text(label: &str) -> String {
format!("{} {}", label, "cacheable prompt chunk ".repeat(300))
@@ -756,7 +997,138 @@ mod tests {
}
#[test]
fn tracker_supports_prefix_hits_without_extending_expiry() {
fn profile_reads_top_level_automatic_cache_control() {
let request = serde_json::json!({
"model": "claude-sonnet-4.6",
"cache_control": {"type": "ephemeral"},
"messages": [{
"role": "user",
"content": long_text("automatic cached turn")
}]
});
let profile =
build_kiro_prompt_cache_profile(&request, estimate_kiro_prompt_input_tokens(&request))
.expect("top-level cache_control should create an automatic cache profile");
let tracker = KiroPromptCacheTracker::default();
let usage = tracker.compute_and_update("cred".to_string(), &profile);
assert_eq!(profile.breakpoints.len(), 1);
assert!(usage.cache_creation_input_tokens > 0);
}
#[test]
fn profile_does_not_create_message_end_breakpoints_from_explicit_cache_control() {
let request = serde_json::json!({
"model": "claude-sonnet-4.6",
"system": [{
"type": "text",
"text": long_text("explicit cached system"),
"cache_control": {"type": "ephemeral"}
}],
"messages": [{
"role": "user",
"content": long_text("uncached later turn")
}]
});
let profile =
build_kiro_prompt_cache_profile(&request, estimate_kiro_prompt_input_tokens(&request))
.expect("explicit cache_control should create a cache profile");
assert_eq!(profile.breakpoints.len(), 1);
assert!(profile.breakpoints[0].cumulative_tokens < profile.total_input_tokens);
}
#[tokio::test]
async fn runtime_state_tracker_reads_cached_prefix_across_calls() {
let request = serde_json::json!({
"model": "claude-sonnet-4.6",
"system": [{
"type": "text",
"text": long_text("runtime shared system"),
"cache_control": {"type": "ephemeral"}
}],
"messages": [{"role": "user", "content": "reuse runtime cache"}]
});
let profile =
build_kiro_prompt_cache_profile(&request, estimate_kiro_prompt_input_tokens(&request))
.expect("cacheable request should create a cache profile");
let runtime = RuntimeState::memory(MemoryRuntimeStateConfig::default());
let first =
compute_kiro_prompt_cache_usage(&runtime, "runtime-cred".to_string(), &profile).await;
let second =
compute_kiro_prompt_cache_usage(&runtime, "runtime-cred".to_string(), &profile).await;
assert!(first.cache_creation_input_tokens > 0);
assert_eq!(first.cache_read_input_tokens, 0);
assert_eq!(second.cache_creation_input_tokens, 0);
assert!(second.cache_read_input_tokens > 0);
}
#[tokio::test]
async fn runtime_state_tracker_trims_oldest_entries_to_capacity() {
let runtime = RuntimeState::memory(MemoryRuntimeStateConfig::default());
let now_ms = current_unix_ms();
let keys = [
"kiro:prompt-cache:test-oldest".to_string(),
"kiro:prompt-cache:test-middle".to_string(),
"kiro:prompt-cache:test-newest".to_string(),
];
for (index, key) in keys.iter().enumerate() {
runtime
.kv_set(
key,
encode_kiro_prompt_cache_runtime_entry(KiroPromptCacheRuntimeEntry {
token_count: 100 + index as u64,
ttl_secs: 120,
}),
Some(Duration::from_secs(120)),
)
.await
.expect("cache entry should store");
runtime
.score_set(
KIRO_PROMPT_CACHE_INDEX_KEY,
key,
now_ms.saturating_add(60_000 + index as u64 * 1_000) as f64,
)
.await
.expect("cache index should store");
}
trim_kiro_prompt_cache_runtime_state(&runtime, 2).await;
assert_eq!(
runtime
.kv_get(&keys[0])
.await
.expect("oldest entry should read"),
None
);
assert!(runtime
.kv_get(&keys[1])
.await
.expect("middle entry should read")
.is_some());
assert!(runtime
.kv_get(&keys[2])
.await
.expect("newest entry should read")
.is_some());
assert_eq!(
runtime
.score_range_by_min(KIRO_PROMPT_CACHE_INDEX_KEY, 0.0)
.await
.expect("cache index should read"),
vec![keys[1].clone(), keys[2].clone()]
);
}
#[test]
fn tracker_refreshes_cached_prefix_ttl_on_read() {
let base = serde_json::json!({
"model": "claude-sonnet-4.6",
"system": [{
@@ -799,13 +1171,13 @@ mod tests {
);
assert!(hit.cache_read_input_tokens > 0);
let expired = tracker.compute_and_update_at(
let refreshed = tracker.compute_and_update_at(
"cred".to_string(),
&base_profile,
start + Duration::from_secs(301),
);
assert!(expired.cache_creation_input_tokens > 0);
assert_eq!(expired.cache_read_input_tokens, 0);
assert_eq!(refreshed.cache_creation_input_tokens, 0);
assert!(refreshed.cache_read_input_tokens > 0);
}
#[test]
@@ -867,6 +1239,146 @@ mod tests {
assert!(hit.cache_creation_input_tokens > 0);
}
#[test]
fn tracker_reads_cached_prefix_within_prompt_cache_lookback_window() {
let first = serde_json::json!({
"model": "claude-sonnet-4.6",
"messages": [{
"role": "user",
"content": [{
"type": "text",
"text": long_text("shared first turn"),
"cache_control": {"type": "ephemeral"}
}]
}]
});
let mut second_messages = vec![serde_json::json!({
"role": "user",
"content": [{
"type": "text",
"text": long_text("shared first turn")
}]
})];
for index in 0..12 {
second_messages.push(serde_json::json!({
"role": if index % 2 == 0 { "assistant" } else { "user" },
"content": format!("intermediate turn {index}")
}));
}
second_messages.push(serde_json::json!({
"role": "user",
"content": [{
"type": "text",
"text": long_text("new tail turn"),
"cache_control": {"type": "ephemeral"}
}]
}));
let second = serde_json::json!({
"model": "claude-sonnet-4.6",
"messages": second_messages
});
let first_profile =
build_kiro_prompt_cache_profile(&first, estimate_kiro_prompt_input_tokens(&first))
.expect("first request should be cacheable");
let second_profile =
build_kiro_prompt_cache_profile(&second, estimate_kiro_prompt_input_tokens(&second))
.expect("second request should be cacheable");
let tracker = KiroPromptCacheTracker::default();
let start = Instant::now();
let created = tracker.compute_and_update_at("cred".to_string(), &first_profile, start);
assert!(created.cache_creation_input_tokens > 0);
assert_eq!(created.cache_read_input_tokens, 0);
let hit = tracker.compute_and_update_at(
"cred".to_string(),
&second_profile,
start + Duration::from_secs(60),
);
assert!(hit.cache_read_input_tokens > 0);
assert!(hit.cache_creation_input_tokens > 0);
}
#[test]
fn tracker_does_not_read_cached_prefix_outside_prompt_cache_lookback_window() {
let first = serde_json::json!({
"model": "claude-sonnet-4.6",
"messages": [{
"role": "user",
"content": [{
"type": "text",
"text": long_text("shared first turn"),
"cache_control": {"type": "ephemeral"}
}]
}]
});
let mut second_messages = vec![serde_json::json!({
"role": "user",
"content": [{
"type": "text",
"text": long_text("shared first turn")
}]
})];
for index in 0..20 {
second_messages.push(serde_json::json!({
"role": if index % 2 == 0 { "assistant" } else { "user" },
"content": format!("intermediate turn {index}")
}));
}
second_messages.push(serde_json::json!({
"role": "user",
"content": [{
"type": "text",
"text": long_text("new tail turn"),
"cache_control": {"type": "ephemeral"}
}]
}));
let second = serde_json::json!({
"model": "claude-sonnet-4.6",
"messages": second_messages
});
let first_profile =
build_kiro_prompt_cache_profile(&first, estimate_kiro_prompt_input_tokens(&first))
.expect("first request should be cacheable");
let second_profile =
build_kiro_prompt_cache_profile(&second, estimate_kiro_prompt_input_tokens(&second))
.expect("second request should be cacheable");
let tracker = KiroPromptCacheTracker::default();
let start = Instant::now();
let created = tracker.compute_and_update_at("cred".to_string(), &first_profile, start);
assert!(created.cache_creation_input_tokens > 0);
assert_eq!(created.cache_read_input_tokens, 0);
let miss = tracker.compute_and_update_at(
"cred".to_string(),
&second_profile,
start + Duration::from_secs(60),
);
assert!(miss.cache_creation_input_tokens > 0);
assert_eq!(miss.cache_read_input_tokens, 0);
}
#[test]
fn profile_reads_message_level_cache_control() {
let request = serde_json::json!({
"model": "claude-sonnet-4.6",
"messages": [{
"role": "system",
"content": long_text("message level cached system"),
"cache_control": {"type": "ephemeral"}
}]
});
let profile =
build_kiro_prompt_cache_profile(&request, estimate_kiro_prompt_input_tokens(&request))
.expect("message-level cache_control should create a cache profile");
let tracker = KiroPromptCacheTracker::default();
let usage = tracker.compute_and_update("cred".to_string(), &profile);
assert!(usage.cache_creation_input_tokens > 0);
}
#[test]
fn billed_input_tokens_subtracts_cache_usage() {
assert_eq!(
@@ -15,8 +15,8 @@ use tracing::{debug, warn};
use uuid::Uuid;
use crate::execution_runtime::kiro_cache::{
billed_input_tokens, build_kiro_prompt_cache_profile, estimate_kiro_prompt_input_tokens,
kiro_prompt_cache_tracker, kiro_simulated_cache_enabled_from_provider_config,
billed_input_tokens, build_kiro_prompt_cache_profile, compute_kiro_prompt_cache_usage,
estimate_kiro_prompt_input_tokens, kiro_simulated_cache_enabled_from_provider_config,
KiroPromptCacheProfile, KiroPromptCacheUsage,
};
use crate::execution_runtime::ndjson::encode_stream_frame_ndjson;
@@ -160,14 +160,17 @@ pub(crate) async fn maybe_execute_kiro_web_search_stream(
let search_results = parse_mcp_search_results(&mcp_execution.result);
let cache_usage = if kiro_simulated_cache_enabled(state, plan).await {
request
.cache_profile
.as_ref()
.map(|profile| {
kiro_prompt_cache_tracker()
.compute_and_update(kiro_cache_credential_id(plan), profile)
})
.unwrap_or_default()
match request.cache_profile.as_ref() {
Some(profile) => {
compute_kiro_prompt_cache_usage(
state.runtime_state(),
kiro_cache_credential_id(plan),
profile,
)
.await
}
None => KiroPromptCacheUsage::default(),
}
} else {
KiroPromptCacheUsage::default()
};
@@ -65,7 +65,7 @@ use crate::execution_runtime::chatgpt_web_image::maybe_execute_chatgpt_web_image
use crate::execution_runtime::grok::maybe_execute_grok_stream;
use crate::execution_runtime::kiro_cache::{
billed_input_tokens as kiro_billed_input_tokens, build_kiro_prompt_cache_profile,
estimate_kiro_prompt_input_tokens, kiro_prompt_cache_tracker,
compute_kiro_prompt_cache_usage, estimate_kiro_prompt_input_tokens,
kiro_simulated_cache_enabled_from_provider_config,
kiro_simulated_cache_enabled_from_report_context, KiroPromptCacheUsage,
KIRO_SIMULATED_CACHE_ENABLED_CONTEXT_FIELD,
@@ -402,7 +402,8 @@ async fn seed_kiro_simulated_cache_enabled(
}
}
fn seed_kiro_report_context_prompt_cache_usage(
async fn seed_kiro_report_context_prompt_cache_usage(
state: &AppState,
plan: &ExecutionPlan,
report_context: &mut Option<Value>,
) {
@@ -450,8 +451,12 @@ fn seed_kiro_report_context_prompt_cache_usage(
return;
};
let cache_usage = kiro_prompt_cache_tracker()
.compute_and_update(kiro_stream_cache_credential_id(plan), &profile);
let cache_usage = compute_kiro_prompt_cache_usage(
state.runtime_state(),
kiro_stream_cache_credential_id(plan),
&profile,
)
.await;
if cache_usage.cache_creation_input_tokens == 0 && cache_usage.cache_read_input_tokens == 0 {
return;
}
@@ -494,7 +499,8 @@ fn kiro_cache_usage_from_report_context(report_context: &Value) -> Option<KiroPr
.and_then(kiro_cache_usage_from_context_object)
}
fn maybe_apply_kiro_prompt_cache_usage_to_stream_summary(
async fn maybe_apply_kiro_prompt_cache_usage_to_stream_summary(
state: &AppState,
plan: &ExecutionPlan,
report_context: Option<&Value>,
summary: &mut Option<ExecutionStreamTerminalSummary>,
@@ -572,8 +578,12 @@ fn maybe_apply_kiro_prompt_cache_usage_to_stream_summary(
return;
};
let cache_usage = kiro_prompt_cache_tracker()
.compute_and_update(kiro_stream_cache_credential_id(plan), &profile);
let cache_usage = compute_kiro_prompt_cache_usage(
state.runtime_state(),
kiro_stream_cache_credential_id(plan),
&profile,
)
.await;
if cache_usage.cache_creation_input_tokens == 0 && cache_usage.cache_read_input_tokens == 0 {
return;
}
@@ -831,7 +841,6 @@ pub(crate) async fn execute_execution_runtime_stream(
) -> Result<Option<Response<Body>>, GatewayError> {
let stream_started_at = Instant::now();
ensure_execution_request_candidate_slot(state, &mut plan, &mut report_context).await;
seed_kiro_report_context_input_tokens(&plan, &mut report_context);
let lifecycle_seed = build_lifecycle_usage_seed(&plan, report_context.as_ref());
let request_candidate_status_snapshot =
snapshot_local_request_candidate_status(&plan, report_context.as_ref());
@@ -1981,10 +1990,12 @@ async fn execute_stream_from_frame_stream(
};
let mut report_context =
attach_provider_response_headers_to_report_context(report_context, &headers);
seed_kiro_report_context_input_tokens(&plan, &mut report_context);
if status_code == 200 {
seed_kiro_simulated_cache_enabled(state, &plan, &mut report_context).await;
seed_kiro_report_context_prompt_cache_usage(&plan, &mut report_context);
if kiro_simulated_cache_enabled_from_report_context(report_context.as_ref()) {
seed_kiro_report_context_input_tokens(&plan, &mut report_context);
}
seed_kiro_report_context_prompt_cache_usage(state, &plan, &mut report_context).await;
}
let mut buffered_frames = VecDeque::new();
let mut stream_terminal_summary: Option<ExecutionStreamTerminalSummary> = None;
@@ -3357,7 +3368,13 @@ async fn execute_stream_from_frame_stream(
"gateway skipped client stream flush after downstream disconnect"
);
}
if let Some(normalizer) = private_stream_normalizer.as_mut() {
// Buffered stream state is partial after a terminal failure; normal
// finish paths may synthesize successful terminal events.
let should_finish_stream_rewriters = terminal_failure.is_none();
if let Some(normalizer) = private_stream_normalizer
.as_mut()
.filter(|_| should_finish_stream_rewriters)
{
match normalizer.finish() {
Ok(normalized_chunk) if !normalized_chunk.is_empty() => {
let provider_private_error_body_json =
@@ -3391,13 +3408,12 @@ async fn execute_stream_from_frame_stream(
error = ?err,
"gateway failed to rewrite normalized private stream chunk during flush"
);
terminal_failure.get_or_insert_with(|| {
build_stream_failure_report(
"execution_runtime_stream_rewrite_flush_error",
format!("failed to rewrite normalized private stream chunk during flush: {err:?}"),
502,
)
});
let failure = build_stream_failure_report(
"execution_runtime_stream_rewrite_flush_error",
format!("failed to rewrite normalized private stream chunk during flush: {err:?}"),
502,
);
terminal_failure.get_or_insert(failure);
Vec::new()
}
}
@@ -3472,7 +3488,7 @@ async fn execute_stream_from_frame_stream(
}
}
}
if !downstream_dropped {
if !downstream_dropped && terminal_failure.is_none() {
if let Some(rewriter) = local_stream_rewriter.as_mut() {
match rewriter.finish() {
Ok(flushed_chunk) if !flushed_chunk.is_empty() => {
@@ -3700,10 +3716,12 @@ async fn execute_stream_from_frame_stream(
}
maybe_apply_kiro_prompt_cache_usage_to_stream_summary(
&state_for_report,
&plan_for_report,
report_context_owned.as_ref(),
&mut stream_terminal_summary,
);
)
.await;
let requires_observed_terminal_event = stream_requires_observed_terminal_event(
plan_for_report.provider_api_format.as_str(),
stream_usage_report_context.as_ref(),
@@ -3898,8 +3916,9 @@ mod tests {
use std::time::{Duration, Instant};
use aether_contracts::{
ExecutionPlan, ExecutionStreamTerminalSummary, ExecutionTimeouts, RequestBody,
StandardizedUsage,
ExecutionError, ExecutionErrorKind, ExecutionPhase, ExecutionPlan,
ExecutionStreamTerminalSummary, ExecutionTimeouts, RequestBody, StandardizedUsage,
StreamFrame, StreamFramePayload, StreamFrameType,
};
use aether_data::repository::candidates::InMemoryRequestCandidateRepository;
use aether_data::repository::usage::InMemoryUsageReadRepository;
@@ -3946,6 +3965,10 @@ mod tests {
.with_execution_runtime_candidate(true)
}
fn test_state() -> AppState {
AppState::new().expect("gateway state should build")
}
#[test]
fn detects_client_visible_sse_terminal_events() {
assert!(stream_chunk_contains_sse_done(b"data: [DONE]\n\n"));
@@ -3996,6 +4019,12 @@ mod tests {
out
}
fn ndjson_frame(frame: StreamFrame) -> Bytes {
let mut bytes = serde_json::to_vec(&frame).expect("stream frame should serialize");
bytes.push(b'\n');
Bytes::from(bytes)
}
#[test]
fn merge_stream_terminal_summary_prefers_more_complete_observed_usage() {
let mut runtime_usage = StandardizedUsage::new();
@@ -4124,8 +4153,8 @@ mod tests {
));
}
#[test]
fn kiro_stream_summary_applies_prompt_cache_usage_from_original_request() {
#[tokio::test]
async fn kiro_stream_summary_applies_prompt_cache_usage_from_original_request() {
let request_body = json!({
"model": "claude-opus-4-7",
"system": [
@@ -4173,6 +4202,7 @@ mod tests {
transport_profile: None,
timeouts: None,
};
let state = test_state();
let mut first_summary = Some(ExecutionStreamTerminalSummary {
standardized_usage: Some(StandardizedUsage {
@@ -4183,10 +4213,12 @@ mod tests {
..ExecutionStreamTerminalSummary::default()
});
maybe_apply_kiro_prompt_cache_usage_to_stream_summary(
&state,
&plan,
Some(&report_context),
&mut first_summary,
);
)
.await;
let first_usage = first_summary
.as_ref()
.and_then(|summary| summary.standardized_usage.as_ref())
@@ -4203,10 +4235,12 @@ mod tests {
..ExecutionStreamTerminalSummary::default()
});
maybe_apply_kiro_prompt_cache_usage_to_stream_summary(
&state,
&plan,
Some(&report_context),
&mut second_summary,
);
)
.await;
let second_usage = second_summary
.as_ref()
.and_then(|summary| summary.standardized_usage.as_ref())
@@ -4217,8 +4251,126 @@ mod tests {
assert_eq!(second_usage.output_tokens, 19);
}
#[test]
fn kiro_stream_summary_seeds_input_tokens_without_cache_control() {
#[tokio::test]
async fn kiro_stream_summary_reads_cached_prefix_within_prompt_cache_lookback_window() {
let first_request_body = json!({
"model": "claude-sonnet-4.6",
"messages": [{
"role": "user",
"content": [{
"type": "text",
"text": "shared first turn ".repeat(600),
"cache_control": {"type": "ephemeral"}
}]
}]
});
let mut second_messages = vec![json!({
"role": "user",
"content": [{
"type": "text",
"text": "shared first turn ".repeat(600)
}]
})];
for index in 0..12 {
second_messages.push(json!({
"role": if index % 2 == 0 { "assistant" } else { "user" },
"content": format!("intermediate stream turn {index}")
}));
}
second_messages.push(json!({
"role": "user",
"content": [{
"type": "text",
"text": "new tail turn ".repeat(600),
"cache_control": {"type": "ephemeral"}
}]
}));
let second_request_body = json!({
"model": "claude-sonnet-4.6",
"messages": second_messages
});
let plan = ExecutionPlan {
request_id: "req-kiro-cache-stream-long-tail".into(),
candidate_id: Some("cand-kiro-cache-stream-long-tail".into()),
provider_name: Some("Kiro".into()),
provider_id: "provider-kiro-cache-stream-long-tail".into(),
endpoint_id: "endpoint-kiro-cache-stream-long-tail".into(),
key_id: "key-kiro-cache-stream-long-tail".into(),
method: "POST".into(),
url: "https://q.us-east-1.amazonaws.com/generateAssistantResponse?beta=true".into(),
headers: BTreeMap::new(),
content_type: Some("application/json".into()),
content_encoding: None,
body: RequestBody::from_json(json!({"conversationState": {}})),
stream: true,
client_api_format: "claude:messages".into(),
provider_api_format: "claude:messages".into(),
model_name: Some("claude-sonnet-4.6".into()),
proxy: None,
transport_profile: None,
timeouts: None,
};
let first_report_context = json!({
"original_request_body": first_request_body,
"kiro_simulated_cache_enabled": true,
});
let second_report_context = json!({
"original_request_body": second_request_body,
"kiro_simulated_cache_enabled": true,
});
let state = test_state();
let mut first_summary = Some(ExecutionStreamTerminalSummary {
standardized_usage: Some(StandardizedUsage {
input_tokens: 4_000,
output_tokens: 17,
..StandardizedUsage::new()
}),
..ExecutionStreamTerminalSummary::default()
});
maybe_apply_kiro_prompt_cache_usage_to_stream_summary(
&state,
&plan,
Some(&first_report_context),
&mut first_summary,
)
.await;
let first_usage = first_summary
.as_ref()
.and_then(|summary| summary.standardized_usage.as_ref())
.expect("first usage should exist");
assert!(first_usage.cache_creation_tokens > 0);
assert_eq!(first_usage.cache_read_tokens, 0);
let mut second_summary = Some(ExecutionStreamTerminalSummary {
standardized_usage: Some(StandardizedUsage {
input_tokens: 8_000,
output_tokens: 19,
..StandardizedUsage::new()
}),
..ExecutionStreamTerminalSummary::default()
});
maybe_apply_kiro_prompt_cache_usage_to_stream_summary(
&state,
&plan,
Some(&second_report_context),
&mut second_summary,
)
.await;
let second_usage = second_summary
.as_ref()
.and_then(|summary| summary.standardized_usage.as_ref())
.expect("second usage should exist");
assert!(
second_usage.cache_read_tokens > 0,
"stream summary should reuse the far earlier cached prefix"
);
assert!(second_usage.cache_creation_tokens > 0);
assert_eq!(second_usage.output_tokens, 19);
}
#[tokio::test]
async fn kiro_stream_summary_seeds_input_tokens_without_cache_control() {
let request_body = json!({
"model": "claude-opus-4-7",
"system": [
@@ -4264,6 +4416,7 @@ mod tests {
transport_profile: None,
timeouts: None,
};
let state = test_state();
let mut summary = Some(ExecutionStreamTerminalSummary {
standardized_usage: Some(StandardizedUsage {
@@ -4275,10 +4428,12 @@ mod tests {
});
maybe_apply_kiro_prompt_cache_usage_to_stream_summary(
&state,
&plan,
Some(&report_context),
&mut summary,
);
)
.await;
let usage = summary
.as_ref()
@@ -4291,8 +4446,8 @@ mod tests {
assert_eq!(usage.output_tokens, 13);
}
#[test]
fn kiro_stream_summary_bills_existing_cache_usage_when_input_is_zero() {
#[tokio::test]
async fn kiro_stream_summary_bills_existing_cache_usage_when_input_is_zero() {
let request_body = json!({
"model": "claude-opus-4-7",
"system": [
@@ -4338,6 +4493,7 @@ mod tests {
transport_profile: None,
timeouts: None,
};
let state = test_state();
let mut summary = Some(ExecutionStreamTerminalSummary {
standardized_usage: Some(StandardizedUsage {
@@ -4350,10 +4506,12 @@ mod tests {
});
maybe_apply_kiro_prompt_cache_usage_to_stream_summary(
&state,
&plan,
Some(&report_context),
&mut summary,
);
)
.await;
let usage = summary
.as_ref()
@@ -4366,8 +4524,8 @@ mod tests {
assert_eq!(usage.output_tokens, 23);
}
#[test]
fn kiro_stream_summary_clears_cache_usage_when_simulated_cache_disabled() {
#[tokio::test]
async fn kiro_stream_summary_clears_cache_usage_when_simulated_cache_disabled() {
let request_body = json!({
"model": "claude-opus-4-7",
"system": [
@@ -4414,6 +4572,7 @@ mod tests {
transport_profile: None,
timeouts: None,
};
let state = test_state();
let mut summary = Some(ExecutionStreamTerminalSummary {
standardized_usage: Some(StandardizedUsage {
@@ -4427,10 +4586,12 @@ mod tests {
});
maybe_apply_kiro_prompt_cache_usage_to_stream_summary(
&state,
&plan,
Some(&report_context),
&mut summary,
);
)
.await;
let usage = summary
.as_ref()
@@ -4443,8 +4604,8 @@ mod tests {
assert_eq!(usage.output_tokens, 23);
}
#[test]
fn kiro_stream_summary_does_not_subtract_cache_from_already_billed_input() {
#[tokio::test]
async fn kiro_stream_summary_does_not_subtract_cache_from_already_billed_input() {
let request_body = json!({
"model": "claude-opus-4-7",
"messages": [
@@ -4492,6 +4653,7 @@ mod tests {
transport_profile: None,
timeouts: None,
};
let state = test_state();
let mut summary = Some(ExecutionStreamTerminalSummary {
standardized_usage: Some(StandardizedUsage {
@@ -4505,10 +4667,12 @@ mod tests {
});
maybe_apply_kiro_prompt_cache_usage_to_stream_summary(
&state,
&plan,
Some(&report_context),
&mut summary,
);
)
.await;
let usage = summary
.as_ref()
@@ -4628,9 +4792,11 @@ mod tests {
"original_request_body": request_body,
"kiro_simulated_cache_enabled": true,
}));
let state = AppState::new().expect("gateway state should build");
super::seed_kiro_report_context_input_tokens(&plan, &mut report_context);
super::seed_kiro_report_context_prompt_cache_usage(&plan, &mut report_context);
super::seed_kiro_report_context_prompt_cache_usage(&state, &plan, &mut report_context)
.await;
let context = report_context.as_ref().expect("context should exist");
assert!(context
@@ -4697,9 +4863,11 @@ mod tests {
let mut report_context = Some(json!({
"original_request_body": request_body,
}));
let state = AppState::new().expect("gateway state should build");
super::seed_kiro_report_context_input_tokens(&plan, &mut report_context);
super::seed_kiro_report_context_prompt_cache_usage(&plan, &mut report_context);
super::seed_kiro_report_context_prompt_cache_usage(&state, &plan, &mut report_context)
.await;
let context = report_context.as_ref().expect("context should exist");
assert!(context
@@ -5017,6 +5185,154 @@ mod tests {
assert_eq!(first.as_ref(), b": aether-keepalive\n\n");
}
#[tokio::test]
async fn execute_stream_from_frame_stream_does_not_finalize_rewritten_tool_call_after_midstream_error(
) {
let usage_repository = Arc::new(InMemoryUsageReadRepository::default());
let request_candidate_repository = Arc::new(InMemoryRequestCandidateRepository::default());
let state = AppState::new()
.expect("app state should build")
.with_data_state_for_tests(
crate::data::GatewayDataState::with_request_candidate_and_usage_repository_for_tests(
Arc::clone(&request_candidate_repository),
Arc::clone(&usage_repository),
),
)
.with_usage_runtime_for_tests(UsageRuntimeConfig {
enabled: true,
..UsageRuntimeConfig::default()
});
let plan = ExecutionPlan {
request_id: "req-responses-tool-midstream-error".into(),
candidate_id: Some("cand-responses-tool-midstream-error".into()),
provider_name: Some("openai".into()),
provider_id: "provider-openai-responses".into(),
endpoint_id: "endpoint-openai-responses".into(),
key_id: "key-openai-responses".into(),
method: "POST".into(),
url: "https://api.openai.com/v1/responses".into(),
headers: BTreeMap::from([
("content-type".into(), "application/json".into()),
("accept".into(), "text/event-stream".into()),
]),
content_type: Some("application/json".into()),
content_encoding: None,
body: RequestBody::from_json(json!({
"model": "gpt-5.5",
"input": [],
"stream": true
})),
stream: true,
client_api_format: "claude:messages".into(),
provider_api_format: "openai:responses".into(),
model_name: Some("gpt-5.5".into()),
proxy: None,
transport_profile: None,
timeouts: None,
};
let upstream_chunk = concat!(
"event: response.created\n",
"data: {\"type\":\"response.created\",\"response\":{\"id\":\"resp_midstream_error\",\"model\":\"gpt-5.5\",\"status\":\"in_progress\"}}\n\n",
"event: response.output_item.added\n",
"data: {\"type\":\"response.output_item.added\",\"output_index\":0,\"item\":{\"type\":\"function_call\",\"id\":\"fc_1\",\"call_id\":\"call_1\",\"name\":\"lookup\",\"arguments\":\"\",\"status\":\"in_progress\"}}\n\n",
"event: response.function_call_arguments.delta\n",
"data: {\"type\":\"response.function_call_arguments.delta\",\"output_index\":0,\"item_id\":\"fc_1\",\"call_id\":\"call_1\",\"delta\":\"{\\\"query\\\":\\\"abc\"}\n\n"
);
let frame_stream = stream! {
yield Ok::<Bytes, std::io::Error>(ndjson_frame(StreamFrame {
frame_type: StreamFrameType::Headers,
payload: StreamFramePayload::Headers {
status_code: 200,
headers: BTreeMap::from([(
"content-type".to_string(),
"text/event-stream".to_string(),
)]),
},
}));
yield Ok::<Bytes, std::io::Error>(ndjson_frame(StreamFrame {
frame_type: StreamFrameType::Data,
payload: StreamFramePayload::Data {
chunk_b64: None,
text: Some(upstream_chunk.to_string()),
},
}));
yield Ok::<Bytes, std::io::Error>(ndjson_frame(StreamFrame {
frame_type: StreamFrameType::Error,
payload: StreamFramePayload::Error {
error: ExecutionError {
kind: ExecutionErrorKind::Internal,
phase: ExecutionPhase::StreamRead,
message: "error reading a body from connection: stream error received: unexpected internal error encountered".to_string(),
upstream_status: Some(200),
retryable: false,
failover_recommended: false,
},
},
}));
}
.boxed();
let response = execute_stream_from_frame_stream(
&state,
plan,
"trace-responses-tool-midstream-error",
&test_decision(),
"openai_responses_stream",
Some("openai_responses_stream_success".to_string()),
Some(json!({
"request_id": "req-responses-tool-midstream-error",
"candidate_id": "cand-responses-tool-midstream-error",
"candidate_index": 0,
"retry_index": 0,
"provider_api_format": "openai:responses",
"client_api_format": "claude:messages",
"needs_conversion": true,
})),
crate::clock::current_unix_ms(),
Instant::now(),
frame_stream,
None,
)
.await
.expect("execution should succeed")
.expect("execution should return a client response");
let body = to_bytes(response.into_body(), usize::MAX)
.await
.expect("response body should read");
let body_text = String::from_utf8(body.to_vec()).expect("body should be utf8");
assert!(body_text.contains("event: content_block_start"));
assert!(body_text.contains("event: content_block_delta"));
assert!(body_text.contains("\"type\":\"tool_use\""));
assert!(!body_text.contains("event: content_block_stop"));
assert!(!body_text.contains("event: message_delta"));
assert!(!body_text.contains("event: message_stop"));
assert!(!body_text.contains("\"stop_reason\":\"tool_use\""));
assert!(body_text.contains("\"error\""));
assert!(body_text.contains("unexpected internal error encountered"));
assert!(body_text.contains("data: [DONE]"));
let candidates = tokio::time::timeout(Duration::from_secs(1), async {
loop {
let candidates = request_candidate_repository
.list_by_request_id("req-responses-tool-midstream-error")
.await
.expect("request candidates should read");
if candidates
.first()
.is_some_and(|candidate| candidate.status == RequestCandidateStatus::Failed)
{
break candidates;
}
tokio::time::sleep(Duration::from_millis(10)).await;
}
})
.await
.expect("candidate should be marked failed");
assert_eq!(candidates[0].status_code, Some(200));
assert_eq!(candidates[0].error_type.as_deref(), Some("internal"));
}
#[tokio::test]
async fn openai_image_stream_ignores_plan_total_timeout() {
let state = AppState::new().expect("app state should build");
@@ -41,6 +41,12 @@ use crate::clock::current_unix_ms as current_request_candidate_unix_ms;
use crate::control::GatewayControlDecision;
use crate::execution_runtime::chatgpt_web_image::maybe_execute_chatgpt_web_image_sync;
use crate::execution_runtime::grok::maybe_execute_grok_sync;
use crate::execution_runtime::kiro_cache::{
build_kiro_prompt_cache_profile, compute_kiro_prompt_cache_usage,
estimate_kiro_prompt_input_tokens, kiro_simulated_cache_enabled_from_provider_config,
kiro_simulated_cache_enabled_from_report_context, KiroPromptCacheUsage,
KIRO_SIMULATED_CACHE_ENABLED_CONTEXT_FIELD,
};
use crate::execution_runtime::oauth_retry::refresh_oauth_plan_auth_for_retry;
#[cfg(test)]
use crate::execution_runtime::remote_compat::post_sync_plan_to_remote_execution_runtime;
@@ -355,6 +361,180 @@ fn build_sync_report_payload(
}
}
fn seed_kiro_sync_report_context_input_tokens(
plan: &ExecutionPlan,
report_context: &mut Option<Value>,
) {
if !plan
.provider_name
.as_deref()
.is_some_and(|provider_name| provider_name.eq_ignore_ascii_case("Kiro"))
{
return;
}
let Some(context) = report_context.as_mut().and_then(Value::as_object_mut) else {
return;
};
if context
.get("input_tokens")
.and_then(Value::as_u64)
.is_some_and(|input_tokens| input_tokens > 0)
{
return;
}
let Some(original_request_body) = context.get("original_request_body").cloned() else {
return;
};
let estimated_input_tokens = estimate_kiro_prompt_input_tokens(&original_request_body);
context.insert(
"input_tokens".to_string(),
Value::from(estimated_input_tokens),
);
}
async fn seed_kiro_sync_simulated_cache_enabled(
state: &AppState,
plan: &ExecutionPlan,
report_context: &mut Option<Value>,
) {
if !plan
.provider_name
.as_deref()
.is_some_and(|provider_name| provider_name.eq_ignore_ascii_case("Kiro"))
{
return;
}
let enabled = match state
.read_provider_catalog_providers_by_ids(std::slice::from_ref(&plan.provider_id))
.await
{
Ok(providers) => providers
.iter()
.find(|provider| provider.id == plan.provider_id)
.filter(|provider| provider.provider_type.eq_ignore_ascii_case("kiro"))
.is_some_and(|provider| {
kiro_simulated_cache_enabled_from_provider_config(provider.config.as_ref())
}),
Err(err) => {
warn!(
event_name = "kiro_simulated_cache_config_read_failed",
log_type = "event",
request_id = %plan.request_id,
provider_id = %plan.provider_id,
error = ?err,
"failed to read Kiro simulated cache provider config; defaulting disabled"
);
false
}
};
let Some(context) = report_context.as_mut().and_then(Value::as_object_mut) else {
return;
};
if enabled {
context.insert(
KIRO_SIMULATED_CACHE_ENABLED_CONTEXT_FIELD.to_string(),
Value::Bool(true),
);
} else {
context.remove(KIRO_SIMULATED_CACHE_ENABLED_CONTEXT_FIELD);
}
}
async fn seed_kiro_sync_report_context_prompt_cache_usage(
state: &AppState,
plan: &ExecutionPlan,
report_context: &mut Option<Value>,
) {
if !plan
.provider_name
.as_deref()
.is_some_and(|provider_name| provider_name.eq_ignore_ascii_case("Kiro"))
{
return;
}
let simulated_cache_enabled =
kiro_simulated_cache_enabled_from_report_context(report_context.as_ref());
let Some(context) = report_context.as_mut().and_then(Value::as_object_mut) else {
return;
};
if context
.get("kiro_web_search_mcp")
.and_then(Value::as_bool)
.unwrap_or(false)
{
return;
}
if !simulated_cache_enabled {
return;
}
if kiro_cache_usage_from_context_object(context).is_some() {
return;
}
let Some(original_request_body) = context.get("original_request_body").cloned() else {
return;
};
let input_tokens = context
.get("input_tokens")
.and_then(Value::as_u64)
.filter(|value| *value > 0)
.unwrap_or_else(|| {
let estimated = estimate_kiro_prompt_input_tokens(&original_request_body);
context.insert("input_tokens".to_string(), Value::from(estimated));
estimated
});
let Some(profile) = build_kiro_prompt_cache_profile(&original_request_body, input_tokens)
else {
return;
};
let cache_usage = compute_kiro_prompt_cache_usage(
state.runtime_state(),
kiro_sync_cache_credential_id(plan),
&profile,
)
.await;
if cache_usage.cache_creation_input_tokens == 0 && cache_usage.cache_read_input_tokens == 0 {
return;
}
context.insert(
"cache_creation_input_tokens".to_string(),
Value::from(cache_usage.cache_creation_input_tokens),
);
context.insert(
"cache_read_input_tokens".to_string(),
Value::from(cache_usage.cache_read_input_tokens),
);
}
fn kiro_sync_cache_credential_id(plan: &ExecutionPlan) -> String {
format!("{}:{}:{}", plan.provider_id, plan.endpoint_id, plan.key_id)
}
fn kiro_cache_usage_from_context_object(
context: &serde_json::Map<String, Value>,
) -> Option<KiroPromptCacheUsage> {
let cache_creation_input_tokens = context
.get("cache_creation_input_tokens")
.and_then(Value::as_u64)
.unwrap_or(0);
let cache_read_input_tokens = context
.get("cache_read_input_tokens")
.and_then(Value::as_u64)
.unwrap_or(0);
(cache_creation_input_tokens > 0 || cache_read_input_tokens > 0).then_some(
KiroPromptCacheUsage {
cache_creation_input_tokens,
cache_read_input_tokens,
},
)
}
fn invalid_gemini_provider_success_message(
plan: &ExecutionPlan,
report_context: Option<&Value>,
@@ -1182,12 +1362,18 @@ fn build_json_whitespace_heartbeat_stream(
}
}
pub(crate) fn build_openai_image_sync_json_whitespace_heartbeat_stream(
pub(crate) fn build_sync_json_whitespace_heartbeat_stream(
rx: mpsc::Receiver<Result<Bytes, IoError>>,
) -> impl futures_util::Stream<Item = Result<Bytes, IoError>> + Send + 'static {
build_json_whitespace_heartbeat_stream(rx, OPENAI_IMAGE_SYNC_JSON_HEARTBEAT_INTERVAL, None)
}
pub(crate) fn build_openai_image_sync_json_whitespace_heartbeat_stream(
rx: mpsc::Receiver<Result<Bytes, IoError>>,
) -> impl futures_util::Stream<Item = Result<Bytes, IoError>> + Send + 'static {
build_sync_json_whitespace_heartbeat_stream(rx)
}
async fn openai_image_sync_json_heartbeat_final_bytes(
result: Result<Option<Response<Body>>, GatewayError>,
) -> Vec<u8> {
@@ -1934,8 +2120,15 @@ async fn execute_execution_runtime_sync_impl(
}
let status_code = result.status_code;
let has_body_bytes = body_base64.is_some();
let report_context =
let mut report_context =
attach_provider_response_headers_to_report_context(report_context, &headers);
if (200..300).contains(&status_code) {
seed_kiro_sync_simulated_cache_enabled(state, &plan, &mut report_context).await;
if kiro_simulated_cache_enabled_from_report_context(report_context.as_ref()) {
seed_kiro_sync_report_context_input_tokens(&plan, &mut report_context);
}
seed_kiro_sync_report_context_prompt_cache_usage(state, &plan, &mut report_context).await;
}
let mut client_headers = headers.clone();
apply_endpoint_response_header_rules(state, &plan, &mut client_headers, body_json.as_ref())
.await?;
@@ -2504,6 +2697,45 @@ mod tests {
plan
}
fn test_kiro_sync_plan() -> ExecutionPlan {
ExecutionPlan {
request_id: "req-kiro-sync-cache-1".to_string(),
candidate_id: Some("candidate-kiro-sync-cache-1".to_string()),
provider_name: Some("Kiro".to_string()),
provider_id: "provider-kiro-sync-1".to_string(),
endpoint_id: "endpoint-kiro-sync-1".to_string(),
key_id: "key-kiro-sync-1".to_string(),
method: "POST".to_string(),
url: "https://kiro.example/generateAssistantResponse".to_string(),
headers: BTreeMap::new(),
content_type: Some("application/json".to_string()),
content_encoding: None,
body: aether_contracts::RequestBody::from_json(json!({
"model": "claude-sonnet-4",
"messages": [{"role": "user", "content": "hello kiro"}],
})),
stream: false,
client_api_format: "claude:messages".to_string(),
provider_api_format: "claude:messages".to_string(),
model_name: Some("claude-sonnet-4".to_string()),
proxy: None,
transport_profile: None,
timeouts: None,
}
}
fn test_kiro_sync_cacheable_request_body() -> serde_json::Value {
json!({
"model": "claude-sonnet-4",
"system": [{
"type": "text",
"text": format!("sync cacheable prompt {}", "cacheable prompt chunk ".repeat(300)),
"cache_control": {"type": "ephemeral"}
}],
"messages": [{"role": "user", "content": "reuse this Kiro prompt"}]
})
}
#[test]
fn invalid_gemini_provider_success_uses_plan_format_when_context_is_missing() {
let plan = test_gemini_chat_plan();
@@ -2659,6 +2891,68 @@ mod tests {
);
}
#[test]
fn kiro_sync_report_context_seeds_input_tokens_from_original_request_body() {
let plan = test_kiro_sync_plan();
let mut report_context = Some(json!({
"original_request_body": test_kiro_sync_cacheable_request_body(),
}));
seed_kiro_sync_report_context_input_tokens(&plan, &mut report_context);
assert!(report_context
.as_ref()
.and_then(|value| value.get("input_tokens"))
.and_then(Value::as_u64)
.is_some_and(|tokens| tokens > 0));
}
#[tokio::test]
async fn kiro_sync_report_context_applies_prompt_cache_usage_from_tracker() {
let state = AppState::new().expect("gateway state should build");
let plan = test_kiro_sync_plan();
let mut first_report_context = Some(json!({
"original_request_body": test_kiro_sync_cacheable_request_body(),
"kiro_simulated_cache_enabled": true,
}));
seed_kiro_sync_report_context_input_tokens(&plan, &mut first_report_context);
seed_kiro_sync_report_context_prompt_cache_usage(&state, &plan, &mut first_report_context)
.await;
let first_creation = first_report_context
.as_ref()
.and_then(|value| value.get("cache_creation_input_tokens"))
.and_then(Value::as_u64)
.unwrap_or_default();
let first_read = first_report_context
.as_ref()
.and_then(|value| value.get("cache_read_input_tokens"))
.and_then(Value::as_u64)
.unwrap_or_default();
assert!(first_creation > 0);
assert_eq!(first_read, 0);
let mut second_report_context = Some(json!({
"original_request_body": test_kiro_sync_cacheable_request_body(),
"kiro_simulated_cache_enabled": true,
}));
seed_kiro_sync_report_context_input_tokens(&plan, &mut second_report_context);
seed_kiro_sync_report_context_prompt_cache_usage(&state, &plan, &mut second_report_context)
.await;
let second_creation = second_report_context
.as_ref()
.and_then(|value| value.get("cache_creation_input_tokens"))
.and_then(Value::as_u64)
.unwrap_or_default();
let second_read = second_report_context
.as_ref()
.and_then(|value| value.get("cache_read_input_tokens"))
.and_then(Value::as_u64)
.unwrap_or_default();
assert_eq!(second_creation, 0);
assert!(second_read > 0);
}
#[tokio::test]
async fn json_whitespace_heartbeat_stream_prefixes_final_json() {
let (tx, rx) = mpsc::channel::<Result<Bytes, IoError>>(1);
@@ -1,7 +1,8 @@
mod execution;
pub(crate) use execution::{
build_openai_image_sync_json_whitespace_heartbeat_stream, execute_execution_runtime_sync,
build_openai_image_sync_json_whitespace_heartbeat_stream,
build_sync_json_whitespace_heartbeat_stream, execute_execution_runtime_sync,
};
#[allow(unused_imports)]
@@ -9,6 +9,7 @@ use serde_json::{json, Value};
use tokio::sync::mpsc;
use crate::ai_serving::api::{
build_core_error_body_for_client_format,
build_local_gemini_files_stream_attempt_source_for_kind,
build_local_gemini_files_sync_attempt_source_for_kind,
build_local_image_stream_attempt_source_for_kind,
@@ -29,17 +30,18 @@ use crate::ai_serving::api::{
resolve_gemini_sync_spec, resolve_local_same_format_stream_spec,
resolve_local_same_format_sync_spec, set_local_openai_chat_execution_exhausted_diagnostic,
set_local_openai_image_execution_exhausted_diagnostic, AiStreamAttempt, AiSyncAttempt,
LocalStandardSpec, EXECUTION_RUNTIME_STREAM_DECISION_ACTION,
LocalCoreSyncErrorKind, LocalStandardSpec, EXECUTION_RUNTIME_STREAM_DECISION_ACTION,
EXECUTION_RUNTIME_SYNC_DECISION_ACTION,
};
use crate::ai_serving::LocalExecutionAttemptSource;
use crate::api::response::{
attach_control_metadata_headers, build_client_response_from_parts_with_mutator,
};
use crate::constants::EXECUTION_PATH_LOCAL_EXECUTION_RUNTIME_MISS;
use crate::constants::{CONTROL_CANDIDATE_ID_HEADER, EXECUTION_PATH_LOCAL_EXECUTION_RUNTIME_MISS};
use crate::control::GatewayControlDecision;
use crate::execution_runtime::sync::{
build_openai_image_sync_json_whitespace_heartbeat_stream, execute_execution_runtime_sync,
build_openai_image_sync_json_whitespace_heartbeat_stream,
build_sync_json_whitespace_heartbeat_stream, execute_execution_runtime_sync,
};
use crate::executor::candidate_loop::{
execute_stream_attempt_source, execute_sync_attempt_source, execute_sync_plan_and_reports,
@@ -47,15 +49,19 @@ use crate::executor::candidate_loop::{
};
use crate::executor::{
build_local_execution_exhaustion, record_failed_usage_for_exhausted_request,
LocalExecutionRequestOutcome,
LocalExecutionExhaustion, LocalExecutionRequestOutcome,
};
use crate::handlers::shared::system_config_bool;
use crate::{AiExecutionDecision, AppState, GatewayError};
const ENABLE_OPENAI_IMAGE_SYNC_HEARTBEAT_CONFIG_KEY: &str = "enable_openai_image_sync_heartbeat";
const ENABLE_STANDARD_TEXT_SYNC_HEARTBEAT_CONFIG_KEY: &str = "enable_standard_text_sync_heartbeat";
const OPENAI_IMAGE_SYNC_HEARTBEAT_INTERNAL_ERROR_STATUS: u16 = 502;
const OPENAI_IMAGE_SYNC_HEARTBEAT_EXHAUSTED_STATUS: u16 = 503;
const OPENAI_IMAGE_SYNC_HEARTBEAT_ERROR_MESSAGE_LIMIT: usize = 4096;
const STANDARD_TEXT_SYNC_HEARTBEAT_INTERNAL_ERROR_STATUS: u16 = 502;
const STANDARD_TEXT_SYNC_HEARTBEAT_EXHAUSTED_STATUS: u16 = 503;
const STANDARD_TEXT_SYNC_HEARTBEAT_ERROR_MESSAGE_LIMIT: usize = 4096;
pub(crate) async fn maybe_execute_sync_local_path(
state: &AppState,
@@ -95,6 +101,65 @@ pub(crate) async fn maybe_execute_sync_via_local_decision(
return Ok(LocalExecutionRequestOutcome::NoPath);
};
if standard_text_sync_heartbeat_should_wrap(state, plan_kind).await {
let parts_for_task = parts.clone();
let body_json_for_task = body_json.clone();
return Ok(LocalExecutionRequestOutcome::responded(
build_standard_text_sync_heartbeat_shell_response(
state.clone(),
parts_for_task,
trace_id.to_string(),
decision.clone(),
plan_kind.to_string(),
move |state, parts, trace_id, decision, plan_kind, started_at| async move {
let Some((attempt_source, candidate_count)) =
build_local_openai_chat_sync_attempt_source_for_kind(
&state,
&parts,
trace_id.as_str(),
&decision,
&body_json_for_task,
plan_kind.as_str(),
)
.await?
else {
return Ok(LocalExecutionRequestOutcome::NoPath);
};
let outcome = execute_sync_attempt_source::<AiSyncAttempt, _>(
&state,
&parts,
trace_id.as_str(),
&decision,
plan_kind.as_str(),
attempt_source,
)
.await?;
match outcome {
LocalExecutionRequestOutcome::Exhausted(exhaustion) => {
set_local_openai_chat_execution_exhausted_diagnostic(
&state,
trace_id.as_str(),
&decision,
plan_kind.as_str(),
&body_json_for_task,
candidate_count,
);
record_standard_text_sync_heartbeat_exhaustion(
&state,
exhaustion,
&started_at,
)
.await;
Ok(LocalExecutionRequestOutcome::NoPath)
}
outcome => Ok(outcome),
}
},
)?,
));
}
let outcome = execute_sync_attempt_source::<AiSyncAttempt, _>(
state,
parts,
@@ -176,6 +241,57 @@ pub(crate) async fn maybe_execute_sync_via_local_openai_responses_decision(
return Ok(LocalExecutionRequestOutcome::NoPath);
};
if standard_text_sync_heartbeat_should_wrap(state, plan_kind).await {
let parts_for_task = parts.clone();
let body_json_for_task = body_json.clone();
return Ok(LocalExecutionRequestOutcome::responded(
build_standard_text_sync_heartbeat_shell_response(
state.clone(),
parts_for_task,
trace_id.to_string(),
decision.clone(),
plan_kind.to_string(),
move |state, parts, trace_id, decision, plan_kind, started_at| async move {
let Some((attempt_source, _candidate_count)) =
build_local_openai_responses_sync_attempt_source_for_kind(
&state,
&parts,
trace_id.as_str(),
&decision,
&body_json_for_task,
plan_kind.as_str(),
)
.await?
else {
return Ok(LocalExecutionRequestOutcome::NoPath);
};
let outcome = execute_sync_attempt_source::<AiSyncAttempt, _>(
&state,
&parts,
trace_id.as_str(),
&decision,
plan_kind.as_str(),
attempt_source,
)
.await?;
match outcome {
LocalExecutionRequestOutcome::Exhausted(exhaustion) => {
record_standard_text_sync_heartbeat_exhaustion(
&state,
exhaustion,
&started_at,
)
.await;
Ok(LocalExecutionRequestOutcome::NoPath)
}
outcome => Ok(outcome),
}
},
)?,
));
}
execute_sync_attempt_source::<AiSyncAttempt, _>(
state,
parts,
@@ -235,6 +351,57 @@ pub(crate) async fn maybe_execute_sync_via_standard_family_decision(
return Ok(LocalExecutionRequestOutcome::NoPath);
};
if standard_text_sync_heartbeat_should_wrap(state, plan_kind).await {
let parts_for_task = parts.clone();
let body_json_for_task = body_json.clone();
return Ok(LocalExecutionRequestOutcome::responded(
build_standard_text_sync_heartbeat_shell_response(
state.clone(),
parts_for_task,
trace_id.to_string(),
decision.clone(),
plan_kind.to_string(),
move |state, parts, trace_id, decision, plan_kind, started_at| async move {
let Some((attempt_source, _candidate_count)) =
build_standard_family_sync_attempt_source(
&state,
&parts,
trace_id.as_str(),
&decision,
&body_json_for_task,
spec,
)
.await?
else {
return Ok(LocalExecutionRequestOutcome::NoPath);
};
let outcome = execute_sync_attempt_source::<AiSyncAttempt, _>(
&state,
&parts,
trace_id.as_str(),
&decision,
plan_kind.as_str(),
attempt_source,
)
.await?;
match outcome {
LocalExecutionRequestOutcome::Exhausted(exhaustion) => {
record_standard_text_sync_heartbeat_exhaustion(
&state,
exhaustion,
&started_at,
)
.await;
Ok(LocalExecutionRequestOutcome::NoPath)
}
outcome => Ok(outcome),
}
},
)?,
));
}
execute_sync_attempt_source::<AiSyncAttempt, _>(
state,
parts,
@@ -399,6 +566,57 @@ pub(crate) async fn maybe_execute_sync_via_local_same_format_provider_decision(
return Ok(LocalExecutionRequestOutcome::NoPath);
};
if standard_text_sync_heartbeat_should_wrap(state, plan_kind).await {
let parts_for_task = parts.clone();
let body_json_for_task = body_json.clone();
return Ok(LocalExecutionRequestOutcome::responded(
build_standard_text_sync_heartbeat_shell_response(
state.clone(),
parts_for_task,
trace_id.to_string(),
decision.clone(),
plan_kind.to_string(),
move |state, parts, trace_id, decision, plan_kind, started_at| async move {
let Some((attempt_source, _candidate_count)) =
build_local_same_format_sync_attempt_source(
&state,
&parts,
trace_id.as_str(),
&decision,
&body_json_for_task,
spec,
)
.await?
else {
return Ok(LocalExecutionRequestOutcome::NoPath);
};
let outcome = execute_sync_attempt_source::<AiSyncAttempt, _>(
&state,
&parts,
trace_id.as_str(),
&decision,
plan_kind.as_str(),
attempt_source,
)
.await?;
match outcome {
LocalExecutionRequestOutcome::Exhausted(exhaustion) => {
record_standard_text_sync_heartbeat_exhaustion(
&state,
exhaustion,
&started_at,
)
.await;
Ok(LocalExecutionRequestOutcome::NoPath)
}
outcome => Ok(outcome),
}
},
)?,
));
}
execute_sync_attempt_source::<AiSyncAttempt, _>(
state,
parts,
@@ -495,6 +713,341 @@ async fn openai_image_sync_heartbeat_enabled(state: &AppState) -> bool {
}
}
async fn standard_text_sync_heartbeat_enabled(state: &AppState) -> bool {
match state
.read_system_config_json_value(ENABLE_STANDARD_TEXT_SYNC_HEARTBEAT_CONFIG_KEY)
.await
{
Ok(value) => system_config_bool(value.as_ref(), false),
Err(err) => {
tracing::warn!(
event_name = "standard_text_sync_heartbeat_config_read_failed",
log_type = "ops",
error = ?err,
"gateway failed to read standard text sync heartbeat config; defaulting disabled"
);
false
}
}
}
fn standard_text_sync_heartbeat_applies_to_plan_kind(plan_kind: &str) -> bool {
matches!(
plan_kind,
"openai_chat_sync"
| "openai_responses_sync"
| "openai_responses_compact_sync"
| "claude_chat_sync"
| "claude_cli_sync"
| "gemini_chat_sync"
| "gemini_cli_sync"
)
}
async fn standard_text_sync_heartbeat_should_wrap(state: &AppState, plan_kind: &str) -> bool {
standard_text_sync_heartbeat_applies_to_plan_kind(plan_kind)
&& standard_text_sync_heartbeat_enabled(state).await
}
fn standard_text_sync_heartbeat_client_api_format_for_plan_kind(plan_kind: &str) -> &'static str {
match plan_kind {
"openai_responses_sync" => "openai:responses",
"openai_responses_compact_sync" => "openai:responses:compact",
"claude_chat_sync" | "claude_cli_sync" => "claude:messages",
"gemini_chat_sync" | "gemini_cli_sync" => "gemini:generate_content",
_ => "openai:chat",
}
}
fn build_standard_text_sync_heartbeat_shell_response<F, Fut>(
state: AppState,
parts: http::request::Parts,
trace_id: String,
decision: GatewayControlDecision,
plan_kind: String,
execute: F,
) -> Result<Response<Body>, GatewayError>
where
F: FnOnce(
AppState,
http::request::Parts,
String,
GatewayControlDecision,
String,
Instant,
) -> Fut
+ Send
+ 'static,
Fut: std::future::Future<Output = Result<LocalExecutionRequestOutcome, GatewayError>>
+ Send
+ 'static,
{
let request_id = (!trace_id.trim().is_empty()).then(|| trace_id.clone());
let client_api_format =
standard_text_sync_heartbeat_client_api_format_for_plan_kind(plan_kind.as_str())
.to_string();
let redaction_slot = parts
.extensions
.get::<crate::privacy::RedactionSessionSlot>()
.cloned();
let trace_id_for_response = trace_id.clone();
let decision_for_response = decision.clone();
let started_at = Instant::now();
let (tx, rx) = mpsc::channel::<Result<Bytes, IoError>>(1);
tokio::spawn(async move {
let bytes = standard_text_sync_heartbeat_final_bytes(
client_api_format.as_str(),
redaction_slot.as_ref(),
execute(state, parts, trace_id, decision, plan_kind, started_at).await,
)
.await;
let _ = tx.send(Ok(Bytes::from(bytes))).await;
});
let headers = BTreeMap::from([(
CONTENT_TYPE.as_str().to_string(),
"application/json".to_string(),
)]);
let response = build_client_response_from_parts_with_mutator(
StatusCode::OK.as_u16(),
&headers,
Body::from_stream(build_sync_json_whitespace_heartbeat_stream(rx)),
trace_id_for_response.as_str(),
Some(&decision_for_response),
|headers| {
headers.remove(CONTENT_LENGTH);
headers.remove(CONTENT_ENCODING);
headers.insert(
CACHE_CONTROL,
HeaderValue::from_static("no-cache, no-transform"),
);
headers.insert(
HeaderName::from_static("x-accel-buffering"),
HeaderValue::from_static("no"),
);
Ok(())
},
)?;
attach_control_metadata_headers(response, request_id.as_deref(), None)
}
async fn record_standard_text_sync_heartbeat_exhaustion(
state: &AppState,
exhaustion: LocalExecutionExhaustion,
started_at: &Instant,
) {
record_failed_usage_for_exhausted_request(
state,
exhaustion,
started_at,
"Standard text sync heartbeat exhausted all local candidates",
EXECUTION_PATH_LOCAL_EXECUTION_RUNTIME_MISS,
None,
)
.await;
}
async fn standard_text_sync_heartbeat_final_bytes(
client_api_format: &str,
redaction_slot: Option<&crate::privacy::RedactionSessionSlot>,
result: Result<LocalExecutionRequestOutcome, GatewayError>,
) -> Vec<u8> {
match result {
Ok(LocalExecutionRequestOutcome::Responded(response)) => {
standard_text_sync_heartbeat_response_body_bytes(
client_api_format,
redaction_slot,
response,
)
.await
}
Ok(LocalExecutionRequestOutcome::Exhausted(_))
| Ok(LocalExecutionRequestOutcome::NoPath) => standard_text_sync_heartbeat_error_body(
client_api_format,
STANDARD_TEXT_SYNC_HEARTBEAT_EXHAUSTED_STATUS,
"standard text sync exhausted all local candidates",
),
Err(err) => standard_text_sync_heartbeat_error_body(
client_api_format,
STANDARD_TEXT_SYNC_HEARTBEAT_INTERNAL_ERROR_STATUS,
&format!("{err:?}"),
),
}
}
async fn standard_text_sync_heartbeat_response_body_bytes(
client_api_format: &str,
redaction_slot: Option<&crate::privacy::RedactionSessionSlot>,
response: Response<Body>,
) -> Vec<u8> {
let status_code = response.status().as_u16();
let (parts, body) = response.into_parts();
match to_bytes(body, usize::MAX).await {
Ok(bytes) => {
let body = match standard_text_sync_heartbeat_restore_response_body(
redaction_slot,
&parts.headers,
bytes.as_ref(),
) {
Ok(body) => body,
Err(err) => {
return standard_text_sync_heartbeat_error_body(
client_api_format,
STANDARD_TEXT_SYNC_HEARTBEAT_INTERNAL_ERROR_STATUS,
&format!("{err:?}"),
);
}
};
if (200..300).contains(&status_code) && !body.is_empty() {
return body;
}
if !(200..300).contains(&status_code) {
return standard_text_sync_heartbeat_error_body_from_response(
client_api_format,
status_code,
body.as_ref(),
);
}
standard_text_sync_heartbeat_error_body(
client_api_format,
STANDARD_TEXT_SYNC_HEARTBEAT_INTERNAL_ERROR_STATUS,
"empty standard text sync response",
)
}
Err(err) => standard_text_sync_heartbeat_error_body(
client_api_format,
STANDARD_TEXT_SYNC_HEARTBEAT_INTERNAL_ERROR_STATUS,
&err.to_string(),
),
}
}
fn standard_text_sync_heartbeat_restore_response_body(
redaction_slot: Option<&crate::privacy::RedactionSessionSlot>,
headers: &http::HeaderMap,
body: &[u8],
) -> Result<Vec<u8>, GatewayError> {
let Some(redaction_slot) = redaction_slot else {
return Ok(body.to_vec());
};
let candidate_id = headers
.get(CONTROL_CANDIDATE_ID_HEADER)
.and_then(|value| value.to_str().ok())
.map(str::trim)
.filter(|value| !value.is_empty());
let Some(session) = redaction_slot.take_for_candidate(candidate_id) else {
return Ok(body.to_vec());
};
let mut header_values = headers
.iter()
.map(|(name, value)| {
(
name.as_str().to_string(),
value.to_str().unwrap_or_default().to_string(),
)
})
.collect::<BTreeMap<_, _>>();
crate::privacy::restore_sync_response_body(&mut header_values, body, &session)
.map(|restored| restored.body)
}
fn standard_text_sync_heartbeat_error_body_from_response(
client_api_format: &str,
status_code: u16,
body: &[u8],
) -> Vec<u8> {
if let Ok(mut value) = serde_json::from_slice::<Value>(body) {
if standard_text_sync_heartbeat_insert_upstream_status(&mut value, status_code) {
return serde_json::to_vec(&value).unwrap_or_else(|_| {
standard_text_sync_heartbeat_error_body(
client_api_format,
status_code,
&format!("upstream returned status {status_code}"),
)
});
}
}
let message = standard_text_sync_heartbeat_error_message_from_body(status_code, body);
standard_text_sync_heartbeat_error_body(client_api_format, status_code, message.as_str())
}
fn standard_text_sync_heartbeat_insert_upstream_status(
value: &mut Value,
status_code: u16,
) -> bool {
let Some(error) = value.get_mut("error").and_then(Value::as_object_mut) else {
return false;
};
error.insert("upstream_status".to_string(), Value::from(status_code));
error
.entry("message".to_string())
.or_insert_with(|| Value::String(format!("upstream returned status {status_code}")));
true
}
fn standard_text_sync_heartbeat_error_message_from_body(status_code: u16, body: &[u8]) -> String {
let text = String::from_utf8_lossy(body).trim().to_string();
if text.is_empty() {
return format!("upstream returned status {status_code}");
}
text.chars()
.take(STANDARD_TEXT_SYNC_HEARTBEAT_ERROR_MESSAGE_LIMIT)
.collect()
}
fn standard_text_sync_heartbeat_error_body(
client_api_format: &str,
status_code: u16,
message: &str,
) -> Vec<u8> {
let mut body = build_core_error_body_for_client_format(
client_api_format,
message,
Some("upstream_error"),
standard_text_sync_heartbeat_error_kind(status_code),
)
.unwrap_or_else(|| {
json!({
"error": {
"type": "upstream_error",
"message": message,
"code": status_code,
}
})
});
if !standard_text_sync_heartbeat_insert_upstream_status(&mut body, status_code) {
body = json!({
"error": {
"type": "upstream_error",
"message": message,
"code": status_code,
"upstream_status": status_code,
}
});
}
serde_json::to_vec(&body).unwrap_or_else(|_| {
format!(
"{{\"error\":{{\"type\":\"upstream_error\",\"code\":{status_code},\"upstream_status\":{status_code}}}}}"
)
.into_bytes()
})
}
fn standard_text_sync_heartbeat_error_kind(status_code: u16) -> LocalCoreSyncErrorKind {
match status_code {
400 => LocalCoreSyncErrorKind::InvalidRequest,
401 => LocalCoreSyncErrorKind::Authentication,
403 => LocalCoreSyncErrorKind::PermissionDenied,
404 => LocalCoreSyncErrorKind::NotFound,
413 => LocalCoreSyncErrorKind::ContextLengthExceeded,
429 => LocalCoreSyncErrorKind::RateLimit,
503 => LocalCoreSyncErrorKind::Overloaded,
_ => LocalCoreSyncErrorKind::ServerError,
}
}
fn build_openai_image_sync_heartbeat_shell_response(
state: AppState,
request_path: String,
@@ -944,10 +1497,35 @@ pub(crate) fn decision_payload_is_direct_execution(payload: &AiExecutionDecision
#[cfg(test)]
mod tests {
use super::*;
use futures_util::StreamExt;
use std::sync::atomic::{AtomicUsize, Ordering};
use std::sync::Arc;
const TEST_OPENAI_IMAGE_SYNC_PLAN_KIND: &str = "openai_image_sync";
const TEST_STANDARD_TEXT_SYNC_PLAN_KIND: &str = "openai_responses_compact_sync";
struct TestSyncAttemptSource {
attempts: VecDeque<AiSyncAttempt>,
}
impl TestSyncAttemptSource {
fn new(attempts: Vec<AiSyncAttempt>) -> Self {
Self {
attempts: VecDeque::from(attempts),
}
}
}
#[async_trait::async_trait]
impl LocalExecutionAttemptSource<AiSyncAttempt> for TestSyncAttemptSource {
async fn next_execution_attempt(&mut self) -> Result<Option<AiSyncAttempt>, GatewayError> {
Ok(self.attempts.pop_front())
}
async fn drain_execution_attempts(&mut self) -> Result<Vec<AiSyncAttempt>, GatewayError> {
Ok(self.attempts.drain(..).collect())
}
}
fn test_openai_image_heartbeat_decision() -> GatewayControlDecision {
GatewayControlDecision::synthetic(
@@ -1028,6 +1606,63 @@ mod tests {
}
}
fn test_standard_text_heartbeat_decision() -> GatewayControlDecision {
GatewayControlDecision::synthetic(
"/v1/responses",
Some("ai_public".to_string()),
Some("openai".to_string()),
Some("responses".to_string()),
Some("openai:responses:compact".to_string()),
)
.with_execution_runtime_candidate(true)
}
fn test_standard_text_heartbeat_plan(
endpoint_id: &str,
candidate_id: &str,
client_api_format: &str,
) -> aether_contracts::ExecutionPlan {
aether_contracts::ExecutionPlan {
request_id: "trace-standard-text-heartbeat-retry".to_string(),
candidate_id: Some(candidate_id.to_string()),
provider_name: Some("OpenAI".to_string()),
provider_id: "provider-openai".to_string(),
endpoint_id: endpoint_id.to_string(),
key_id: "key-openai".to_string(),
method: "POST".to_string(),
url: "https://api.openai.com/v1/responses".to_string(),
headers: BTreeMap::new(),
content_type: Some("application/json".to_string()),
content_encoding: None,
body: aether_contracts::RequestBody::from_json(json!({"model": "gpt-5"})),
stream: false,
client_api_format: client_api_format.to_string(),
provider_api_format: client_api_format.to_string(),
model_name: Some("gpt-5".to_string()),
proxy: None,
transport_profile: None,
timeouts: None,
}
}
fn test_standard_text_heartbeat_attempt(
candidate_index: u32,
endpoint_id: &str,
candidate_id: &str,
client_api_format: &str,
) -> AiSyncAttempt {
AiSyncAttempt {
plan: test_standard_text_heartbeat_plan(endpoint_id, candidate_id, client_api_format),
report_kind: None,
report_context: Some(json!({
"candidate_index": candidate_index,
"retry_index": 0,
"client_api_format": client_api_format,
"provider_api_format": client_api_format,
})),
}
}
#[tokio::test]
async fn openai_image_sync_heartbeat_success_body_is_unchanged() {
let response = Response::builder()
@@ -1133,4 +1768,239 @@ mod tests {
assert_eq!(call_count.load(Ordering::SeqCst), 2);
assert_eq!(body, json!({"data": [{"b64_json": "second-candidate"}]}));
}
#[tokio::test]
async fn standard_text_sync_heartbeat_missing_config_defaults_disabled() {
let state = AppState::new().expect("state should build");
assert!(!standard_text_sync_heartbeat_enabled(&state).await);
}
#[tokio::test]
async fn standard_text_sync_heartbeat_no_local_candidates_preserves_no_path() {
let state = AppState::new()
.expect("state should build")
.with_data_state_for_tests(
crate::data::GatewayDataState::disabled().with_system_config_values_for_tests([(
ENABLE_STANDARD_TEXT_SYNC_HEARTBEAT_CONFIG_KEY.to_string(),
json!(true),
)]),
);
let (parts, _) = http::Request::builder()
.method(http::Method::POST)
.uri("/v1/responses")
.body(())
.expect("request should build")
.into_parts();
let outcome = maybe_execute_sync_via_local_openai_responses_decision(
&state,
&parts,
"trace-standard-text-heartbeat-no-path",
&test_standard_text_heartbeat_decision(),
&json!({"model": "missing-local-candidate"}),
TEST_STANDARD_TEXT_SYNC_PLAN_KIND,
)
.await
.expect("heartbeat no-path check should execute");
assert!(matches!(outcome, LocalExecutionRequestOutcome::NoPath));
}
#[tokio::test]
async fn standard_text_sync_heartbeat_success_body_is_unchanged() {
let response = Response::builder()
.status(StatusCode::OK)
.body(Body::from(r#"{"id":"resp_123","output":[]}"#))
.expect("response should build");
let bytes = standard_text_sync_heartbeat_response_body_bytes(
"openai:responses:compact",
None,
response,
)
.await;
let body: Value = serde_json::from_slice(&bytes).expect("body should decode");
assert_eq!(body, json!({"id": "resp_123", "output": []}));
}
#[tokio::test]
async fn standard_text_sync_heartbeat_claude_error_body_includes_upstream_status() {
let response = Response::builder()
.status(StatusCode::TOO_MANY_REQUESTS)
.body(Body::from(
r#"{"type":"error","error":{"type":"rate_limit_error","message":"slow down"}}"#,
))
.expect("response should build");
let bytes =
standard_text_sync_heartbeat_response_body_bytes("claude:messages", None, response)
.await;
let body: Value = serde_json::from_slice(&bytes).expect("body should decode");
assert_eq!(body["type"], json!("error"));
assert_eq!(body["error"]["type"], json!("rate_limit_error"));
assert_eq!(body["error"]["message"], json!("slow down"));
assert_eq!(body["error"]["upstream_status"], json!(429));
}
#[test]
fn standard_text_sync_heartbeat_applies_to_chat_and_cli_plan_kinds() {
assert!(standard_text_sync_heartbeat_applies_to_plan_kind(
"claude_chat_sync"
));
assert!(standard_text_sync_heartbeat_applies_to_plan_kind(
"claude_cli_sync"
));
assert!(standard_text_sync_heartbeat_applies_to_plan_kind(
"gemini_chat_sync"
));
assert!(standard_text_sync_heartbeat_applies_to_plan_kind(
"gemini_cli_sync"
));
assert!(!standard_text_sync_heartbeat_applies_to_plan_kind(
"openai_embedding_sync"
));
}
#[tokio::test]
async fn standard_text_sync_heartbeat_redirect_status_is_wrapped_as_error() {
let response = Response::builder()
.status(StatusCode::TEMPORARY_REDIRECT)
.body(Body::from(r#"{"location":"https://upstream.example"}"#))
.expect("response should build");
let bytes = standard_text_sync_heartbeat_response_body_bytes(
"openai:responses:compact",
None,
response,
)
.await;
let body: Value = serde_json::from_slice(&bytes).expect("body should decode");
assert_eq!(body["error"]["type"], json!("server_error"));
assert_eq!(body["error"]["upstream_status"], json!(307));
}
#[tokio::test]
async fn standard_text_sync_heartbeat_shell_sends_whitespace_before_background_finishes() {
let state = AppState::new().expect("state should build");
let (parts, _) = http::Request::builder()
.method(http::Method::POST)
.uri("/v1/responses")
.body(())
.expect("request should build")
.into_parts();
let (release_tx, release_rx) = tokio::sync::oneshot::channel::<()>();
let response = build_standard_text_sync_heartbeat_shell_response(
state,
parts,
"trace-standard-text-heartbeat-shell".to_string(),
test_standard_text_heartbeat_decision(),
TEST_STANDARD_TEXT_SYNC_PLAN_KIND.to_string(),
move |_state, _parts, _trace_id, _decision, _plan_kind, _started_at| async move {
let _ = release_rx.await;
Ok(LocalExecutionRequestOutcome::responded(
Response::builder()
.status(StatusCode::OK)
.body(Body::from(r#"{"id":"resp_done","output":[]}"#))
.expect("response should build"),
))
},
)
.expect("heartbeat shell should build");
let mut body_stream = response.into_body().into_data_stream();
let first = body_stream
.next()
.await
.expect("heartbeat stream should yield")
.expect("heartbeat chunk should be ok");
assert_eq!(first.as_ref(), b"\n");
let _ = release_tx.send(());
}
#[test]
fn standard_text_sync_heartbeat_compact_non_json_error_body_is_wrapped_in_client_format() {
let bytes = standard_text_sync_heartbeat_error_body_from_response(
"openai:responses:compact",
502,
b"bad gateway from upstream",
);
let body: Value = serde_json::from_slice(&bytes).expect("body should decode");
assert_eq!(body["error"]["type"], json!("server_error"));
assert_eq!(body["error"]["message"], json!("bad gateway from upstream"));
assert_eq!(body["error"]["upstream_status"], json!(502));
}
#[tokio::test]
async fn standard_text_sync_heartbeat_attempts_retry_first_candidate_then_return_second() {
let call_count = Arc::new(AtomicUsize::new(0));
let call_count_for_override = Arc::clone(&call_count);
let state = AppState::new()
.expect("state should build")
.with_execution_runtime_sync_override_for_tests(move |plan| {
call_count_for_override.fetch_add(1, Ordering::SeqCst);
if plan.endpoint_id == "endpoint-retry" {
Ok(test_openai_image_execution_result(
plan,
StatusCode::TOO_MANY_REQUESTS.as_u16(),
json!({"error": {"message": "retry this candidate"}}),
))
} else {
Ok(test_openai_image_execution_result(
plan,
StatusCode::OK.as_u16(),
json!({"id": "resp_second_candidate", "output": []}),
))
}
});
let attempts = vec![
test_standard_text_heartbeat_attempt(
0,
"endpoint-retry",
"candidate-retry",
"openai:responses:compact",
),
test_standard_text_heartbeat_attempt(
1,
"endpoint-success",
"candidate-success",
"openai:responses:compact",
),
];
let (parts, _) = http::Request::builder()
.method(http::Method::POST)
.uri("/v1/responses")
.body(())
.expect("request should build")
.into_parts();
let outcome = execute_sync_attempt_source::<AiSyncAttempt, _>(
&state,
&parts,
"trace-standard-text-heartbeat-retry",
&test_standard_text_heartbeat_decision(),
TEST_STANDARD_TEXT_SYNC_PLAN_KIND,
TestSyncAttemptSource::new(attempts),
)
.await
.expect("heartbeat attempts should execute");
let LocalExecutionRequestOutcome::Responded(response) = outcome else {
panic!("second candidate should return a response");
};
let bytes = standard_text_sync_heartbeat_response_body_bytes(
"openai:responses:compact",
None,
response,
)
.await;
let body: Value = serde_json::from_slice(&bytes).expect("body should decode");
assert_eq!(call_count.load(Ordering::SeqCst), 2);
assert_eq!(body, json!({"id": "resp_second_candidate", "output": []}));
}
}
@@ -13,6 +13,7 @@ const EMBEDDING_API_FORMATS: &[&str] = &[
"jina:embedding",
"gemini:embedding",
"doubao:embedding",
"aliyun:multimodal_embedding",
];
fn json_value_contains_string(value: &serde_json::Value, expected: &str) -> bool {
@@ -6,6 +6,7 @@ const EMBEDDING_API_FORMATS: &[&str] = &[
"jina:embedding",
"gemini:embedding",
"doubao:embedding",
"aliyun:multimodal_embedding",
];
pub(crate) fn model_tiered_pricing_first_tier_value(
@@ -94,7 +94,7 @@ fn split_admin_monitoring_api_format_and_model(
fn is_known_admin_monitoring_api_format_family(value: &str) -> bool {
matches!(
value.trim().to_ascii_lowercase().as_str(),
"openai" | "claude" | "gemini" | "jina" | "doubao"
"openai" | "claude" | "gemini" | "jina" | "doubao" | "aliyun"
)
}
@@ -116,6 +116,7 @@ fn is_known_admin_monitoring_api_format(value: &str) -> bool {
| "jina:embedding"
| "jina:rerank"
| "doubao:embedding"
| "aliyun:multimodal_embedding"
)
}
@@ -597,6 +597,22 @@ fn provider_query_build_test_request_body_for_api_format(
let message = provider_query_extract_message(payload)
.unwrap_or_else(|| DEFAULT_PROVIDER_QUERY_TEST_MESSAGE.to_string());
match client_api_format.as_str() {
"openai:embedding" => json!({
"model": model,
"input": message,
}),
"openai:rerank" => json!({
"model": model,
"query": message,
"documents": [
"apple",
"banana",
"fruit",
"vegetable"
],
"return_documents": true,
"top_n": 4,
}),
"openai:responses" | "openai:responses:compact" => json!({
"model": model,
"input": message,
@@ -649,6 +665,21 @@ fn provider_query_insert_default_test_conversation(
let message = provider_query_extract_message(payload)
.unwrap_or_else(|| DEFAULT_PROVIDER_QUERY_TEST_MESSAGE.to_string());
match client_api_format {
"openai:embedding" => {
object.insert("input".to_string(), Value::String(message));
}
"openai:rerank" => {
object.insert("query".to_string(), Value::String(message));
object
.entry("documents".to_string())
.or_insert_with(|| json!(["apple", "banana", "fruit", "vegetable"]));
object
.entry("return_documents".to_string())
.or_insert(Value::Bool(true));
object
.entry("top_n".to_string())
.or_insert_with(|| Value::from(4_u64));
}
"openai:responses" | "openai:responses:compact" => {
object.insert("input".to_string(), Value::String(message));
}
@@ -2920,8 +2951,13 @@ async fn provider_query_execute_standard_test_candidate(
);
provider_request_body
}
"openai:embedding" | "gemini:embedding" | "jina:embedding" | "doubao:embedding"
| "openai:rerank" | "jina:rerank" => {
"openai:embedding"
| "gemini:embedding"
| "jina:embedding"
| "doubao:embedding"
| "aliyun:multimodal_embedding"
| "openai:rerank"
| "jina:rerank" => {
let Some(mut provider_request_body) =
crate::ai_serving::build_standard_request_body_with_model_directives_and_request_headers(
&request_body,
@@ -3051,6 +3087,7 @@ async fn provider_query_execute_standard_test_candidate(
| "gemini:embedding"
| "jina:embedding"
| "doubao:embedding"
| "aliyun:multimodal_embedding"
| "openai:rerank"
| "jina:rerank" => state.resolve_local_oauth_header_auth(&transport).await?,
_ => None,
@@ -3062,6 +3099,7 @@ async fn provider_query_execute_standard_test_candidate(
| "openai:embedding"
| "jina:embedding"
| "doubao:embedding"
| "aliyun:multimodal_embedding"
| "openai:rerank"
| "jina:rerank" => {
crate::provider_transport::auth::resolve_local_openai_bearer_auth(&transport)
@@ -74,6 +74,7 @@ pub(super) fn provider_query_standard_test_unsupported_reason(
| "openai:embedding"
| "jina:embedding"
| "doubao:embedding"
| "aliyun:multimodal_embedding"
| "openai:rerank"
| "jina:rerank" => {
crate::provider_transport::policy::local_standard_transport_unsupported_reason_with_network(
@@ -253,6 +254,7 @@ pub(super) fn provider_query_test_adapter_for_provider_api_format(
| "gemini:embedding"
| "jina:embedding"
| "doubao:embedding"
| "aliyun:multimodal_embedding"
| "openai:rerank"
| "jina:rerank"
) {
@@ -352,6 +354,7 @@ pub(super) fn provider_query_transport_supports_model_test_execution(
| "openai:embedding"
| "jina:embedding"
| "doubao:embedding"
| "aliyun:multimodal_embedding"
| "openai:rerank"
| "jina:rerank" => {
crate::provider_transport::policy::supports_local_standard_transport_with_network(
@@ -138,6 +138,12 @@ fn provider_query_endpoint_route_payload(
"embeddings",
"openai_batch",
),
"aliyun:multimodal_embedding" => (
"Aliyun DashScope",
"dashscope_native",
"multimodal-embedding",
"dashscope_contents",
),
"openai:chat" if is_vertex && is_openai_compat => (
"Vertex AI OpenAI-compatible",
"openai_compatible",
@@ -470,6 +470,26 @@ fn provider_query_compact_test_request_body_defaults_to_responses_input() {
assert!(body.get("messages").is_none());
}
#[test]
fn provider_query_embedding_test_request_body_defaults_to_embedding_input() {
let payload = json!({"message": "hello from embedding"});
let client_api_format =
provider_query_standard_test_client_api_format("aliyun:multimodal_embedding");
let body = provider_query_build_test_request_body_for_api_format(
&payload,
"qwen3-vl-embedding",
"/api/admin/provider-query/test-model",
client_api_format,
);
assert_eq!(client_api_format, "openai:embedding");
assert_eq!(body["model"], json!("qwen3-vl-embedding"));
assert_eq!(body["input"], json!("hello from embedding"));
assert!(body.get("messages").is_none());
assert!(body.get("stream").is_none());
}
#[test]
fn provider_query_compact_test_request_body_promotes_prompt_to_input() {
let payload = json!({
@@ -658,6 +678,13 @@ fn provider_query_test_adapter_routes_fixed_provider_endpoint_types() {
provider_query_test_adapter_for_provider_api_format("custom", "gemini:embedding"),
Some(ProviderQueryTestAdapter::Standard)
);
assert_eq!(
provider_query_test_adapter_for_provider_api_format(
"aliyun",
"aliyun:multimodal_embedding"
),
Some(ProviderQueryTestAdapter::Standard)
);
assert_eq!(
provider_query_test_adapter_for_provider_api_format("jina", "jina:rerank"),
Some(ProviderQueryTestAdapter::Standard)
@@ -278,6 +278,9 @@ pub(crate) fn normalize_admin_user_api_formats(
if item.is_empty() {
return Err("allowed_api_formats 不能为空".to_string());
}
if !looks_like_admin_api_format_signature(item) {
return Err(format!("allowed_api_formats 格式无效: {item}"));
}
let Some(normalized_item) = crate::api::ai::normalize_admin_endpoint_signature(item) else {
return Err(format!("allowed_api_formats 格式无效: {item}"));
};
@@ -289,6 +292,12 @@ pub(crate) fn normalize_admin_user_api_formats(
Ok(Some(normalized))
}
fn looks_like_admin_api_format_signature(value: &str) -> bool {
value
.split_once(':')
.is_some_and(|(family, kind)| !family.trim().is_empty() && !kind.trim().is_empty())
}
pub(crate) fn normalize_admin_user_ip_rules(
value: Option<Vec<String>>,
) -> Result<Option<Vec<String>>, String> {
@@ -522,12 +522,45 @@ fn embedding_array_input_is_non_empty(items: &[Value]) -> bool {
item.as_array()
.is_some_and(|items| embedding_token_array_is_non_empty(items))
})
|| items.iter().all(embedding_multimodal_content_is_non_empty)
}
fn embedding_token_array_is_non_empty(items: &[Value]) -> bool {
!items.is_empty() && items.iter().all(|item| item.as_u64().is_some())
}
fn embedding_multimodal_content_is_non_empty(value: &Value) -> bool {
let Some(object) = value.as_object() else {
return false;
};
let valid_text = object
.get("text")
.map(|value| value.as_str().is_some_and(|text| !text.trim().is_empty()));
let valid_image = object
.get("image")
.map(|value| value.as_str().is_some_and(|image| !image.trim().is_empty()));
let valid_video = object
.get("video")
.map(|value| value.as_str().is_some_and(|video| !video.trim().is_empty()));
let valid_multi_images = object.get("multi_images").map(|value| {
value.as_array().is_some_and(|items| {
!items.is_empty()
&& items
.iter()
.all(|item| item.as_str().is_some_and(|image| !image.trim().is_empty()))
})
});
[valid_text, valid_image, valid_video, valid_multi_images]
.into_iter()
.flatten()
.all(|valid| valid)
&& [valid_text, valid_image, valid_video, valid_multi_images]
.into_iter()
.flatten()
.any(|valid| valid)
}
fn image_request_count(value: &Value) -> Option<u64> {
value
.as_u64()
@@ -327,6 +327,32 @@ pub(super) async fn handle_billing_plan_checkout(
false,
);
}
match state
.find_pending_plan_purchase_order_by_user_id(&auth.user.id, &plan.id)
.await
{
Ok(Some(order)) => {
return build_auth_json_response(
http::StatusCode::OK,
json!({
"order": payment_order_payload(&order, &plan),
"payment_instructions": sanitize_wallet_gateway_response(
order.gateway_response.clone()
),
"reused_pending_order": true,
}),
None,
)
}
Ok(None) => {}
Err(err) => {
return build_auth_error_response(
http::StatusCode::INTERNAL_SERVER_ERROR,
format!("pending billing checkout lookup failed: {err:?}"),
false,
)
}
}
let now = Utc::now();
let order_no = billing_order_no(now);
let expires_at = now + chrono::Duration::minutes(30);
@@ -25,6 +25,7 @@ pub(crate) fn models_api_format(request_context: &GatewayPublicRequestContext) -
"jina:embedding" => Some("jina:embedding"),
"jina:rerank" => Some("jina:rerank"),
"doubao:embedding" => Some("doubao:embedding"),
"aliyun:multimodal_embedding" => Some("aliyun:multimodal_embedding"),
_ => None,
}
}
@@ -43,6 +44,7 @@ const MODELS_EMBEDDING_QUERY_API_FORMATS: &[&str] = &[
"jina:embedding",
"gemini:embedding",
"doubao:embedding",
"aliyun:multimodal_embedding",
];
const MODELS_RERANK_QUERY_API_FORMATS: &[&str] = &["openai:rerank", "jina:rerank"];
@@ -54,9 +56,11 @@ pub(super) fn models_query_api_formats(api_format: &str) -> &'static [&'static s
| "claude:messages"
| "gemini:generate_content" => MODELS_CROSS_FORMAT_QUERY_API_FORMATS,
"openai:image" => &["openai:image"],
"openai:embedding" | "jina:embedding" | "gemini:embedding" | "doubao:embedding" => {
MODELS_EMBEDDING_QUERY_API_FORMATS
}
"openai:embedding"
| "jina:embedding"
| "gemini:embedding"
| "doubao:embedding"
| "aliyun:multimodal_embedding" => MODELS_EMBEDDING_QUERY_API_FORMATS,
"openai:rerank" | "jina:rerank" => MODELS_RERANK_QUERY_API_FORMATS,
_ => &[],
}
@@ -138,6 +138,18 @@ impl AppState {
.map_err(|err| GatewayError::Internal(err.to_string()))
}
pub(crate) async fn find_pending_plan_purchase_order_by_user_id(
&self,
user_id: &str,
product_id: &str,
) -> Result<Option<aether_data::repository::wallet::StoredAdminPaymentOrder>, GatewayError>
{
self.data
.find_pending_plan_purchase_order_by_user_id(user_id, product_id)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))
}
pub(crate) async fn find_wallet_refund(
&self,
wallet_id: &str,
@@ -24,8 +24,39 @@ use aether_data_contracts::repository::provider_catalog::{
};
use sha2::{Digest, Sha256};
#[tokio::test]
async fn gateway_executes_openai_chat_sync_upstream_stream_via_local_finalize_response() {
const OPENAI_CHAT_FINALIZE_TEST_STACK_BYTES: usize = 16 * 1024 * 1024;
fn run_openai_chat_finalize_test<F, Fut>(test_name: &'static str, make_future: F)
where
F: FnOnce() -> Fut + Send + 'static,
Fut: std::future::Future<Output = ()> + 'static,
{
let handle = std::thread::Builder::new()
.name(test_name.to_string())
.stack_size(OPENAI_CHAT_FINALIZE_TEST_STACK_BYTES)
.spawn(move || {
let runtime = tokio::runtime::Builder::new_current_thread()
.enable_all()
.build()
.expect("test runtime should build");
runtime.block_on(make_future());
})
.expect("openai chat finalize test thread should spawn");
if let Err(payload) = handle.join() {
std::panic::resume_unwind(payload);
}
}
#[test]
fn gateway_executes_openai_chat_sync_upstream_stream_via_local_finalize_response() {
run_openai_chat_finalize_test(
"gateway_executes_openai_chat_sync_upstream_stream_via_local_finalize_response",
gateway_executes_openai_chat_sync_upstream_stream_via_local_finalize_response_impl,
);
}
async fn gateway_executes_openai_chat_sync_upstream_stream_via_local_finalize_response_impl() {
use base64::Engine as _;
#[derive(Debug, Clone)]
@@ -521,8 +552,16 @@ async fn gateway_executes_openai_chat_sync_upstream_stream_via_local_finalize_re
upstream_handle.abort();
}
#[tokio::test]
async fn gateway_executes_openai_chat_cross_format_upstream_stream_via_local_finalize_response() {
#[test]
fn gateway_executes_openai_chat_cross_format_upstream_stream_via_local_finalize_response() {
run_openai_chat_finalize_test(
"gateway_executes_openai_chat_cross_format_upstream_stream_via_local_finalize_response",
gateway_executes_openai_chat_cross_format_upstream_stream_via_local_finalize_response_impl,
);
}
async fn gateway_executes_openai_chat_cross_format_upstream_stream_via_local_finalize_response_impl(
) {
use base64::Engine as _;
#[derive(Debug, Clone)]
struct SeenRemoteExecutionRuntimeRequest {
@@ -955,8 +994,16 @@ async fn gateway_executes_openai_chat_cross_format_upstream_stream_via_local_fin
upstream_handle.abort();
}
#[tokio::test]
async fn gateway_executes_openai_chat_cross_format_tool_use_upstream_stream_via_local_finalize_response(
#[test]
fn gateway_executes_openai_chat_cross_format_tool_use_upstream_stream_via_local_finalize_response()
{
run_openai_chat_finalize_test(
"gateway_executes_openai_chat_cross_format_tool_use_upstream_stream_via_local_finalize_response",
gateway_executes_openai_chat_cross_format_tool_use_upstream_stream_via_local_finalize_response_impl,
);
}
async fn gateway_executes_openai_chat_cross_format_tool_use_upstream_stream_via_local_finalize_response_impl(
) {
use base64::Engine as _;
@@ -1782,8 +1829,15 @@ async fn gateway_skips_openai_chat_antigravity_cross_format_sync_candidate_as_tr
upstream_handle.abort();
}
#[tokio::test]
async fn gateway_executes_openai_chat_cross_format_claude_upstream_sync_via_local_finalize_response(
#[test]
fn gateway_executes_openai_chat_cross_format_claude_upstream_sync_via_local_finalize_response() {
run_openai_chat_finalize_test(
"gateway_executes_openai_chat_cross_format_claude_upstream_sync_via_local_finalize_response",
gateway_executes_openai_chat_cross_format_claude_upstream_sync_via_local_finalize_response_impl,
);
}
async fn gateway_executes_openai_chat_cross_format_claude_upstream_sync_via_local_finalize_response_impl(
) {
fn hash_api_key(value: &str) -> String {
let mut hasher = Sha256::new();
@@ -2134,8 +2188,15 @@ async fn gateway_executes_openai_chat_cross_format_claude_upstream_sync_via_loca
upstream_handle.abort();
}
#[tokio::test]
async fn gateway_executes_openai_chat_cross_format_gemini_upstream_sync_via_local_finalize_response(
#[test]
fn gateway_executes_openai_chat_cross_format_gemini_upstream_sync_via_local_finalize_response() {
run_openai_chat_finalize_test(
"gateway_executes_openai_chat_cross_format_gemini_upstream_sync_via_local_finalize_response",
gateway_executes_openai_chat_cross_format_gemini_upstream_sync_via_local_finalize_response_impl,
);
}
async fn gateway_executes_openai_chat_cross_format_gemini_upstream_sync_via_local_finalize_response_impl(
) {
fn hash_api_key(value: &str) -> String {
let mut hasher = Sha256::new();
@@ -8,6 +8,9 @@ use super::{
StoredProviderCatalogKey, StoredProviderCatalogProvider, StoredProviderModelMapping,
DEVELOPMENT_ENCRYPTION_KEY, TRACE_ID_HEADER,
};
use aether_data::repository::usage::InMemoryUsageReadRepository;
use aether_data_contracts::repository::usage::{StoredRequestUsageAudit, UsageReadRepository};
use aether_usage_runtime::UsageRuntimeConfig;
const KIRO_CLAUDE_CLI_SYNC_TEST_STACK_BYTES: usize = 16 * 1024 * 1024;
@@ -33,6 +36,30 @@ where
}
}
async fn wait_for_completed_usage<T>(repository: &T, request_id: &str) -> StoredRequestUsageAudit
where
T: UsageReadRepository + ?Sized,
{
let timeout = std::time::Duration::from_secs(60);
let deadline = tokio::time::Instant::now() + timeout;
loop {
if let Some(usage) = repository
.find_by_request_id(request_id)
.await
.expect("usage should read")
{
if usage.status == "completed" {
return usage;
}
}
assert!(
tokio::time::Instant::now() < deadline,
"usage {request_id} should complete within {timeout:?}"
);
tokio::time::sleep(std::time::Duration::from_millis(10)).await;
}
}
#[test]
fn gateway_executes_kiro_claude_cli_sync_via_local_provider_catalog_candidate() {
run_kiro_claude_cli_sync_test(
@@ -194,7 +221,7 @@ async fn gateway_executes_kiro_claude_cli_sync_via_local_provider_catalog_candid
Some(serde_json::json!({"url":"http://provider-proxy.internal:8080"})),
Some(20.0),
None,
None,
Some(serde_json::json!({"kiro": {"simulated_cache_enabled": true}})),
)
}
@@ -354,14 +381,15 @@ async fn gateway_executes_kiro_claude_cli_sync_via_local_provider_catalog_candid
let raw_body = to_bytes(body, usize::MAX).await.expect("body should read");
let payload: serde_json::Value =
serde_json::from_slice(&raw_body).expect("execution runtime payload should parse");
let trace_id = parts
.headers
.get(TRACE_ID_HEADER)
.and_then(|value| value.to_str().ok())
.unwrap_or_default()
.to_string();
*seen_execution_runtime_inner.lock().expect("mutex should lock") =
Some(SeenExecutionRuntimeSyncRequest {
trace_id: parts
.headers
.get(TRACE_ID_HEADER)
.and_then(|value| value.to_str().ok())
.unwrap_or_default()
.to_string(),
trace_id: trace_id.clone(),
url: payload
.get("url")
.and_then(|value| value.as_str())
@@ -453,7 +481,7 @@ async fn gateway_executes_kiro_claude_cli_sync_via_local_provider_catalog_candid
.concat();
Json(json!({
"request_id": "trace-kiro-cli-local-sync-123",
"request_id": trace_id,
"status_code": 200,
"headers": {
"content-type": "application/vnd.amazon.eventstream"
@@ -481,6 +509,7 @@ async fn gateway_executes_kiro_claude_cli_sync_via_local_provider_catalog_candid
sample_candidate_row(),
]));
let request_candidate_repository = Arc::new(InMemoryRequestCandidateRepository::default());
let usage_repository = Arc::new(InMemoryUsageReadRepository::default());
let provider_catalog_repository = Arc::new(InMemoryProviderCatalogReadRepository::seed(
vec![sample_provider_catalog_provider()],
vec![sample_provider_catalog_endpoint()],
@@ -491,34 +520,51 @@ async fn gateway_executes_kiro_claude_cli_sync_via_local_provider_catalog_candid
let (execution_runtime_url, execution_runtime_handle) = start_server(execution_runtime).await;
let gateway_state = build_state_with_execution_runtime_override(execution_runtime_url.clone())
.with_data_state_for_tests(
crate::data::GatewayDataState::with_auth_candidate_selection_provider_catalog_and_request_candidate_repository_for_tests(
crate::data::GatewayDataState::with_auth_candidate_selection_provider_catalog_request_candidates_and_usage_for_tests(
auth_repository,
candidate_selection_repository,
provider_catalog_repository,
Arc::clone(&request_candidate_repository),
Arc::clone(&usage_repository),
DEVELOPMENT_ENCRYPTION_KEY,
),
);
)
.with_usage_runtime_for_tests(UsageRuntimeConfig {
enabled: true,
..UsageRuntimeConfig::default()
});
let gateway = build_router_with_state(gateway_state);
let (gateway_url, gateway_handle) = start_server(gateway).await;
let response = reqwest::Client::new()
.post(format!("{gateway_url}/v1/messages"))
.header(http::header::CONTENT_TYPE, "application/json")
.header(
http::header::AUTHORIZATION,
"Bearer sk-client-kiro-cli-local-sync",
)
.header(TRACE_ID_HEADER, "trace-kiro-cli-local-sync-123")
.body(
"{\"model\":\"claude-sonnet-4\",\"messages\":[{\"role\":\"user\",\"content\":\"hello kiro\"}],\"thinking\":{\"type\":\"enabled\",\"budget_tokens\":64}}",
)
.send()
.await
.expect("request should succeed");
async fn send_kiro_request(
gateway_url: &str,
trace_id: &str,
body: String,
) -> (StatusCode, String) {
let response = reqwest::Client::new()
.post(format!("{gateway_url}/v1/messages"))
.header(http::header::CONTENT_TYPE, "application/json")
.header(
http::header::AUTHORIZATION,
"Bearer sk-client-kiro-cli-local-sync",
)
.header(TRACE_ID_HEADER, trace_id)
.body(body)
.send()
.await
.expect("request should succeed");
let status = response.status();
let response_body = response.text().await.expect("body should read");
let status = response.status();
let response_body = response.text().await.expect("body should read");
(status, response_body)
}
let (status, response_body) = send_kiro_request(
&gateway_url,
"trace-kiro-cli-local-sync-123",
"{\"model\":\"claude-sonnet-4\",\"messages\":[{\"role\":\"user\",\"content\":\"hello kiro\"}],\"thinking\":{\"type\":\"enabled\",\"budget_tokens\":64}}".to_string(),
)
.await;
assert!(
status == StatusCode::OK,
"unexpected status={status} body={response_body} decision_hits={} plan_hits={} public_hits={}",
@@ -598,6 +644,58 @@ async fn gateway_executes_kiro_claude_cli_sync_via_local_provider_catalog_candid
"report-sync should stay local when request candidate persistence is available"
);
let cacheable_request_body = serde_json::json!({
"model": "claude-sonnet-4",
"system": [{
"type": "text",
"text": format!("sync cacheable prompt {}", "cacheable prompt chunk ".repeat(300)),
"cache_control": {"type": "ephemeral"}
}],
"messages": [{"role": "user", "content": "reuse this Kiro prompt"}]
})
.to_string();
let (first_cache_status, first_cache_body) = send_kiro_request(
&gateway_url,
"trace-kiro-cli-local-sync-cache-1",
cacheable_request_body.clone(),
)
.await;
assert!(
first_cache_status == StatusCode::OK,
"unexpected first cache status={first_cache_status} body={first_cache_body}"
);
let first_usage = wait_for_completed_usage(
usage_repository.as_ref(),
"trace-kiro-cli-local-sync-cache-1",
)
.await;
assert!(
first_usage.cache_creation_input_tokens > 0,
"first Kiro sync cacheable request should create simulated cache"
);
assert_eq!(first_usage.cache_read_input_tokens, 0);
let (second_cache_status, second_cache_body) = send_kiro_request(
&gateway_url,
"trace-kiro-cli-local-sync-cache-2",
cacheable_request_body,
)
.await;
assert!(
second_cache_status == StatusCode::OK,
"unexpected second cache status={second_cache_status} body={second_cache_body}"
);
let second_usage = wait_for_completed_usage(
usage_repository.as_ref(),
"trace-kiro-cli-local-sync-cache-2",
)
.await;
assert!(
second_usage.cache_read_input_tokens > 0,
"second Kiro sync cacheable request should read simulated cache"
);
assert_eq!(second_usage.cache_creation_input_tokens, 0);
assert_eq!(*decision_hits.lock().expect("mutex should lock"), 0);
assert_eq!(*plan_hits.lock().expect("mutex should lock"), 0);
assert_eq!(*public_hits.lock().expect("mutex should lock"), 0);
@@ -26,8 +26,40 @@ use aether_data_contracts::repository::provider_catalog::{
use base64::Engine as _;
use sha2::{Digest, Sha256};
#[tokio::test]
async fn gateway_executes_openai_responses_sync_via_local_decision_gate_with_local_sync_decision() {
const CLI_SYNC_TEST_STACK_BYTES: usize = 16 * 1024 * 1024;
fn run_cli_sync_test<F, Fut>(test_name: &'static str, make_future: F)
where
F: FnOnce() -> Fut + Send + 'static,
Fut: std::future::Future<Output = ()> + 'static,
{
let handle = std::thread::Builder::new()
.name(test_name.to_string())
.stack_size(CLI_SYNC_TEST_STACK_BYTES)
.spawn(move || {
let runtime = tokio::runtime::Builder::new_current_thread()
.enable_all()
.build()
.expect("test runtime should build");
runtime.block_on(make_future());
})
.expect("cli sync test thread should spawn");
if let Err(payload) = handle.join() {
std::panic::resume_unwind(payload);
}
}
#[test]
fn gateway_executes_openai_responses_sync_via_local_decision_gate_with_local_sync_decision() {
run_cli_sync_test(
"gateway_executes_openai_responses_sync_via_local_decision_gate_with_local_sync_decision",
gateway_executes_openai_responses_sync_via_local_decision_gate_with_local_sync_decision_impl,
);
}
async fn gateway_executes_openai_responses_sync_via_local_decision_gate_with_local_sync_decision_impl(
) {
#[derive(Debug, Clone)]
struct SeenExecutionRuntimeSyncRequest {
trace_id: String,
@@ -546,8 +578,15 @@ async fn gateway_executes_openai_responses_sync_via_local_decision_gate_with_loc
upstream_handle.abort();
}
#[tokio::test]
async fn gateway_waits_for_api_key_concurrency_slot_then_executes_openai_responses_sync() {
#[test]
fn gateway_waits_for_api_key_concurrency_slot_then_executes_openai_responses_sync() {
run_cli_sync_test(
"gateway_waits_for_api_key_concurrency_slot_then_executes_openai_responses_sync",
gateway_waits_for_api_key_concurrency_slot_then_executes_openai_responses_sync_impl,
);
}
async fn gateway_waits_for_api_key_concurrency_slot_then_executes_openai_responses_sync_impl() {
fn hash_api_key(value: &str) -> String {
let mut hasher = Sha256::new();
hasher.update(value.as_bytes());
@@ -929,8 +968,16 @@ async fn gateway_waits_for_api_key_concurrency_slot_then_executes_openai_respons
upstream_handle.abort();
}
#[tokio::test]
async fn gateway_returns_concurrency_limited_after_wait_budget_expires_for_openai_responses_sync() {
#[test]
fn gateway_returns_concurrency_limited_after_wait_budget_expires_for_openai_responses_sync() {
run_cli_sync_test(
"gateway_returns_concurrency_limited_after_wait_budget_expires_for_openai_responses_sync",
gateway_returns_concurrency_limited_after_wait_budget_expires_for_openai_responses_sync_impl,
);
}
async fn gateway_returns_concurrency_limited_after_wait_budget_expires_for_openai_responses_sync_impl(
) {
fn hash_api_key(value: &str) -> String {
let mut hasher = Sha256::new();
hasher.update(value.as_bytes());
@@ -1278,8 +1325,15 @@ async fn gateway_returns_concurrency_limited_after_wait_budget_expires_for_opena
upstream_handle.abort();
}
#[tokio::test]
async fn gateway_returns_openai_responses_error_for_local_sync_failure() {
#[test]
fn gateway_returns_openai_responses_error_for_local_sync_failure() {
run_cli_sync_test(
"gateway_returns_openai_responses_error_for_local_sync_failure",
gateway_returns_openai_responses_error_for_local_sync_failure_impl,
);
}
async fn gateway_returns_openai_responses_error_for_local_sync_failure_impl() {
fn hash_api_key(value: &str) -> String {
let mut hasher = Sha256::new();
hasher.update(value.as_bytes());
@@ -1571,8 +1625,16 @@ async fn gateway_returns_openai_responses_error_for_local_sync_failure() {
upstream_handle.abort();
}
#[tokio::test]
async fn gateway_returns_openai_responses_error_for_local_cross_format_gemini_cli_sync_failure() {
#[test]
fn gateway_returns_openai_responses_error_for_local_cross_format_gemini_cli_sync_failure() {
run_cli_sync_test(
"gateway_returns_openai_responses_error_for_local_cross_format_gemini_cli_sync_failure",
gateway_returns_openai_responses_error_for_local_cross_format_gemini_cli_sync_failure_impl,
);
}
async fn gateway_returns_openai_responses_error_for_local_cross_format_gemini_cli_sync_failure_impl(
) {
#[derive(Debug, Clone)]
struct SeenExecutionRuntimeSyncRequest {
trace_id: String,
@@ -1972,8 +2034,15 @@ async fn gateway_returns_openai_responses_error_for_local_cross_format_gemini_cl
upstream_handle.abort();
}
#[tokio::test]
async fn gateway_returns_openai_responses_error_for_local_cross_format_claude_sync_failure() {
#[test]
fn gateway_returns_openai_responses_error_for_local_cross_format_claude_sync_failure() {
run_cli_sync_test(
"gateway_returns_openai_responses_error_for_local_cross_format_claude_sync_failure",
gateway_returns_openai_responses_error_for_local_cross_format_claude_sync_failure_impl,
);
}
async fn gateway_returns_openai_responses_error_for_local_cross_format_claude_sync_failure_impl() {
#[derive(Debug, Clone)]
struct SeenExecutionRuntimeSyncRequest {
trace_id: String,
@@ -2349,8 +2418,16 @@ async fn gateway_returns_openai_responses_error_for_local_cross_format_claude_sy
upstream_handle.abort();
}
#[tokio::test]
async fn gateway_returns_openai_responses_error_for_local_cross_format_claude_chat_sync_failure() {
#[test]
fn gateway_returns_openai_responses_error_for_local_cross_format_claude_chat_sync_failure() {
run_cli_sync_test(
"gateway_returns_openai_responses_error_for_local_cross_format_claude_chat_sync_failure",
gateway_returns_openai_responses_error_for_local_cross_format_claude_chat_sync_failure_impl,
);
}
async fn gateway_returns_openai_responses_error_for_local_cross_format_claude_chat_sync_failure_impl(
) {
#[derive(Debug, Clone)]
struct SeenExecutionRuntimeSyncRequest {
trace_id: String,
@@ -2729,8 +2806,16 @@ async fn gateway_returns_openai_responses_error_for_local_cross_format_claude_ch
upstream_handle.abort();
}
#[tokio::test]
async fn gateway_returns_openai_responses_error_for_local_cross_format_gemini_chat_sync_failure() {
#[test]
fn gateway_returns_openai_responses_error_for_local_cross_format_gemini_chat_sync_failure() {
run_cli_sync_test(
"gateway_returns_openai_responses_error_for_local_cross_format_gemini_chat_sync_failure",
gateway_returns_openai_responses_error_for_local_cross_format_gemini_chat_sync_failure_impl,
);
}
async fn gateway_returns_openai_responses_error_for_local_cross_format_gemini_chat_sync_failure_impl(
) {
#[derive(Debug, Clone)]
struct SeenExecutionRuntimeSyncRequest {
trace_id: String,
@@ -3109,8 +3194,15 @@ async fn gateway_returns_openai_responses_error_for_local_cross_format_gemini_ch
upstream_handle.abort();
}
#[tokio::test]
async fn gateway_executes_codex_cli_sync_via_local_decision_gate_after_oauth_refresh() {
#[test]
fn gateway_executes_codex_cli_sync_via_local_decision_gate_after_oauth_refresh() {
run_cli_sync_test(
"gateway_executes_codex_cli_sync_via_local_decision_gate_after_oauth_refresh",
gateway_executes_codex_cli_sync_via_local_decision_gate_after_oauth_refresh_impl,
);
}
async fn gateway_executes_codex_cli_sync_via_local_decision_gate_after_oauth_refresh_impl() {
#[derive(Debug, Clone)]
struct SeenExecutionRuntimeSyncRequest {
trace_id: String,
@@ -226,19 +226,24 @@ impl StandardFormat {
}
}
#[tokio::test]
async fn ai_execute_openai_responses_pii_redaction_round_trip_same_format() {
let (response_json, seen) = run_redaction_format_case(RedactionFormatCase {
test_id: "openai-responses-pii-redaction-same-format",
trace_id: "trace-openai-responses-pii-redaction-same-format",
client_format: StandardFormat::OpenAiResponses,
provider_format: StandardFormat::OpenAiResponses,
})
.await;
#[test]
fn ai_execute_openai_responses_pii_redaction_round_trip_same_format() {
run_async_test_on_large_stack(
"ai_execute_openai_responses_pii_redaction_round_trip_same_format",
async {
let (response_json, seen) = run_redaction_format_case(RedactionFormatCase {
test_id: "openai-responses-pii-redaction-same-format",
trace_id: "trace-openai-responses-pii-redaction-same-format",
client_format: StandardFormat::OpenAiResponses,
provider_format: StandardFormat::OpenAiResponses,
})
.await;
assert_provider_request_redacted(&seen, StandardFormat::OpenAiResponses);
assert!(seen.body.get("input").is_some());
assert_restored_response(&response_json, StandardFormat::OpenAiResponses);
assert_provider_request_redacted(&seen, StandardFormat::OpenAiResponses);
assert!(seen.body.get("input").is_some());
assert_restored_response(&response_json, StandardFormat::OpenAiResponses);
},
);
}
#[test]
@@ -261,42 +266,52 @@ fn ai_execute_claude_messages_pii_redaction_round_trip_same_format() {
);
}
#[tokio::test]
async fn ai_execute_openai_chat_pii_redaction_before_claude_conversion() {
let (response_json, seen) = run_redaction_format_case(RedactionFormatCase {
test_id: "openai-chat-pii-redaction-before-claude-conversion",
trace_id: "trace-openai-chat-pii-redaction-before-claude-conversion",
client_format: StandardFormat::OpenAiChat,
provider_format: StandardFormat::ClaudeMessages,
})
.await;
#[test]
fn ai_execute_openai_chat_pii_redaction_before_claude_conversion() {
run_async_test_on_large_stack(
"ai_execute_openai_chat_pii_redaction_before_claude_conversion",
async {
let (response_json, seen) = run_redaction_format_case(RedactionFormatCase {
test_id: "openai-chat-pii-redaction-before-claude-conversion",
trace_id: "trace-openai-chat-pii-redaction-before-claude-conversion",
client_format: StandardFormat::OpenAiChat,
provider_format: StandardFormat::ClaudeMessages,
})
.await;
assert_provider_request_redacted(&seen, StandardFormat::ClaudeMessages);
assert!(seen.body.get("messages").is_some());
assert_eq!(
seen.body["model"],
StandardFormat::ClaudeMessages.provider_model()
assert_provider_request_redacted(&seen, StandardFormat::ClaudeMessages);
assert!(seen.body.get("messages").is_some());
assert_eq!(
seen.body["model"],
StandardFormat::ClaudeMessages.provider_model()
);
assert_restored_response(&response_json, StandardFormat::OpenAiChat);
},
);
assert_restored_response(&response_json, StandardFormat::OpenAiChat);
}
#[tokio::test]
async fn ai_execute_openai_responses_pii_redaction_before_claude_conversion() {
let (response_json, seen) = run_redaction_format_case(RedactionFormatCase {
test_id: "openai-responses-pii-redaction-before-claude-conversion",
trace_id: "trace-openai-responses-pii-redaction-before-claude-conversion",
client_format: StandardFormat::OpenAiResponses,
provider_format: StandardFormat::ClaudeMessages,
})
.await;
#[test]
fn ai_execute_openai_responses_pii_redaction_before_claude_conversion() {
run_async_test_on_large_stack(
"ai_execute_openai_responses_pii_redaction_before_claude_conversion",
async {
let (response_json, seen) = run_redaction_format_case(RedactionFormatCase {
test_id: "openai-responses-pii-redaction-before-claude-conversion",
trace_id: "trace-openai-responses-pii-redaction-before-claude-conversion",
client_format: StandardFormat::OpenAiResponses,
provider_format: StandardFormat::ClaudeMessages,
})
.await;
assert_provider_request_redacted(&seen, StandardFormat::ClaudeMessages);
assert!(seen.body.get("messages").is_some());
assert_eq!(
seen.body["model"],
StandardFormat::ClaudeMessages.provider_model()
assert_provider_request_redacted(&seen, StandardFormat::ClaudeMessages);
assert!(seen.body.get("messages").is_some());
assert_eq!(
seen.body["model"],
StandardFormat::ClaudeMessages.provider_model()
);
assert_restored_response(&response_json, StandardFormat::OpenAiResponses);
},
);
assert_restored_response(&response_json, StandardFormat::OpenAiResponses);
}
#[test]
@@ -28,7 +28,7 @@ fn admin_system_build_version_contract_uses_explicit_local_build_arg() {
let deploy = read_workspace_file("deploy.sh");
for pattern in [
"detect_build_version()",
"git describe --tags --always --dirty",
"git describe --tags --match 'v[0-9]*' --always --dirty",
"AETHER_BUILD_VERSION=\"${AETHER_BUILD_VERSION:-$(detect_build_version)}\"",
"--build-arg \"AETHER_BUILD_VERSION=$AETHER_BUILD_VERSION\"",
">>> AETHER_BUILD_VERSION",
@@ -43,6 +43,8 @@ fn admin_system_build_version_contract_uses_explicit_local_build_arg() {
for pattern in [
"process.env.AETHER_BUILD_VERSION",
"process.env.AETHER_VERSION",
"git describe --tags --match \"v[0-9]*\" --always --dirty",
"trimmed.startsWith('tunnel-v')",
] {
assert!(
vite_config.contains(pattern),
@@ -60,6 +62,18 @@ fn admin_system_build_version_contract_uses_explicit_local_build_arg() {
"api/core.rs should expose build version pattern {pattern}"
);
}
let build_rs = read_workspace_file("apps/aether-gateway/build.rs");
for pattern in [
"\"--match\"",
"\"v[0-9]*\"",
"trimmed.starts_with(\"tunnel-v\")",
] {
assert!(
build_rs.contains(pattern),
"apps/aether-gateway/build.rs should ignore tunnel release tags for gateway version pattern {pattern}"
);
}
}
#[test]
@@ -1452,6 +1452,14 @@ async fn gateway_handles_admin_system_api_formats_locally_with_trusted_admin_pri
assert!(formats.iter().any(|item| item["value"] == "jina:embedding"));
assert!(formats.iter().any(|item| item["value"] == "jina:rerank"));
assert!(formats.iter().any(|item| item["value"] == "gemini:video"));
let aliyun_embedding = formats
.iter()
.find(|item| item["value"] == "aliyun:multimodal_embedding")
.expect("aliyun multimodal embedding format should exist");
assert_eq!(
aliyun_embedding["default_path"],
"/api/v1/services/embeddings/multimodal-embedding/multimodal-embedding"
);
assert_eq!(*upstream_hits.lock().expect("mutex should lock"), 0);
gateway_handle.abort();
@@ -179,6 +179,129 @@ fn vertex_gemini_embedding_success_state(execution_runtime_url: String) -> AppSt
.with_data_state_for_tests(data_state)
}
fn aliyun_embedding_success_state(execution_runtime_url: String) -> AppState {
let mut snapshot = sample_currently_usable_auth_snapshot(
"key-aliyun-embedding-success",
"user-aliyun-embedding-success",
);
snapshot.user_allowed_providers = None;
snapshot.api_key_allowed_providers = None;
snapshot.user_allowed_api_formats = Some(vec!["openai:embedding".to_string()]);
snapshot.api_key_allowed_api_formats = Some(vec!["openai:embedding".to_string()]);
snapshot.user_allowed_models = Some(vec!["qwen3-vl-embedding".to_string()]);
snapshot.api_key_allowed_models = Some(vec!["qwen3-vl-embedding".to_string()]);
let auth_repository = Arc::new(InMemoryAuthApiKeySnapshotRepository::seed(vec![(
Some(hash_api_key("sk-aliyun-embedding-success")),
snapshot,
)]));
let candidate_repository =
Arc::new(InMemoryMinimalCandidateSelectionReadRepository::seed(vec![
aliyun_embedding_candidate_row(),
]));
let mut provider = sample_provider("provider-aliyun-embedding", "Aliyun DashScope", 1);
provider.provider_type = "aliyun".to_string();
let provider_catalog_repository = Arc::new(InMemoryProviderCatalogReadRepository::seed(
vec![provider],
vec![sample_endpoint(
"endpoint-aliyun-embedding",
"provider-aliyun-embedding",
"aliyun:multimodal_embedding",
"https://dashscope.aliyuncs.com",
)],
vec![sample_key(
"key-upstream-aliyun-embedding",
"provider-aliyun-embedding",
"aliyun:multimodal_embedding",
"sk-upstream-aliyun-embedding",
)],
));
let data_state =
GatewayDataState::with_provider_catalog_and_minimal_candidate_selection_for_tests(
provider_catalog_repository,
candidate_repository,
)
.with_auth_api_key_reader(auth_repository)
.with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY);
build_state_with_execution_runtime_override(execution_runtime_url)
.with_data_state_for_tests(data_state)
}
fn mixed_embedding_success_state(execution_runtime_url: String) -> AppState {
let mut snapshot = sample_currently_usable_auth_snapshot(
"key-mixed-embedding-success",
"user-mixed-embedding-success",
);
snapshot.user_allowed_providers = None;
snapshot.api_key_allowed_providers = None;
snapshot.user_allowed_api_formats = Some(vec!["openai:embedding".to_string()]);
snapshot.api_key_allowed_api_formats = Some(vec!["openai:embedding".to_string()]);
snapshot.user_allowed_models = Some(vec!["qwen3-vl-embedding".to_string()]);
snapshot.api_key_allowed_models = Some(vec!["qwen3-vl-embedding".to_string()]);
let auth_repository = Arc::new(InMemoryAuthApiKeySnapshotRepository::seed(vec![(
Some(hash_api_key("sk-mixed-embedding-success")),
snapshot,
)]));
let mut openai_candidate = embedding_candidate_row();
openai_candidate.model_id = "model-openai-qwen-vl-embedding".to_string();
openai_candidate.global_model_id = "global-qwen3-vl-embedding".to_string();
openai_candidate.global_model_name = "qwen3-vl-embedding".to_string();
openai_candidate.model_provider_model_name = "openai-qwen-fallback".to_string();
let candidate_repository =
Arc::new(InMemoryMinimalCandidateSelectionReadRepository::seed(vec![
openai_candidate,
aliyun_embedding_candidate_row(),
]));
let mut aliyun_provider = sample_provider("provider-aliyun-embedding", "Aliyun DashScope", 1);
aliyun_provider.provider_type = "aliyun".to_string();
let provider_catalog_repository = Arc::new(InMemoryProviderCatalogReadRepository::seed(
vec![
sample_provider("provider-embedding", "OpenAI Embeddings", 1),
aliyun_provider,
],
vec![
sample_endpoint(
"endpoint-embedding",
"provider-embedding",
"openai:embedding",
"https://api.openai.example",
),
sample_endpoint(
"endpoint-aliyun-embedding",
"provider-aliyun-embedding",
"aliyun:multimodal_embedding",
"https://dashscope.aliyuncs.com",
),
],
vec![
sample_key(
"key-upstream-embedding",
"provider-embedding",
"openai:embedding",
"sk-upstream-embedding",
),
sample_key(
"key-upstream-aliyun-embedding",
"provider-aliyun-embedding",
"aliyun:multimodal_embedding",
"sk-upstream-aliyun-embedding",
),
],
));
let data_state =
GatewayDataState::with_provider_catalog_and_minimal_candidate_selection_for_tests(
provider_catalog_repository,
candidate_repository,
)
.with_auth_api_key_reader(auth_repository)
.with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY);
build_state_with_execution_runtime_override(execution_runtime_url)
.with_data_state_for_tests(data_state)
}
fn gemini_embedding_conversion_execution_runtime() -> Router {
Router::new().route(
"/v1/execute/sync",
@@ -189,6 +312,29 @@ fn gemini_embedding_conversion_execution_runtime() -> Router {
)
}
fn aliyun_embedding_conversion_execution_runtime(
expected_contents: serde_json::Value,
expected_parameters: Option<serde_json::Value>,
) -> Router {
let expected_contents = Arc::new(expected_contents);
let expected_parameters = Arc::new(expected_parameters);
Router::new().route(
"/v1/execute/sync",
any(move |Json(plan): Json<ExecutionPlan>| {
let expected_contents = Arc::clone(&expected_contents);
let expected_parameters = Arc::clone(&expected_parameters);
async move {
assert_openai_to_aliyun_embedding_execution_plan(
&plan,
&expected_contents,
expected_parameters.as_ref().as_ref(),
);
Json(aliyun_embedding_execution_result(&plan))
}
}),
)
}
fn vertex_gemini_embedding_conversion_execution_runtime() -> Router {
Router::new().route(
"/v1/execute/sync",
@@ -300,6 +446,40 @@ fn vertex_gemini_embedding_candidate_row() -> StoredMinimalCandidateSelectionRow
row
}
fn aliyun_embedding_candidate_row() -> StoredMinimalCandidateSelectionRow {
StoredMinimalCandidateSelectionRow {
provider_id: "provider-aliyun-embedding".to_string(),
provider_name: "Aliyun DashScope".to_string(),
provider_type: "aliyun".to_string(),
provider_priority: 1,
provider_is_active: true,
endpoint_id: "endpoint-aliyun-embedding".to_string(),
endpoint_api_format: "aliyun:multimodal_embedding".to_string(),
endpoint_api_family: Some("aliyun".to_string()),
endpoint_kind: Some("multimodal_embedding".to_string()),
endpoint_is_active: true,
key_id: "key-upstream-aliyun-embedding".to_string(),
key_name: "default".to_string(),
key_auth_type: "api_key".to_string(),
key_is_active: true,
key_api_formats: Some(vec!["aliyun:multimodal_embedding".to_string()]),
key_allowed_models: None,
key_capabilities: None,
key_internal_priority: 50,
key_global_priority_by_format: None,
model_id: "model-qwen3-vl-embedding".to_string(),
global_model_id: "global-qwen3-vl-embedding".to_string(),
global_model_name: "qwen3-vl-embedding".to_string(),
global_model_mappings: None,
global_model_supports_streaming: Some(false),
model_provider_model_name: "qwen3-vl-embedding".to_string(),
model_provider_model_mappings: None,
model_supports_streaming: Some(false),
model_is_active: true,
model_is_available: true,
}
}
fn assert_embedding_execution_plan(plan: &ExecutionPlan) {
assert_eq!(plan.client_api_format, "openai:embedding");
assert_eq!(plan.provider_api_format, "openai:embedding");
@@ -311,6 +491,34 @@ fn assert_embedding_execution_plan(plan: &ExecutionPlan) {
assert!(body.get("input").is_some());
}
fn assert_openai_to_aliyun_embedding_execution_plan(
plan: &ExecutionPlan,
expected_contents: &serde_json::Value,
expected_parameters: Option<&serde_json::Value>,
) {
assert_eq!(plan.client_api_format, "openai:embedding");
assert_eq!(plan.provider_api_format, "aliyun:multimodal_embedding");
assert_eq!(plan.method, "POST");
assert_eq!(
plan.url,
"https://dashscope.aliyuncs.com/api/v1/services/embeddings/multimodal-embedding/multimodal-embedding"
);
assert_eq!(
plan.headers.get("authorization").map(String::as_str),
Some("Bearer sk-upstream-aliyun-embedding")
);
assert_eq!(plan.model_name.as_deref(), Some("qwen3-vl-embedding"));
assert!(!plan.stream);
let body = plan.body.json_body.as_ref().expect("json request body");
assert_eq!(body["model"], "qwen3-vl-embedding");
assert_eq!(&body["input"]["contents"], expected_contents);
match expected_parameters {
Some(expected) => assert_eq!(&body["parameters"], expected),
None => assert!(body.get("parameters").is_none()),
}
assert!(body.get("messages").is_none());
}
fn assert_openai_to_gemini_embedding_execution_plan(plan: &ExecutionPlan) {
assert_eq!(plan.client_api_format, "openai:embedding");
assert_eq!(plan.provider_api_format, "gemini:embedding");
@@ -503,6 +711,41 @@ fn gemini_batch_embedding_execution_result(plan: &ExecutionPlan) -> ExecutionRes
}
}
fn aliyun_embedding_execution_result(plan: &ExecutionPlan) -> ExecutionResult {
ExecutionResult {
request_id: plan.request_id.clone(),
candidate_id: plan.candidate_id.clone(),
status_code: 200,
headers: BTreeMap::from([("content-type".to_string(), "application/json".to_string())]),
body: Some(ResponseBody {
json_body: Some(json!({
"output": {
"embeddings": [
{
"index": 0,
"embedding": [0.1, 0.2, 0.3],
"type": "fusion"
}
]
},
"usage": {
"input_tokens": 432,
"input_tokens_details": {
"image_tokens": 402,
"text_tokens": 30
},
"output_tokens": 1,
"total_tokens": 433
},
"request_id": "aliyun-request-1"
})),
body_bytes_b64: None,
}),
telemetry: None,
error: None,
}
}
#[tokio::test]
async fn embeddings_route_accepts_openai_payload() {
let (execution_runtime_url, execution_runtime_handle) =
@@ -714,6 +957,183 @@ async fn embeddings_route_converts_openai_batch_payload_to_gemini_batch_endpoint
execution_runtime_handle.abort();
}
#[tokio::test]
async fn embeddings_route_converts_text_payload_to_aliyun_embedding_provider() {
let (execution_runtime_url, execution_runtime_handle) =
start_server(aliyun_embedding_conversion_execution_runtime(
json!([{ "text": "hello" }]),
Some(json!({ "dimension": 1024 })),
))
.await;
let gateway = build_router_with_state(aliyun_embedding_success_state(execution_runtime_url));
let (gateway_url, gateway_handle) = start_server(gateway).await;
let response = reqwest::Client::new()
.post(format!("{gateway_url}/v1/embeddings"))
.header(
http::header::AUTHORIZATION,
"Bearer sk-aliyun-embedding-success",
)
.json(&json!({
"model": "qwen3-vl-embedding",
"input": "hello",
"dimensions": 1024
}))
.send()
.await
.expect("request should succeed");
assert_eq!(response.status(), StatusCode::OK);
assert_eq!(
response
.headers()
.get(CONTROL_ENDPOINT_SIGNATURE_HEADER)
.and_then(|value| value.to_str().ok()),
Some("openai:embedding")
);
let payload: serde_json::Value = response.json().await.expect("body should parse");
assert_eq!(payload["object"], "list");
assert_eq!(payload["request_id"], "aliyun-request-1");
assert_eq!(payload["model"], "qwen3-vl-embedding");
assert_eq!(payload["data"][0]["object"], "embedding");
assert_eq!(payload["data"][0]["embedding"], json!([0.1, 0.2, 0.3]));
assert_eq!(payload["data"][0]["type"], "fusion");
assert_eq!(payload["usage"]["prompt_tokens"], json!(432));
assert_eq!(payload["usage"]["completion_tokens"], json!(1));
assert_eq!(payload["usage"]["total_tokens"], json!(433));
gateway_handle.abort();
execution_runtime_handle.abort();
}
#[tokio::test]
async fn embeddings_route_converts_multimodal_payload_to_aliyun_embedding_provider() {
let expected_contents = json!([
{ "text": "white running shoes" },
{ "image": "https://example.com/shoe.png" },
{ "multi_images": ["https://example.com/a.png", "https://example.com/b.png"] }
]);
let (execution_runtime_url, execution_runtime_handle) =
start_server(aliyun_embedding_conversion_execution_runtime(
expected_contents.clone(),
Some(json!({ "res_level": 2, "max_video_frames": 64 })),
))
.await;
let gateway = build_router_with_state(aliyun_embedding_success_state(execution_runtime_url));
let (gateway_url, gateway_handle) = start_server(gateway).await;
let response = reqwest::Client::new()
.post(format!("{gateway_url}/v1/embeddings"))
.header(
http::header::AUTHORIZATION,
"Bearer sk-aliyun-embedding-success",
)
.json(&json!({
"model": "qwen3-vl-embedding",
"input": expected_contents,
"parameters": {
"res_level": 2,
"max_video_frames": 64
}
}))
.send()
.await
.expect("request should succeed");
assert_eq!(response.status(), StatusCode::OK);
let payload: serde_json::Value = response.json().await.expect("body should parse");
assert_eq!(payload["data"][0]["embedding"], json!([0.1, 0.2, 0.3]));
assert_eq!(payload["data"][0]["type"], "fusion");
gateway_handle.abort();
execution_runtime_handle.abort();
}
#[tokio::test]
async fn embeddings_route_skips_openai_candidate_for_multimodal_payload() {
let expected_contents = json!([
{ "text": "white running shoes" },
{ "image": "https://example.com/shoe.png" }
]);
let (execution_runtime_url, execution_runtime_handle) =
start_server(aliyun_embedding_conversion_execution_runtime(
expected_contents.clone(),
Some(json!({ "enable_fusion": true })),
))
.await;
let gateway = build_router_with_state(mixed_embedding_success_state(execution_runtime_url));
let (gateway_url, gateway_handle) = start_server(gateway).await;
let response = reqwest::Client::new()
.post(format!("{gateway_url}/v1/embeddings"))
.header(
http::header::AUTHORIZATION,
"Bearer sk-mixed-embedding-success",
)
.json(&json!({
"model": "qwen3-vl-embedding",
"input": expected_contents,
"parameters": {
"enable_fusion": true
}
}))
.send()
.await
.expect("request should succeed");
assert_eq!(response.status(), StatusCode::OK);
let payload: serde_json::Value = response.json().await.expect("body should parse");
assert_eq!(payload["data"][0]["embedding"], json!([0.1, 0.2, 0.3]));
assert_eq!(payload["data"][0]["type"], "fusion");
gateway_handle.abort();
execution_runtime_handle.abort();
}
#[tokio::test]
async fn embeddings_route_converts_fusion_payload_to_aliyun_embedding_provider() {
let expected_contents = json!([
{
"text": "white running shoes",
"image": "https://example.com/shoe.png"
},
{ "video": "https://example.com/demo.mp4" }
]);
let (execution_runtime_url, execution_runtime_handle) =
start_server(aliyun_embedding_conversion_execution_runtime(
expected_contents.clone(),
Some(json!({ "enable_fusion": true })),
))
.await;
let gateway = build_router_with_state(aliyun_embedding_success_state(execution_runtime_url));
let (gateway_url, gateway_handle) = start_server(gateway).await;
let response = reqwest::Client::new()
.post(format!("{gateway_url}/v1/embeddings"))
.header(
http::header::AUTHORIZATION,
"Bearer sk-aliyun-embedding-success",
)
.json(&json!({
"model": "qwen3-vl-embedding",
"input": expected_contents,
"parameters": {
"enable_fusion": true
}
}))
.send()
.await
.expect("request should succeed");
assert_eq!(response.status(), StatusCode::OK);
let payload: serde_json::Value = response.json().await.expect("body should parse");
assert_eq!(payload["data"][0]["embedding"], json!([0.1, 0.2, 0.3]));
assert_eq!(payload["data"][0]["type"], "fusion");
gateway_handle.abort();
execution_runtime_handle.abort();
}
#[tokio::test]
async fn gemini_embed_content_route_uses_native_gemini_embedding_provider() {
let (execution_runtime_url, execution_runtime_handle) =
@@ -836,6 +1256,14 @@ async fn embeddings_route_rejects_invalid_local_payloads() {
r#"{"model":"text-embedding-3-small","input":[[1],[]]}"#,
"Embedding request input is required",
),
(
r#"{"model":"text-embedding-3-small","input":[{}]}"#,
"Embedding request input is required",
),
(
r#"{"model":"text-embedding-3-small","input":[{"image":" "} ]}"#,
"Embedding request input is required",
),
(
r#"{"model":"text-embedding-3-small","input":"hello","stream":true}"#,
"Embedding requests do not support streaming",
@@ -24,6 +24,7 @@ use aether_data::repository::auth::{
use aether_data::repository::auth_modules::{
InMemoryAuthModuleReadRepository, StoredLdapModuleConfig, StoredOAuthProviderModuleConfig,
};
use aether_data::repository::billing::InMemoryBillingReadRepository;
use aether_data::repository::management_tokens::{
InMemoryManagementTokenRepository, StoredManagementToken, StoredManagementTokenUserSummary,
StoredManagementTokenWithUser,
@@ -36,6 +37,10 @@ use aether_data::repository::users::{
use aether_data::repository::wallet::{
InMemoryWalletRepository, StoredWalletSnapshot, WalletWriteRepository,
};
use aether_data_contracts::repository::billing::{
AdminBillingMutationOutcome, BillingPlanWriteInput, BillingReadRepository,
PaymentGatewayConfigWriteInput,
};
use aether_data_contracts::repository::global_models::StoredProviderActiveGlobalModel;
use aether_data_contracts::repository::provider_catalog::ProviderCatalogReadRepository;
use aether_data_contracts::repository::usage::{StoredRequestUsageAudit, UsageRepository};
@@ -3901,6 +3906,209 @@ async fn gateway_creates_wallet_recharge_orders_locally_without_proxying_upstrea
upstream_handle.abort();
}
#[tokio::test]
async fn gateway_reuses_pending_billing_plan_checkout_order_without_proxying_upstream() {
let now = Utc::now();
let user = StoredUserAuthRecord::new(
"user-billing-checkout-reuse".to_string(),
Some("[email protected]".to_string()),
true,
"billing_checkout_reuse_user".to_string(),
Some("$2y$10$.OBQfixAECpsb8V/VS3csOMf00x2E/jD/gnud20t6RG0yiQosyOZ2".to_string()),
"user".to_string(),
"local".to_string(),
Some(json!(["openai"])),
Some(json!(["openai:chat"])),
Some(json!(["gpt-5"])),
true,
false,
Some(now),
Some(now),
)
.expect("auth user should build");
let wallet = StoredWalletSnapshot::new(
"wallet-billing-checkout-reuse".to_string(),
Some(user.id.clone()),
None,
12.5,
3.0,
"finite".to_string(),
"USD".to_string(),
"active".to_string(),
20.0,
4.5,
0.0,
0.0,
now.timestamp(),
)
.expect("wallet should build");
let access_token = build_test_auth_token(
"access",
serde_json::Map::from_iter([
("user_id".to_string(), json!(user.id.clone())),
("role".to_string(), json!(user.role.clone())),
(
"created_at".to_string(),
json!(user.created_at.map(|value| value.to_rfc3339())),
),
(
"session_id".to_string(),
json!("session-billing-checkout-reuse"),
),
]),
now + chrono::Duration::hours(1),
);
let billing_repository = Arc::new(InMemoryBillingReadRepository::seed(Vec::new()));
let encrypted_key = encrypt_python_fernet_plaintext(DEVELOPMENT_ENCRYPTION_KEY, "epay-secret")
.expect("merchant key should encrypt");
let AdminBillingMutationOutcome::Applied(_) = billing_repository
.upsert_payment_gateway_config(&PaymentGatewayConfigWriteInput {
provider: "epay".to_string(),
enabled: true,
endpoint_url: "https://pay.example.com/".to_string(),
callback_base_url: Some("https://app.example.com".to_string()),
merchant_id: "merchant-1".to_string(),
merchant_key_encrypted: Some(encrypted_key),
preserve_existing_secret: false,
pay_currency: "CNY".to_string(),
usd_exchange_rate: 7.25,
min_recharge_usd: 1.0,
channels_json: json!([
{
"channel": "alipay",
"display_name": "支付宝"
}
]),
})
.await
.expect("gateway config should create")
else {
panic!("gateway config should apply");
};
let plan = match billing_repository
.create_billing_plan(&BillingPlanWriteInput {
title: "每日额度月卡".to_string(),
description: Some("测试套餐".to_string()),
price_amount: 100.0,
price_currency: "CNY".to_string(),
duration_unit: "month".to_string(),
duration_value: 1,
enabled: true,
sort_order: 1,
max_active_per_user: 1,
purchase_limit_scope: "active_period".to_string(),
entitlements_json: json!([
{
"type": "daily_quota",
"daily_quota_usd": 50.0,
"reset_timezone": "Asia/Shanghai",
"allow_wallet_overage": false
}
]),
})
.await
.expect("billing plan should create")
{
AdminBillingMutationOutcome::Applied(plan) => plan,
other => panic!("billing plan should apply, got {other:?}"),
};
let upstream_hits = Arc::new(Mutex::new(0usize));
let upstream_hits_clone = Arc::clone(&upstream_hits);
let upstream = Router::new().route(
"/{*path}",
any(move |_request: Request| {
let upstream_hits_inner = Arc::clone(&upstream_hits_clone);
async move {
*upstream_hits_inner.lock().expect("mutex should lock") += 1;
(StatusCode::OK, Body::from("proxied"))
}
}),
);
let (_upstream_url, upstream_handle) = start_server(upstream).await;
let user_repository = Arc::new(InMemoryUserReadRepository::seed_auth_users(vec![user]));
let wallet_repository = Arc::new(InMemoryWalletRepository::seed(vec![wallet]));
let state = AppState::new()
.expect("gateway should build")
.with_data_state_for_tests(GatewayDataState::with_user_billing_and_wallet_for_tests(
user_repository,
billing_repository,
wallet_repository,
))
.with_auth_sessions_for_tests([sample_auth_session(
"user-billing-checkout-reuse",
"session-billing-checkout-reuse",
"device-billing-checkout-reuse",
"refresh-token-placeholder",
now,
)]);
let gateway = build_router_with_state(state);
let (gateway_url, gateway_handle) = start_server(gateway).await;
let client = reqwest::Client::new();
let checkout_body = json!({
"payment_provider": "epay",
"payment_method": "epay",
"payment_channel": "alipay",
});
let first_response = client
.post(format!(
"{gateway_url}/api/billing/plans/{}/checkout",
plan.id
))
.header("authorization", format!("Bearer {access_token}"))
.header("x-client-device-id", "device-billing-checkout-reuse")
.header("user-agent", "AetherTest/1.0")
.json(&checkout_body)
.send()
.await
.expect("first checkout request should succeed");
assert_eq!(first_response.status(), StatusCode::OK);
let first_payload: serde_json::Value = first_response
.json()
.await
.expect("first checkout json should parse");
let first_order_id = first_payload["order"]["id"]
.as_str()
.expect("first order id should exist")
.to_string();
assert_eq!(first_payload["order"]["status"], "pending");
assert_eq!(first_payload["order"]["product_id"], plan.id);
assert_eq!(
first_payload["reused_pending_order"],
serde_json::Value::Null
);
let second_response = client
.post(format!(
"{gateway_url}/api/billing/plans/{}/checkout",
plan.id
))
.header("authorization", format!("Bearer {access_token}"))
.header("x-client-device-id", "device-billing-checkout-reuse")
.header("user-agent", "AetherTest/1.0")
.json(&checkout_body)
.send()
.await
.expect("second checkout request should succeed");
assert_eq!(second_response.status(), StatusCode::OK);
let second_payload: serde_json::Value = second_response
.json()
.await
.expect("second checkout json should parse");
assert_eq!(second_payload["order"]["id"], first_order_id);
assert_eq!(second_payload["reused_pending_order"], true);
assert_eq!(
second_payload["payment_instructions"],
first_payload["payment_instructions"]
);
assert_eq!(*upstream_hits.lock().expect("mutex should lock"), 0);
gateway_handle.abort();
upstream_handle.abort();
}
#[tokio::test]
async fn gateway_creates_wallet_refunds_locally_without_proxying_upstream() {
let now = Utc::now();
+128 -125
View File
@@ -154,74 +154,79 @@ async fn gateway_records_usage_for_execution_runtime_sync_when_runtime_enabled()
#[test]
fn gateway_records_pending_usage_before_execution_runtime_sync_result_arrives() {
run_async_test_on_large_stack("pending-usage-sync-before-runtime-result", async move {
let usage_repository = Arc::new(InMemoryUsageReadRepository::default());
let request_candidate_repository = Arc::new(InMemoryRequestCandidateRepository::default());
let execution_request_started = Arc::new(tokio::sync::Notify::new());
let allow_execution_response = Arc::new(tokio::sync::Notify::new());
run_async_test_on_large_stack(
"gateway_records_pending_usage_before_execution_runtime_sync_result_arrives",
gateway_records_pending_usage_before_execution_runtime_sync_result_arrives_impl(),
);
}
let upstream = Router::new().route(
"/api/internal/gateway/report-sync",
any(|_request: Request| async move { Json(json!({"ok": true})) }),
);
async fn gateway_records_pending_usage_before_execution_runtime_sync_result_arrives_impl() {
let usage_repository = Arc::new(InMemoryUsageReadRepository::default());
let request_candidate_repository = Arc::new(InMemoryRequestCandidateRepository::default());
let execution_request_started = Arc::new(tokio::sync::Notify::new());
let allow_execution_response = Arc::new(tokio::sync::Notify::new());
let execution_runtime = Router::new().route(
"/v1/execute/sync",
any({
let upstream = Router::new().route(
"/api/internal/gateway/report-sync",
any(|_request: Request| async move { Json(json!({"ok": true})) }),
);
let execution_runtime = Router::new().route(
"/v1/execute/sync",
any({
let execution_request_started = Arc::clone(&execution_request_started);
let allow_execution_response = Arc::clone(&allow_execution_response);
move |_request: Request| {
let execution_request_started = Arc::clone(&execution_request_started);
let allow_execution_response = Arc::clone(&allow_execution_response);
move |_request: Request| {
let execution_request_started = Arc::clone(&execution_request_started);
let allow_execution_response = Arc::clone(&allow_execution_response);
async move {
execution_request_started.notify_one();
allow_execution_response.notified().await;
Json(json!({
"request_id": "req-usage-sync-pending-123",
"status_code": 200,
"headers": {
"content-type": "application/json"
},
"body": {
"json_body": {
"id": "chatcmpl-usage-sync-pending-123",
"usage": {
"input_tokens": 3,
"output_tokens": 5,
"total_tokens": 8
}
async move {
execution_request_started.notify_one();
allow_execution_response.notified().await;
Json(json!({
"request_id": "req-usage-sync-pending-123",
"status_code": 200,
"headers": {
"content-type": "application/json"
},
"body": {
"json_body": {
"id": "chatcmpl-usage-sync-pending-123",
"usage": {
"input_tokens": 3,
"output_tokens": 5,
"total_tokens": 8
}
},
"telemetry": {
"elapsed_ms": 45
}
}))
}
},
"telemetry": {
"elapsed_ms": 45
}
}))
}
}),
);
}
}),
);
let auth_repository = Arc::new(InMemoryAuthApiKeySnapshotRepository::seed(vec![(
Some(hash_api_key("sk-client-openai-usage-sync-pending")),
sample_local_openai_auth_snapshot(
"api-key-usage-sync-pending-123",
"user-usage-sync-pending-123",
),
)]));
let candidate_selection_repository =
Arc::new(InMemoryMinimalCandidateSelectionReadRepository::seed(vec![
sample_local_openai_candidate_row(),
]));
let provider_catalog_repository = Arc::new(InMemoryProviderCatalogReadRepository::seed(
vec![sample_local_openai_provider()],
vec![sample_local_openai_endpoint()],
vec![sample_local_openai_key()],
));
let auth_repository = Arc::new(InMemoryAuthApiKeySnapshotRepository::seed(vec![(
Some(hash_api_key("sk-client-openai-usage-sync-pending")),
sample_local_openai_auth_snapshot(
"api-key-usage-sync-pending-123",
"user-usage-sync-pending-123",
),
)]));
let candidate_selection_repository =
Arc::new(InMemoryMinimalCandidateSelectionReadRepository::seed(vec![
sample_local_openai_candidate_row(),
]));
let provider_catalog_repository = Arc::new(InMemoryProviderCatalogReadRepository::seed(
vec![sample_local_openai_provider()],
vec![sample_local_openai_endpoint()],
vec![sample_local_openai_key()],
));
let (upstream_url, upstream_handle) = start_server(upstream).await;
let (execution_runtime_url, execution_runtime_handle) =
start_server(execution_runtime).await;
let gateway_state =
let (upstream_url, upstream_handle) = start_server(upstream).await;
let (execution_runtime_url, execution_runtime_handle) = start_server(execution_runtime).await;
let gateway_state =
build_state_with_execution_runtime_override(execution_runtime_url)
.with_data_state_for_tests(
GatewayDataState::with_auth_candidate_selection_provider_catalog_request_candidates_and_usage_for_tests(
@@ -237,76 +242,74 @@ fn gateway_records_pending_usage_before_execution_runtime_sync_result_arrives()
enabled: true,
..UsageRuntimeConfig::default()
});
let gateway = build_router_with_state(gateway_state);
let (gateway_url, gateway_handle) = start_server(gateway).await;
let gateway = build_router_with_state(gateway_state);
let (gateway_url, gateway_handle) = start_server(gateway).await;
let request_task = tokio::spawn({
let gateway_url = gateway_url.clone();
async move {
let response = reqwest::Client::new()
.post(format!("{gateway_url}/v1/chat/completions"))
.header(http::header::CONTENT_TYPE, "application/json")
.header(
http::header::AUTHORIZATION,
"Bearer sk-client-openai-usage-sync-pending",
)
.header(TRACE_ID_HEADER, "req-usage-sync-pending-123")
.body("{\"model\":\"gpt-5\",\"messages\":[]}")
.send()
.await
.expect("request should succeed");
let status = response.status();
let body = response.text().await.expect("body should read");
(status, body)
}
});
execution_request_started.notified().await;
let mut pending = None;
for _ in 0..50 {
pending = usage_repository
.find_by_request_id("req-usage-sync-pending-123")
let request_task = tokio::spawn({
let gateway_url = gateway_url.clone();
async move {
let response = reqwest::Client::new()
.post(format!("{gateway_url}/v1/chat/completions"))
.header(http::header::CONTENT_TYPE, "application/json")
.header(
http::header::AUTHORIZATION,
"Bearer sk-client-openai-usage-sync-pending",
)
.header(TRACE_ID_HEADER, "req-usage-sync-pending-123")
.body("{\"model\":\"gpt-5\",\"messages\":[]}")
.send()
.await
.expect("usage lookup should succeed");
if pending
.as_ref()
.is_some_and(|stored| stored.status == "pending")
{
break;
}
tokio::time::sleep(std::time::Duration::from_millis(10)).await;
.expect("request should succeed");
let status = response.status();
let body = response.text().await.expect("body should read");
(status, body)
}
let pending =
pending.expect("pending usage should be recorded before sync result resolves");
assert_eq!(pending.status, "pending");
assert_eq!(pending.billing_status, "pending");
assert_eq!(pending.response_time_ms, None);
allow_execution_response.notify_one();
let (status, _body) = request_task.await.expect("request task should join");
assert_eq!(status, StatusCode::OK);
let mut stored = None;
for _ in 0..50 {
stored = usage_repository
.find_by_request_id("req-usage-sync-pending-123")
.await
.expect("usage lookup should succeed");
if stored.as_ref().is_some_and(|row| row.status == "completed") {
break;
}
tokio::time::sleep(std::time::Duration::from_millis(10)).await;
}
let stored = stored.expect("usage should be finalized");
assert_eq!(stored.status, "completed");
assert_eq!(stored.response_time_ms, Some(45));
gateway_handle.abort();
execution_runtime_handle.abort();
upstream_handle.abort();
});
execution_request_started.notified().await;
let mut pending = None;
for _ in 0..50 {
pending = usage_repository
.find_by_request_id("req-usage-sync-pending-123")
.await
.expect("usage lookup should succeed");
if pending
.as_ref()
.is_some_and(|stored| stored.status == "pending")
{
break;
}
tokio::time::sleep(std::time::Duration::from_millis(10)).await;
}
let pending = pending.expect("pending usage should be recorded before sync result resolves");
assert_eq!(pending.status, "pending");
assert_eq!(pending.billing_status, "pending");
assert_eq!(pending.response_time_ms, None);
allow_execution_response.notify_one();
let (status, _body) = request_task.await.expect("request task should join");
assert_eq!(status, StatusCode::OK);
let mut stored = None;
for _ in 0..50 {
stored = usage_repository
.find_by_request_id("req-usage-sync-pending-123")
.await
.expect("usage lookup should succeed");
if stored.as_ref().is_some_and(|row| row.status == "completed") {
break;
}
tokio::time::sleep(std::time::Duration::from_millis(10)).await;
}
let stored = stored.expect("usage should be finalized");
assert_eq!(stored.status, "completed");
assert_eq!(stored.response_time_ms, Some(45));
gateway_handle.abort();
execution_runtime_handle.abort();
upstream_handle.abort();
}
#[tokio::test]