Merge remote-tracking branch 'origin/aether-rust-pioneer' into aether-rust-pioneer

# Conflicts:
#	crates/aether-data-contracts/src/repository/usage/mod.rs
#	crates/aether-data/src/repository/global_models/postgres.rs
#	crates/aether-data/src/repository/usage/postgres/mod.rs
This commit is contained in:
fawney19
2026-05-05 18:53:14 +08:00
107 changed files with 7399 additions and 244 deletions

View File

@@ -4,7 +4,7 @@ use aether_data_contracts::repository::{
usage::{
StoredRequestUsageAudit, StoredUsageCostSavingsSummary, StoredUsageErrorDistributionRow,
StoredUsageLeaderboardSummary, StoredUsagePerformancePercentilesRow,
StoredUsageTimeSeriesBucket,
StoredUsageProviderPerformance, StoredUsageTimeSeriesBucket,
},
};
use axum::{
@@ -584,6 +584,20 @@ pub fn round_to(value: f64, decimals: u32) -> f64 {
(value * factor).round() / factor
}
fn rounded_option(value: Option<f64>, decimals: u32) -> serde_json::Value {
value
.map(|value| json!(round_to(value, decimals)))
.unwrap_or(serde_json::Value::Null)
}
fn success_rate(request_count: u64, success_count: u64) -> f64 {
if request_count == 0 {
0.0
} else {
round_to(success_count as f64 / request_count as f64 * 100.0, 2)
}
}
pub fn admin_stats_provider_quota_usage_empty_response() -> Response<Body> {
Json(json!({
"providers": [],
@@ -652,6 +666,21 @@ pub fn admin_stats_performance_percentiles_empty_response() -> Response<Body> {
Json(json!([])).into_response()
}
pub fn admin_stats_provider_performance_empty_response() -> Response<Body> {
Json(json!({
"summary": {
"request_count": 0,
"success_rate": 0.0,
"avg_output_tps": serde_json::Value::Null,
"avg_first_byte_time_ms": serde_json::Value::Null,
"avg_response_time_ms": serde_json::Value::Null,
},
"providers": [],
"timeline": [],
}))
.into_response()
}
pub fn admin_stats_cost_savings_empty_response() -> Response<Body> {
Json(json!({
"cache_read_tokens": 0,
@@ -1061,6 +1090,64 @@ pub fn build_admin_stats_performance_percentiles_response_from_summaries(
Json(serde_json::Value::Array(payload)).into_response()
}
pub fn build_admin_stats_provider_performance_response(
performance: &StoredUsageProviderPerformance,
) -> Response<Body> {
let summary = &performance.summary;
let providers = performance
.providers
.iter()
.map(|row| {
json!({
"provider_id": row.provider_id.as_str(),
"provider": row.provider.as_str(),
"request_count": row.request_count,
"success_count": row.success_count,
"error_count": row.request_count.saturating_sub(row.success_count),
"success_rate": success_rate(row.request_count, row.success_count),
"output_tokens": row.output_tokens,
"avg_output_tps": rounded_option(row.avg_output_tps, 2),
"avg_first_byte_time_ms": rounded_option(row.avg_first_byte_time_ms, 2),
"avg_response_time_ms": rounded_option(row.avg_response_time_ms, 2),
"p90_response_time_ms": row.p90_response_time_ms,
"p90_first_byte_time_ms": row.p90_first_byte_time_ms,
"tps_sample_count": row.tps_sample_count,
"first_byte_sample_count": row.first_byte_sample_count,
})
})
.collect::<Vec<_>>();
let timeline = performance
.timeline
.iter()
.map(|row| {
json!({
"date": row.date.as_str(),
"provider_id": row.provider_id.as_str(),
"provider": row.provider.as_str(),
"request_count": row.request_count,
"output_tokens": row.output_tokens,
"avg_output_tps": rounded_option(row.avg_output_tps, 2),
"avg_first_byte_time_ms": rounded_option(row.avg_first_byte_time_ms, 2),
"avg_response_time_ms": rounded_option(row.avg_response_time_ms, 2),
"success_rate": success_rate(row.request_count, row.success_count),
})
})
.collect::<Vec<_>>();
Json(json!({
"summary": {
"request_count": summary.request_count,
"success_rate": success_rate(summary.request_count, summary.success_count),
"avg_output_tps": rounded_option(summary.avg_output_tps, 2),
"avg_first_byte_time_ms": rounded_option(summary.avg_first_byte_time_ms, 2),
"avg_response_time_ms": rounded_option(summary.avg_response_time_ms, 2),
},
"providers": providers,
"timeline": timeline,
}))
.into_response()
}
pub fn build_admin_stats_time_series_response(
time_range: &AdminStatsTimeRange,
granularity: AdminStatsGranularity,

View File

@@ -3,6 +3,13 @@ use chrono::{SecondsFormat, Utc};
use serde_json::{json, Map, Value};
use std::collections::BTreeMap;
const EMBEDDING_API_FORMATS: &[&str] = &[
"openai:embedding",
"jina:embedding",
"gemini:embedding",
"doubao:embedding",
];
fn unix_secs_to_rfc3339(unix_secs: u64) -> Option<String> {
let timestamp = i64::try_from(unix_secs).ok()?;
Some(
@@ -36,6 +43,50 @@ fn model_effective_capability(
})
}
fn value_contains_string(value: &Value, expected: &str) -> bool {
match value {
Value::String(value) => value.trim().eq_ignore_ascii_case(expected),
Value::Array(values) => values
.iter()
.any(|value| value_contains_string(value, expected)),
Value::Object(object) => object
.values()
.any(|value| value_contains_string(value, expected)),
_ => false,
}
}
fn value_has_true_key(value: &Value, key: &str) -> bool {
value
.as_object()
.and_then(|object| object.get(key))
.and_then(Value::as_bool)
.unwrap_or(false)
}
fn value_contains_embedding_metadata(value: &Value) -> bool {
value_has_true_key(value, "embedding")
|| value_contains_string(value, "embedding")
|| EMBEDDING_API_FORMATS
.iter()
.any(|api_format| value_contains_string(value, api_format))
}
fn model_effective_embedding_capability(model: &StoredAdminProviderModel) -> bool {
model
.config
.as_ref()
.is_some_and(value_contains_embedding_metadata)
|| model
.global_model_supported_capabilities
.as_ref()
.is_some_and(value_contains_embedding_metadata)
|| model
.global_model_config
.as_ref()
.is_some_and(value_contains_embedding_metadata)
}
fn merge_json_values(base: &mut Value, overlay: Value) {
match (base, overlay) {
(Value::Object(base_map), Value::Object(overlay_map)) => {
@@ -128,6 +179,7 @@ pub fn admin_provider_model_effective_capability(
model.global_model_config.as_ref(),
"image_generation",
),
"embedding" => model_effective_embedding_capability(model),
_ => false,
}
}
@@ -161,6 +213,7 @@ pub fn build_admin_provider_model_response(
"supports_streaming": model.supports_streaming,
"supports_extended_thinking": model.supports_extended_thinking,
"supports_image_generation": model.supports_image_generation,
"supports_embedding": model_effective_embedding_capability(model),
"effective_supports_vision": admin_provider_model_effective_capability(model, "vision"),
"effective_supports_function_calling": admin_provider_model_effective_capability(
model,
@@ -175,6 +228,10 @@ pub fn build_admin_provider_model_response(
model,
"image_generation",
),
"effective_supports_embedding": admin_provider_model_effective_capability(
model,
"embedding",
),
"is_active": model.is_active,
"is_available": model.is_available,
"config": model.config.clone(),
@@ -214,6 +271,7 @@ pub fn build_admin_provider_available_source_models_payload(
"supports_vision": admin_provider_model_effective_capability(&model, "vision"),
"supports_function_calling": admin_provider_model_effective_capability(&model, "function_calling"),
"supports_streaming": admin_provider_model_effective_capability(&model, "streaming"),
"supports_embedding": admin_provider_model_effective_capability(&model, "embedding"),
}),
"is_active": model.is_active,
})

View File

@@ -545,6 +545,18 @@ const ADMIN_API_FORMAT_DEFINITIONS: &[AdminApiFormatDefinition] = &[
default_path: "/v1/responses/compact",
aliases: &["responses_compact"],
},
AdminApiFormatDefinition {
value: "openai:embedding",
label: "OpenAI Embedding",
default_path: "/v1/embeddings",
aliases: &["openai_embedding", "embeddings"],
},
AdminApiFormatDefinition {
value: "openai:rerank",
label: "OpenAI Rerank",
default_path: "/v1/rerank",
aliases: &["openai_rerank", "rerank"],
},
AdminApiFormatDefinition {
value: "openai:image",
label: "OpenAI Image",
@@ -569,12 +581,36 @@ const ADMIN_API_FORMAT_DEFINITIONS: &[AdminApiFormatDefinition] = &[
default_path: "/v1beta/models/{model}:{action}",
aliases: &["gemini", "google", "vertex"],
},
AdminApiFormatDefinition {
value: "gemini:embedding",
label: "Gemini Embedding",
default_path: "/v1/embeddings",
aliases: &["gemini_embedding"],
},
AdminApiFormatDefinition {
value: "gemini:video",
label: "Gemini Video",
default_path: "/v1beta/models/{model}:predictLongRunning",
aliases: &["gemini_video", "veo"],
},
AdminApiFormatDefinition {
value: "jina:embedding",
label: "Jina Embedding",
default_path: "/v1/embeddings",
aliases: &["jina_embedding"],
},
AdminApiFormatDefinition {
value: "jina:rerank",
label: "Jina Rerank",
default_path: "/v1/rerank",
aliases: &["jina_rerank"],
},
AdminApiFormatDefinition {
value: "doubao:embedding",
label: "Doubao Embedding",
default_path: "/v1/embeddings",
aliases: &["doubao_embedding"],
},
];
pub fn build_admin_system_check_update_payload(current_version: String) -> serde_json::Value {

View File

@@ -23,9 +23,10 @@ pub use crate::contracts::{
OPENAI_CHAT_STREAM_PLAN_KIND, OPENAI_CHAT_STREAM_SUCCESS_REPORT_KIND,
OPENAI_CHAT_SYNC_ERROR_REPORT_KIND, OPENAI_CHAT_SYNC_FINALIZE_REPORT_KIND,
OPENAI_CHAT_SYNC_PLAN_KIND, OPENAI_CHAT_SYNC_SUCCESS_REPORT_KIND,
OPENAI_IMAGE_STREAM_PLAN_KIND, OPENAI_IMAGE_STREAM_SUCCESS_REPORT_KIND,
OPENAI_IMAGE_SYNC_FINALIZE_REPORT_KIND, OPENAI_IMAGE_SYNC_PLAN_KIND,
OPENAI_IMAGE_SYNC_SUCCESS_REPORT_KIND, OPENAI_RESPONSES_COMPACT_STREAM_PLAN_KIND,
OPENAI_EMBEDDING_SYNC_PLAN_KIND, OPENAI_IMAGE_STREAM_PLAN_KIND,
OPENAI_IMAGE_STREAM_SUCCESS_REPORT_KIND, OPENAI_IMAGE_SYNC_FINALIZE_REPORT_KIND,
OPENAI_IMAGE_SYNC_PLAN_KIND, OPENAI_IMAGE_SYNC_SUCCESS_REPORT_KIND,
OPENAI_RERANK_SYNC_PLAN_KIND, OPENAI_RESPONSES_COMPACT_STREAM_PLAN_KIND,
OPENAI_RESPONSES_COMPACT_STREAM_SUCCESS_REPORT_KIND,
OPENAI_RESPONSES_COMPACT_SYNC_ERROR_REPORT_KIND,
OPENAI_RESPONSES_COMPACT_SYNC_FINALIZE_REPORT_KIND, OPENAI_RESPONSES_COMPACT_SYNC_PLAN_KIND,

View File

@@ -18,12 +18,12 @@ pub use plan_kinds::{
GEMINI_FILES_DOWNLOAD_PLAN_KIND, GEMINI_FILES_GET_PLAN_KIND, GEMINI_FILES_LIST_PLAN_KIND,
GEMINI_FILES_UPLOAD_PLAN_KIND, GEMINI_VIDEO_CANCEL_SYNC_PLAN_KIND,
GEMINI_VIDEO_CREATE_SYNC_PLAN_KIND, OPENAI_CHAT_STREAM_PLAN_KIND, OPENAI_CHAT_SYNC_PLAN_KIND,
OPENAI_IMAGE_STREAM_PLAN_KIND, OPENAI_IMAGE_SYNC_PLAN_KIND,
OPENAI_RESPONSES_COMPACT_STREAM_PLAN_KIND, OPENAI_RESPONSES_COMPACT_SYNC_PLAN_KIND,
OPENAI_RESPONSES_STREAM_PLAN_KIND, OPENAI_RESPONSES_SYNC_PLAN_KIND,
OPENAI_VIDEO_CANCEL_SYNC_PLAN_KIND, OPENAI_VIDEO_CONTENT_PLAN_KIND,
OPENAI_VIDEO_CREATE_SYNC_PLAN_KIND, OPENAI_VIDEO_DELETE_SYNC_PLAN_KIND,
OPENAI_VIDEO_REMIX_SYNC_PLAN_KIND,
OPENAI_EMBEDDING_SYNC_PLAN_KIND, OPENAI_IMAGE_STREAM_PLAN_KIND, OPENAI_IMAGE_SYNC_PLAN_KIND,
OPENAI_RERANK_SYNC_PLAN_KIND, OPENAI_RESPONSES_COMPACT_STREAM_PLAN_KIND,
OPENAI_RESPONSES_COMPACT_SYNC_PLAN_KIND, OPENAI_RESPONSES_STREAM_PLAN_KIND,
OPENAI_RESPONSES_SYNC_PLAN_KIND, OPENAI_VIDEO_CANCEL_SYNC_PLAN_KIND,
OPENAI_VIDEO_CONTENT_PLAN_KIND, OPENAI_VIDEO_CREATE_SYNC_PLAN_KIND,
OPENAI_VIDEO_DELETE_SYNC_PLAN_KIND, OPENAI_VIDEO_REMIX_SYNC_PLAN_KIND,
};
pub use report_kinds::{
core_error_background_report_kind, core_error_default_client_api_format,

View File

@@ -20,6 +20,8 @@ pub const CLAUDE_CLI_STREAM_PLAN_KIND: &str = "claude_cli_stream";
pub const GEMINI_CLI_STREAM_PLAN_KIND: &str = "gemini_cli_stream";
pub const OPENAI_VIDEO_CREATE_SYNC_PLAN_KIND: &str = "openai_video_create_sync";
pub const OPENAI_CHAT_SYNC_PLAN_KIND: &str = "openai_chat_sync";
pub const OPENAI_EMBEDDING_SYNC_PLAN_KIND: &str = "openai_embedding_sync";
pub const OPENAI_RERANK_SYNC_PLAN_KIND: &str = "openai_rerank_sync";
pub const OPENAI_RESPONSES_SYNC_PLAN_KIND: &str = "openai_responses_sync";
pub const OPENAI_RESPONSES_COMPACT_SYNC_PLAN_KIND: &str = "openai_responses_compact_sync";
pub const CLAUDE_CHAT_SYNC_PLAN_KIND: &str = "claude_chat_sync";

View File

@@ -9,19 +9,21 @@ pub mod response;
pub use protocol::canonical::{
canonical_request_unknown_block_count, canonical_response_unknown_block_count,
canonical_to_claude_request, canonical_to_claude_response, canonical_to_gemini_request,
canonical_to_gemini_response, canonical_to_openai_chat_request,
canonical_to_claude_request, canonical_to_claude_response, canonical_to_embedding_response,
canonical_to_gemini_request, canonical_to_gemini_response, canonical_to_openai_chat_request,
canonical_to_openai_chat_response, canonical_to_openai_responses_compact_request,
canonical_to_openai_responses_compact_response, canonical_to_openai_responses_request,
canonical_to_openai_responses_response, canonical_unknown_block_count,
from_claude_to_canonical_request, from_claude_to_canonical_response,
from_gemini_to_canonical_request, from_gemini_to_canonical_response,
from_openai_chat_to_canonical_request, from_openai_chat_to_canonical_response,
from_openai_responses_to_canonical_request, from_openai_responses_to_canonical_response,
CanonicalContentBlock, CanonicalGenerationConfig, CanonicalInstruction, CanonicalMessage,
CanonicalRequest, CanonicalResponse, CanonicalResponseFormat, CanonicalResponseOutput,
CanonicalRole, CanonicalStopReason, CanonicalStreamEvent, CanonicalStreamFrame,
CanonicalThinkingConfig, CanonicalToolChoice, CanonicalToolDefinition, CanonicalUsage,
from_embedding_to_canonical_response, from_gemini_to_canonical_request,
from_gemini_to_canonical_response, from_openai_chat_to_canonical_request,
from_openai_chat_to_canonical_response, from_openai_responses_to_canonical_request,
from_openai_responses_to_canonical_response, CanonicalContentBlock, CanonicalEmbedding,
CanonicalEmbeddingInput, CanonicalEmbeddingRequest, CanonicalEmbeddingResponse,
CanonicalGenerationConfig, CanonicalInstruction, CanonicalMessage, CanonicalRequest,
CanonicalResponse, CanonicalResponseFormat, CanonicalResponseOutput, CanonicalRole,
CanonicalStopReason, CanonicalStreamEvent, CanonicalStreamFrame, CanonicalThinkingConfig,
CanonicalToolChoice, CanonicalToolDefinition, CanonicalUsage,
};
pub use protocol::context::{FormatContext, FormatError};
pub use protocol::formats::{
@@ -30,10 +32,11 @@ pub use protocol::formats::{
FormatFamily, FormatId, FormatProfile,
};
pub use protocol::matrix::{
request_candidate_api_format_preference, request_candidate_api_formats,
request_conversion_kind, request_conversion_requires_enable_flag,
sync_chat_response_conversion_kind, sync_cli_response_conversion_kind, RequestConversionKind,
SyncChatResponseConversionKind, SyncCliResponseConversionKind,
is_embedding_api_format, is_rerank_api_format, request_candidate_api_format_preference,
request_candidate_api_formats, request_conversion_kind,
request_conversion_requires_enable_flag, sync_chat_response_conversion_kind,
sync_cli_response_conversion_kind, RequestConversionKind, SyncChatResponseConversionKind,
SyncCliResponseConversionKind,
};
pub use protocol::registry::{build_stream_transcoder, convert_request, convert_response};
pub use request::model_directives::{

View File

@@ -222,6 +222,94 @@ pub struct CanonicalUsage {
pub extensions: BTreeMap<String, Value>,
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
#[serde(untagged)]
pub enum CanonicalEmbeddingInput {
String(String),
StringArray(Vec<String>),
TokenArray(Vec<i64>),
TokenArrayArray(Vec<Vec<i64>>),
}
impl CanonicalEmbeddingInput {
fn is_empty(&self) -> bool {
match self {
Self::String(value) => value.trim().is_empty(),
Self::StringArray(values) => {
values.is_empty() || values.iter().any(|value| value.trim().is_empty())
}
Self::TokenArray(values) => values.is_empty(),
Self::TokenArrayArray(values) => values.is_empty() || values.iter().any(Vec::is_empty),
}
}
fn as_string_items(&self) -> Option<Vec<&str>> {
match self {
Self::String(value) => Some(vec![value.as_str()]),
Self::StringArray(values) => Some(values.iter().map(String::as_str).collect()),
Self::TokenArray(_) | Self::TokenArrayArray(_) => None,
}
}
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub struct CanonicalEmbeddingRequest {
pub input: CanonicalEmbeddingInput,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub encoding_format: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub dimensions: Option<u64>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub task: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub user: Option<String>,
#[serde(default, skip_serializing_if = "BTreeMap::is_empty")]
pub extensions: BTreeMap<String, Value>,
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub struct CanonicalRerankRequest {
pub query: String,
#[serde(default)]
pub documents: Vec<Value>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub top_n: Option<u64>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub return_documents: Option<bool>,
#[serde(default, skip_serializing_if = "BTreeMap::is_empty")]
pub extensions: BTreeMap<String, Value>,
}
impl CanonicalRerankRequest {
fn is_empty(&self) -> bool {
self.query.trim().is_empty()
|| self.documents.is_empty()
|| self.documents.iter().any(rerank_document_is_empty)
}
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub struct CanonicalEmbedding {
#[serde(default)]
pub index: usize,
#[serde(default)]
pub embedding: Vec<f64>,
#[serde(default, skip_serializing_if = "BTreeMap::is_empty")]
pub extensions: BTreeMap<String, Value>,
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub struct CanonicalEmbeddingResponse {
pub id: String,
pub model: String,
#[serde(default)]
pub embeddings: Vec<CanonicalEmbedding>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub usage: Option<CanonicalUsage>,
#[serde(default, skip_serializing_if = "BTreeMap::is_empty")]
pub extensions: BTreeMap<String, Value>,
}
#[derive(Debug, Clone, Default, PartialEq, Serialize, Deserialize)]
pub struct CanonicalRequest {
#[serde(default)]
@@ -232,6 +320,10 @@ pub struct CanonicalRequest {
pub system: Option<String>,
#[serde(default)]
pub messages: Vec<CanonicalMessage>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub embedding: Option<CanonicalEmbeddingRequest>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub rerank: Option<CanonicalRerankRequest>,
#[serde(default)]
pub generation: CanonicalGenerationConfig,
#[serde(default)]
@@ -373,6 +465,47 @@ pub fn canonical_to_gemini_request(
)
}
pub(crate) fn from_embedding_to_canonical_request(
body_json: &Value,
namespace: &str,
) -> Option<CanonicalRequest> {
embedding_request_from_raw(body_json, namespace)
}
pub(crate) fn canonical_to_embedding_request(
canonical: &CanonicalRequest,
mapped_model: &str,
namespace: &str,
) -> Option<Value> {
match namespace {
"openai" => canonical_to_openai_embedding_request(canonical, mapped_model),
"jina" => canonical_to_jina_embedding_request(canonical, mapped_model),
"gemini" => canonical_to_gemini_embedding_request(canonical, mapped_model),
"doubao" => canonical_to_doubao_embedding_request(canonical, mapped_model),
_ => None,
}
}
pub(crate) fn from_rerank_to_canonical_request(
body_json: &Value,
namespace: &str,
) -> Option<CanonicalRequest> {
rerank_request_from_raw(body_json, namespace)
}
pub(crate) fn canonical_to_rerank_request(
canonical: &CanonicalRequest,
mapped_model: &str,
namespace: &str,
) -> Option<Value> {
match namespace {
"openai" | "jina" => {
canonical_to_openai_like_rerank_request(canonical, mapped_model, namespace)
}
_ => None,
}
}
pub fn from_openai_chat_to_canonical_response(body_json: &Value) -> Option<CanonicalResponse> {
crate::protocol::formats::openai_chat::response::from_raw(body_json)
}
@@ -556,6 +689,23 @@ pub fn canonical_to_gemini_response(
crate::protocol::formats::gemini_generate_content::response::to_raw(canonical, report_context)
}
pub fn from_embedding_to_canonical_response(
body_json: &Value,
namespace: &str,
) -> Option<CanonicalEmbeddingResponse> {
embedding_response_from_raw(body_json, namespace)
}
pub fn canonical_to_embedding_response(
canonical: &CanonicalEmbeddingResponse,
namespace: &str,
) -> Option<Value> {
match namespace {
"openai" | "jina" => Some(canonical_to_openai_embedding_response(canonical, namespace)),
_ => None,
}
}
pub fn canonical_unknown_block_count(blocks: &[CanonicalContentBlock]) -> usize {
blocks
.iter()
@@ -4123,6 +4273,411 @@ pub(crate) fn strip_claude_billing_header(text: &str) -> String {
remainder.trim_start_matches('\n').trim().to_string()
}
fn embedding_request_from_raw(body_json: &Value, namespace: &str) -> Option<CanonicalRequest> {
let request = body_json.as_object()?;
let model = request
.get("model")
.and_then(Value::as_str)
.map(str::trim)
.filter(|value| !value.is_empty())?
.to_string();
let input =
serde_json::from_value::<CanonicalEmbeddingInput>(request.get("input")?.clone()).ok()?;
if input.is_empty() {
return None;
}
let embedding = CanonicalEmbeddingRequest {
input,
encoding_format: request
.get("encoding_format")
.and_then(Value::as_str)
.map(ToOwned::to_owned),
dimensions: request.get("dimensions").and_then(Value::as_u64),
task: request
.get("task")
.and_then(Value::as_str)
.map(ToOwned::to_owned),
user: request
.get("user")
.and_then(Value::as_str)
.map(ToOwned::to_owned),
extensions: namespace_extensions(
namespace,
request,
&[
"model",
"input",
"encoding_format",
"dimensions",
"task",
"user",
],
),
};
Some(CanonicalRequest {
model,
embedding: Some(embedding),
..CanonicalRequest::default()
})
}
fn rerank_request_from_raw(body_json: &Value, namespace: &str) -> Option<CanonicalRequest> {
let request = body_json.as_object()?;
let model = request
.get("model")
.and_then(Value::as_str)
.map(str::trim)
.filter(|value| !value.is_empty())?
.to_string();
let query = request
.get("query")
.and_then(Value::as_str)
.map(str::trim)
.filter(|value| !value.is_empty())?
.to_string();
let documents = request.get("documents").and_then(Value::as_array)?.to_vec();
let rerank = CanonicalRerankRequest {
query,
documents,
top_n: request
.get("top_n")
.or_else(|| request.get("topN"))
.and_then(Value::as_u64),
return_documents: request
.get("return_documents")
.or_else(|| request.get("returnDocuments"))
.and_then(Value::as_bool),
extensions: namespace_extensions(
namespace,
request,
&[
"model",
"query",
"documents",
"top_n",
"topN",
"return_documents",
"returnDocuments",
],
),
};
if rerank.is_empty() || rerank.top_n == Some(0) {
return None;
}
Some(CanonicalRequest {
model,
rerank: Some(rerank),
..CanonicalRequest::default()
})
}
fn canonical_to_openai_like_rerank_request(
canonical: &CanonicalRequest,
mapped_model: &str,
namespace: &str,
) -> Option<Value> {
let rerank = canonical.rerank.as_ref()?;
if rerank.is_empty() || rerank.top_n == Some(0) {
return None;
}
let mut output = Map::new();
output.insert(
"model".to_string(),
Value::String(mapped_rerank_model(canonical, mapped_model)),
);
output.insert("query".to_string(), Value::String(rerank.query.clone()));
output.insert(
"documents".to_string(),
Value::Array(rerank.documents.clone()),
);
if let Some(value) = rerank.top_n {
output.insert("top_n".to_string(), Value::from(value));
}
if let Some(value) = rerank.return_documents {
output.insert("return_documents".to_string(), Value::Bool(value));
}
output.extend(namespace_extension_object(
&rerank.extensions,
namespace,
&output,
));
Some(Value::Object(output))
}
fn rerank_document_is_empty(value: &Value) -> bool {
match value {
Value::String(text) => text.trim().is_empty(),
Value::Object(object) => object
.get("text")
.and_then(Value::as_str)
.is_some_and(|text| text.trim().is_empty()),
Value::Null => true,
_ => false,
}
}
fn mapped_rerank_model(canonical: &CanonicalRequest, mapped_model: &str) -> String {
mapped_model
.trim()
.chars()
.next()
.map(|_| mapped_model.trim().to_string())
.unwrap_or_else(|| canonical.model.clone())
}
fn canonical_to_openai_embedding_request(
canonical: &CanonicalRequest,
mapped_model: &str,
) -> Option<Value> {
canonical_to_openai_like_embedding_request(canonical, mapped_model, "openai", false)
}
fn canonical_to_jina_embedding_request(
canonical: &CanonicalRequest,
mapped_model: &str,
) -> Option<Value> {
canonical_to_openai_like_embedding_request(canonical, mapped_model, "jina", true)
}
fn canonical_to_openai_like_embedding_request(
canonical: &CanonicalRequest,
mapped_model: &str,
namespace: &str,
default_task: bool,
) -> Option<Value> {
let embedding = canonical.embedding.as_ref()?;
if embedding.input.is_empty() {
return None;
}
let mut output = Map::new();
output.insert(
"model".to_string(),
Value::String(mapped_embedding_model(canonical, mapped_model)),
);
output.insert(
"input".to_string(),
serde_json::to_value(&embedding.input).ok()?,
);
if let Some(value) = &embedding.encoding_format {
output.insert("encoding_format".to_string(), Value::String(value.clone()));
}
if let Some(value) = embedding.dimensions {
output.insert("dimensions".to_string(), Value::from(value));
}
if let Some(value) = &embedding.user {
output.insert("user".to_string(), Value::String(value.clone()));
}
if let Some(task) = embedding
.task
.as_ref()
.filter(|value| !value.trim().is_empty())
{
output.insert("task".to_string(), Value::String(task.clone()));
} else if default_task {
output.insert(
"task".to_string(),
Value::String("text-matching".to_string()),
);
}
output.extend(namespace_extension_object(
&embedding.extensions,
namespace,
&output,
));
Some(Value::Object(output))
}
fn canonical_to_gemini_embedding_request(
canonical: &CanonicalRequest,
mapped_model: &str,
) -> Option<Value> {
let embedding = canonical.embedding.as_ref()?;
let items = embedding.input.as_string_items()?;
if items.is_empty() || items.iter().any(|value| value.trim().is_empty()) {
return None;
}
let model = mapped_embedding_model(canonical, mapped_model);
if items.len() == 1 {
return Some(json!({
"model": model,
"content": {
"parts": [{"text": items[0]}]
}
}));
}
Some(json!({
"model": model,
"requests": items.into_iter().map(|text| {
json!({
"model": model,
"content": {
"parts": [{"text": text}]
}
})
}).collect::<Vec<_>>()
}))
}
fn canonical_to_doubao_embedding_request(
canonical: &CanonicalRequest,
mapped_model: &str,
) -> Option<Value> {
let embedding = canonical.embedding.as_ref()?;
let items = embedding.input.as_string_items()?;
if items.is_empty() || items.iter().any(|value| value.trim().is_empty()) {
return None;
}
let mut output = Map::new();
output.insert(
"model".to_string(),
Value::String(mapped_embedding_model(canonical, mapped_model)),
);
output.insert(
"input".to_string(),
Value::Array(
items
.into_iter()
.map(|text| json!({"type": "text", "text": text}))
.collect(),
),
);
if let Some(dimensions) = embedding.dimensions {
output.insert("dimensions".to_string(), Value::from(dimensions));
}
output.extend(namespace_extension_object(
&embedding.extensions,
"doubao",
&output,
));
Some(Value::Object(output))
}
fn embedding_response_from_raw(
body_json: &Value,
namespace: &str,
) -> Option<CanonicalEmbeddingResponse> {
let body = body_json.as_object()?;
if body.contains_key("error") {
return None;
}
let data = body.get("data")?.as_array()?;
let mut embeddings = Vec::new();
for (fallback_index, item) in data.iter().enumerate() {
let item_object = item.as_object()?;
let values = item_object.get("embedding")?.as_array()?;
let embedding = values
.iter()
.map(Value::as_f64)
.collect::<Option<Vec<_>>>()?;
embeddings.push(CanonicalEmbedding {
index: item_object
.get("index")
.and_then(Value::as_u64)
.and_then(|value| usize::try_from(value).ok())
.unwrap_or(fallback_index),
embedding,
extensions: namespace_extensions(
namespace,
item_object,
&["object", "index", "embedding"],
),
});
}
Some(CanonicalEmbeddingResponse {
id: body
.get("id")
.and_then(Value::as_str)
.unwrap_or("embd-unknown")
.to_string(),
model: body
.get("model")
.and_then(Value::as_str)
.unwrap_or("unknown")
.to_string(),
embeddings,
usage: openai_usage_to_canonical(body.get("usage")),
extensions: namespace_extensions(
namespace,
body,
&["id", "object", "model", "data", "usage"],
),
})
}
fn canonical_to_openai_embedding_response(
canonical: &CanonicalEmbeddingResponse,
namespace: &str,
) -> Value {
let mut response = Map::new();
response.insert("object".to_string(), Value::String("list".to_string()));
if !canonical.model.trim().is_empty() && canonical.model != "unknown" {
response.insert("model".to_string(), Value::String(canonical.model.clone()));
}
response.insert(
"data".to_string(),
Value::Array(
canonical
.embeddings
.iter()
.map(|embedding| {
let mut item = Map::new();
item.insert("object".to_string(), Value::String("embedding".to_string()));
item.insert("index".to_string(), Value::from(embedding.index as u64));
item.insert("embedding".to_string(), json!(embedding.embedding));
item.extend(namespace_extension_object(
&embedding.extensions,
namespace,
&item,
));
Value::Object(item)
})
.collect(),
),
);
if let Some(usage) = &canonical.usage {
response.insert("usage".to_string(), canonical_usage_to_openai(usage));
}
response.extend(namespace_extension_object(
&canonical.extensions,
namespace,
&response,
));
Value::Object(response)
}
fn mapped_embedding_model(canonical: &CanonicalRequest, mapped_model: &str) -> String {
let mapped_model = mapped_model.trim();
if mapped_model.is_empty() {
canonical.model.clone()
} else {
mapped_model.to_string()
}
}
fn namespace_extensions(
namespace: &str,
object: &Map<String, Value>,
handled_keys: &[&str],
) -> BTreeMap<String, Value> {
let handled = handled_keys
.iter()
.copied()
.collect::<std::collections::BTreeSet<_>>();
let raw = object
.iter()
.filter(|(key, _)| !handled.contains(key.as_str()))
.map(|(key, value)| (key.clone(), value.clone()))
.collect::<Map<String, Value>>();
if raw.is_empty() {
BTreeMap::new()
} else {
BTreeMap::from([(namespace.to_string(), Value::Object(raw))])
}
}
#[cfg(test)]
mod tests {
use super::{
@@ -4135,10 +4690,328 @@ mod tests {
from_gemini_to_canonical_request, from_gemini_to_canonical_response,
from_openai_chat_to_canonical_request, from_openai_chat_to_canonical_response,
from_openai_responses_to_canonical_request, from_openai_responses_to_canonical_response,
CanonicalContentBlock, CanonicalRole,
CanonicalContentBlock, CanonicalEmbedding, CanonicalEmbeddingInput,
CanonicalEmbeddingRequest, CanonicalRole, CanonicalUsage,
};
use serde_json::{json, Value};
#[test]
fn canonical_embedding_request_accepts_axonhub_input_shapes() {
for input in [
json!("hello"),
json!(["hello", "world"]),
json!([1, 2, 3]),
json!([[1, 2], [3, 4]]),
] {
let request = json!({
"model": "text-embedding-3-small",
"embedding": {
"input": input,
"encoding_format": "float",
"dimensions": 3
}
});
let canonical = serde_json::from_value::<super::CanonicalRequest>(request)
.expect("embedding request should deserialize");
assert!(canonical.embedding.is_some());
assert!(canonical.messages.is_empty());
let encoded = serde_json::to_value(&canonical).expect("serialize");
assert!(encoded.get("messages").is_some());
assert!(encoded.get("embedding").is_some());
}
}
#[test]
fn embedding_wire_request_accepts_all_axonhub_input_shapes() {
let cases = [
(
json!("hello"),
"single string",
CanonicalEmbeddingInput::String("hello".to_string()),
),
(
json!(["hello", "world"]),
"string array",
CanonicalEmbeddingInput::StringArray(vec![
"hello".to_string(),
"world".to_string(),
]),
),
(
json!([1, 2, 3]),
"token array",
CanonicalEmbeddingInput::TokenArray(vec![1, 2, 3]),
),
(
json!([[1, 2], [3, 4]]),
"nested token array",
CanonicalEmbeddingInput::TokenArrayArray(vec![vec![1, 2], vec![3, 4]]),
),
];
for (input, label, expected_input) in cases {
let body = json!({
"model": "text-embedding-3-small",
"input": input
});
let canonical = super::from_embedding_to_canonical_request(&body, "openai")
.unwrap_or_else(|| panic!("{label} should parse"));
assert_eq!(
canonical.embedding.expect("embedding request").input,
expected_input,
"{label} should preserve its canonical variant"
);
assert!(canonical.messages.is_empty());
}
}
#[test]
fn embedding_wire_request_rejects_empty_invalid_or_chat_payloads() {
for body in [
json!({"model": "text-embedding-3-small", "input": " "}),
json!({"model": "text-embedding-3-small", "input": []}),
json!({"model": "text-embedding-3-small", "input": [1, "two"]}),
json!({"model": "text-embedding-3-small", "input": [[1], []]}),
json!({"model": "", "input": "hello"}),
json!({"input": "hello"}),
json!({"model": "text-embedding-3-small", "messages": []}),
] {
assert!(
super::from_embedding_to_canonical_request(&body, "openai").is_none(),
"invalid embedding payload should be rejected: {body}"
);
}
}
#[test]
fn embedding_openai_request_response_roundtrip_stays_non_chat() {
let body = json!({
"model": "text-embedding-3-small",
"input": ["hello", "world"],
"encoding_format": "float",
"dimensions": 2,
"user": "user-1",
"extra": true
});
let canonical =
super::from_embedding_to_canonical_request(&body, "openai").expect("embedding request");
assert_eq!(canonical.model, "text-embedding-3-small");
assert!(canonical.messages.is_empty());
assert!(matches!(
canonical.embedding.as_ref().map(|embedding| &embedding.input),
Some(CanonicalEmbeddingInput::StringArray(values)) if values == &vec!["hello".to_string(), "world".to_string()]
));
let rebuilt =
super::canonical_to_embedding_request(&canonical, "upstream-embedding", "openai")
.expect("openai embedding request");
assert_eq!(rebuilt["model"], "upstream-embedding");
assert_eq!(rebuilt["input"], json!(["hello", "world"]));
assert!(rebuilt.get("messages").is_none());
let response = json!({
"object": "list",
"model": "upstream-embedding",
"data": [
{"object": "embedding", "index": 0, "embedding": [0.1, 0.2]},
{"object": "embedding", "index": 1, "embedding": [0.3, 0.4]}
],
"usage": {"prompt_tokens": 4, "total_tokens": 4}
});
let canonical_response = super::from_embedding_to_canonical_response(&response, "openai")
.expect("embedding response");
assert_eq!(canonical_response.embeddings.len(), 2);
let emitted = super::canonical_to_embedding_response(&canonical_response, "openai")
.expect("embedding response output");
assert_eq!(emitted["data"][0]["embedding"], json!([0.1, 0.2]));
assert!(emitted.get("choices").is_none());
assert!(emitted.get("messages").is_none());
}
#[test]
fn embedding_provider_request_emitters_preserve_provider_contracts() {
let canonical = super::CanonicalRequest {
model: "text-embedding-3-small".to_string(),
embedding: Some(CanonicalEmbeddingRequest {
input: CanonicalEmbeddingInput::StringArray(vec![
"alpha".to_string(),
"beta".to_string(),
]),
encoding_format: Some("float".to_string()),
dimensions: Some(2),
task: None,
user: None,
extensions: Default::default(),
}),
..Default::default()
};
let jina = super::canonical_to_embedding_request(&canonical, "jina-embeddings-v3", "jina")
.expect("jina embedding request");
assert_eq!(jina["task"], "text-matching");
assert_eq!(jina["input"], json!(["alpha", "beta"]));
let gemini =
super::canonical_to_embedding_request(&canonical, "gemini-embedding-001", "gemini")
.expect("gemini embedding request");
assert_eq!(
gemini["requests"][0]["content"]["parts"][0]["text"],
"alpha"
);
assert!(gemini.get("messages").is_none());
let doubao =
super::canonical_to_embedding_request(&canonical, "doubao-embedding-vision", "doubao")
.expect("doubao embedding request");
assert_eq!(doubao["input"][0], json!({"type": "text", "text": "alpha"}));
assert!(doubao.get("messages").is_none());
}
#[test]
fn embedding_provider_request_emitters_cover_golden_payload_variants() {
let single = super::CanonicalRequest {
model: "text-embedding-3-small".to_string(),
embedding: Some(CanonicalEmbeddingRequest {
input: CanonicalEmbeddingInput::String("alpha".to_string()),
encoding_format: Some("float".to_string()),
dimensions: Some(1536),
task: Some("retrieval.passage".to_string()),
user: Some("user-1".to_string()),
extensions: Default::default(),
}),
..Default::default()
};
let openai =
super::canonical_to_embedding_request(&single, "text-embedding-3-large", "openai")
.expect("openai embedding request");
assert_eq!(openai["model"], "text-embedding-3-large");
assert_eq!(openai["input"], "alpha");
assert_eq!(openai["encoding_format"], "float");
assert_eq!(openai["dimensions"], 1536);
assert_eq!(openai["user"], "user-1");
assert_eq!(openai["task"], "retrieval.passage");
let jina = super::canonical_to_embedding_request(&single, "jina-embeddings-v3", "jina")
.expect("jina embedding request");
assert_eq!(jina["task"], "retrieval.passage");
assert_eq!(jina["input"], "alpha");
let gemini =
super::canonical_to_embedding_request(&single, "gemini-embedding-001", "gemini")
.expect("gemini single embedding request");
assert_eq!(gemini["model"], "gemini-embedding-001");
assert_eq!(gemini["content"]["parts"][0]["text"], "alpha");
assert!(gemini.get("requests").is_none());
let doubao =
super::canonical_to_embedding_request(&single, "doubao-embedding-vision", "doubao")
.expect("doubao embedding request");
assert_eq!(doubao["model"], "doubao-embedding-vision");
assert_eq!(doubao["input"], json!([{"type": "text", "text": "alpha"}]));
assert_eq!(doubao["dimensions"], 1536);
}
#[test]
fn gemini_and_doubao_embedding_emitters_reject_token_inputs() {
for input in [
CanonicalEmbeddingInput::TokenArray(vec![1, 2, 3]),
CanonicalEmbeddingInput::TokenArrayArray(vec![vec![1, 2], vec![3, 4]]),
] {
let canonical = super::CanonicalRequest {
model: "token-model".to_string(),
embedding: Some(CanonicalEmbeddingRequest {
input,
encoding_format: None,
dimensions: None,
task: None,
user: None,
extensions: Default::default(),
}),
..Default::default()
};
assert!(super::canonical_to_embedding_request(
&canonical,
"gemini-embedding-001",
"gemini"
)
.is_none());
assert!(super::canonical_to_embedding_request(
&canonical,
"doubao-embedding",
"doubao"
)
.is_none());
}
}
#[test]
fn embedding_response_parser_rejects_error_and_malformed_vectors() {
for body in [
json!({"error": {"message": "bad"}}),
json!({"object": "list"}),
json!({"data": [{"object": "embedding", "embedding": [0.1, "bad"]}]}),
json!({"data": [{"object": "embedding"}]}),
] {
assert!(
super::from_embedding_to_canonical_response(&body, "openai").is_none(),
"malformed embedding response should be rejected: {body}"
);
}
}
#[test]
fn embedding_response_parser_uses_openai_fallback_fields() {
let body = json!({
"object": "list",
"data": [
{"object": "embedding", "embedding": [0.1, 0.2]},
{"object": "embedding", "index": 7, "embedding": [0.3, 0.4]}
]
});
let canonical = super::from_embedding_to_canonical_response(&body, "openai")
.expect("fallback embedding response");
assert_eq!(canonical.id, "embd-unknown");
assert_eq!(canonical.model, "unknown");
assert_eq!(canonical.embeddings[0].index, 0);
assert_eq!(canonical.embeddings[1].index, 7);
assert!(super::canonical_to_embedding_response(&canonical, "jina").is_some());
assert!(super::canonical_to_embedding_response(&canonical, "gemini").is_none());
assert!(super::canonical_to_embedding_response(&canonical, "doubao").is_none());
}
#[test]
fn embedding_response_contract_serializes_vectors_without_chat_outputs() {
let response = super::CanonicalEmbeddingResponse {
id: "embd-1".to_string(),
model: "text-embedding-3-small".to_string(),
embeddings: vec![CanonicalEmbedding {
index: 0,
embedding: vec![0.1, 0.2, 0.3],
extensions: Default::default(),
}],
usage: Some(CanonicalUsage {
input_tokens: 3,
total_tokens: 3,
..Default::default()
}),
extensions: Default::default(),
};
let encoded = serde_json::to_value(&response).expect("serialize");
assert_eq!(
encoded["embeddings"][0]["embedding"],
json!([0.1, 0.2, 0.3])
);
assert!(encoded.get("choices").is_none());
let decoded = serde_json::from_value::<super::CanonicalEmbeddingResponse>(encoded)
.expect("deserialize");
assert_eq!(decoded, response);
}
#[test]
fn canonical_request_preserves_openai_multimodal_tools_and_extensions() {
let request = json!({

View File

@@ -16,6 +16,8 @@ pub enum FormatFamily {
OpenAi,
Claude,
Gemini,
Jina,
Doubao,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord, Hash)]
@@ -29,8 +31,14 @@ pub enum FormatId {
OpenAiChat,
OpenAiResponses,
OpenAiResponsesCompact,
OpenAiEmbedding,
OpenAiRerank,
ClaudeMessages,
GeminiGenerateContent,
GeminiEmbedding,
JinaEmbedding,
JinaRerank,
DoubaoEmbedding,
}
impl FormatId {
@@ -44,11 +52,15 @@ impl FormatId {
pub fn family(self) -> FormatFamily {
match self {
Self::OpenAiChat | Self::OpenAiResponses | Self::OpenAiResponsesCompact => {
FormatFamily::OpenAi
}
Self::OpenAiChat
| Self::OpenAiResponses
| Self::OpenAiResponsesCompact
| Self::OpenAiEmbedding
| Self::OpenAiRerank => FormatFamily::OpenAi,
Self::ClaudeMessages => FormatFamily::Claude,
Self::GeminiGenerateContent => FormatFamily::Gemini,
Self::GeminiGenerateContent | Self::GeminiEmbedding => FormatFamily::Gemini,
Self::JinaEmbedding | Self::JinaRerank => FormatFamily::Jina,
Self::DoubaoEmbedding => FormatFamily::Doubao,
}
}
@@ -64,8 +76,14 @@ impl FormatId {
Self::OpenAiChat => "openai:chat",
Self::OpenAiResponses => "openai:responses",
Self::OpenAiResponsesCompact => "openai:responses:compact",
Self::OpenAiEmbedding => "openai:embedding",
Self::OpenAiRerank => "openai:rerank",
Self::ClaudeMessages => "claude:messages",
Self::GeminiGenerateContent => "gemini:generate_content",
Self::GeminiEmbedding => "gemini:embedding",
Self::JinaEmbedding => "jina:embedding",
Self::JinaRerank => "jina:rerank",
Self::DoubaoEmbedding => "doubao:embedding",
}
}
}
@@ -86,8 +104,14 @@ impl FromStr for FormatId {
"openai:responses:compact" | "/v1/responses/compact" => {
Ok(Self::OpenAiResponsesCompact)
}
"openai:embedding" | "/v1/embeddings" => Ok(Self::OpenAiEmbedding),
"openai:rerank" | "/v1/rerank" => Ok(Self::OpenAiRerank),
"claude:messages" | "/v1/messages" => Ok(Self::ClaudeMessages),
"gemini:generate_content" => Ok(Self::GeminiGenerateContent),
"gemini:embedding" => Ok(Self::GeminiEmbedding),
"jina:embedding" | "/jina/v1/embeddings" => Ok(Self::JinaEmbedding),
"jina:rerank" | "/jina/v1/rerank" => Ok(Self::JinaRerank),
"doubao:embedding" => Ok(Self::DoubaoEmbedding),
_ => Err(()),
}
}
@@ -136,6 +160,90 @@ mod tests {
assert_eq!(FormatId::parse("gemini:cli"), None);
}
#[test]
fn parses_embedding_api_formats() {
assert_eq!(
FormatId::parse("openai:embedding"),
Some(FormatId::OpenAiEmbedding)
);
assert_eq!(
FormatId::parse("/v1/embeddings"),
Some(FormatId::OpenAiEmbedding)
);
assert_eq!(
FormatId::parse("gemini:embedding"),
Some(FormatId::GeminiEmbedding)
);
assert_eq!(
FormatId::parse("jina:embedding"),
Some(FormatId::JinaEmbedding)
);
assert_eq!(
FormatId::parse("/jina/v1/embeddings"),
Some(FormatId::JinaEmbedding)
);
assert_eq!(
FormatId::parse("doubao:embedding"),
Some(FormatId::DoubaoEmbedding)
);
assert_eq!(FormatId::OpenAiEmbedding.to_string(), "openai:embedding");
}
#[test]
fn embedding_format_ids_keep_provider_family_and_default_profile() {
use super::{FormatFamily, FormatProfile};
for (format, family) in [
(FormatId::OpenAiEmbedding, FormatFamily::OpenAi),
(FormatId::GeminiEmbedding, FormatFamily::Gemini),
(FormatId::JinaEmbedding, FormatFamily::Jina),
(FormatId::DoubaoEmbedding, FormatFamily::Doubao),
] {
assert_eq!(format.family(), family);
assert_eq!(format.profile(), FormatProfile::Default);
assert_eq!(FormatId::parse(format.as_str()), Some(format));
assert_eq!(format.to_string(), format.as_str());
}
}
#[test]
fn parses_rerank_api_formats() {
assert_eq!(
FormatId::parse("openai:rerank"),
Some(FormatId::OpenAiRerank)
);
assert_eq!(FormatId::parse("/v1/rerank"), Some(FormatId::OpenAiRerank));
assert_eq!(FormatId::parse("jina:rerank"), Some(FormatId::JinaRerank));
assert_eq!(
FormatId::parse("/jina/v1/rerank"),
Some(FormatId::JinaRerank)
);
assert_eq!(FormatId::OpenAiRerank.to_string(), "openai:rerank");
}
#[test]
fn rerank_format_ids_keep_provider_family_and_default_profile() {
use super::{FormatFamily, FormatProfile};
for (format, family) in [
(FormatId::OpenAiRerank, FormatFamily::OpenAi),
(FormatId::JinaRerank, FormatFamily::Jina),
] {
assert_eq!(format.family(), family);
assert_eq!(format.profile(), FormatProfile::Default);
assert_eq!(FormatId::parse(format.as_str()), Some(format));
assert_eq!(format.to_string(), format.as_str());
}
}
#[test]
fn rejects_unknown_embedding_format() {
assert_eq!(FormatId::parse("embedding"), None);
assert_eq!(FormatId::parse("openai:embeddings"), None);
assert_eq!(FormatId::parse("claude:embedding"), None);
assert_eq!(FormatId::parse("gemini:embed_content"), None);
}
#[test]
fn normalizes_api_format_aliases() {
assert_eq!(
@@ -154,6 +262,10 @@ mod tests {
normalize_api_format_alias("GEMINI:GENERATE_CONTENT"),
"gemini:generate_content"
);
assert_eq!(
normalize_api_format_alias("OPENAI:EMBEDDING"),
"openai:embedding"
);
assert_eq!(normalize_api_format_alias("openai:image"), "openai:image");
assert_eq!(normalize_api_format_alias("openai:video"), "openai:video");
assert_eq!(normalize_api_format_alias("gemini:video"), "gemini:video");
@@ -184,5 +296,21 @@ mod tests {
api_format_storage_aliases("gemini:generate_content"),
vec!["gemini:generate_content".to_string()]
);
assert_eq!(
api_format_storage_aliases("openai:embedding"),
vec!["openai:embedding".to_string()]
);
assert_eq!(
api_format_storage_aliases("gemini:embedding"),
vec!["gemini:embedding".to_string()]
);
assert_eq!(
api_format_storage_aliases("jina:embedding"),
vec!["jina:embedding".to_string()]
);
assert_eq!(
api_format_storage_aliases("doubao:embedding"),
vec!["doubao:embedding".to_string()]
);
}
}

View File

@@ -37,6 +37,13 @@ const STANDARD_API_FORMAT_ORDER: &[&str] = &[
"claude:messages",
"gemini:generate_content",
];
const EMBEDDING_CANDIDATE_API_FORMATS: &[&str] = &[
"openai:embedding",
"gemini:embedding",
"jina:embedding",
"doubao:embedding",
];
const RERANK_CANDIDATE_API_FORMATS: &[&str] = &["openai:rerank", "jina:rerank"];
pub fn request_candidate_api_format_preference(
client_api_format: &str,
@@ -48,6 +55,26 @@ pub fn request_candidate_api_format_preference(
if client_api_format == "openai:responses:compact" {
return (provider_api_format == "openai:responses:compact").then_some((0, 0));
}
if is_embedding_api_format(client_api_format.as_str()) {
return is_embedding_api_format(provider_api_format.as_str()).then_some((
if client_api_format == provider_api_format {
0
} else {
1
},
embedding_api_format_priority(provider_api_format.as_str()),
));
}
if is_rerank_api_format(client_api_format.as_str()) {
return is_rerank_api_format(provider_api_format.as_str()).then_some((
if client_api_format == provider_api_format {
0
} else {
1
},
rerank_api_format_priority(provider_api_format.as_str()),
));
}
let (client_family, client_kind) =
parse_non_compact_standard_api_format(client_api_format.as_str())?;
@@ -77,6 +104,22 @@ pub fn request_candidate_api_formats(
if client_api_format == "openai:responses:compact" {
return vec!["openai:responses:compact"];
}
if is_embedding_api_format(client_api_format.as_str()) {
let mut candidate_api_formats = EMBEDDING_CANDIDATE_API_FORMATS.to_vec();
candidate_api_formats.sort_by_key(|provider_api_format| {
request_candidate_api_format_preference(client_api_format.as_str(), provider_api_format)
.unwrap_or((u8::MAX, u8::MAX))
});
return candidate_api_formats;
}
if is_rerank_api_format(client_api_format.as_str()) {
let mut candidate_api_formats = RERANK_CANDIDATE_API_FORMATS.to_vec();
candidate_api_formats.sort_by_key(|provider_api_format| {
request_candidate_api_format_preference(client_api_format.as_str(), provider_api_format)
.unwrap_or((u8::MAX, u8::MAX))
});
return candidate_api_formats;
}
if parse_non_compact_standard_api_format(client_api_format.as_str()).is_none() {
return Vec::new();
}
@@ -192,6 +235,20 @@ pub fn is_standard_api_format(api_format: &str) -> bool {
)
}
pub fn is_embedding_api_format(api_format: &str) -> bool {
matches!(
normalize_api_format_alias(api_format).as_str(),
"openai:embedding" | "gemini:embedding" | "jina:embedding" | "doubao:embedding"
)
}
pub fn is_rerank_api_format(api_format: &str) -> bool {
matches!(
normalize_api_format_alias(api_format).as_str(),
"openai:rerank" | "jina:rerank"
)
}
pub fn parse_non_compact_standard_api_format(
api_format: &str,
) -> Option<(&'static str, &'static str)> {
@@ -210,6 +267,10 @@ pub fn api_data_format_id(api_format: &str) -> Option<&'static str> {
"gemini:generate_content" => Some("gemini"),
"openai:chat" => Some("openai_chat"),
"openai:responses" | "openai:responses:compact" => Some("openai_responses"),
"openai:embedding" | "gemini:embedding" | "jina:embedding" | "doubao:embedding" => {
Some("embedding")
}
"openai:rerank" | "jina:rerank" => Some("rerank"),
_ => None,
}
}
@@ -226,9 +287,26 @@ fn standard_api_format_priority(api_format: &str) -> u8 {
.unwrap_or(STANDARD_API_FORMAT_ORDER.len()) as u8
}
fn embedding_api_format_priority(api_format: &str) -> u8 {
let api_format = normalize_api_format_alias(api_format);
EMBEDDING_CANDIDATE_API_FORMATS
.iter()
.position(|candidate| *candidate == api_format)
.unwrap_or(EMBEDDING_CANDIDATE_API_FORMATS.len()) as u8
}
fn rerank_api_format_priority(api_format: &str) -> u8 {
let api_format = normalize_api_format_alias(api_format);
RERANK_CANDIDATE_API_FORMATS
.iter()
.position(|candidate| *candidate == api_format)
.unwrap_or(RERANK_CANDIDATE_API_FORMATS.len()) as u8
}
#[cfg(test)]
mod tests {
use super::{
api_data_format_id, is_embedding_api_format, is_rerank_api_format,
request_candidate_api_format_preference, request_candidate_api_formats,
request_conversion_kind, request_conversion_requires_enable_flag,
sync_chat_response_conversion_kind, sync_cli_response_conversion_kind,
@@ -355,6 +433,162 @@ mod tests {
);
}
#[test]
fn embedding_candidate_registry_excludes_chat_generation_formats() {
assert_eq!(
request_candidate_api_formats("openai:embedding", false),
vec![
"openai:embedding",
"gemini:embedding",
"jina:embedding",
"doubao:embedding",
]
);
assert_eq!(
request_candidate_api_formats("jina:embedding", false),
vec![
"jina:embedding",
"openai:embedding",
"gemini:embedding",
"doubao:embedding",
]
);
assert!(!request_candidate_api_formats("openai:embedding", false).contains(&"openai:chat"));
assert!(!request_candidate_api_formats("openai:embedding", false)
.contains(&"gemini:generate_content"));
assert_eq!(
request_conversion_kind("openai:embedding", "jina:embedding"),
None
);
assert_eq!(
request_conversion_kind("openai:embedding", "openai:chat"),
None
);
assert!(!request_conversion_requires_enable_flag(
"openai:embedding",
"jina:embedding"
));
}
#[test]
fn embedding_candidate_registry_covers_all_provider_orderings() {
assert_eq!(
request_candidate_api_formats("gemini:embedding", true),
vec![
"gemini:embedding",
"openai:embedding",
"jina:embedding",
"doubao:embedding",
]
);
assert_eq!(
request_candidate_api_formats("doubao:embedding", false),
vec![
"doubao:embedding",
"openai:embedding",
"gemini:embedding",
"jina:embedding",
]
);
let embedding_formats = [
"openai:embedding",
"gemini:embedding",
"jina:embedding",
"doubao:embedding",
];
for client_api_format in embedding_formats {
for provider_api_format in embedding_formats {
assert!(
request_candidate_api_format_preference(client_api_format, provider_api_format)
.is_some(),
"{client_api_format} should consider {provider_api_format} as embedding candidate"
);
assert_eq!(
request_conversion_kind(client_api_format, provider_api_format),
None,
"embedding pair should not use chat/generation conversion kind"
);
}
}
}
#[test]
fn embedding_candidate_registry_never_crosses_chat_generation_boundary() {
let embedding_formats = [
"openai:embedding",
"gemini:embedding",
"jina:embedding",
"doubao:embedding",
];
let standard_formats = [
"openai:chat",
"openai:responses",
"claude:messages",
"gemini:generate_content",
];
for embedding_api_format in embedding_formats {
assert!(is_embedding_api_format(embedding_api_format));
assert_eq!(api_data_format_id(embedding_api_format), Some("embedding"));
for standard_api_format in standard_formats {
assert_eq!(
request_candidate_api_format_preference(
embedding_api_format,
standard_api_format
),
None
);
assert_eq!(
request_candidate_api_format_preference(
standard_api_format,
embedding_api_format
),
None
);
assert_eq!(
request_conversion_kind(embedding_api_format, standard_api_format),
None
);
assert_eq!(
request_conversion_kind(standard_api_format, embedding_api_format),
None
);
}
}
}
#[test]
fn rerank_candidate_registry_excludes_chat_and_embedding_formats() {
assert_eq!(
request_candidate_api_formats("openai:rerank", false),
vec!["openai:rerank", "jina:rerank"]
);
assert_eq!(
request_candidate_api_formats("jina:rerank", false),
vec!["jina:rerank", "openai:rerank"]
);
assert_eq!(api_data_format_id("openai:rerank"), Some("rerank"));
assert!(is_rerank_api_format("jina:rerank"));
assert!(!is_embedding_api_format("openai:rerank"));
assert_eq!(
request_candidate_api_format_preference("openai:rerank", "openai:embedding"),
None
);
assert_eq!(
request_candidate_api_format_preference("openai:rerank", "openai:chat"),
None
);
assert_eq!(
request_conversion_kind("openai:rerank", "jina:rerank"),
None
);
assert!(!request_conversion_requires_enable_flag(
"openai:rerank",
"jina:rerank"
));
}
#[test]
fn request_candidate_registry_prefers_same_kind_before_same_family_fallbacks() {
assert_eq!(

View File

@@ -1,7 +1,11 @@
use serde_json::Value;
use crate::{
protocol::canonical::{CanonicalRequest, CanonicalResponse},
protocol::canonical::{
canonical_to_embedding_request, canonical_to_rerank_request,
from_embedding_to_canonical_request, from_rerank_to_canonical_request, CanonicalRequest,
CanonicalResponse,
},
protocol::formats::{
claude_messages, gemini_generate_content, openai_chat, openai_responses, FormatId,
},
@@ -22,6 +26,11 @@ pub fn parse_request(
}
FormatId::ClaudeMessages => claude_messages::request::from(body, ctx),
FormatId::GeminiGenerateContent => gemini_generate_content::request::from(body, ctx),
FormatId::OpenAiEmbedding => from_embedding_to_canonical_request(body, "openai"),
FormatId::JinaEmbedding => from_embedding_to_canonical_request(body, "jina"),
FormatId::OpenAiRerank => from_rerank_to_canonical_request(body, "openai"),
FormatId::JinaRerank => from_rerank_to_canonical_request(body, "jina"),
FormatId::GeminiEmbedding | FormatId::DoubaoEmbedding => None,
}
.ok_or_else(|| FormatError::RequestParseFailed {
format: source.as_str().to_string(),
@@ -48,6 +57,36 @@ pub fn emit_request(
FormatId::OpenAiResponsesCompact => openai_responses::request::to_compact(&request, ctx),
FormatId::ClaudeMessages => claude_messages::request::to(&request, ctx),
FormatId::GeminiGenerateContent => gemini_generate_content::request::to(&request, ctx),
FormatId::OpenAiEmbedding => canonical_to_embedding_request(
&request,
ctx.mapped_model_or(request.model.as_str()),
"openai",
),
FormatId::JinaEmbedding => canonical_to_embedding_request(
&request,
ctx.mapped_model_or(request.model.as_str()),
"jina",
),
FormatId::OpenAiRerank => canonical_to_rerank_request(
&request,
ctx.mapped_model_or(request.model.as_str()),
"openai",
),
FormatId::JinaRerank => canonical_to_rerank_request(
&request,
ctx.mapped_model_or(request.model.as_str()),
"jina",
),
FormatId::GeminiEmbedding => canonical_to_embedding_request(
&request,
ctx.mapped_model_or(request.model.as_str()),
"gemini",
),
FormatId::DoubaoEmbedding => canonical_to_embedding_request(
&request,
ctx.mapped_model_or(request.model.as_str()),
"doubao",
),
}
.ok_or_else(|| FormatError::RequestEmitFailed {
format: target.as_str().to_string(),
@@ -77,6 +116,12 @@ pub fn parse_response(
}
FormatId::ClaudeMessages => claude_messages::response::from(body, ctx),
FormatId::GeminiGenerateContent => gemini_generate_content::response::from(body, ctx),
FormatId::OpenAiEmbedding
| FormatId::JinaEmbedding
| FormatId::OpenAiRerank
| FormatId::JinaRerank
| FormatId::GeminiEmbedding
| FormatId::DoubaoEmbedding => None,
}
.ok_or_else(|| FormatError::ResponseParseFailed {
format: source.as_str().to_string(),
@@ -95,6 +140,12 @@ pub fn emit_response(
FormatId::OpenAiResponsesCompact => openai_responses::response::to_compact(response, ctx),
FormatId::ClaudeMessages => claude_messages::response::to(response, ctx),
FormatId::GeminiGenerateContent => gemini_generate_content::response::to(response, ctx),
FormatId::OpenAiEmbedding
| FormatId::JinaEmbedding
| FormatId::OpenAiRerank
| FormatId::JinaRerank
| FormatId::GeminiEmbedding
| FormatId::DoubaoEmbedding => None,
}
.ok_or_else(|| FormatError::ResponseEmitFailed {
format: target.as_str().to_string(),
@@ -169,6 +220,116 @@ mod tests {
assert_eq!(converted["input"][0]["content"][0]["type"], "input_text");
}
#[test]
fn converts_openai_embedding_to_jina_without_chat_fields() {
let body = json!({
"model": "text-embedding-3-small",
"input": ["alpha", "beta"],
"dimensions": 2
});
let ctx = FormatContext::default().with_mapped_model("jina-embeddings-v3");
let converted = convert_request("openai:embedding", "jina:embedding", &body, &ctx)
.expect("embedding request conversion should succeed");
assert_eq!(converted["model"], "jina-embeddings-v3");
assert_eq!(converted["task"], "text-matching");
assert_eq!(converted["input"], json!(["alpha", "beta"]));
assert!(converted.get("messages").is_none());
}
#[test]
fn converts_openai_embedding_to_gemini_and_doubao_payload_shapes() {
let body = json!({
"model": "text-embedding-3-small",
"input": ["alpha", "beta"],
"dimensions": 2
});
let gemini = convert_request(
"openai:embedding",
"gemini:embedding",
&body,
&FormatContext::default().with_mapped_model("gemini-embedding-001"),
)
.expect("gemini embedding conversion should succeed");
assert_eq!(gemini["model"], "gemini-embedding-001");
assert_eq!(
gemini["requests"][0]["content"]["parts"][0]["text"],
"alpha"
);
assert!(gemini.get("messages").is_none());
let doubao = convert_request(
"openai:embedding",
"doubao:embedding",
&body,
&FormatContext::default().with_mapped_model("doubao-embedding-vision"),
)
.expect("doubao embedding conversion should succeed");
assert_eq!(doubao["model"], "doubao-embedding-vision");
assert_eq!(doubao["input"][0], json!({"type": "text", "text": "alpha"}));
assert!(doubao.get("messages").is_none());
}
#[test]
fn embedding_registry_keeps_gemini_and_doubao_emit_only() {
let body = json!({
"model": "gemini-embedding-001",
"content": {"parts": [{"text": "alpha"}]}
});
let ctx = FormatContext::default();
assert!(convert_request("gemini:embedding", "openai:embedding", &body, &ctx).is_err());
assert!(convert_request("doubao:embedding", "openai:embedding", &body, &ctx).is_err());
}
#[test]
fn embedding_registry_rejects_chat_payload_for_embedding_format() {
let body = json!({
"model": "gpt-5",
"messages": [{"role": "user", "content": "hello"}]
});
let ctx = FormatContext::default();
assert!(convert_request("openai:embedding", "jina:embedding", &body, &ctx).is_err());
}
#[test]
fn converts_openai_rerank_to_jina_without_chat_fields() {
let body = json!({
"model": "rerank-source",
"query": "best document",
"documents": ["alpha", {"text": "beta"}],
"top_n": 1,
"return_documents": true
});
let ctx = FormatContext::default().with_mapped_model("jina-reranker-v2-base-multilingual");
let converted = convert_request("openai:rerank", "jina:rerank", &body, &ctx)
.expect("rerank request conversion should succeed");
assert_eq!(converted["model"], "jina-reranker-v2-base-multilingual");
assert_eq!(converted["query"], "best document");
assert_eq!(converted["documents"], json!(["alpha", {"text": "beta"}]));
assert_eq!(converted["top_n"], 1);
assert_eq!(converted["return_documents"], true);
assert!(converted.get("messages").is_none());
}
#[test]
fn rerank_registry_rejects_invalid_payloads() {
let ctx = FormatContext::default();
for body in [
json!({"model": "rerank", "documents": ["alpha"]}),
json!({"model": "rerank", "query": "q", "documents": []}),
json!({"model": "rerank", "query": "q", "documents": [""]}),
json!({"model": "rerank", "query": "q", "documents": ["alpha"], "top_n": 0}),
] {
assert!(convert_request("openai:rerank", "jina:rerank", &body, &ctx).is_err());
}
}
#[test]
fn registry_does_not_call_wire_specific_canonical_functions_directly() {
let implementation = include_str!("registry.rs")

View File

@@ -1,7 +1,8 @@
use crate::contracts::{
CLAUDE_CHAT_STREAM_PLAN_KIND, CLAUDE_CHAT_SYNC_PLAN_KIND, CLAUDE_CLI_STREAM_PLAN_KIND,
CLAUDE_CLI_SYNC_PLAN_KIND, GEMINI_CHAT_STREAM_PLAN_KIND, GEMINI_CHAT_SYNC_PLAN_KIND,
GEMINI_CLI_STREAM_PLAN_KIND, GEMINI_CLI_SYNC_PLAN_KIND,
GEMINI_CLI_STREAM_PLAN_KIND, GEMINI_CLI_SYNC_PLAN_KIND, OPENAI_EMBEDDING_SYNC_PLAN_KIND,
OPENAI_RERANK_SYNC_PLAN_KIND,
};
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
@@ -49,6 +50,20 @@ pub fn resolve_sync_spec(plan_kind: &str) -> Option<LocalSameFormatProviderSpec>
family: LocalSameFormatProviderFamily::Gemini,
require_streaming: false,
}),
OPENAI_EMBEDDING_SYNC_PLAN_KIND => Some(LocalSameFormatProviderSpec {
api_format: "openai:embedding",
decision_kind: OPENAI_EMBEDDING_SYNC_PLAN_KIND,
report_kind: "openai_embedding_sync_success",
family: LocalSameFormatProviderFamily::Standard,
require_streaming: false,
}),
OPENAI_RERANK_SYNC_PLAN_KIND => Some(LocalSameFormatProviderSpec {
api_format: "openai:rerank",
decision_kind: OPENAI_RERANK_SYNC_PLAN_KIND,
report_kind: "openai_rerank_sync_success",
family: LocalSameFormatProviderFamily::Standard,
require_streaming: false,
}),
_ => None,
}
}
@@ -106,4 +121,20 @@ mod tests {
assert_eq!(spec.report_kind, "gemini_cli_stream_success");
assert!(spec.require_streaming);
}
#[test]
fn resolves_openai_embedding_sync_same_format_spec() {
let spec = resolve_sync_spec("openai_embedding_sync").expect("spec");
assert_eq!(spec.api_format, "openai:embedding");
assert_eq!(spec.report_kind, "openai_embedding_sync_success");
assert!(!spec.require_streaming);
}
#[test]
fn resolves_openai_rerank_sync_same_format_spec() {
let spec = resolve_sync_spec("openai_rerank_sync").expect("spec");
assert_eq!(spec.api_format, "openai:rerank");
assert_eq!(spec.report_kind, "openai_rerank_sync_success");
assert!(!spec.require_streaming);
}
}

View File

@@ -7,12 +7,12 @@ use crate::contracts::{
GEMINI_FILES_DOWNLOAD_PLAN_KIND, GEMINI_FILES_GET_PLAN_KIND, GEMINI_FILES_LIST_PLAN_KIND,
GEMINI_FILES_UPLOAD_PLAN_KIND, GEMINI_VIDEO_CANCEL_SYNC_PLAN_KIND,
GEMINI_VIDEO_CREATE_SYNC_PLAN_KIND, OPENAI_CHAT_STREAM_PLAN_KIND, OPENAI_CHAT_SYNC_PLAN_KIND,
OPENAI_IMAGE_STREAM_PLAN_KIND, OPENAI_IMAGE_SYNC_PLAN_KIND,
OPENAI_RESPONSES_COMPACT_STREAM_PLAN_KIND, OPENAI_RESPONSES_COMPACT_SYNC_PLAN_KIND,
OPENAI_RESPONSES_STREAM_PLAN_KIND, OPENAI_RESPONSES_SYNC_PLAN_KIND,
OPENAI_VIDEO_CANCEL_SYNC_PLAN_KIND, OPENAI_VIDEO_CONTENT_PLAN_KIND,
OPENAI_VIDEO_CREATE_SYNC_PLAN_KIND, OPENAI_VIDEO_DELETE_SYNC_PLAN_KIND,
OPENAI_VIDEO_REMIX_SYNC_PLAN_KIND,
OPENAI_EMBEDDING_SYNC_PLAN_KIND, OPENAI_IMAGE_STREAM_PLAN_KIND, OPENAI_IMAGE_SYNC_PLAN_KIND,
OPENAI_RERANK_SYNC_PLAN_KIND, OPENAI_RESPONSES_COMPACT_STREAM_PLAN_KIND,
OPENAI_RESPONSES_COMPACT_SYNC_PLAN_KIND, OPENAI_RESPONSES_STREAM_PLAN_KIND,
OPENAI_RESPONSES_SYNC_PLAN_KIND, OPENAI_VIDEO_CANCEL_SYNC_PLAN_KIND,
OPENAI_VIDEO_CONTENT_PLAN_KIND, OPENAI_VIDEO_CREATE_SYNC_PLAN_KIND,
OPENAI_VIDEO_DELETE_SYNC_PLAN_KIND, OPENAI_VIDEO_REMIX_SYNC_PLAN_KIND,
};
use crate::request::specialized::image::is_openai_image_stream_request;
@@ -173,6 +173,22 @@ pub fn resolve_execution_runtime_sync_plan_kind(
return Some(OPENAI_CHAT_SYNC_PLAN_KIND);
}
if route_family == Some("openai")
&& route_kind == Some("embedding")
&& *method == Method::POST
&& path == "/v1/embeddings"
{
return Some(OPENAI_EMBEDDING_SYNC_PLAN_KIND);
}
if route_family == Some("openai")
&& route_kind == Some("rerank")
&& *method == Method::POST
&& path == "/v1/rerank"
{
return Some(OPENAI_RERANK_SYNC_PLAN_KIND);
}
if route_family == Some("openai")
&& route_kind == Some("image")
&& *method == Method::POST
@@ -327,6 +343,8 @@ pub fn supports_sync_execution_decision_kind(plan_kind: &str) -> bool {
matches!(
plan_kind,
OPENAI_CHAT_SYNC_PLAN_KIND
| OPENAI_EMBEDDING_SYNC_PLAN_KIND
| OPENAI_RERANK_SYNC_PLAN_KIND
| OPENAI_IMAGE_SYNC_PLAN_KIND
| OPENAI_RESPONSES_SYNC_PLAN_KIND
| OPENAI_RESPONSES_COMPACT_SYNC_PLAN_KIND
@@ -377,7 +395,8 @@ mod tests {
CLAUDE_CHAT_STREAM_PLAN_KIND, CLAUDE_CHAT_SYNC_PLAN_KIND, CLAUDE_CLI_STREAM_PLAN_KIND,
CLAUDE_CLI_SYNC_PLAN_KIND, GEMINI_CHAT_STREAM_PLAN_KIND, GEMINI_CHAT_SYNC_PLAN_KIND,
GEMINI_CLI_STREAM_PLAN_KIND, GEMINI_CLI_SYNC_PLAN_KIND, OPENAI_CHAT_STREAM_PLAN_KIND,
OPENAI_CHAT_SYNC_PLAN_KIND, OPENAI_IMAGE_STREAM_PLAN_KIND, OPENAI_IMAGE_SYNC_PLAN_KIND,
OPENAI_CHAT_SYNC_PLAN_KIND, OPENAI_EMBEDDING_SYNC_PLAN_KIND, OPENAI_IMAGE_STREAM_PLAN_KIND,
OPENAI_IMAGE_SYNC_PLAN_KIND, OPENAI_RERANK_SYNC_PLAN_KIND,
OPENAI_RESPONSES_COMPACT_STREAM_PLAN_KIND, OPENAI_RESPONSES_COMPACT_SYNC_PLAN_KIND,
OPENAI_RESPONSES_STREAM_PLAN_KIND, OPENAI_RESPONSES_SYNC_PLAN_KIND,
};
@@ -650,6 +669,42 @@ mod tests {
));
}
#[test]
fn resolves_openai_embedding_sync_plan_kind() {
assert_eq!(
resolve_execution_runtime_sync_plan_kind(
Some("ai_public"),
Some("openai"),
Some("embedding"),
None,
&Method::POST,
"/v1/embeddings",
),
Some(OPENAI_EMBEDDING_SYNC_PLAN_KIND)
);
assert!(supports_sync_execution_decision_kind(
OPENAI_EMBEDDING_SYNC_PLAN_KIND
));
}
#[test]
fn resolves_openai_rerank_sync_plan_kind() {
assert_eq!(
resolve_execution_runtime_sync_plan_kind(
Some("ai_public"),
Some("openai"),
Some("rerank"),
None,
&Method::POST,
"/v1/rerank",
),
Some(OPENAI_RERANK_SYNC_PLAN_KIND)
);
assert!(supports_sync_execution_decision_kind(
OPENAI_RERANK_SYNC_PLAN_KIND
));
}
#[test]
fn resolves_openai_image_stream_plan_kind() {
assert_eq!(

View File

@@ -8,6 +8,7 @@ pub enum LocalStandardSourceFamily {
pub enum LocalStandardSourceMode {
Chat,
Cli,
Embedding,
}
#[derive(Debug, Clone, Copy)]

View File

@@ -210,6 +210,12 @@ impl ProviderStreamParser {
}
FormatId::ClaudeMessages => Self::Claude(ClaudeProviderState::default()),
FormatId::GeminiGenerateContent => Self::Gemini(GeminiProviderState::default()),
FormatId::OpenAiEmbedding
| FormatId::OpenAiRerank
| FormatId::GeminiEmbedding
| FormatId::JinaEmbedding
| FormatId::JinaRerank
| FormatId::DoubaoEmbedding => return None,
})
}
@@ -294,6 +300,12 @@ impl ClientStreamEmitter {
}
FormatId::ClaudeMessages => Self::Claude(ClaudeClientEmitter::default()),
FormatId::GeminiGenerateContent => Self::Gemini(GeminiClientEmitter::default()),
FormatId::OpenAiEmbedding
| FormatId::OpenAiRerank
| FormatId::GeminiEmbedding
| FormatId::JinaEmbedding
| FormatId::JinaRerank
| FormatId::DoubaoEmbedding => return None,
})
}
@@ -352,6 +364,12 @@ fn parse_provider_error(
}
FormatId::ClaudeMessages => parse_claude_error(payload),
FormatId::GeminiGenerateContent => parse_gemini_error(payload),
FormatId::OpenAiEmbedding
| FormatId::OpenAiRerank
| FormatId::GeminiEmbedding
| FormatId::JinaEmbedding
| FormatId::JinaRerank
| FormatId::DoubaoEmbedding => None,
}
}

View File

@@ -1,6 +1,147 @@
use async_trait::async_trait;
use serde_json::Value;
const EMBEDDING_CAPABILITY: &str = "embedding";
const EMBEDDING_API_FORMATS: &[&str] = &[
"openai:embedding",
"gemini:embedding",
"jina:embedding",
"doubao:embedding",
"/v1/embeddings",
"/jina/v1/embeddings",
];
fn validate_optional_price(
field_name: &str,
value: Option<f64>,
) -> Result<(), crate::DataLayerError> {
if value.is_some_and(|price| !price.is_finite() || price < 0.0) {
return Err(crate::DataLayerError::UnexpectedValue(format!(
"{field_name} must be a non-negative finite number"
)));
}
Ok(())
}
fn validate_embedding_global_billing(
default_price_per_request: Option<f64>,
default_tiered_pricing: Option<&Value>,
supported_capabilities: Option<&Value>,
config: Option<&Value>,
) -> Result<(), crate::DataLayerError> {
validate_optional_price(
"global_models.default_price_per_request",
default_price_per_request,
)?;
if !has_embedding_metadata(supported_capabilities, config) {
return Ok(());
}
if has_request_or_input_token_pricing(default_price_per_request, default_tiered_pricing) {
return Ok(());
}
Err(crate::DataLayerError::UnexpectedValue(
"embedding global model requires default_price_per_request or default_tiered_pricing.tiers[].input_price_per_1m".to_string(),
))
}
fn validate_provider_model_pricing(
price_per_request: Option<f64>,
) -> Result<(), crate::DataLayerError> {
validate_optional_price("models.price_per_request", price_per_request)
}
fn has_request_or_input_token_pricing(
price_per_request: Option<f64>,
tiered_pricing: Option<&Value>,
) -> bool {
price_per_request.is_some_and(|price| price.is_finite() && price >= 0.0)
|| tiered_pricing
.and_then(|value| value.get("tiers"))
.and_then(Value::as_array)
.is_some_and(|tiers| {
tiers.iter().any(|tier| {
tier.get("input_price_per_1m")
.and_then(Value::as_f64)
.is_some_and(|price| price.is_finite() && price >= 0.0)
})
})
}
fn has_embedding_metadata(supported_capabilities: Option<&Value>, config: Option<&Value>) -> bool {
supported_capabilities.is_some_and(value_contains_embedding_capability)
|| config.is_some_and(value_contains_embedding_metadata)
}
fn value_contains_embedding_capability(value: &Value) -> bool {
match value {
Value::String(value) => value.trim().eq_ignore_ascii_case(EMBEDDING_CAPABILITY),
Value::Array(values) => values.iter().any(value_contains_embedding_capability),
Value::Object(object) => {
object
.get(EMBEDDING_CAPABILITY)
.and_then(Value::as_bool)
.unwrap_or(false)
|| [
"capability",
"model_type",
"type",
"task_type",
"request_type",
]
.iter()
.any(|key| {
object
.get(*key)
.is_some_and(value_contains_embedding_capability)
})
|| ["capabilities", "supported_capabilities"]
.iter()
.any(|key| {
object
.get(*key)
.is_some_and(value_contains_embedding_capability)
})
}
_ => false,
}
}
fn value_contains_embedding_metadata(value: &Value) -> bool {
match value {
Value::String(value) => {
value.trim().eq_ignore_ascii_case(EMBEDDING_CAPABILITY)
|| is_known_embedding_api_format(value)
}
Value::Array(values) => values.iter().any(value_contains_embedding_metadata),
Value::Object(object) => {
value_contains_embedding_capability(value)
|| ["api_format", "client_api_format", "provider_api_format"]
.iter()
.any(|key| {
object
.get(*key)
.and_then(Value::as_str)
.is_some_and(is_known_embedding_api_format)
})
|| ["api_formats", "client_api_formats", "provider_api_formats"]
.iter()
.any(|key| {
object
.get(*key)
.is_some_and(value_contains_embedding_metadata)
})
}
_ => false,
}
}
fn is_known_embedding_api_format(value: &str) -> bool {
let normalized = value.trim().to_ascii_lowercase();
EMBEDDING_API_FORMATS
.iter()
.any(|api_format| normalized == *api_format)
}
#[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)]
pub struct StoredPublicGlobalModel {
pub id: String,
@@ -37,6 +178,12 @@ impl StoredPublicGlobalModel {
"global_models.name is empty".to_string(),
));
}
validate_embedding_global_billing(
default_price_per_request,
default_tiered_pricing.as_ref(),
supported_capabilities.as_ref(),
config.as_ref(),
)?;
Ok(Self {
id,
@@ -77,6 +224,7 @@ pub struct StoredPublicCatalogModel {
pub supports_vision: Option<bool>,
pub supports_function_calling: Option<bool>,
pub supports_streaming: Option<bool>,
pub supports_embedding: Option<bool>,
pub is_active: bool,
}
@@ -98,6 +246,7 @@ impl StoredPublicCatalogModel {
supports_vision: Option<bool>,
supports_function_calling: Option<bool>,
supports_streaming: Option<bool>,
supports_embedding: Option<bool>,
is_active: bool,
) -> Result<Self, crate::DataLayerError> {
if id.trim().is_empty() {
@@ -147,6 +296,7 @@ impl StoredPublicCatalogModel {
supports_vision,
supports_function_calling,
supports_streaming,
supports_embedding,
is_active,
})
}
@@ -223,6 +373,12 @@ impl StoredAdminGlobalModel {
"global_models.display_name is empty".to_string(),
));
}
validate_embedding_global_billing(
default_price_per_request,
default_tiered_pricing.as_ref(),
supported_capabilities.as_ref(),
config.as_ref(),
)?;
Ok(Self {
id,
@@ -273,6 +429,7 @@ pub struct StoredAdminProviderModel {
pub global_model_display_name: Option<String>,
pub global_model_default_price_per_request: Option<f64>,
pub global_model_default_tiered_pricing: Option<Value>,
pub global_model_supported_capabilities: Option<Value>,
pub global_model_config: Option<Value>,
}
@@ -300,6 +457,7 @@ impl StoredAdminProviderModel {
global_model_display_name: Option<String>,
global_model_default_price_per_request: Option<f64>,
global_model_default_tiered_pricing: Option<Value>,
global_model_supported_capabilities: Option<Value>,
global_model_config: Option<Value>,
) -> Result<Self, crate::DataLayerError> {
if id.trim().is_empty() {
@@ -322,6 +480,7 @@ impl StoredAdminProviderModel {
"models.provider_model_name is empty".to_string(),
));
}
validate_provider_model_pricing(price_per_request)?;
Ok(Self {
id,
@@ -345,6 +504,7 @@ impl StoredAdminProviderModel {
global_model_display_name,
global_model_default_price_per_request,
global_model_default_tiered_pricing,
global_model_supported_capabilities,
global_model_config,
})
}
@@ -408,6 +568,7 @@ impl UpsertAdminProviderModelRecord {
"models.provider_model_name is empty".to_string(),
));
}
validate_provider_model_pricing(price_per_request)?;
Ok(Self {
id,
@@ -468,6 +629,12 @@ impl CreateAdminGlobalModelRecord {
"global_models.display_name is empty".to_string(),
));
}
validate_embedding_global_billing(
default_price_per_request,
default_tiered_pricing.as_ref(),
supported_capabilities.as_ref(),
config.as_ref(),
)?;
Ok(Self {
id,
@@ -514,6 +681,12 @@ impl UpdateAdminGlobalModelRecord {
"global_models.display_name is empty".to_string(),
));
}
validate_embedding_global_billing(
default_price_per_request,
default_tiered_pricing.as_ref(),
supported_capabilities.as_ref(),
config.as_ref(),
)?;
Ok(Self {
id,
@@ -695,3 +868,99 @@ pub trait GlobalModelWriteRepository: Send + Sync {
global_model_id: &str,
) -> Result<bool, crate::DataLayerError>;
}
#[cfg(test)]
mod tests {
use serde_json::json;
use super::{CreateAdminGlobalModelRecord, UpsertAdminProviderModelRecord};
#[test]
fn embedding_missing_billing_config_rejected() {
let err = CreateAdminGlobalModelRecord::new(
"gm-embedding".to_string(),
"text-embedding-3-small".to_string(),
"Text Embedding 3 Small".to_string(),
true,
None,
None,
Some(json!(["embedding"])),
None,
)
.expect_err("embedding model without explicit billing should be rejected");
assert!(err.to_string().contains(
"embedding global model requires default_price_per_request or default_tiered_pricing.tiers[].input_price_per_1m"
));
}
#[test]
fn embedding_api_format_config_requires_billing_config() {
let err = CreateAdminGlobalModelRecord::new(
"gm-embedding".to_string(),
"jina-embeddings-v3".to_string(),
"Jina Embeddings v3".to_string(),
true,
None,
Some(json!({"tiers": []})),
None,
Some(json!({"api_formats": ["jina:embedding"]})),
)
.expect_err("embedding API format without price should be rejected");
assert!(err.to_string().contains("embedding global model requires"));
}
#[test]
fn embedding_input_token_or_request_pricing_is_accepted() {
CreateAdminGlobalModelRecord::new(
"gm-embedding-input".to_string(),
"text-embedding-3-small".to_string(),
"Text Embedding 3 Small".to_string(),
true,
None,
Some(json!({"tiers":[{"up_to":null,"input_price_per_1m":0.02}]})),
Some(json!(["embedding"])),
Some(json!({"dimensions": 1536})),
)
.expect("input-token pricing should satisfy embedding billing");
CreateAdminGlobalModelRecord::new(
"gm-embedding-request".to_string(),
"custom-embedding".to_string(),
"Custom Embedding".to_string(),
true,
Some(0.0),
None,
Some(json!(["embedding"])),
None,
)
.expect("explicit request pricing should satisfy embedding billing");
}
#[test]
fn provider_model_negative_request_price_rejected() {
let err = UpsertAdminProviderModelRecord::new(
"model-1".to_string(),
"provider-1".to_string(),
"global-model-1".to_string(),
"text-embedding-3-small".to_string(),
None,
Some(-0.01),
None,
None,
None,
None,
None,
None,
true,
true,
None,
)
.expect_err("negative provider model request price should be rejected");
assert!(err
.to_string()
.contains("models.price_per_request must be a non-negative finite number"));
}
}

View File

@@ -9,7 +9,9 @@ pub use types::{
StoredUsageDailySummary, StoredUsageDashboardDailyBreakdownRow,
StoredUsageDashboardProviderCount, StoredUsageDashboardSummary,
StoredUsageErrorDistributionRow, StoredUsageLeaderboardSummary,
StoredUsagePerformancePercentilesRow, StoredUsageSettledCostSummary,
StoredUsagePerformancePercentilesRow, StoredUsageProviderPerformance,
StoredUsageProviderPerformanceProviderRow, StoredUsageProviderPerformanceSummary,
StoredUsageProviderPerformanceTimelineRow, StoredUsageSettledCostSummary,
StoredUsageTimeSeriesBucket, StoredUsageUserTotals, UpsertUsageRecord,
UsageAuditAggregationGroupBy, UsageAuditAggregationQuery, UsageAuditKeywordSearchQuery,
UsageAuditListQuery, UsageAuditSummaryQuery, UsageBodyCaptureResult, UsageBodyCaptureState,
@@ -20,7 +22,7 @@ pub use types::{
UsageDashboardDailyBreakdownQuery, UsageDashboardProviderCountsQuery,
UsageDashboardSummaryQuery, UsageErrorDistributionQuery, UsageLeaderboardGroupBy,
UsageLeaderboardQuery, UsageMonitoringErrorCountQuery, UsageMonitoringErrorListQuery,
UsagePerformancePercentilesQuery, UsageReadRepository, UsageRepository,
UsageSettledCostSummaryQuery, UsageTimeSeriesGranularity, UsageTimeSeriesQuery,
UsageWriteRepository,
UsagePerformancePercentilesQuery, UsageProviderPerformanceQuery, UsageReadRepository,
UsageRepository, UsageSettledCostSummaryQuery, UsageTimeSeriesGranularity,
UsageTimeSeriesQuery, UsageWriteRepository,
};

View File

@@ -950,6 +950,60 @@ pub struct StoredUsagePerformancePercentilesRow {
pub p99_first_byte_time_ms: Option<u64>,
}
#[derive(Debug, Clone, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
pub struct UsageProviderPerformanceQuery {
pub created_from_unix_secs: u64,
pub created_until_unix_secs: u64,
pub granularity: UsageTimeSeriesGranularity,
pub tz_offset_minutes: i32,
pub limit: usize,
}
#[derive(Debug, Clone, PartialEq, Default, serde::Serialize, serde::Deserialize)]
pub struct StoredUsageProviderPerformanceSummary {
pub request_count: u64,
pub success_count: u64,
pub avg_output_tps: Option<f64>,
pub avg_first_byte_time_ms: Option<f64>,
pub avg_response_time_ms: Option<f64>,
}
#[derive(Debug, Clone, PartialEq, Default, serde::Serialize, serde::Deserialize)]
pub struct StoredUsageProviderPerformanceProviderRow {
pub provider_id: String,
pub provider: String,
pub request_count: u64,
pub success_count: u64,
pub output_tokens: u64,
pub avg_output_tps: Option<f64>,
pub avg_first_byte_time_ms: Option<f64>,
pub avg_response_time_ms: Option<f64>,
pub p90_response_time_ms: Option<u64>,
pub p90_first_byte_time_ms: Option<u64>,
pub tps_sample_count: u64,
pub first_byte_sample_count: u64,
}
#[derive(Debug, Clone, PartialEq, Default, serde::Serialize, serde::Deserialize)]
pub struct StoredUsageProviderPerformanceTimelineRow {
pub date: String,
pub provider_id: String,
pub provider: String,
pub request_count: u64,
pub success_count: u64,
pub output_tokens: u64,
pub avg_output_tps: Option<f64>,
pub avg_first_byte_time_ms: Option<f64>,
pub avg_response_time_ms: Option<f64>,
}
#[derive(Debug, Clone, PartialEq, Default, serde::Serialize, serde::Deserialize)]
pub struct StoredUsageProviderPerformance {
pub summary: StoredUsageProviderPerformanceSummary,
pub providers: Vec<StoredUsageProviderPerformanceProviderRow>,
pub timeline: Vec<StoredUsageProviderPerformanceTimelineRow>,
}
#[derive(Debug, Clone, PartialEq, Eq, Default, serde::Serialize, serde::Deserialize)]
pub struct UsageCostSavingsSummaryQuery {
pub created_from_unix_secs: u64,
@@ -1365,6 +1419,11 @@ pub trait UsageReadRepository: Send + Sync {
query: &UsagePerformancePercentilesQuery,
) -> Result<Vec<StoredUsagePerformancePercentilesRow>, crate::DataLayerError>;
async fn summarize_usage_provider_performance(
&self,
query: &UsageProviderPerformanceQuery,
) -> Result<StoredUsageProviderPerformance, crate::DataLayerError>;
async fn summarize_usage_cost_savings(
&self,
query: &UsageCostSavingsSummaryQuery,

View File

@@ -497,6 +497,7 @@ impl GlobalModelWriteRepository for InMemoryGlobalModelReadRepository {
Some(global_model.display_name.clone()),
global_model.default_price_per_request,
global_model.default_tiered_pricing.clone(),
global_model.supported_capabilities.clone(),
global_model.config.clone(),
)?;
self.admin_provider_model_items
@@ -542,6 +543,7 @@ impl GlobalModelWriteRepository for InMemoryGlobalModelReadRepository {
existing.global_model_display_name = Some(global_model.display_name.clone());
existing.global_model_default_price_per_request = global_model.default_price_per_request;
existing.global_model_default_tiered_pricing = global_model.default_tiered_pricing.clone();
existing.global_model_supported_capabilities = global_model.supported_capabilities.clone();
existing.global_model_config = global_model.config.clone();
Ok(Some(existing.clone()))
}
@@ -639,8 +641,9 @@ mod tests {
use super::InMemoryGlobalModelReadRepository;
use crate::repository::global_models::{
GlobalModelReadRepository, PublicCatalogModelListQuery, PublicCatalogModelSearchQuery,
PublicGlobalModelQuery, StoredPublicCatalogModel, StoredPublicGlobalModel,
CreateAdminGlobalModelRecord, GlobalModelReadRepository, GlobalModelWriteRepository,
PublicCatalogModelListQuery, PublicCatalogModelSearchQuery, PublicGlobalModelQuery,
StoredPublicCatalogModel, StoredPublicGlobalModel,
};
fn sample_model(
@@ -687,11 +690,83 @@ mod tests {
Some(true),
Some(true),
Some(true),
Some(false),
true,
)
.expect("public catalog model should build")
}
#[tokio::test]
async fn embedding_model_metadata_roundtrip() {
let repository =
InMemoryGlobalModelReadRepository::seed(Vec::<StoredPublicGlobalModel>::new());
let record = CreateAdminGlobalModelRecord::new(
"gm-embedding".to_string(),
"text-embedding-3-small".to_string(),
"Text Embedding 3 Small".to_string(),
true,
None,
Some(json!({"tiers":[{"up_to":null,"input_price_per_1m":0.02}]})),
Some(json!(["embedding"])),
Some(json!({
"api_formats": ["openai:embedding"],
"dimensions": 1536
})),
)
.expect("embedding global model should validate");
repository
.create_admin_global_model(&record)
.await
.expect("embedding global model should persist")
.expect("embedding global model should be returned");
let stored = repository
.get_admin_global_model_by_name("text-embedding-3-small")
.await
.expect("embedding global model should read")
.expect("embedding global model should exist");
assert_eq!(stored.supported_capabilities, Some(json!(["embedding"])));
assert_eq!(
stored
.config
.as_ref()
.and_then(|value| value.get("dimensions")),
Some(&json!(1536))
);
assert_eq!(
stored
.default_tiered_pricing
.as_ref()
.and_then(|value| value.get("tiers"))
.and_then(serde_json::Value::as_array)
.and_then(|tiers| tiers.first())
.and_then(|tier| tier.get("input_price_per_1m"))
.and_then(serde_json::Value::as_f64),
Some(0.02)
);
}
#[tokio::test]
async fn embedding_missing_billing_config_rejected() {
let error = CreateAdminGlobalModelRecord::new(
"gm-embedding".to_string(),
"text-embedding-3-small".to_string(),
"Text Embedding 3 Small".to_string(),
true,
None,
None,
Some(json!(["embedding"])),
None,
)
.expect_err("embedding metadata without billing should fail closed");
assert!(error
.to_string()
.contains("embedding global model requires"));
}
#[tokio::test]
async fn defaults_to_active_models_only() {
let repository = InMemoryGlobalModelReadRepository::seed(vec![
@@ -791,6 +866,52 @@ mod tests {
assert_eq!(items[0].name, "gpt-5");
}
#[tokio::test]
async fn public_catalog_preserves_embedding_capability_without_contaminating_chat_models() {
let mut embedding_model = sample_public_catalog_model(
"model-embedding",
"provider-openai",
"openai",
"text-embedding-3-small",
"text-embedding-3-small",
"Text Embedding 3 Small",
);
embedding_model.supports_embedding = Some(true);
embedding_model.supports_streaming = Some(false);
let chat_model = sample_public_catalog_model(
"model-chat",
"provider-openai",
"openai",
"gpt-5-upstream",
"gpt-5",
"GPT 5",
);
let repository =
InMemoryGlobalModelReadRepository::seed(Vec::<StoredPublicGlobalModel>::new())
.with_public_catalog_models(vec![embedding_model, chat_model]);
let items = repository
.list_public_catalog_models(&PublicCatalogModelListQuery {
provider_id: Some("provider-openai".to_string()),
offset: 0,
limit: 50,
})
.await
.expect("catalog should list");
let embedding = items
.iter()
.find(|item| item.name == "text-embedding-3-small")
.expect("embedding model should be listed");
let chat = items
.iter()
.find(|item| item.name == "gpt-5")
.expect("chat model should be listed");
assert_eq!(embedding.supports_embedding, Some(true));
assert_eq!(embedding.supports_streaming, Some(false));
assert_eq!(chat.supports_embedding, Some(false));
}
#[tokio::test]
async fn searches_public_catalog_models_by_provider_and_display_name() {
let repository =

View File

@@ -3,6 +3,8 @@ mod mysql;
mod postgres;
mod sqlite;
use serde_json::Value;
#[allow(unused_imports)]
pub(crate) use aether_data_contracts::repository::global_models::{
AdminGlobalModelListQuery, AdminProviderModelListQuery, CreateAdminGlobalModelRecord,
@@ -16,3 +18,95 @@ pub use memory::InMemoryGlobalModelReadRepository;
pub use mysql::MysqlGlobalModelReadRepository;
pub use postgres::SqlxGlobalModelReadRepository;
pub use sqlite::SqliteGlobalModelReadRepository;
const EMBEDDING_CAPABILITY: &str = "embedding";
const EMBEDDING_API_FORMATS: &[&str] = &[
"openai:embedding",
"gemini:embedding",
"jina:embedding",
"doubao:embedding",
"/v1/embeddings",
"/jina/v1/embeddings",
];
pub(super) fn metadata_supports_embedding(
supported_capabilities: Option<&Value>,
global_config: Option<&Value>,
model_config: Option<&Value>,
) -> Option<bool> {
Some(
supported_capabilities.is_some_and(value_contains_embedding_capability)
|| global_config.is_some_and(value_contains_embedding_metadata)
|| model_config.is_some_and(value_contains_embedding_metadata),
)
}
fn value_contains_embedding_capability(value: &Value) -> bool {
match value {
Value::String(value) => value.trim().eq_ignore_ascii_case(EMBEDDING_CAPABILITY),
Value::Array(values) => values.iter().any(value_contains_embedding_capability),
Value::Object(object) => {
object
.get(EMBEDDING_CAPABILITY)
.and_then(Value::as_bool)
.unwrap_or(false)
|| [
"capability",
"model_type",
"type",
"task_type",
"request_type",
]
.iter()
.any(|key| {
object
.get(*key)
.is_some_and(value_contains_embedding_capability)
})
|| ["capabilities", "supported_capabilities"]
.iter()
.any(|key| {
object
.get(*key)
.is_some_and(value_contains_embedding_capability)
})
}
_ => false,
}
}
fn value_contains_embedding_metadata(value: &Value) -> bool {
match value {
Value::String(value) => {
value.trim().eq_ignore_ascii_case(EMBEDDING_CAPABILITY)
|| is_known_embedding_api_format(value)
}
Value::Array(values) => values.iter().any(value_contains_embedding_metadata),
Value::Object(object) => {
value_contains_embedding_capability(value)
|| ["api_format", "client_api_format", "provider_api_format"]
.iter()
.any(|key| {
object
.get(*key)
.and_then(Value::as_str)
.is_some_and(is_known_embedding_api_format)
})
|| ["api_formats", "client_api_formats", "provider_api_formats"]
.iter()
.any(|key| {
object
.get(*key)
.is_some_and(value_contains_embedding_metadata)
})
}
_ => false,
}
}
fn is_known_embedding_api_format(value: &str) -> bool {
let normalized = value.trim().to_ascii_lowercase();
EMBEDDING_API_FORMATS
.iter()
.any(|format| normalized == *format || normalized.ends_with(*format))
}

View File

@@ -2,13 +2,13 @@ use async_trait::async_trait;
use sqlx::{mysql::MySqlRow, Row};
use super::{
AdminGlobalModelListQuery, AdminProviderModelListQuery, CreateAdminGlobalModelRecord,
GlobalModelReadRepository, GlobalModelWriteRepository, InMemoryGlobalModelReadRepository,
PublicCatalogModelListQuery, PublicCatalogModelSearchQuery, PublicGlobalModelQuery,
StoredAdminGlobalModel, StoredAdminGlobalModelPage, StoredAdminProviderModel,
StoredProviderActiveGlobalModel, StoredProviderModelStats, StoredPublicCatalogModel,
StoredPublicGlobalModel, StoredPublicGlobalModelPage, UpdateAdminGlobalModelRecord,
UpsertAdminProviderModelRecord,
metadata_supports_embedding, AdminGlobalModelListQuery, AdminProviderModelListQuery,
CreateAdminGlobalModelRecord, GlobalModelReadRepository, GlobalModelWriteRepository,
InMemoryGlobalModelReadRepository, PublicCatalogModelListQuery, PublicCatalogModelSearchQuery,
PublicGlobalModelQuery, StoredAdminGlobalModel, StoredAdminGlobalModelPage,
StoredAdminProviderModel, StoredProviderActiveGlobalModel, StoredProviderModelStats,
StoredPublicCatalogModel, StoredPublicGlobalModel, StoredPublicGlobalModelPage,
UpdateAdminGlobalModelRecord, UpsertAdminProviderModelRecord,
};
use crate::driver::mysql::MysqlPool;
use crate::error::SqlResultExt;
@@ -104,6 +104,7 @@ SELECT
gm.display_name AS global_model_display_name,
gm.default_price_per_request AS global_model_default_price_per_request,
gm.default_tiered_pricing AS global_model_default_tiered_pricing,
gm.supported_capabilities AS global_model_supported_capabilities,
gm.config AS global_model_config
FROM models m
LEFT JOIN global_models gm ON gm.id = m.global_model_id
@@ -130,6 +131,8 @@ SELECT
COALESCE(gm.name, m.provider_model_name) AS name,
COALESCE(NULLIF(gm.display_name, ''), m.provider_model_name) AS display_name,
gm.config AS global_model_config,
gm.supported_capabilities AS global_model_supported_capabilities,
m.config AS model_config,
m.tiered_pricing,
gm.default_tiered_pricing,
m.supports_vision,
@@ -781,6 +784,11 @@ fn map_admin_provider_model_row(
.map_sql_err()?,
"global_models.default_tiered_pricing",
)?,
optional_json_from_string(
row.try_get("global_model_supported_capabilities")
.map_sql_err()?,
"global_models.supported_capabilities",
)?,
optional_json_from_string(
row.try_get("global_model_config").map_sql_err()?,
"global_models.config",
@@ -795,6 +803,13 @@ fn map_public_catalog_model_row(
row.try_get("global_model_config").map_sql_err()?,
"global_models.config",
)?;
let global_model_supported_capabilities = optional_json_from_string(
row.try_get("global_model_supported_capabilities")
.map_sql_err()?,
"global_models.supported_capabilities",
)?;
let model_config =
optional_json_from_string(row.try_get("model_config").map_sql_err()?, "models.config")?;
let tiered_pricing = optional_json_from_string(
row.try_get("tiered_pricing").map_sql_err()?,
"models.tiered_pricing",
@@ -835,6 +850,11 @@ fn map_public_catalog_model_row(
row.try_get("supports_vision").map_sql_err()?,
row.try_get("supports_function_calling").map_sql_err()?,
row.try_get("supports_streaming").map_sql_err()?,
metadata_supports_embedding(
global_model_supported_capabilities.as_ref(),
global_model_config.as_ref(),
model_config.as_ref(),
),
model_is_active && provider_is_active && global_model_is_active,
)
}

View File

@@ -60,6 +60,27 @@ SELECT
COALESCE(m.supports_vision, CAST(gm.config->>'vision' AS BOOLEAN), FALSE) AS supports_vision,
COALESCE(m.supports_function_calling, CAST(gm.config->>'function_calling' AS BOOLEAN), FALSE) AS supports_function_calling,
COALESCE(m.supports_streaming, CAST(gm.config->>'streaming' AS BOOLEAN), TRUE) AS supports_streaming,
(
COALESCE(gm.supported_capabilities::jsonb @> '["embedding"]'::jsonb, FALSE)
OR LOWER(COALESCE(gm.config->>'embedding', 'false')) = 'true'
OR LOWER(COALESCE(gm.config->>'model_type', '')) = 'embedding'
OR LOWER(COALESCE(gm.config->>'type', '')) = 'embedding'
OR COALESCE(gm.config->'capabilities' @> '["embedding"]'::jsonb, FALSE)
OR COALESCE(gm.config->'supported_capabilities' @> '["embedding"]'::jsonb, FALSE)
OR COALESCE(gm.config->'api_formats' @> '["openai:embedding"]'::jsonb, FALSE)
OR COALESCE(gm.config->'api_formats' @> '["jina:embedding"]'::jsonb, FALSE)
OR COALESCE(gm.config->'api_formats' @> '["gemini:embedding"]'::jsonb, FALSE)
OR COALESCE(gm.config->'api_formats' @> '["doubao:embedding"]'::jsonb, FALSE)
OR LOWER(COALESCE(m.config->>'embedding', 'false')) = 'true'
OR LOWER(COALESCE(m.config->>'model_type', '')) = 'embedding'
OR LOWER(COALESCE(m.config->>'type', '')) = 'embedding'
OR COALESCE(m.config::jsonb->'capabilities' @> '["embedding"]'::jsonb, FALSE)
OR COALESCE(m.config::jsonb->'supported_capabilities' @> '["embedding"]'::jsonb, FALSE)
OR COALESCE(m.config::jsonb->'api_formats' @> '["openai:embedding"]'::jsonb, FALSE)
OR COALESCE(m.config::jsonb->'api_formats' @> '["jina:embedding"]'::jsonb, FALSE)
OR COALESCE(m.config::jsonb->'api_formats' @> '["gemini:embedding"]'::jsonb, FALSE)
OR COALESCE(m.config::jsonb->'api_formats' @> '["doubao:embedding"]'::jsonb, FALSE)
) AS supports_embedding,
m.is_active
FROM models m
JOIN providers p ON p.id = m.provider_id
@@ -98,6 +119,7 @@ SELECT
gm.display_name AS global_model_display_name,
CAST(gm.default_price_per_request AS DOUBLE PRECISION) AS global_model_default_price_per_request,
gm.default_tiered_pricing AS global_model_default_tiered_pricing,
gm.supported_capabilities AS global_model_supported_capabilities,
gm.config AS global_model_config
FROM models m
LEFT JOIN global_models gm ON gm.id = m.global_model_id
@@ -302,6 +324,7 @@ SELECT
gm.display_name AS global_model_display_name,
CAST(gm.default_price_per_request AS DOUBLE PRECISION) AS global_model_default_price_per_request,
gm.default_tiered_pricing AS global_model_default_tiered_pricing,
gm.supported_capabilities AS global_model_supported_capabilities,
gm.config AS global_model_config
FROM models m
LEFT JOIN global_models gm ON gm.id = m.global_model_id
@@ -348,6 +371,7 @@ SELECT
gm.display_name AS global_model_display_name,
CAST(gm.default_price_per_request AS DOUBLE PRECISION) AS global_model_default_price_per_request,
gm.default_tiered_pricing AS global_model_default_tiered_pricing,
gm.supported_capabilities AS global_model_supported_capabilities,
gm.config AS global_model_config
FROM models m
JOIN global_models gm ON gm.id = m.global_model_id
@@ -487,6 +511,7 @@ SELECT
gm.display_name AS global_model_display_name,
CAST(gm.default_price_per_request AS DOUBLE PRECISION) AS global_model_default_price_per_request,
gm.default_tiered_pricing AS global_model_default_tiered_pricing,
gm.supported_capabilities AS global_model_supported_capabilities,
gm.config AS global_model_config
FROM models m
LEFT JOIN global_models gm ON gm.id = m.global_model_id
@@ -1020,7 +1045,7 @@ fn apply_public_catalog_model_filters(
provider_id: Option<&str>,
search: Option<&str>,
) {
builder.push(" WHERE m.is_active = TRUE AND p.is_active = TRUE");
builder.push(" WHERE m.is_active = TRUE AND COALESCE(m.is_available, TRUE) = TRUE AND p.is_active = TRUE AND COALESCE(gm.is_active, TRUE) = TRUE");
if let Some(provider_id) = provider_id.map(str::trim).filter(|value| !value.is_empty()) {
builder
@@ -1060,6 +1085,7 @@ fn map_public_catalog_model_row(row: &PgRow) -> Result<StoredPublicCatalogModel,
row.try_get("supports_function_calling")
.map_postgres_err()?,
row.try_get("supports_streaming").map_postgres_err()?,
row.try_get("supports_embedding").map_postgres_err()?,
row.try_get("is_active").map_postgres_err()?,
)
}
@@ -1101,6 +1127,8 @@ fn map_admin_provider_model_row(row: &PgRow) -> Result<StoredAdminProviderModel,
.map_postgres_err()?,
row.try_get("global_model_default_tiered_pricing")
.map_postgres_err()?,
row.try_get("global_model_supported_capabilities")
.map_postgres_err()?,
row.try_get("global_model_config").map_postgres_err()?,
)
}
@@ -1177,9 +1205,42 @@ fn map_provider_active_global_model_row(
#[cfg(test)]
mod tests {
use super::SqlxGlobalModelReadRepository;
use super::{SqlxGlobalModelReadRepository, LIST_ADMIN_PROVIDER_MODELS_PREFIX};
use crate::driver::postgres::{PostgresPoolConfig, PostgresPoolFactory};
const ADMIN_PROVIDER_MODEL_REQUIRED_COLUMNS: &[&str] = &[
"global_model_default_tiered_pricing",
"global_model_supported_capabilities",
"global_model_config",
];
fn assert_admin_provider_model_projection_has_required_columns(sql: &str) {
for column in ADMIN_PROVIDER_MODEL_REQUIRED_COLUMNS {
assert!(
sql.contains(column),
"admin provider model SQL projection should include {column}"
);
}
}
#[test]
fn admin_provider_model_sql_projections_include_supported_capabilities() {
assert_admin_provider_model_projection_has_required_columns(
LIST_ADMIN_PROVIDER_MODELS_PREFIX,
);
assert_admin_provider_model_projection_has_required_columns(include_str!("postgres.rs"));
let supported_capabilities_projection = format!(
"{} AS {}",
"gm.supported_capabilities", "global_model_supported_capabilities"
);
assert_eq!(
include_str!("postgres.rs")
.matches(&supported_capabilities_projection)
.count(),
4
);
}
#[tokio::test]
async fn repository_constructs_from_lazy_pool() {
let factory = PostgresPoolFactory::new(PostgresPoolConfig {

View File

@@ -2,13 +2,13 @@ use async_trait::async_trait;
use sqlx::{sqlite::SqliteRow, Row};
use super::{
AdminGlobalModelListQuery, AdminProviderModelListQuery, CreateAdminGlobalModelRecord,
GlobalModelReadRepository, GlobalModelWriteRepository, InMemoryGlobalModelReadRepository,
PublicCatalogModelListQuery, PublicCatalogModelSearchQuery, PublicGlobalModelQuery,
StoredAdminGlobalModel, StoredAdminGlobalModelPage, StoredAdminProviderModel,
StoredProviderActiveGlobalModel, StoredProviderModelStats, StoredPublicCatalogModel,
StoredPublicGlobalModel, StoredPublicGlobalModelPage, UpdateAdminGlobalModelRecord,
UpsertAdminProviderModelRecord,
metadata_supports_embedding, AdminGlobalModelListQuery, AdminProviderModelListQuery,
CreateAdminGlobalModelRecord, GlobalModelReadRepository, GlobalModelWriteRepository,
InMemoryGlobalModelReadRepository, PublicCatalogModelListQuery, PublicCatalogModelSearchQuery,
PublicGlobalModelQuery, StoredAdminGlobalModel, StoredAdminGlobalModelPage,
StoredAdminProviderModel, StoredProviderActiveGlobalModel, StoredProviderModelStats,
StoredPublicCatalogModel, StoredPublicGlobalModel, StoredPublicGlobalModelPage,
UpdateAdminGlobalModelRecord, UpsertAdminProviderModelRecord,
};
use crate::driver::sqlite::SqlitePool;
use crate::error::SqlResultExt;
@@ -117,6 +117,7 @@ SELECT
gm.display_name AS global_model_display_name,
gm.default_price_per_request AS global_model_default_price_per_request,
gm.default_tiered_pricing AS global_model_default_tiered_pricing,
gm.supported_capabilities AS global_model_supported_capabilities,
gm.config AS global_model_config
FROM models m
LEFT JOIN global_models gm ON gm.id = m.global_model_id
@@ -143,6 +144,8 @@ SELECT
COALESCE(gm.name, m.provider_model_name) AS name,
COALESCE(NULLIF(gm.display_name, ''), m.provider_model_name) AS display_name,
gm.config AS global_model_config,
gm.supported_capabilities AS global_model_supported_capabilities,
m.config AS model_config,
m.tiered_pricing,
gm.default_tiered_pricing,
m.supports_vision,
@@ -794,6 +797,11 @@ fn map_admin_provider_model_row(
.map_sql_err()?,
"global_models.default_tiered_pricing",
)?,
optional_json_from_string(
row.try_get("global_model_supported_capabilities")
.map_sql_err()?,
"global_models.supported_capabilities",
)?,
optional_json_from_string(
row.try_get("global_model_config").map_sql_err()?,
"global_models.config",
@@ -808,6 +816,13 @@ fn map_public_catalog_model_row(
row.try_get("global_model_config").map_sql_err()?,
"global_models.config",
)?;
let global_model_supported_capabilities = optional_json_from_string(
row.try_get("global_model_supported_capabilities")
.map_sql_err()?,
"global_models.supported_capabilities",
)?;
let model_config =
optional_json_from_string(row.try_get("model_config").map_sql_err()?, "models.config")?;
let tiered_pricing = optional_json_from_string(
row.try_get("tiered_pricing").map_sql_err()?,
"models.tiered_pricing",
@@ -848,6 +863,11 @@ fn map_public_catalog_model_row(
row.try_get("supports_vision").map_sql_err()?,
row.try_get("supports_function_calling").map_sql_err()?,
row.try_get("supports_streaming").map_sql_err()?,
metadata_supports_embedding(
global_model_supported_capabilities.as_ref(),
global_model_config.as_ref(),
model_config.as_ref(),
),
model_is_active && provider_is_active && global_model_is_active,
)
}

View File

@@ -1270,57 +1270,58 @@ INSERT INTO provider_api_keys (
$23,
$24,
$25,
CASE
WHEN $26::double precision IS NULL THEN NULL
ELSE TO_TIMESTAMP($26::double precision)
END,
$26,
CASE
WHEN $27::double precision IS NULL THEN NULL
ELSE TO_TIMESTAMP($27::double precision)
END,
$28,
$29,
COALESCE($30, 0),
COALESCE($31, 0),
CASE
WHEN $32::double precision IS NULL THEN NULL
ELSE TO_TIMESTAMP($32::double precision)
WHEN $28::double precision IS NULL THEN NULL
ELSE TO_TIMESTAMP($28::double precision)
END,
$29,
$30,
COALESCE($31, 0),
COALESCE($32, 0),
CASE
WHEN $33::double precision IS NULL THEN NULL
ELSE TO_TIMESTAMP($33::double precision)
END,
$33,
$34,
$35,
$36,
CASE
WHEN $36::double precision IS NULL THEN NULL
ELSE TO_TIMESTAMP($36::double precision)
WHEN $37::double precision IS NULL THEN NULL
ELSE TO_TIMESTAMP($37::double precision)
END,
$37,
COALESCE($38, 0),
COALESCE($39, 0),
COALESCE($40, 0),
COALESCE($41, 0),
COALESCE($42, 0),
COALESCE($43, 0),
CASE
WHEN $44::double precision IS NULL THEN NULL
ELSE TO_TIMESTAMP($44::double precision)
END,
COALESCE($44, 0),
CASE
WHEN $45::double precision IS NULL THEN NULL
ELSE TO_TIMESTAMP($45::double precision)
END,
$46,
CASE
WHEN $46::double precision IS NULL THEN NULL
ELSE TO_TIMESTAMP($46::double precision)
END,
$47,
$48,
$49,
CASE
WHEN $50::double precision IS NULL THEN NOW()
ELSE TO_TIMESTAMP($50::double precision)
END,
$50,
CASE
WHEN $51::double precision IS NULL THEN NOW()
ELSE TO_TIMESTAMP($51::double precision)
END,
$52
CASE
WHEN $52::double precision IS NULL THEN NOW()
ELSE TO_TIMESTAMP($52::double precision)
END,
$53
)
"#,
)
@@ -2613,4 +2614,19 @@ mod tests {
assert!(source.contains(".bind(&key.allow_auth_channel_mismatch_formats)"));
assert!(source.contains("row.try_get(\"allow_auth_channel_mismatch_formats\").ok()"));
}
#[test]
fn provider_api_keys_create_key_insert_placeholders_match_bind_order() {
let source = include_str!("postgres.rs");
assert!(source.contains(
" $24,\n $25,\n $26,\n CASE\n WHEN $27::double precision IS NULL THEN NULL"
));
assert!(source.contains(" $29,\n $30,\n COALESCE($31, 0),"));
assert!(source.contains(
" COALESCE($42, 0),\n COALESCE($43, 0),\n COALESCE($44, 0),\n CASE\n WHEN $45::double precision IS NULL THEN NULL"
));
assert!(source.contains(
" CASE\n WHEN $52::double precision IS NULL THEN NOW()\n ELSE TO_TIMESTAMP($52::double precision)\n END,\n $53"
));
}
}

View File

@@ -8,7 +8,9 @@ use aether_data_contracts::repository::usage::{
StoredUsageCacheAffinityIntervalRow, StoredUsageCacheHitSummary, StoredUsageCostSavingsSummary,
StoredUsageDashboardDailyBreakdownRow, StoredUsageDashboardProviderCount,
StoredUsageDashboardSummary, StoredUsageErrorDistributionRow, StoredUsageLeaderboardSummary,
StoredUsagePerformancePercentilesRow, StoredUsageSettledCostSummary,
StoredUsagePerformancePercentilesRow, StoredUsageProviderPerformance,
StoredUsageProviderPerformanceProviderRow, StoredUsageProviderPerformanceSummary,
StoredUsageProviderPerformanceTimelineRow, StoredUsageSettledCostSummary,
StoredUsageTimeSeriesBucket, StoredUsageUserTotals, UsageAuditAggregationGroupBy,
UsageAuditAggregationQuery, UsageAuditKeywordSearchQuery, UsageAuditSummaryQuery,
UsageBodyField, UsageBreakdownGroupBy, UsageBreakdownSummaryQuery,
@@ -17,8 +19,8 @@ use aether_data_contracts::repository::usage::{
UsageDashboardDailyBreakdownQuery, UsageDashboardProviderCountsQuery,
UsageDashboardSummaryQuery, UsageErrorDistributionQuery, UsageLeaderboardGroupBy,
UsageLeaderboardQuery, UsageMonitoringErrorCountQuery, UsageMonitoringErrorListQuery,
UsagePerformancePercentilesQuery, UsageSettledCostSummaryQuery, UsageTimeSeriesGranularity,
UsageTimeSeriesQuery,
UsagePerformancePercentilesQuery, UsageProviderPerformanceQuery, UsageSettledCostSummaryQuery,
UsageTimeSeriesGranularity, UsageTimeSeriesQuery,
};
use async_trait::async_trait;
use chrono::Utc;
@@ -587,6 +589,38 @@ fn usage_matches_performance_percentiles_query(
&& item.status == "completed"
}
fn usage_provider_performance_identity(item: &StoredRequestUsageAudit) -> Option<(String, String)> {
let provider_id = item.provider_id.as_deref()?.trim();
let provider_id_status = provider_id.to_ascii_lowercase();
if provider_id.is_empty() || matches!(provider_id_status.as_str(), "unknown" | "pending") {
return None;
}
let provider_name = item.provider_name.trim();
let provider_name_status = provider_name.to_ascii_lowercase();
if matches!(provider_name_status.as_str(), "unknown" | "pending") {
return None;
}
let display_name = if provider_name.is_empty() {
provider_id
} else {
provider_name
};
Some((provider_id.to_string(), display_name.to_string()))
}
fn usage_matches_provider_performance_query(
item: &StoredRequestUsageAudit,
query: &UsageProviderPerformanceQuery,
) -> Option<(String, String)> {
if item.created_at_unix_ms < query.created_from_unix_secs
|| item.created_at_unix_ms >= query.created_until_unix_secs
|| matches!(item.status.as_str(), "pending" | "streaming")
{
return None;
}
usage_provider_performance_identity(item)
}
fn usage_matches_cost_savings_query(
item: &StoredRequestUsageAudit,
query: &UsageCostSavingsSummaryQuery,
@@ -1650,6 +1684,218 @@ impl UsageReadRepository for InMemoryUsageReadRepository {
.collect())
}
async fn summarize_usage_provider_performance(
&self,
query: &UsageProviderPerformanceQuery,
) -> Result<StoredUsageProviderPerformance, DataLayerError> {
#[derive(Default)]
struct ProviderPerformanceBucket {
provider: String,
request_count: u64,
success_count: u64,
output_tokens: u64,
tps_output_tokens: u64,
tps_response_time_ms_sum: u64,
tps_sample_count: u64,
first_byte_time_ms_sum: u64,
first_byte_sample_count: u64,
response_time_ms_sum: u64,
response_time_sample_count: u64,
response_times: Vec<u64>,
first_byte_times: Vec<u64>,
}
impl ProviderPerformanceBucket {
fn add(&mut self, item: &StoredRequestUsageAudit) {
self.request_count = self.request_count.saturating_add(1);
self.output_tokens = self.output_tokens.saturating_add(item.output_tokens);
if !usage_is_success(item) {
return;
}
self.success_count = self.success_count.saturating_add(1);
if let Some(response_time_ms) = item.response_time_ms {
self.response_time_ms_sum =
self.response_time_ms_sum.saturating_add(response_time_ms);
self.response_time_sample_count =
self.response_time_sample_count.saturating_add(1);
self.response_times.push(response_time_ms);
if response_time_ms > 0 && item.output_tokens > 0 {
self.tps_output_tokens =
self.tps_output_tokens.saturating_add(item.output_tokens);
self.tps_response_time_ms_sum = self
.tps_response_time_ms_sum
.saturating_add(response_time_ms);
self.tps_sample_count = self.tps_sample_count.saturating_add(1);
}
}
if let Some(first_byte_time_ms) = item.first_byte_time_ms {
self.first_byte_time_ms_sum = self
.first_byte_time_ms_sum
.saturating_add(first_byte_time_ms);
self.first_byte_sample_count = self.first_byte_sample_count.saturating_add(1);
self.first_byte_times.push(first_byte_time_ms);
}
}
}
fn avg(sum: u64, samples: u64) -> Option<f64> {
if samples == 0 {
None
} else {
Some(sum as f64 / samples as f64)
}
}
fn avg_tps(tokens: u64, response_time_ms_sum: u64) -> Option<f64> {
if response_time_ms_sum == 0 {
None
} else {
Some(tokens as f64 * 1000.0 / response_time_ms_sum as f64)
}
}
let usage = self
.by_request_id
.read()
.expect("usage repository lock")
.values()
.cloned()
.collect::<Vec<_>>();
let mut grouped = BTreeMap::<String, ProviderPerformanceBucket>::new();
let mut summary_bucket = ProviderPerformanceBucket::default();
for item in &usage {
let Some((provider_id, provider)) =
usage_matches_provider_performance_query(item, query)
else {
continue;
};
summary_bucket.add(item);
let bucket = grouped.entry(provider_id).or_default();
if bucket.provider.is_empty() {
bucket.provider = provider;
}
bucket.add(item);
}
let summary = StoredUsageProviderPerformanceSummary {
request_count: summary_bucket.request_count,
success_count: summary_bucket.success_count,
avg_output_tps: avg_tps(
summary_bucket.tps_output_tokens,
summary_bucket.tps_response_time_ms_sum,
),
avg_first_byte_time_ms: avg(
summary_bucket.first_byte_time_ms_sum,
summary_bucket.first_byte_sample_count,
),
avg_response_time_ms: avg(
summary_bucket.response_time_ms_sum,
summary_bucket.response_time_sample_count,
),
};
let mut providers = grouped
.into_iter()
.map(|(provider_id, mut bucket)| {
let p90_response_time_ms = usage_percentile_cont(&mut bucket.response_times, 0.9);
let p90_first_byte_time_ms =
usage_percentile_cont(&mut bucket.first_byte_times, 0.9);
StoredUsageProviderPerformanceProviderRow {
provider_id,
provider: bucket.provider,
request_count: bucket.request_count,
success_count: bucket.success_count,
output_tokens: bucket.output_tokens,
avg_output_tps: avg_tps(
bucket.tps_output_tokens,
bucket.tps_response_time_ms_sum,
),
avg_first_byte_time_ms: avg(
bucket.first_byte_time_ms_sum,
bucket.first_byte_sample_count,
),
avg_response_time_ms: avg(
bucket.response_time_ms_sum,
bucket.response_time_sample_count,
),
p90_response_time_ms,
p90_first_byte_time_ms,
tps_sample_count: bucket.tps_sample_count,
first_byte_sample_count: bucket.first_byte_sample_count,
}
})
.collect::<Vec<_>>();
providers.sort_by(|left, right| {
right
.request_count
.cmp(&left.request_count)
.then_with(|| left.provider_id.cmp(&right.provider_id))
});
providers.truncate(query.limit.max(1));
let top_provider_ids = providers
.iter()
.map(|row| row.provider_id.clone())
.collect::<Vec<_>>();
let mut timeline_grouped = BTreeMap::<(String, String), ProviderPerformanceBucket>::new();
for item in &usage {
let Some((provider_id, provider)) =
usage_matches_provider_performance_query(item, query)
else {
continue;
};
if !top_provider_ids.iter().any(|value| value == &provider_id) {
continue;
}
let Some(bucket_key) =
usage_time_series_bucket_key(item, query.granularity, query.tz_offset_minutes)
else {
continue;
};
let bucket = timeline_grouped
.entry((bucket_key, provider_id))
.or_default();
if bucket.provider.is_empty() {
bucket.provider = provider;
}
bucket.add(item);
}
let timeline = timeline_grouped
.into_iter()
.map(
|((date, provider_id), bucket)| StoredUsageProviderPerformanceTimelineRow {
date,
provider_id,
provider: bucket.provider,
request_count: bucket.request_count,
success_count: bucket.success_count,
output_tokens: bucket.output_tokens,
avg_output_tps: avg_tps(
bucket.tps_output_tokens,
bucket.tps_response_time_ms_sum,
),
avg_first_byte_time_ms: avg(
bucket.first_byte_time_ms_sum,
bucket.first_byte_sample_count,
),
avg_response_time_ms: avg(
bucket.response_time_ms_sum,
bucket.response_time_sample_count,
),
},
)
.collect();
Ok(StoredUsageProviderPerformance {
summary,
providers,
timeline,
})
}
async fn summarize_usage_cost_savings(
&self,
query: &UsageCostSavingsSummaryQuery,
@@ -2579,7 +2825,9 @@ mod tests {
StoredProviderUsageWindow, StoredRequestUsageAudit, UpsertUsageRecord, UsageReadRepository,
UsageWriteRepository,
};
use aether_data_contracts::repository::usage::{usage_body_ref, UsageBodyField};
use aether_data_contracts::repository::usage::{
usage_body_ref, UsageBodyField, UsageProviderPerformanceQuery, UsageTimeSeriesGranularity,
};
use serde_json::json;
fn sample_usage(request_id: &str, created_at_unix_ms: i64) -> StoredRequestUsageAudit {
@@ -4562,4 +4810,76 @@ mod tests {
assert_eq!(key.total_tokens, 300);
assert_eq!(key.total_cost_usd, 0.24);
}
#[tokio::test]
async fn summarize_usage_provider_performance_computes_tps_and_top_provider_timeline() {
let mut first = sample_usage("req-provider-perf-1", 1_711_000_000);
first.output_tokens = 60;
first.response_time_ms = Some(3000);
first.first_byte_time_ms = Some(100);
let mut second = sample_usage("req-provider-perf-2", 1_711_000_300);
second.output_tokens = 40;
second.response_time_ms = Some(1000);
second.first_byte_time_ms = Some(200);
let mut failed = sample_usage("req-provider-perf-failed", 1_711_000_400);
failed.output_tokens = 999;
failed.response_time_ms = Some(10);
failed.first_byte_time_ms = Some(1);
failed.status = "failed".to_string();
failed.status_code = Some(500);
let mut other_provider = sample_usage("req-provider-perf-other", 1_711_003_600);
other_provider.provider_id = Some("provider-2".to_string());
other_provider.provider_name = "Anthropic".to_string();
other_provider.output_tokens = 30;
other_provider.response_time_ms = Some(3000);
other_provider.first_byte_time_ms = None;
let repository =
InMemoryUsageReadRepository::seed(vec![first, second, failed, other_provider]);
let summary = repository
.summarize_usage_provider_performance(&UsageProviderPerformanceQuery {
created_from_unix_secs: 1_711_000_000,
created_until_unix_secs: 1_711_010_000,
granularity: UsageTimeSeriesGranularity::Hour,
tz_offset_minutes: 0,
limit: 1,
})
.await
.expect("provider performance should summarize");
assert_eq!(summary.summary.request_count, 4);
assert_eq!(summary.summary.success_count, 3);
assert!((summary.summary.avg_output_tps.expect("summary tps") - 18.571_428).abs() < 0.001);
assert_eq!(summary.summary.avg_first_byte_time_ms, Some(150.0));
assert!(
(summary
.summary
.avg_response_time_ms
.expect("summary response")
- 2333.333)
.abs()
< 0.001
);
assert_eq!(summary.providers.len(), 1);
let provider = &summary.providers[0];
assert_eq!(provider.provider_id, "provider-1");
assert_eq!(provider.request_count, 3);
assert_eq!(provider.success_count, 2);
assert_eq!(provider.output_tokens, 1099);
assert_eq!(provider.avg_output_tps, Some(25.0));
assert_eq!(provider.avg_first_byte_time_ms, Some(150.0));
assert_eq!(provider.avg_response_time_ms, Some(2000.0));
assert_eq!(provider.p90_response_time_ms, None);
assert_eq!(provider.tps_sample_count, 2);
assert_eq!(provider.first_byte_sample_count, 2);
assert_eq!(summary.timeline.len(), 1);
assert_eq!(summary.timeline[0].date, "2024-03-21T05:00:00+00:00");
assert_eq!(summary.timeline[0].provider_id, "provider-1");
assert_eq!(summary.timeline[0].avg_output_tps, Some(25.0));
}
}

View File

@@ -241,6 +241,17 @@ macro_rules! impl_materialized_usage_read_repository {
<$crate::repository::usage::InMemoryUsageReadRepository as $crate::repository::usage::UsageReadRepository>::summarize_usage_performance_percentiles(&repository, query).await
}
async fn summarize_usage_provider_performance(
&self,
query: &$crate::repository::usage::UsageProviderPerformanceQuery,
) -> Result<
$crate::repository::usage::StoredUsageProviderPerformance,
$crate::DataLayerError,
> {
let repository = self.materialize_read_model().await?;
<$crate::repository::usage::InMemoryUsageReadRepository as $crate::repository::usage::UsageReadRepository>::summarize_usage_provider_performance(&repository, query).await
}
async fn summarize_usage_cost_savings(
&self,
query: &$crate::repository::usage::UsageCostSavingsSummaryQuery,
@@ -348,7 +359,9 @@ pub(crate) use aether_data_contracts::repository::usage::{
StoredUsageDailySummary, StoredUsageDashboardDailyBreakdownRow,
StoredUsageDashboardProviderCount, StoredUsageDashboardSummary,
StoredUsageErrorDistributionRow, StoredUsageLeaderboardSummary,
StoredUsagePerformancePercentilesRow, StoredUsageSettledCostSummary,
StoredUsagePerformancePercentilesRow, StoredUsageProviderPerformance,
StoredUsageProviderPerformanceProviderRow, StoredUsageProviderPerformanceSummary,
StoredUsageProviderPerformanceTimelineRow, StoredUsageSettledCostSummary,
StoredUsageTimeSeriesBucket, StoredUsageUserTotals, UpsertUsageRecord,
UsageAuditAggregationGroupBy, UsageAuditAggregationQuery, UsageAuditKeywordSearchQuery,
UsageAuditListQuery, UsageAuditSummaryQuery, UsageBreakdownGroupBy, UsageBreakdownSummaryQuery,
@@ -358,9 +371,9 @@ pub(crate) use aether_data_contracts::repository::usage::{
UsageDashboardDailyBreakdownQuery, UsageDashboardProviderCountsQuery,
UsageDashboardSummaryQuery, UsageErrorDistributionQuery, UsageLeaderboardGroupBy,
UsageLeaderboardQuery, UsageMonitoringErrorCountQuery, UsageMonitoringErrorListQuery,
UsagePerformancePercentilesQuery, UsageReadRepository, UsageRepository,
UsageSettledCostSummaryQuery, UsageTimeSeriesGranularity, UsageTimeSeriesQuery,
UsageWriteRepository,
UsagePerformancePercentilesQuery, UsageProviderPerformanceQuery, UsageReadRepository,
UsageRepository, UsageSettledCostSummaryQuery, UsageTimeSeriesGranularity,
UsageTimeSeriesQuery, UsageWriteRepository,
};
pub mod cleanup {
pub use super::postgres::cleanup::*;

View File

@@ -4,7 +4,9 @@ use aether_data_contracts::repository::usage::{
StoredUsageCacheAffinityIntervalRow, StoredUsageCacheHitSummary, StoredUsageCostSavingsSummary,
StoredUsageDashboardDailyBreakdownRow, StoredUsageDashboardProviderCount,
StoredUsageDashboardSummary, StoredUsageErrorDistributionRow, StoredUsageLeaderboardSummary,
StoredUsagePerformancePercentilesRow, StoredUsageSettledCostSummary,
StoredUsagePerformancePercentilesRow, StoredUsageProviderPerformance,
StoredUsageProviderPerformanceProviderRow, StoredUsageProviderPerformanceSummary,
StoredUsageProviderPerformanceTimelineRow, StoredUsageSettledCostSummary,
StoredUsageTimeSeriesBucket, StoredUsageUserTotals, UsageAuditAggregationGroupBy,
UsageAuditAggregationQuery, UsageAuditKeywordSearchQuery, UsageAuditSummaryQuery,
UsageBodyCaptureState, UsageBodyField, UsageBreakdownGroupBy, UsageBreakdownSummaryQuery,
@@ -13,8 +15,8 @@ use aether_data_contracts::repository::usage::{
UsageCleanupWindow, UsageCostSavingsSummaryQuery, UsageDashboardDailyBreakdownQuery,
UsageDashboardProviderCountsQuery, UsageDashboardSummaryQuery, UsageErrorDistributionQuery,
UsageLeaderboardGroupBy, UsageLeaderboardQuery, UsageMonitoringErrorCountQuery,
UsageMonitoringErrorListQuery, UsagePerformancePercentilesQuery, UsageSettledCostSummaryQuery,
UsageTimeSeriesGranularity, UsageTimeSeriesQuery,
UsageMonitoringErrorListQuery, UsagePerformancePercentilesQuery, UsageProviderPerformanceQuery,
UsageSettledCostSummaryQuery, UsageTimeSeriesGranularity, UsageTimeSeriesQuery,
};
use async_trait::async_trait;
use chrono::{DateTime, Utc};
@@ -806,6 +808,107 @@ fn decode_usage_performance_percentiles_row(
})
}
fn decode_usage_provider_performance_summary(
row: &PgRow,
) -> Result<StoredUsageProviderPerformanceSummary, DataLayerError> {
Ok(StoredUsageProviderPerformanceSummary {
request_count: row
.try_get::<i64, _>("request_count")
.map_postgres_err()?
.max(0) as u64,
success_count: row
.try_get::<i64, _>("success_count")
.map_postgres_err()?
.max(0) as u64,
avg_output_tps: row
.try_get::<Option<f64>, _>("avg_output_tps")
.map_postgres_err()?,
avg_first_byte_time_ms: row
.try_get::<Option<f64>, _>("avg_first_byte_time_ms")
.map_postgres_err()?,
avg_response_time_ms: row
.try_get::<Option<f64>, _>("avg_response_time_ms")
.map_postgres_err()?,
})
}
fn decode_usage_provider_performance_provider_row(
row: &PgRow,
) -> Result<StoredUsageProviderPerformanceProviderRow, DataLayerError> {
Ok(StoredUsageProviderPerformanceProviderRow {
provider_id: row.try_get::<String, _>("provider_id").map_postgres_err()?,
provider: row.try_get::<String, _>("provider").map_postgres_err()?,
request_count: row
.try_get::<i64, _>("request_count")
.map_postgres_err()?
.max(0) as u64,
success_count: row
.try_get::<i64, _>("success_count")
.map_postgres_err()?
.max(0) as u64,
output_tokens: row
.try_get::<i64, _>("output_tokens")
.map_postgres_err()?
.max(0) as u64,
avg_output_tps: row
.try_get::<Option<f64>, _>("avg_output_tps")
.map_postgres_err()?,
avg_first_byte_time_ms: row
.try_get::<Option<f64>, _>("avg_first_byte_time_ms")
.map_postgres_err()?,
avg_response_time_ms: row
.try_get::<Option<f64>, _>("avg_response_time_ms")
.map_postgres_err()?,
p90_response_time_ms: row
.try_get::<Option<i64>, _>("p90_response_time_ms")
.map_postgres_err()?
.map(|value| value.max(0) as u64),
p90_first_byte_time_ms: row
.try_get::<Option<i64>, _>("p90_first_byte_time_ms")
.map_postgres_err()?
.map(|value| value.max(0) as u64),
tps_sample_count: row
.try_get::<i64, _>("tps_sample_count")
.map_postgres_err()?
.max(0) as u64,
first_byte_sample_count: row
.try_get::<i64, _>("first_byte_sample_count")
.map_postgres_err()?
.max(0) as u64,
})
}
fn decode_usage_provider_performance_timeline_row(
row: &PgRow,
) -> Result<StoredUsageProviderPerformanceTimelineRow, DataLayerError> {
Ok(StoredUsageProviderPerformanceTimelineRow {
date: row.try_get::<String, _>("date").map_postgres_err()?,
provider_id: row.try_get::<String, _>("provider_id").map_postgres_err()?,
provider: row.try_get::<String, _>("provider").map_postgres_err()?,
request_count: row
.try_get::<i64, _>("request_count")
.map_postgres_err()?
.max(0) as u64,
success_count: row
.try_get::<i64, _>("success_count")
.map_postgres_err()?
.max(0) as u64,
output_tokens: row
.try_get::<i64, _>("output_tokens")
.map_postgres_err()?
.max(0) as u64,
avg_output_tps: row
.try_get::<Option<f64>, _>("avg_output_tps")
.map_postgres_err()?,
avg_first_byte_time_ms: row
.try_get::<Option<f64>, _>("avg_first_byte_time_ms")
.map_postgres_err()?,
avg_response_time_ms: row
.try_get::<Option<f64>, _>("avg_response_time_ms")
.map_postgres_err()?,
})
}
fn decode_usage_time_series_bucket_row(
row: &PgRow,
) -> Result<StoredUsageTimeSeriesBucket, DataLayerError> {
@@ -4467,6 +4570,306 @@ ORDER BY date ASC
Ok(items)
}
async fn summarize_usage_provider_performance_summary(
&self,
query: &UsageProviderPerformanceQuery,
) -> Result<StoredUsageProviderPerformanceSummary, DataLayerError> {
let mut builder = QueryBuilder::<Postgres>::new(
r#"
WITH filtered_usage AS (
SELECT
GREATEST(COALESCE("usage".output_tokens, 0), 0) AS output_tokens,
GREATEST(COALESCE("usage".response_time_ms, 0), 0) AS response_time_ms,
GREATEST(COALESCE("usage".first_byte_time_ms, 0), 0) AS first_byte_time_ms,
"usage".response_time_ms IS NOT NULL AS has_response_time,
"usage".first_byte_time_ms IS NOT NULL AS has_first_byte_time,
CASE
WHEN lower(COALESCE("usage".status, '')) IN ('completed', 'success', 'ok', 'billed', 'settled')
AND ("usage".status_code IS NULL OR "usage".status_code < 400)
THEN 1
ELSE 0
END AS success_flag
FROM usage_billing_facts AS "usage"
WHERE "usage".created_at >= TO_TIMESTAMP("#,
);
builder.push_bind(query.created_from_unix_secs as f64);
builder.push(
r#"::double precision)
AND "usage".created_at < TO_TIMESTAMP("#,
);
builder.push_bind(query.created_until_unix_secs as f64);
builder.push(
r#"::double precision)
AND COALESCE("usage".status, '') NOT IN ('pending', 'streaming')
AND NULLIF(BTRIM(COALESCE("usage".provider_id, '')), '') IS NOT NULL
AND lower(BTRIM(COALESCE("usage".provider_id, ''))) NOT IN ('unknown', 'pending')
AND lower(BTRIM(COALESCE("usage".provider_name, ''))) NOT IN ('unknown', 'pending')
)
SELECT
COUNT(*)::BIGINT AS request_count,
COALESCE(SUM(success_flag), 0)::BIGINT AS success_count,
CASE
WHEN COALESCE(SUM(CASE
WHEN success_flag = 1 AND response_time_ms > 0 AND output_tokens > 0
THEN response_time_ms
ELSE 0
END), 0) > 0
THEN COALESCE(SUM(CASE
WHEN success_flag = 1 AND response_time_ms > 0 AND output_tokens > 0
THEN output_tokens
ELSE 0
END), 0)::DOUBLE PRECISION * 1000.0 / COALESCE(SUM(CASE
WHEN success_flag = 1 AND response_time_ms > 0 AND output_tokens > 0
THEN response_time_ms
ELSE 0
END), 0)::DOUBLE PRECISION
ELSE NULL
END AS avg_output_tps,
AVG(first_byte_time_ms::DOUBLE PRECISION)
FILTER (WHERE success_flag = 1 AND has_first_byte_time) AS avg_first_byte_time_ms,
AVG(response_time_ms::DOUBLE PRECISION)
FILTER (WHERE success_flag = 1 AND has_response_time) AS avg_response_time_ms
FROM filtered_usage
"#,
);
let row = builder
.build()
.fetch_one(&self.pool)
.await
.map_postgres_err()?;
decode_usage_provider_performance_summary(&row)
}
async fn summarize_usage_provider_performance_providers(
&self,
query: &UsageProviderPerformanceQuery,
) -> Result<Vec<StoredUsageProviderPerformanceProviderRow>, DataLayerError> {
let mut builder = QueryBuilder::<Postgres>::new(
r#"
WITH filtered_usage AS (
SELECT
COALESCE("usage".provider_id, '') AS provider_id,
COALESCE(NULLIF(BTRIM("usage".provider_name), ''), COALESCE("usage".provider_id, '')) AS provider,
GREATEST(COALESCE("usage".output_tokens, 0), 0) AS output_tokens,
GREATEST(COALESCE("usage".response_time_ms, 0), 0) AS response_time_ms,
GREATEST(COALESCE("usage".first_byte_time_ms, 0), 0) AS first_byte_time_ms,
"usage".response_time_ms IS NOT NULL AS has_response_time,
"usage".first_byte_time_ms IS NOT NULL AS has_first_byte_time,
CASE
WHEN lower(COALESCE("usage".status, '')) IN ('completed', 'success', 'ok', 'billed', 'settled')
AND ("usage".status_code IS NULL OR "usage".status_code < 400)
THEN 1
ELSE 0
END AS success_flag
FROM usage_billing_facts AS "usage"
WHERE "usage".created_at >= TO_TIMESTAMP("#,
);
builder.push_bind(query.created_from_unix_secs as f64);
builder.push(
r#"::double precision)
AND "usage".created_at < TO_TIMESTAMP("#,
);
builder.push_bind(query.created_until_unix_secs as f64);
builder.push(
r#"::double precision)
AND COALESCE("usage".status, '') NOT IN ('pending', 'streaming')
AND NULLIF(BTRIM(COALESCE("usage".provider_id, '')), '') IS NOT NULL
AND lower(BTRIM(COALESCE("usage".provider_id, ''))) NOT IN ('unknown', 'pending')
AND lower(BTRIM(COALESCE("usage".provider_name, ''))) NOT IN ('unknown', 'pending')
)
SELECT
provider_id,
COALESCE(MAX(NULLIF(provider, '')), provider_id) AS provider,
COUNT(*)::BIGINT AS request_count,
COALESCE(SUM(success_flag), 0)::BIGINT AS success_count,
COALESCE(SUM(output_tokens), 0)::BIGINT AS output_tokens,
CASE
WHEN COALESCE(SUM(CASE
WHEN success_flag = 1 AND response_time_ms > 0 AND output_tokens > 0
THEN response_time_ms
ELSE 0
END), 0) > 0
THEN COALESCE(SUM(CASE
WHEN success_flag = 1 AND response_time_ms > 0 AND output_tokens > 0
THEN output_tokens
ELSE 0
END), 0)::DOUBLE PRECISION * 1000.0 / COALESCE(SUM(CASE
WHEN success_flag = 1 AND response_time_ms > 0 AND output_tokens > 0
THEN response_time_ms
ELSE 0
END), 0)::DOUBLE PRECISION
ELSE NULL
END AS avg_output_tps,
AVG(first_byte_time_ms::DOUBLE PRECISION)
FILTER (WHERE success_flag = 1 AND has_first_byte_time) AS avg_first_byte_time_ms,
AVG(response_time_ms::DOUBLE PRECISION)
FILTER (WHERE success_flag = 1 AND has_response_time) AS avg_response_time_ms,
CASE
WHEN COUNT(response_time_ms) FILTER (WHERE success_flag = 1 AND has_response_time) >= 10
THEN FLOOR(PERCENTILE_CONT(0.9) WITHIN GROUP (ORDER BY response_time_ms)
FILTER (WHERE success_flag = 1 AND has_response_time))::BIGINT
ELSE NULL
END AS p90_response_time_ms,
CASE
WHEN COUNT(first_byte_time_ms) FILTER (WHERE success_flag = 1 AND has_first_byte_time) >= 10
THEN FLOOR(PERCENTILE_CONT(0.9) WITHIN GROUP (ORDER BY first_byte_time_ms)
FILTER (WHERE success_flag = 1 AND has_first_byte_time))::BIGINT
ELSE NULL
END AS p90_first_byte_time_ms,
COALESCE(SUM(CASE
WHEN success_flag = 1 AND response_time_ms > 0 AND output_tokens > 0
THEN 1
ELSE 0
END), 0)::BIGINT AS tps_sample_count,
(COUNT(first_byte_time_ms) FILTER (WHERE success_flag = 1 AND has_first_byte_time))::BIGINT
AS first_byte_sample_count
FROM filtered_usage
GROUP BY provider_id
ORDER BY request_count DESC, provider_id ASC
"#,
);
let mut rows = builder.build().fetch(&self.pool);
let mut items = Vec::new();
while let Some(row) = rows.try_next().await.map_postgres_err()? {
items.push(decode_usage_provider_performance_provider_row(&row)?);
}
Ok(items)
}
async fn summarize_usage_provider_performance_timeline(
&self,
query: &UsageProviderPerformanceQuery,
provider_ids: &[String],
) -> Result<Vec<StoredUsageProviderPerformanceTimelineRow>, DataLayerError> {
if provider_ids.is_empty() {
return Ok(Vec::new());
}
let mut builder = QueryBuilder::<Postgres>::new("WITH filtered_usage AS ( SELECT ");
match query.granularity {
UsageTimeSeriesGranularity::Day => {
builder
.push("TO_CHAR(date_trunc('day', \"usage\".created_at + (")
.push_bind(query.tz_offset_minutes)
.push("::integer * INTERVAL '1 minute')), 'YYYY-MM-DD') AS date");
}
UsageTimeSeriesGranularity::Hour => {
builder
.push("TO_CHAR(date_trunc('hour', \"usage\".created_at + (")
.push_bind(query.tz_offset_minutes)
.push("::integer * INTERVAL '1 minute')), 'YYYY-MM-DD\"T\"HH24:00:00+00:00') AS date");
}
}
builder.push(
r#",
COALESCE("usage".provider_id, '') AS provider_id,
COALESCE(NULLIF(BTRIM("usage".provider_name), ''), COALESCE("usage".provider_id, '')) AS provider,
GREATEST(COALESCE("usage".output_tokens, 0), 0) AS output_tokens,
GREATEST(COALESCE("usage".response_time_ms, 0), 0) AS response_time_ms,
GREATEST(COALESCE("usage".first_byte_time_ms, 0), 0) AS first_byte_time_ms,
"usage".response_time_ms IS NOT NULL AS has_response_time,
"usage".first_byte_time_ms IS NOT NULL AS has_first_byte_time,
CASE
WHEN lower(COALESCE("usage".status, '')) IN ('completed', 'success', 'ok', 'billed', 'settled')
AND ("usage".status_code IS NULL OR "usage".status_code < 400)
THEN 1
ELSE 0
END AS success_flag
FROM usage_billing_facts AS "usage"
WHERE "usage".created_at >= TO_TIMESTAMP("#,
);
builder.push_bind(query.created_from_unix_secs as f64);
builder.push(
r#"::double precision)
AND "usage".created_at < TO_TIMESTAMP("#,
);
builder.push_bind(query.created_until_unix_secs as f64);
builder.push(
r#"::double precision)
AND COALESCE("usage".status, '') NOT IN ('pending', 'streaming')
AND NULLIF(BTRIM(COALESCE("usage".provider_id, '')), '') IS NOT NULL
AND lower(BTRIM(COALESCE("usage".provider_id, ''))) NOT IN ('unknown', 'pending')
AND lower(BTRIM(COALESCE("usage".provider_name, ''))) NOT IN ('unknown', 'pending')
AND "usage".provider_id = ANY("#,
);
builder.push_bind(provider_ids.to_vec());
builder.push(
r#")
)
SELECT
date,
provider_id,
COALESCE(MAX(NULLIF(provider, '')), provider_id) AS provider,
COUNT(*)::BIGINT AS request_count,
COALESCE(SUM(success_flag), 0)::BIGINT AS success_count,
COALESCE(SUM(output_tokens), 0)::BIGINT AS output_tokens,
CASE
WHEN COALESCE(SUM(CASE
WHEN success_flag = 1 AND response_time_ms > 0 AND output_tokens > 0
THEN response_time_ms
ELSE 0
END), 0) > 0
THEN COALESCE(SUM(CASE
WHEN success_flag = 1 AND response_time_ms > 0 AND output_tokens > 0
THEN output_tokens
ELSE 0
END), 0)::DOUBLE PRECISION * 1000.0 / COALESCE(SUM(CASE
WHEN success_flag = 1 AND response_time_ms > 0 AND output_tokens > 0
THEN response_time_ms
ELSE 0
END), 0)::DOUBLE PRECISION
ELSE NULL
END AS avg_output_tps,
AVG(first_byte_time_ms::DOUBLE PRECISION)
FILTER (WHERE success_flag = 1 AND has_first_byte_time) AS avg_first_byte_time_ms,
AVG(response_time_ms::DOUBLE PRECISION)
FILTER (WHERE success_flag = 1 AND has_response_time) AS avg_response_time_ms
FROM filtered_usage
GROUP BY date, provider_id
ORDER BY date ASC, provider_id ASC
"#,
);
let mut rows = builder.build().fetch(&self.pool);
let mut items = Vec::new();
while let Some(row) = rows.try_next().await.map_postgres_err()? {
items.push(decode_usage_provider_performance_timeline_row(&row)?);
}
Ok(items)
}
pub async fn summarize_usage_provider_performance(
&self,
query: &UsageProviderPerformanceQuery,
) -> Result<StoredUsageProviderPerformance, DataLayerError> {
if query.created_from_unix_secs >= query.created_until_unix_secs {
return Ok(StoredUsageProviderPerformance::default());
}
let summary = self
.summarize_usage_provider_performance_summary(query)
.await?;
let mut providers = self
.summarize_usage_provider_performance_providers(query)
.await?;
providers.truncate(query.limit.max(1));
let provider_ids = providers
.iter()
.map(|row| row.provider_id.clone())
.collect::<Vec<_>>();
let timeline = self
.summarize_usage_provider_performance_timeline(query, &provider_ids)
.await?;
Ok(StoredUsageProviderPerformance {
summary,
providers,
timeline,
})
}
async fn summarize_usage_cost_savings_raw_from_range(
&self,
start_utc: DateTime<Utc>,
@@ -6999,6 +7402,13 @@ impl UsageReadRepository for SqlxUsageReadRepository {
Self::summarize_usage_performance_percentiles(self, query).await
}
async fn summarize_usage_provider_performance(
&self,
query: &UsageProviderPerformanceQuery,
) -> Result<StoredUsageProviderPerformance, DataLayerError> {
Self::summarize_usage_provider_performance(self, query).await
}
async fn summarize_usage_cost_savings(
&self,
query: &UsageCostSavingsSummaryQuery,

View File

@@ -1,4 +1,5 @@
use super::provider_types::{
provider_type_supports_local_embedding_transport,
provider_type_supports_local_openai_chat_transport,
provider_type_supports_local_same_format_transport,
};
@@ -195,10 +196,200 @@ fn local_same_format_transport_unsupported_reason(
if !provider_type_supported(&transport.provider.provider_type) {
return Some("transport_provider_type_unsupported");
}
if aether_ai_formats::is_embedding_api_format(api_format) {
if !endpoint_kind_allows_embedding(transport.endpoint.endpoint_kind.as_deref()) {
return Some("transport_endpoint_kind_unsupported");
}
if !provider_type_supports_local_embedding_transport(
&transport.provider.provider_type,
api_format,
) {
return Some("transport_provider_type_unsupported");
}
}
if aether_ai_formats::is_rerank_api_format(api_format) {
if !endpoint_kind_allows_rerank(transport.endpoint.endpoint_kind.as_deref()) {
return Some("transport_endpoint_kind_unsupported");
}
if !provider_type_supports_local_embedding_transport(
&transport.provider.provider_type,
api_format,
) {
return Some("transport_provider_type_unsupported");
}
}
None
}
fn endpoint_kind_allows_embedding(endpoint_kind: Option<&str>) -> bool {
endpoint_kind
.map(str::trim)
.filter(|value| !value.is_empty())
.map(|value| {
matches!(
value.to_ascii_lowercase().as_str(),
"embedding" | "embeddings"
)
})
.unwrap_or(true)
}
fn endpoint_kind_allows_rerank(endpoint_kind: Option<&str>) -> bool {
endpoint_kind
.map(str::trim)
.filter(|value| !value.is_empty())
.map(|value| matches!(value.to_ascii_lowercase().as_str(), "rerank" | "reranking"))
.unwrap_or(true)
}
fn same_api_format(left: &str, right: &str) -> bool {
aether_ai_formats::api_format_alias_matches(left, right)
}
#[cfg(test)]
mod tests {
use super::local_standard_transport_unsupported_reason_with_network;
use crate::snapshot::{
GatewayProviderTransportEndpoint, GatewayProviderTransportKey,
GatewayProviderTransportProvider, GatewayProviderTransportSnapshot,
};
fn sample_transport(
provider_type: &str,
api_format: &str,
endpoint_kind: Option<&str>,
) -> GatewayProviderTransportSnapshot {
GatewayProviderTransportSnapshot {
provider: GatewayProviderTransportProvider {
id: "provider-1".to_string(),
name: "provider".to_string(),
provider_type: provider_type.to_string(),
website: None,
is_active: true,
keep_priority_on_conversion: false,
enable_format_conversion: false,
concurrent_limit: None,
max_retries: None,
proxy: None,
request_timeout_secs: None,
stream_first_byte_timeout_secs: None,
config: None,
},
endpoint: GatewayProviderTransportEndpoint {
id: "endpoint-1".to_string(),
provider_id: "provider-1".to_string(),
api_format: api_format.to_string(),
api_family: None,
endpoint_kind: endpoint_kind.map(ToOwned::to_owned),
is_active: true,
base_url: "https://provider.example".to_string(),
header_rules: None,
body_rules: None,
max_retries: None,
custom_path: None,
config: None,
format_acceptance_config: None,
proxy: None,
},
key: GatewayProviderTransportKey {
id: "key-1".to_string(),
provider_id: "provider-1".to_string(),
name: "key".to_string(),
auth_type: "api_key".to_string(),
is_active: true,
api_formats: None,
auth_type_by_format: None,
allow_auth_channel_mismatch_formats: None,
allowed_models: None,
capabilities: None,
rate_multipliers: None,
global_priority_by_format: None,
expires_at_unix_secs: None,
proxy: None,
fingerprint: None,
decrypted_api_key: "sk-test".to_string(),
decrypted_auth_config: None,
},
}
}
#[test]
fn rejects_unsupported_embedding_provider_format_pairs() {
let openai_on_gemini = sample_transport("openai", "gemini:embedding", Some("embedding"));
let gemini_on_openai = sample_transport("gemini", "openai:embedding", Some("embedding"));
let chat_marked_embedding = sample_transport("openai", "openai:embedding", Some("chat"));
assert_eq!(
local_standard_transport_unsupported_reason_with_network(
&openai_on_gemini,
"gemini:embedding"
),
Some("transport_provider_type_unsupported")
);
assert_eq!(
local_standard_transport_unsupported_reason_with_network(
&gemini_on_openai,
"openai:embedding"
),
Some("transport_provider_type_unsupported")
);
assert_eq!(
local_standard_transport_unsupported_reason_with_network(
&chat_marked_embedding,
"openai:embedding"
),
Some("transport_endpoint_kind_unsupported")
);
}
#[test]
fn accepts_supported_embedding_provider_format_pairs() {
for (provider_type, api_format) in [
("openai", "openai:embedding"),
("gemini", "gemini:embedding"),
("google", "gemini:embedding"),
("jina", "jina:embedding"),
("doubao", "doubao:embedding"),
("volcengine", "doubao:embedding"),
("custom", "openai:embedding"),
("custom", "gemini:embedding"),
("custom", "jina:embedding"),
("custom", "doubao:embedding"),
] {
let transport = sample_transport(provider_type, api_format, Some("embedding"));
assert_eq!(
local_standard_transport_unsupported_reason_with_network(&transport, api_format),
None,
"{provider_type} should support {api_format}"
);
}
}
#[test]
fn embedding_policy_accepts_embedding_endpoint_kind_aliases_only() {
for endpoint_kind in [None, Some(""), Some(" embedding "), Some("EMBEDDINGS")] {
let transport = sample_transport("openai", "openai:embedding", endpoint_kind);
assert_eq!(
local_standard_transport_unsupported_reason_with_network(
&transport,
"openai:embedding"
),
None,
"endpoint kind {endpoint_kind:?} should be accepted"
);
}
for endpoint_kind in [Some("chat"), Some("responses"), Some("image")] {
let transport = sample_transport("openai", "openai:embedding", endpoint_kind);
assert_eq!(
local_standard_transport_unsupported_reason_with_network(
&transport,
"openai:embedding"
),
Some("transport_endpoint_kind_unsupported"),
"endpoint kind {endpoint_kind:?} should fail closed"
);
}
}
}

View File

@@ -225,6 +225,24 @@ pub fn provider_type_supports_local_same_format_transport(provider_type: &str) -
)
}
pub fn provider_type_supports_local_embedding_transport(
provider_type: &str,
api_format: &str,
) -> bool {
let provider_type = provider_type.trim().to_ascii_lowercase();
let api_format = aether_ai_formats::normalize_api_format_alias(api_format);
match api_format.as_str() {
"openai:embedding" => matches!(provider_type.as_str(), "custom" | "openai"),
"openai:rerank" => matches!(provider_type.as_str(), "custom" | "openai"),
"gemini:embedding" => matches!(provider_type.as_str(), "custom" | "gemini" | "google"),
"jina:embedding" => matches!(provider_type.as_str(), "custom" | "jina"),
"jina:rerank" => matches!(provider_type.as_str(), "custom" | "jina"),
"doubao:embedding" => matches!(provider_type.as_str(), "custom" | "doubao" | "volcengine"),
_ => false,
}
}
pub fn is_codex_cli_backend_url(url: &str) -> bool {
let url = url.trim().to_ascii_lowercase();
url.contains("/codex") && (url.contains("/backend-api/") || url.contains("/backendapi/"))
@@ -301,7 +319,8 @@ pub const ADMIN_PROVIDER_OAUTH_TEMPLATE_TYPES: &[&str] =
mod tests {
use super::{
fixed_provider_endpoint_template_by_api_format, fixed_provider_key_inherits_api_formats,
fixed_provider_template, FixedProviderEndpointConfigValue,
fixed_provider_template, provider_type_supports_local_embedding_transport,
FixedProviderEndpointConfigValue,
};
#[test]
@@ -355,4 +374,41 @@ mod tests {
"custom", "oauth", None
));
}
#[test]
fn provider_type_supports_only_matching_embedding_formats() {
for (provider_type, api_format) in [
("openai", "openai:embedding"),
("custom", "openai:embedding"),
("gemini", "gemini:embedding"),
("google", "gemini:embedding"),
("jina", "jina:embedding"),
("doubao", "doubao:embedding"),
("volcengine", "doubao:embedding"),
] {
assert!(
provider_type_supports_local_embedding_transport(provider_type, api_format),
"{provider_type} should support {api_format}"
);
}
for (provider_type, api_format) in [
("openai", "gemini:embedding"),
("gemini", "openai:embedding"),
("jina", "doubao:embedding"),
("doubao", "jina:embedding"),
("claude_code", "openai:embedding"),
("openai", "openai:chat"),
] {
assert!(
!provider_type_supports_local_embedding_transport(provider_type, api_format),
"{provider_type} should not support {api_format}"
);
}
assert!(provider_type_supports_local_embedding_transport(
" Google ",
"GEMINI:EMBEDDING"
));
}
}

View File

@@ -78,6 +78,12 @@ pub fn build_transport_request_url(
params.request_query,
true,
)),
"openai:embedding" | "jina:embedding" => {
build_provider_embedding_v1_url(&transport.endpoint.base_url, params.request_query)
}
"openai:rerank" | "jina:rerank" => {
build_provider_rerank_v1_url(&transport.endpoint.base_url, params.request_query)
}
"claude:messages" => Some(build_claude_messages_url(
&transport.endpoint.base_url,
params.request_query,
@@ -88,6 +94,17 @@ pub fn build_transport_request_url(
params.upstream_is_stream,
params.request_query,
),
"gemini:embedding" => build_gemini_embedding_url(
&transport.endpoint.base_url,
params.mapped_model?,
params.request_query,
),
"doubao:embedding" => build_passthrough_path_url(
&transport.endpoint.base_url,
"/embeddings/multimodal",
params.request_query,
&[],
),
_ => None,
}?;
@@ -221,11 +238,8 @@ fn build_transport_hook_url(
));
}
if params
.provider_api_format
.trim()
.to_ascii_lowercase()
.starts_with("gemini:")
if aether_ai_formats::normalize_api_format_alias(params.provider_api_format)
== "gemini:generate_content"
{
if let Some(auth) = resolve_local_vertex_api_key_query_auth(transport) {
return build_vertex_api_key_gemini_content_url(
@@ -266,15 +280,14 @@ fn build_path_params(params: TransportRequestUrlParams<'_>) -> BTreeMap<&'static
{
path_params.insert("model", model);
}
if params
.provider_api_format
.trim()
.to_ascii_lowercase()
.starts_with("gemini:")
{
let provider_api_format =
aether_ai_formats::normalize_api_format_alias(params.provider_api_format);
if provider_api_format.starts_with("gemini:") {
path_params.insert(
"action",
if params.upstream_is_stream {
if provider_api_format == "gemini:embedding" {
"embedContent"
} else if params.upstream_is_stream {
"streamGenerateContent"
} else {
"generateContent"
@@ -284,6 +297,60 @@ fn build_path_params(params: TransportRequestUrlParams<'_>) -> BTreeMap<&'static
path_params
}
fn build_provider_embedding_v1_url(upstream_base_url: &str, query: Option<&str>) -> Option<String> {
build_provider_v1_url(upstream_base_url, "/embeddings", "/v1/embeddings", query)
}
fn build_provider_rerank_v1_url(upstream_base_url: &str, query: Option<&str>) -> Option<String> {
build_provider_v1_url(upstream_base_url, "/rerank", "/v1/rerank", query)
}
fn build_provider_v1_url(
upstream_base_url: &str,
v1_path: &str,
default_path: &str,
query: Option<&str>,
) -> Option<String> {
let base_without_query = upstream_base_url
.trim()
.split_once('?')
.map(|(base, _)| base)
.unwrap_or_else(|| upstream_base_url.trim())
.trim_end_matches('/');
let path = if base_without_query.ends_with("/v1") {
v1_path
} else {
default_path
};
build_passthrough_path_url(upstream_base_url, path, query, &[])
}
fn build_gemini_embedding_url(
upstream_base_url: &str,
model: &str,
query: Option<&str>,
) -> Option<String> {
let trimmed_base_url = upstream_base_url
.trim()
.split_once('?')
.map(|(base, _)| base)
.unwrap_or_else(|| upstream_base_url.trim())
.trim_end_matches('/');
let trimmed_model = model.trim();
if trimmed_base_url.is_empty() || trimmed_model.is_empty() {
return None;
}
let path = if trimmed_base_url.ends_with("/v1beta") {
format!("/models/{trimmed_model}:embedContent")
} else if trimmed_base_url.contains("/v1beta/models/") {
":embedContent".to_string()
} else {
format!("/v1beta/models/{trimmed_model}:embedContent")
};
build_passthrough_path_url(upstream_base_url, &path, query, &["key"])
}
fn expand_custom_path_template(path: &str, params: BTreeMap<&'static str, &str>) -> String {
if params.is_empty() {
return path.to_string();
@@ -320,7 +387,10 @@ fn maybe_add_gemini_stream_alt_sse(
provider_api_format: &str,
upstream_is_stream: bool,
) -> String {
if !provider_api_format.starts_with("gemini:") || !upstream_is_stream {
if aether_ai_formats::normalize_api_format_alias(provider_api_format)
!= "gemini:generate_content"
|| !upstream_is_stream
{
return upstream_url;
}
@@ -548,4 +618,260 @@ mod tests {
));
assert!(url.contains("conversationId=abc"));
}
#[test]
fn embedding_request_url_builds_provider_default_paths() {
let openai = sample_transport(
"openai",
"openai:embedding",
"https://api.openai.example/v1",
None,
);
let jina = sample_transport("jina", "jina:embedding", "https://api.jina.example", None);
let gemini = sample_transport(
"gemini",
"gemini:embedding",
"https://generativelanguage.googleapis.com/v1beta",
None,
);
let doubao = sample_transport(
"doubao",
"doubao:embedding",
"https://ark.volces.example/api/v3",
None,
);
assert_eq!(
build_transport_request_url(
&openai,
TransportRequestUrlParams {
provider_api_format: "openai:embedding",
mapped_model: Some("text-embedding-3-small"),
upstream_is_stream: false,
request_query: Some("tenant=demo"),
kiro_api_region: None,
},
)
.as_deref(),
Some("https://api.openai.example/v1/embeddings?tenant=demo")
);
assert_eq!(
build_transport_request_url(
&jina,
TransportRequestUrlParams {
provider_api_format: "jina:embedding",
mapped_model: Some("jina-embeddings-v3"),
upstream_is_stream: false,
request_query: None,
kiro_api_region: None,
},
)
.as_deref(),
Some("https://api.jina.example/v1/embeddings")
);
assert_eq!(
build_transport_request_url(
&gemini,
TransportRequestUrlParams {
provider_api_format: "gemini:embedding",
mapped_model: Some("gemini-embedding-001"),
upstream_is_stream: false,
request_query: Some("key=client-key&foo=bar"),
kiro_api_region: None,
},
)
.as_deref(),
Some(
"https://generativelanguage.googleapis.com/v1beta/models/gemini-embedding-001:embedContent?foo=bar"
)
);
assert_eq!(
build_transport_request_url(
&doubao,
TransportRequestUrlParams {
provider_api_format: "doubao:embedding",
mapped_model: Some("doubao-embedding-vision"),
upstream_is_stream: false,
request_query: None,
kiro_api_region: None,
},
)
.as_deref(),
Some("https://ark.volces.example/api/v3/embeddings/multimodal")
);
}
#[test]
fn rerank_request_url_builds_provider_default_paths() {
let openai = sample_transport(
"openai",
"openai:rerank",
"https://api.openai.example/v1",
None,
);
let jina = sample_transport("jina", "jina:rerank", "https://api.jina.example", None);
assert_eq!(
build_transport_request_url(
&openai,
TransportRequestUrlParams {
provider_api_format: "openai:rerank",
mapped_model: Some("rerank-1"),
upstream_is_stream: false,
request_query: Some("tenant=demo"),
kiro_api_region: None,
},
)
.as_deref(),
Some("https://api.openai.example/v1/rerank?tenant=demo")
);
assert_eq!(
build_transport_request_url(
&jina,
TransportRequestUrlParams {
provider_api_format: "jina:rerank",
mapped_model: Some("jina-reranker-v2-base-multilingual"),
upstream_is_stream: false,
request_query: None,
kiro_api_region: None,
},
)
.as_deref(),
Some("https://api.jina.example/v1/rerank")
);
}
#[test]
fn embedding_request_url_handles_base_variants_and_queries() {
let openai_without_v1 = sample_transport(
"openai",
"openai:embedding",
"https://api.openai.example/root?tenant=base",
None,
);
let jina_with_v1 = sample_transport(
"jina",
"jina:embedding",
"https://api.jina.example/v1/",
None,
);
let gemini_model_base = sample_transport(
"gemini",
"gemini:embedding",
"https://generativelanguage.googleapis.com/v1beta/models/gemini-embedding-001",
None,
);
let doubao_with_query = sample_transport(
"doubao",
"doubao:embedding",
"https://ark.volces.example/api/v3?tenant=base",
None,
);
assert_eq!(
build_transport_request_url(
&openai_without_v1,
TransportRequestUrlParams {
provider_api_format: "OPENAI:EMBEDDING",
mapped_model: Some("text-embedding-3-small"),
upstream_is_stream: false,
request_query: Some("tenant=request&trace=1"),
kiro_api_region: None,
},
)
.as_deref(),
Some("https://api.openai.example/root/v1/embeddings?tenant=request&trace=1")
);
assert_eq!(
build_transport_request_url(
&jina_with_v1,
TransportRequestUrlParams {
provider_api_format: "jina:embedding",
mapped_model: Some("jina-embeddings-v3"),
upstream_is_stream: false,
request_query: Some("trace=2"),
kiro_api_region: None,
},
)
.as_deref(),
Some("https://api.jina.example/v1/embeddings?trace=2")
);
assert_eq!(
build_transport_request_url(
&gemini_model_base,
TransportRequestUrlParams {
provider_api_format: "gemini:embedding",
mapped_model: Some("gemini-embedding-001"),
upstream_is_stream: false,
request_query: Some("key=client-key&trace=3"),
kiro_api_region: None,
},
)
.as_deref(),
Some("https://generativelanguage.googleapis.com/v1beta/models/gemini-embedding-001:embedContent?trace=3")
);
assert_eq!(
build_transport_request_url(
&doubao_with_query,
TransportRequestUrlParams {
provider_api_format: "doubao:embedding",
mapped_model: Some("doubao-embedding-vision"),
upstream_is_stream: false,
request_query: Some("trace=4"),
kiro_api_region: None,
},
)
.as_deref(),
Some("https://ark.volces.example/api/v3/embeddings/multimodal?tenant=base&trace=4")
);
}
#[test]
fn gemini_embedding_request_url_requires_mapped_model_without_custom_path() {
let transport = sample_transport(
"gemini",
"gemini:embedding",
"https://generativelanguage.googleapis.com/v1beta",
None,
);
assert!(build_transport_request_url(
&transport,
TransportRequestUrlParams {
provider_api_format: "gemini:embedding",
mapped_model: None,
upstream_is_stream: false,
request_query: None,
kiro_api_region: None,
},
)
.is_none());
}
#[test]
fn embedding_request_url_expands_custom_gemini_embed_action() {
let transport = sample_transport(
"custom",
"gemini:embedding",
"https://generativelanguage.googleapis.com",
Some("/v1beta/models/{model}:{action}"),
);
let url = build_transport_request_url(
&transport,
TransportRequestUrlParams {
provider_api_format: "gemini:embedding",
mapped_model: Some("gemini-embedding-001"),
upstream_is_stream: true,
request_query: Some("key=client-key&foo=bar"),
kiro_api_region: None,
},
)
.expect("expanded custom embedding path url");
assert_eq!(
url,
"https://generativelanguage.googleapis.com/v1beta/models/gemini-embedding-001:embedContent?foo=bar"
);
}
}

View File

@@ -52,6 +52,7 @@ pub struct SameFormatProviderRequestBehavior {
pub struct SameFormatProviderRequestBodyInput<'a> {
pub body_json: &'a Value,
pub mapped_model: &'a str,
pub client_api_format: &'a str,
pub provider_api_format: &'a str,
pub source_model: Option<&'a str>,
pub family: SameFormatProviderFamily,
@@ -133,12 +134,27 @@ pub fn build_same_format_provider_request_body(
);
}
let request_body_object = input.body_json.as_object()?;
let mut provider_request_body = serde_json::Map::from_iter(
request_body_object
.iter()
.map(|(key, value)| (key.clone(), value.clone())),
);
let mut provider_request_body = if aether_ai_formats::api_format_alias_matches(
input.client_api_format,
input.provider_api_format,
) {
let request_body_object = input.body_json.as_object()?;
serde_json::Map::from_iter(
request_body_object
.iter()
.map(|(key, value)| (key.clone(), value.clone())),
)
} else {
aether_ai_formats::convert_request(
input.client_api_format,
input.provider_api_format,
input.body_json,
&aether_ai_formats::FormatContext::default().with_mapped_model(input.mapped_model),
)
.ok()?
.as_object()?
.clone()
};
match input.family {
SameFormatProviderFamily::Standard => {
provider_request_body.insert(
@@ -488,6 +504,7 @@ mod tests {
"messages": [{"role": "user", "content": "hello"}]
}),
mapped_model: "upstream-model",
client_api_format: "openai:chat",
provider_api_format: "openai:chat",
source_model: Some("client-model"),
family: SameFormatProviderFamily::Standard,
@@ -512,6 +529,7 @@ mod tests {
"reasoning_effort": "low"
}),
mapped_model: "upstream-model",
client_api_format: "openai:chat",
provider_api_format: "openai:chat",
source_model: Some("gpt-5.4-high"),
family: SameFormatProviderFamily::Standard,