mirror of
https://github.com/fawney19/Aether.git
synced 2026-10-08 10:27:46 +08:00
Merge remote-tracking branch 'upstream/main'
This commit is contained in:
@@ -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(
|
||||
|
||||
@@ -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,3 +1,4 @@
|
||||
mod aliyun;
|
||||
mod claude;
|
||||
mod doubao;
|
||||
mod gemini;
|
||||
|
||||
@@ -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();
|
||||
|
||||
@@ -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]
|
||||
|
||||
Reference in New Issue
Block a user