mirror of
https://github.com/fawney19/Aether.git
synced 2026-10-04 16:37:46 +08:00
feat(gateway): harden provider request execution
Preserve exact request payloads and model client surface and API operation explicitly. Add Anthropic compatibility profiles, bounded stream commitment, and scoped OAuth retry behavior across provider transports.
This commit is contained in:
@@ -287,6 +287,7 @@ pub(super) async fn handle_admin_provider_oauth_complete_key(
|
||||
&ProviderCatalogKeyOAuthRuntimeStateCasUpdate {
|
||||
key_id: key_id.clone(),
|
||||
expected_encrypted_auth_config: state_data.expected_encrypted_auth_config,
|
||||
expected_credential: None,
|
||||
encrypted_auth_config: persisted_encrypted_auth_config.clone(),
|
||||
encrypted_api_key_update: Some(encrypted_api_key),
|
||||
expires_at_unix_secs_update: Some(expires_at),
|
||||
|
||||
@@ -370,6 +370,7 @@ pub(crate) async fn persist_fenced_provider_quota_refresh_state(
|
||||
&ProviderCatalogKeyOAuthRuntimeStateCasUpdate {
|
||||
key_id: key_id.to_string(),
|
||||
expected_encrypted_auth_config: Some(expected_encrypted_auth_config.to_string()),
|
||||
expected_credential: None,
|
||||
encrypted_auth_config: expected_encrypted_auth_config.to_string(),
|
||||
encrypted_api_key_update: None,
|
||||
expires_at_unix_secs_update: None,
|
||||
|
||||
@@ -100,14 +100,6 @@ fn select_provider_oauth_runtime_endpoint(
|
||||
.api_format
|
||||
.trim()
|
||||
.eq_ignore_ascii_case("gemini:generate_content")
|
||||
})
|
||||
.or_else(|| {
|
||||
matching_endpoint(endpoints, include_inactive, |endpoint| {
|
||||
endpoint
|
||||
.api_format
|
||||
.trim()
|
||||
.eq_ignore_ascii_case("claude:messages")
|
||||
})
|
||||
}),
|
||||
_ => matching_endpoint(endpoints, include_inactive, |_| true),
|
||||
}
|
||||
@@ -255,3 +247,44 @@ pub(crate) fn spawn_provider_oauth_account_state_refresh_after_update(
|
||||
.await;
|
||||
});
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::{
|
||||
provider_oauth_maintenance_endpoint_for_provider,
|
||||
provider_oauth_runtime_endpoint_for_provider,
|
||||
};
|
||||
use aether_data_contracts::repository::provider_catalog::StoredProviderCatalogEndpoint;
|
||||
|
||||
fn endpoint(id: &str, api_format: &str, is_active: bool) -> StoredProviderCatalogEndpoint {
|
||||
StoredProviderCatalogEndpoint::new(
|
||||
id.to_string(),
|
||||
"provider-1".to_string(),
|
||||
api_format.to_string(),
|
||||
None,
|
||||
None,
|
||||
is_active,
|
||||
)
|
||||
.expect("endpoint should build")
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn vertex_oauth_runtime_never_falls_back_to_retired_claude_endpoint() {
|
||||
let endpoints = vec![endpoint("claude", "claude:messages", true)];
|
||||
|
||||
assert!(provider_oauth_runtime_endpoint_for_provider("vertex_ai", &endpoints).is_none());
|
||||
assert!(
|
||||
provider_oauth_maintenance_endpoint_for_provider("vertex_ai", &endpoints).is_none()
|
||||
);
|
||||
|
||||
let endpoints = vec![
|
||||
endpoint("claude", "claude:messages", true),
|
||||
endpoint("gemini", "gemini:generate_content", true),
|
||||
];
|
||||
assert_eq!(
|
||||
provider_oauth_runtime_endpoint_for_provider("vertex_ai", &endpoints)
|
||||
.map(|endpoint| endpoint.id),
|
||||
Some("gemini".to_string())
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -2668,6 +2668,7 @@ async fn provider_query_execute_antigravity_test_candidate(
|
||||
upstream_is_stream: false,
|
||||
request_query: parts.uri.query(),
|
||||
kiro_api_region: None,
|
||||
api_operation: None,
|
||||
},
|
||||
);
|
||||
let Some(request_url) = request_url else {
|
||||
@@ -3364,6 +3365,7 @@ async fn provider_query_execute_standard_test_candidate(
|
||||
upstream_is_stream,
|
||||
request_query: parts.uri.query(),
|
||||
kiro_api_region: None,
|
||||
api_operation: None,
|
||||
},
|
||||
Some(&provider_request_body),
|
||||
);
|
||||
|
||||
@@ -195,11 +195,7 @@ pub(crate) fn validate_vertex_api_formats(
|
||||
|
||||
let allowed = match auth_type {
|
||||
"api_key" => &["gemini:generate_content", "gemini:embedding"][..],
|
||||
"service_account" | "vertex_ai" => &[
|
||||
"claude:messages",
|
||||
"gemini:generate_content",
|
||||
"gemini:embedding",
|
||||
][..],
|
||||
"service_account" | "vertex_ai" => &["gemini:generate_content", "gemini:embedding"][..],
|
||||
_ => return Ok(()),
|
||||
};
|
||||
let invalid = api_formats
|
||||
@@ -410,20 +406,11 @@ mod tests {
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn validate_vertex_api_formats_uses_canonical_message_formats() {
|
||||
fn validate_vertex_api_formats_rejects_unimplemented_anthropic_transport() {
|
||||
assert!(validate_vertex_api_formats(
|
||||
"vertex_ai",
|
||||
"service_account",
|
||||
&[
|
||||
"claude:messages".to_string(),
|
||||
"gemini:generate_content".to_string()
|
||||
],
|
||||
)
|
||||
.is_ok());
|
||||
assert!(validate_vertex_api_formats(
|
||||
"vertex_ai",
|
||||
"service_account",
|
||||
&["claude:chat".to_string()],
|
||||
&["claude:messages".to_string()],
|
||||
)
|
||||
.is_err());
|
||||
}
|
||||
@@ -443,7 +430,6 @@ mod tests {
|
||||
"vertex_ai",
|
||||
"service_account",
|
||||
&[
|
||||
"claude:messages".to_string(),
|
||||
"gemini:generate_content".to_string(),
|
||||
"gemini:embedding".to_string()
|
||||
],
|
||||
|
||||
@@ -158,6 +158,8 @@ pub(crate) async fn build_admin_create_provider_record(
|
||||
}
|
||||
}
|
||||
let config = (!config_map.is_empty()).then_some(serde_json::Value::Object(config_map));
|
||||
crate::provider_transport::validate_anthropic_compatibility_profile_config(config.as_ref())
|
||||
.map_err(|_| "无效的 Anthropic compatibility profile".to_string())?;
|
||||
|
||||
let now_unix_secs = SystemTime::now()
|
||||
.duration_since(UNIX_EPOCH)
|
||||
|
||||
@@ -312,6 +312,10 @@ pub(crate) async fn build_admin_update_provider_record(
|
||||
}
|
||||
|
||||
updated.config = (!config_map.is_empty()).then_some(serde_json::Value::Object(config_map));
|
||||
crate::provider_transport::validate_anthropic_compatibility_profile_config(
|
||||
updated.config.as_ref(),
|
||||
)
|
||||
.map_err(|_| "无效的 Anthropic compatibility profile".to_string())?;
|
||||
updated.updated_at_unix_secs = SystemTime::now()
|
||||
.duration_since(UNIX_EPOCH)
|
||||
.ok()
|
||||
|
||||
@@ -259,6 +259,10 @@ impl<'a> AdminAppState<'a> {
|
||||
admin_endpoint_signature_parts(&payload.api_format)
|
||||
.ok_or_else(|| format!("无效的 api_format: {}", payload.api_format))?;
|
||||
validate_admin_endpoint_stream_policy(normalized_api_format, payload.config.as_ref())?;
|
||||
crate::provider_transport::validate_anthropic_compatibility_profile_config(
|
||||
payload.config.as_ref(),
|
||||
)
|
||||
.map_err(|_| "无效的 Anthropic compatibility profile".to_string())?;
|
||||
let base_url = normalize_admin_base_url(&payload.base_url)?;
|
||||
|
||||
let existing_endpoints = self
|
||||
@@ -369,6 +373,10 @@ impl<'a> AdminAppState<'a> {
|
||||
existing_endpoint.api_format.as_str(),
|
||||
updated.config.as_ref(),
|
||||
)?;
|
||||
crate::provider_transport::validate_anthropic_compatibility_profile_config(
|
||||
updated.config.as_ref(),
|
||||
)
|
||||
.map_err(|_| "无效的 Anthropic compatibility profile".to_string())?;
|
||||
}
|
||||
|
||||
if provider_type == "codex"
|
||||
|
||||
@@ -192,6 +192,15 @@ fn normalize_import_endpoint_format(value: &str) -> Result<String, String> {
|
||||
.ok_or_else(|| format!("无效的 api_format: {value}"))
|
||||
}
|
||||
|
||||
fn fixed_provider_import_endpoint_supported(provider_type: &str, api_format: &str) -> bool {
|
||||
crate::provider_transport::provider_types::fixed_provider_template(provider_type).is_none()
|
||||
|| crate::provider_transport::provider_types::fixed_provider_endpoint_template_by_api_format(
|
||||
provider_type,
|
||||
api_format,
|
||||
)
|
||||
.is_some()
|
||||
}
|
||||
|
||||
fn normalize_import_key_formats(
|
||||
item: &ImportedProviderKey,
|
||||
provider_endpoint_formats: &BTreeSet<String>,
|
||||
@@ -1419,6 +1428,12 @@ impl<'a> AdminAppState<'a> {
|
||||
for imported_provider_item in imported_providers {
|
||||
let (raw_provider, imported_provider) = imported_provider_item.into_parts();
|
||||
let provider_name = invalid!(trim_required(&imported_provider.name, "name"));
|
||||
invalid!(
|
||||
crate::provider_transport::validate_anthropic_compatibility_profile_config(
|
||||
imported_provider.config.as_ref(),
|
||||
)
|
||||
.map_err(|_| "无效的 Anthropic compatibility profile".to_string())
|
||||
);
|
||||
let existing_provider = providers_by_name.get(&provider_name).cloned();
|
||||
|
||||
let provider = if let Some(existing) = existing_provider {
|
||||
@@ -1513,6 +1528,42 @@ impl<'a> AdminAppState<'a> {
|
||||
let normalized_api_format = invalid!(normalize_import_endpoint_format(
|
||||
&imported_endpoint.api_format
|
||||
));
|
||||
invalid!(
|
||||
crate::provider_transport::validate_anthropic_compatibility_profile_config(
|
||||
imported_endpoint.config.as_ref(),
|
||||
)
|
||||
.map_err(|_| "无效的 Anthropic compatibility profile".to_string())
|
||||
);
|
||||
if !fixed_provider_import_endpoint_supported(
|
||||
&provider.provider_type,
|
||||
&normalized_api_format,
|
||||
) {
|
||||
let retired = existing_endpoints_by_format.remove(&normalized_api_format);
|
||||
if let Some(mut retired) = retired {
|
||||
if retired.is_active {
|
||||
retired.is_active = false;
|
||||
retired.updated_at_unix_secs = SystemTime::now()
|
||||
.duration_since(UNIX_EPOCH)
|
||||
.ok()
|
||||
.map(|duration| duration.as_secs());
|
||||
let Some(_) = self.update_provider_catalog_endpoint(&retired).await?
|
||||
else {
|
||||
return Ok(Err(invalid_request(format!(
|
||||
"停用 Provider '{provider_name}' 的已移除 Endpoint '{normalized_api_format}' 失败"
|
||||
))));
|
||||
};
|
||||
stats.endpoints.updated += 1;
|
||||
} else {
|
||||
stats.endpoints.skipped += 1;
|
||||
}
|
||||
} else {
|
||||
stats.endpoints.skipped += 1;
|
||||
}
|
||||
stats.errors.push(format!(
|
||||
"固定 Provider '{provider_name}' 不再支持 Endpoint '{normalized_api_format}',已跳过或停用"
|
||||
));
|
||||
continue;
|
||||
}
|
||||
let existing_endpoint = existing_endpoints_by_format
|
||||
.get(&normalized_api_format)
|
||||
.cloned();
|
||||
|
||||
@@ -17,8 +17,8 @@ use crate::ai_serving::api::{
|
||||
};
|
||||
use crate::api::response::{
|
||||
build_client_response, build_client_response_from_parts, build_local_auth_rejection_response,
|
||||
build_local_http_error_response, build_local_overloaded_response,
|
||||
build_local_user_rpm_limited_response,
|
||||
build_local_http_error_response, build_local_http_error_response_with_request_path,
|
||||
build_local_overloaded_response, build_local_user_rpm_limited_response,
|
||||
};
|
||||
use crate::constants::{
|
||||
CONTROL_CANDIDATE_ID_HEADER, DEPENDENCY_REASON_HEADER, EXECUTION_PATH_CONTROL_EXECUTE_STREAM,
|
||||
@@ -935,7 +935,13 @@ async fn proxy_request_inner(
|
||||
limit,
|
||||
})) => {
|
||||
let trace_id = extract_or_generate_trace_id(request.headers());
|
||||
let response = build_local_overloaded_response(&trace_id, None, gate, limit)?;
|
||||
let response = build_local_overloaded_response(
|
||||
&trace_id,
|
||||
None,
|
||||
Some(request.uri().path()),
|
||||
gate,
|
||||
limit,
|
||||
)?;
|
||||
return Ok(finalize_gateway_response(
|
||||
&state,
|
||||
response,
|
||||
@@ -965,7 +971,13 @@ async fn proxy_request_inner(
|
||||
aether_runtime_state::RuntimeSemaphoreError::Unavailable { gate, limit, .. },
|
||||
)) => {
|
||||
let trace_id = extract_or_generate_trace_id(request.headers());
|
||||
let response = build_local_overloaded_response(&trace_id, None, gate, limit)?;
|
||||
let response = build_local_overloaded_response(
|
||||
&trace_id,
|
||||
None,
|
||||
Some(request.uri().path()),
|
||||
gate,
|
||||
limit,
|
||||
)?;
|
||||
return Ok(finalize_gateway_response(
|
||||
&state,
|
||||
response,
|
||||
@@ -999,9 +1011,10 @@ async fn proxy_request_inner(
|
||||
path = %request.uri().path(),
|
||||
"gateway rejected blacklisted client IP"
|
||||
);
|
||||
let response = build_local_http_error_response(
|
||||
let response = build_local_http_error_response_with_request_path(
|
||||
&trace_id,
|
||||
None,
|
||||
Some(request.uri().path()),
|
||||
http::StatusCode::FORBIDDEN,
|
||||
"当前 IP 已被禁止访问",
|
||||
)?;
|
||||
@@ -1057,9 +1070,10 @@ async fn proxy_request_inner(
|
||||
loop_guard_header = EXECUTION_RUNTIME_LOOP_GUARD_HEADER,
|
||||
"gateway rejected execution runtime request loop into frontdoor"
|
||||
);
|
||||
let response = build_local_http_error_response(
|
||||
let response = build_local_http_error_response_with_request_path(
|
||||
&trace_id,
|
||||
None,
|
||||
Some(parts.uri.path()),
|
||||
http::StatusCode::LOOP_DETECTED,
|
||||
LOCAL_EXECUTION_LOOP_DETECTED_DETAIL,
|
||||
)?;
|
||||
@@ -1534,9 +1548,10 @@ async fn proxy_request_inner(
|
||||
}
|
||||
|
||||
if control_decision.is_none() {
|
||||
let response = build_local_http_error_response(
|
||||
let response = build_local_http_error_response_with_request_path(
|
||||
&trace_id,
|
||||
None,
|
||||
Some(request_context.request_path.as_str()),
|
||||
http::StatusCode::NOT_FOUND,
|
||||
LOCAL_ROUTE_NOT_FOUND_DETAIL,
|
||||
)?;
|
||||
|
||||
@@ -1,4 +1,6 @@
|
||||
use crate::ai_serving::normalize_openai_image_quality;
|
||||
use crate::ai_serving::{
|
||||
build_core_error_body_for_client_format, normalize_openai_image_quality, LocalCoreSyncErrorKind,
|
||||
};
|
||||
use crate::async_task::CancelVideoTaskError;
|
||||
use crate::control::GatewayControlDecision;
|
||||
use crate::control::GatewayPublicRequestContext;
|
||||
@@ -13,8 +15,6 @@ use axum::response::IntoResponse;
|
||||
use axum::Json;
|
||||
use serde_json::{json, Value};
|
||||
|
||||
const CLAUDE_COUNT_TOKENS_INVALID_PAYLOAD_DETAIL: &str = "Invalid token count payload";
|
||||
const CLAUDE_COUNT_TOKENS_MISSING_BODY_DETAIL: &str = "请求体不能为空";
|
||||
const GEMINI_VIDEO_TASK_NOT_FOUND_DETAIL: &str = "Video task not found";
|
||||
const AI_PUBLIC_METHOD_NOT_ALLOWED_DETAIL: &str = "Method not allowed";
|
||||
const AI_PUBLIC_UNAUTHORIZED_DETAIL: &str = "Unauthorized";
|
||||
@@ -51,6 +51,10 @@ const OPENAI_RERANK_TOP_N_DETAIL: &str = "Rerank request top_n must be a positiv
|
||||
const OPENAI_RERANK_CHAT_PAYLOAD_DETAIL: &str =
|
||||
"Rerank request must use query/documents, not chat messages";
|
||||
const OPENAI_RERANK_STREAM_UNSUPPORTED_DETAIL: &str = "Rerank requests do not support streaming";
|
||||
const CLAUDE_COUNT_TOKENS_BODY_REQUIRED_DETAIL: &str = "Request body is required";
|
||||
const CLAUDE_COUNT_TOKENS_INVALID_JSON_DETAIL: &str = "Invalid JSON body";
|
||||
const CLAUDE_COUNT_TOKENS_MODEL_REQUIRED_DETAIL: &str = "model: Field required";
|
||||
const CLAUDE_COUNT_TOKENS_MESSAGES_REQUIRED_DETAIL: &str = "messages: Field required";
|
||||
const ANTIGRAVITY_USER_SETTINGS_MISSING_BODY_DETAIL: &str =
|
||||
"Antigravity setUserSettings request body is required";
|
||||
const ANTIGRAVITY_USER_SETTINGS_INVALID_JSON_DETAIL: &str =
|
||||
@@ -135,7 +139,7 @@ pub(crate) async fn maybe_build_local_ai_public_response(
|
||||
}
|
||||
|
||||
if let Some(response) =
|
||||
maybe_build_local_claude_count_tokens_response(request_context, request_body)
|
||||
maybe_build_local_claude_count_tokens_validation_response(request_context, request_body)
|
||||
{
|
||||
return Some(response);
|
||||
}
|
||||
@@ -863,7 +867,7 @@ fn maybe_build_local_ai_public_route_guard_response(
|
||||
None
|
||||
}
|
||||
|
||||
fn maybe_build_local_claude_count_tokens_response(
|
||||
fn maybe_build_local_claude_count_tokens_validation_response(
|
||||
request_context: &GatewayPublicRequestContext,
|
||||
request_body: Option<&Bytes>,
|
||||
) -> Option<Response<Body>> {
|
||||
@@ -876,34 +880,45 @@ fn maybe_build_local_claude_count_tokens_response(
|
||||
return None;
|
||||
}
|
||||
|
||||
let Some(request_body) = request_body else {
|
||||
return Some(build_ai_public_error_response(
|
||||
http::StatusCode::BAD_REQUEST,
|
||||
CLAUDE_COUNT_TOKENS_MISSING_BODY_DETAIL,
|
||||
));
|
||||
};
|
||||
let validation = validate_claude_count_tokens_request(request_body);
|
||||
validation.err().map(build_claude_invalid_request_response)
|
||||
}
|
||||
|
||||
let payload = match serde_json::from_slice::<serde_json::Value>(request_body) {
|
||||
Ok(payload) => payload,
|
||||
Err(_) => {
|
||||
return Some(build_ai_public_error_response(
|
||||
http::StatusCode::BAD_REQUEST,
|
||||
CLAUDE_COUNT_TOKENS_INVALID_PAYLOAD_DETAIL,
|
||||
));
|
||||
}
|
||||
};
|
||||
fn validate_claude_count_tokens_request(request_body: Option<&Bytes>) -> Result<(), &'static str> {
|
||||
let request_body = request_body
|
||||
.filter(|body| !body.is_empty())
|
||||
.ok_or(CLAUDE_COUNT_TOKENS_BODY_REQUIRED_DETAIL)?;
|
||||
let payload = serde_json::from_slice::<Value>(request_body)
|
||||
.map_err(|_| CLAUDE_COUNT_TOKENS_INVALID_JSON_DETAIL)?;
|
||||
let object = payload
|
||||
.as_object()
|
||||
.ok_or(CLAUDE_COUNT_TOKENS_INVALID_JSON_DETAIL)?;
|
||||
|
||||
let input_tokens = match estimate_claude_count_tokens(&payload) {
|
||||
Ok(tokens) => tokens,
|
||||
Err(_) => {
|
||||
return Some(build_ai_public_error_response(
|
||||
http::StatusCode::BAD_REQUEST,
|
||||
CLAUDE_COUNT_TOKENS_INVALID_PAYLOAD_DETAIL,
|
||||
));
|
||||
}
|
||||
};
|
||||
if object
|
||||
.get("model")
|
||||
.and_then(Value::as_str)
|
||||
.map(str::trim)
|
||||
.filter(|model| !model.is_empty())
|
||||
.is_none()
|
||||
{
|
||||
return Err(CLAUDE_COUNT_TOKENS_MODEL_REQUIRED_DETAIL);
|
||||
}
|
||||
if object.get("messages").and_then(Value::as_array).is_none() {
|
||||
return Err(CLAUDE_COUNT_TOKENS_MESSAGES_REQUIRED_DETAIL);
|
||||
}
|
||||
|
||||
Some(Json(json!({ "input_tokens": input_tokens })).into_response())
|
||||
Ok(())
|
||||
}
|
||||
|
||||
fn build_claude_invalid_request_response(detail: &'static str) -> Response<Body> {
|
||||
let body = build_core_error_body_for_client_format(
|
||||
"claude:messages",
|
||||
detail,
|
||||
None,
|
||||
LocalCoreSyncErrorKind::InvalidRequest,
|
||||
)
|
||||
.expect("Claude core error format should be available");
|
||||
(http::StatusCode::BAD_REQUEST, Json(body)).into_response()
|
||||
}
|
||||
|
||||
fn maybe_build_local_antigravity_v1internal_response(
|
||||
@@ -1583,132 +1598,43 @@ fn build_ai_public_error_response(
|
||||
(status, Json(json!({ "detail": detail.into() }))).into_response()
|
||||
}
|
||||
|
||||
fn estimate_claude_count_tokens(payload: &serde_json::Value) -> Result<u64, ()> {
|
||||
let object = payload.as_object().ok_or(())?;
|
||||
let model = object
|
||||
.get("model")
|
||||
.and_then(serde_json::Value::as_str)
|
||||
.ok_or(())?;
|
||||
if model.trim().is_empty() {
|
||||
return Err(());
|
||||
}
|
||||
|
||||
let messages = object
|
||||
.get("messages")
|
||||
.and_then(serde_json::Value::as_array)
|
||||
.ok_or(())?;
|
||||
|
||||
let system_tokens = estimate_claude_system_tokens(object.get("system"))?;
|
||||
let message_tokens = estimate_claude_message_tokens(messages)?;
|
||||
Ok(system_tokens.saturating_add(message_tokens))
|
||||
}
|
||||
|
||||
fn estimate_claude_system_tokens(system: Option<&serde_json::Value>) -> Result<u64, ()> {
|
||||
let Some(system) = system else {
|
||||
return Ok(0);
|
||||
};
|
||||
|
||||
match system {
|
||||
serde_json::Value::Null => Ok(0),
|
||||
serde_json::Value::String(text) => Ok(estimate_text_tokens(text)),
|
||||
serde_json::Value::Array(blocks) => {
|
||||
let mut total = 0_u64;
|
||||
for block in blocks {
|
||||
let block = block.as_object().ok_or(())?;
|
||||
if let Some(text) = block.get("text").and_then(serde_json::Value::as_str) {
|
||||
total = total.saturating_add(estimate_text_tokens(text));
|
||||
}
|
||||
}
|
||||
Ok(total)
|
||||
}
|
||||
serde_json::Value::Object(_) => Ok(0),
|
||||
_ => Err(()),
|
||||
}
|
||||
}
|
||||
|
||||
fn estimate_claude_message_tokens(messages: &[serde_json::Value]) -> Result<u64, ()> {
|
||||
let mut total = 0_u64;
|
||||
|
||||
for message in messages {
|
||||
let message = message.as_object().ok_or(())?;
|
||||
let role = message
|
||||
.get("role")
|
||||
.and_then(serde_json::Value::as_str)
|
||||
.ok_or(())?;
|
||||
if !matches!(role, "user" | "assistant") {
|
||||
return Err(());
|
||||
}
|
||||
|
||||
total = total.saturating_add(4);
|
||||
let content = message.get("content").ok_or(())?;
|
||||
match content {
|
||||
serde_json::Value::String(text) => {
|
||||
total = total.saturating_add(estimate_text_tokens(text));
|
||||
}
|
||||
serde_json::Value::Array(items) => {
|
||||
for item in items {
|
||||
let item = item.as_object().ok_or(())?;
|
||||
if let Some(text) = item.get("text").and_then(serde_json::Value::as_str) {
|
||||
total = total.saturating_add(estimate_text_tokens(text));
|
||||
}
|
||||
}
|
||||
}
|
||||
_ => return Err(()),
|
||||
}
|
||||
}
|
||||
|
||||
Ok(total)
|
||||
}
|
||||
|
||||
fn estimate_text_tokens(text: &str) -> u64 {
|
||||
if text.is_empty() {
|
||||
return 0;
|
||||
}
|
||||
|
||||
let char_count = text.chars().count() as u64;
|
||||
std::cmp::max(1, char_count / 4)
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::{
|
||||
estimate_claude_count_tokens, parse_openai_image_validation_input, validate_openai_image_n,
|
||||
OpenAiImageOperation,
|
||||
parse_openai_image_validation_input, validate_claude_count_tokens_request,
|
||||
validate_openai_image_n, OpenAiImageOperation, CLAUDE_COUNT_TOKENS_BODY_REQUIRED_DETAIL,
|
||||
CLAUDE_COUNT_TOKENS_INVALID_JSON_DETAIL, CLAUDE_COUNT_TOKENS_MESSAGES_REQUIRED_DETAIL,
|
||||
CLAUDE_COUNT_TOKENS_MODEL_REQUIRED_DETAIL,
|
||||
};
|
||||
use axum::body::Bytes;
|
||||
use serde_json::json;
|
||||
|
||||
#[test]
|
||||
fn estimates_claude_count_tokens_from_system_and_messages() {
|
||||
let payload = json!({
|
||||
"model": "claude-sonnet-4-5",
|
||||
"system": [{"type": "text", "text": "abcdefghijklmnop"}],
|
||||
"messages": [
|
||||
{
|
||||
"role": "user",
|
||||
"content": "abcdefghijkl"
|
||||
},
|
||||
{
|
||||
"role": "assistant",
|
||||
"content": [
|
||||
{"type": "text", "text": "abcdefgh"},
|
||||
{"type": "tool_use", "name": "ignored", "input": {"city": "SF"}}
|
||||
]
|
||||
}
|
||||
]
|
||||
});
|
||||
|
||||
assert_eq!(estimate_claude_count_tokens(&payload), Ok(17));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn rejects_invalid_claude_count_tokens_payload() {
|
||||
let payload = json!({
|
||||
"model": "claude-sonnet-4-5",
|
||||
"messages": [{"role": "system", "content": "bad"}]
|
||||
});
|
||||
|
||||
assert_eq!(estimate_claude_count_tokens(&payload), Err(()));
|
||||
fn count_tokens_validation_rejects_only_structurally_invalid_requests() {
|
||||
assert_eq!(
|
||||
validate_claude_count_tokens_request(None),
|
||||
Err(CLAUDE_COUNT_TOKENS_BODY_REQUIRED_DETAIL)
|
||||
);
|
||||
assert_eq!(
|
||||
validate_claude_count_tokens_request(Some(&Bytes::from_static(b"{"))),
|
||||
Err(CLAUDE_COUNT_TOKENS_INVALID_JSON_DETAIL)
|
||||
);
|
||||
assert_eq!(
|
||||
validate_claude_count_tokens_request(Some(&Bytes::from_static(br#"{"messages":[]}"#,))),
|
||||
Err(CLAUDE_COUNT_TOKENS_MODEL_REQUIRED_DETAIL)
|
||||
);
|
||||
assert_eq!(
|
||||
validate_claude_count_tokens_request(Some(&Bytes::from_static(
|
||||
br#"{"model":"claude-sonnet-4-5"}"#,
|
||||
))),
|
||||
Err(CLAUDE_COUNT_TOKENS_MESSAGES_REQUIRED_DETAIL)
|
||||
);
|
||||
assert_eq!(
|
||||
validate_claude_count_tokens_request(Some(&Bytes::from_static(
|
||||
br#"{"model":"claude-sonnet-4-5","messages":[],"tools":[{"name":"x"}]}"#,
|
||||
))),
|
||||
Ok(())
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
|
||||
@@ -253,6 +253,7 @@ pub(super) async fn maybe_build_local_test_connection_route_response(
|
||||
upstream_is_stream: false,
|
||||
request_query: None,
|
||||
kiro_api_region: None,
|
||||
api_operation: None,
|
||||
},
|
||||
);
|
||||
let Some(upstream_url) = upstream_url else {
|
||||
|
||||
Reference in New Issue
Block a user