mirror of
https://github.com/fawney19/Aether.git
synced 2026-10-06 17:37:47 +08:00
Merge pull request #371 from Kayphoon/feature/embedding-model-support
feat: add embedding and rerank support
This commit is contained in:
@@ -11,6 +11,7 @@
|
||||
<p align="center">
|
||||
<a href="#简介">简介</a> •
|
||||
<a href="#部署">部署</a> •
|
||||
<a href="#api-文档">API 文档</a> •
|
||||
<a href="#环境变量">环境变量</a> •
|
||||
<a href="#qa">Q&A</a>
|
||||
</p>
|
||||
@@ -123,6 +124,11 @@ Aether Proxy 是配套的正向代理节点,部署在海外 VPS 上,为墙
|
||||
- 通过 `aether-proxy setup` 完成交互式配置,自动注册为系统服务
|
||||
- 详细文档见 [apps/aether-proxy/README.md](apps/aether-proxy/README.md)
|
||||
|
||||
## API 文档
|
||||
|
||||
- Embeddings: [OpenAI compatible `POST /v1/embeddings`](docs/api/embeddings.md)
|
||||
- Rerank: [OpenAI/Jina compatible `POST /v1/rerank`](docs/api/rerank.md)
|
||||
|
||||
## 环境变量
|
||||
|
||||
部署建议直接参考对应示例文件:
|
||||
|
||||
@@ -53,10 +53,11 @@ pub(crate) use aether_ai_formats::api::{
|
||||
LocalStandardSourceFamily, LocalStandardSourceMode, LocalStandardSpec,
|
||||
StreamingStandardTerminalObserver, EXECUTION_RUNTIME_STREAM_DECISION_ACTION,
|
||||
EXECUTION_RUNTIME_SYNC_DECISION_ACTION, GEMINI_FILES_DOWNLOAD_PLAN_KIND,
|
||||
GEMINI_VIDEO_CANCEL_SYNC_PLAN_KIND, OPENAI_IMAGE_STREAM_PLAN_KIND,
|
||||
OPENAI_IMAGE_SYNC_FINALIZE_REPORT_KIND, OPENAI_IMAGE_SYNC_PLAN_KIND,
|
||||
OPENAI_VIDEO_CANCEL_SYNC_PLAN_KIND, OPENAI_VIDEO_CONTENT_PLAN_KIND,
|
||||
OPENAI_VIDEO_DELETE_SYNC_PLAN_KIND, OPENAI_VIDEO_REMIX_SYNC_PLAN_KIND,
|
||||
GEMINI_VIDEO_CANCEL_SYNC_PLAN_KIND, OPENAI_EMBEDDING_SYNC_PLAN_KIND,
|
||||
OPENAI_IMAGE_STREAM_PLAN_KIND, OPENAI_IMAGE_SYNC_FINALIZE_REPORT_KIND,
|
||||
OPENAI_IMAGE_SYNC_PLAN_KIND, OPENAI_RERANK_SYNC_PLAN_KIND, OPENAI_VIDEO_CANCEL_SYNC_PLAN_KIND,
|
||||
OPENAI_VIDEO_CONTENT_PLAN_KIND, OPENAI_VIDEO_DELETE_SYNC_PLAN_KIND,
|
||||
OPENAI_VIDEO_REMIX_SYNC_PLAN_KIND,
|
||||
};
|
||||
|
||||
pub(crate) fn parse_direct_request_body(
|
||||
|
||||
@@ -13,12 +13,12 @@ pub(crate) use crate::ai_serving::{
|
||||
GEMINI_FILES_DELETE_PLAN_KIND, 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(crate) use aether_ai_serving::AiRequestedModelFamily as RequestedModelFamily;
|
||||
|
||||
@@ -4,12 +4,13 @@ use crate::ai_serving::planner::common::{
|
||||
GEMINI_CLI_STREAM_PLAN_KIND, GEMINI_CLI_SYNC_PLAN_KIND, GEMINI_FILES_DELETE_PLAN_KIND,
|
||||
GEMINI_FILES_DOWNLOAD_PLAN_KIND, GEMINI_FILES_GET_PLAN_KIND, GEMINI_FILES_LIST_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_CHAT_STREAM_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,
|
||||
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::ai_serving::planner::plan_builders::{
|
||||
build_gemini_stream_plan_from_decision, build_gemini_sync_plan_from_decision,
|
||||
@@ -103,7 +104,10 @@ fn build_sync_plan_payload_from_decision(
|
||||
OPENAI_RESPONSES_COMPACT_SYNC_PLAN_KIND => {
|
||||
build_openai_responses_sync_plan_from_decision(parts, body_json, payload, true)?
|
||||
}
|
||||
CLAUDE_CHAT_SYNC_PLAN_KIND | CLAUDE_CLI_SYNC_PLAN_KIND => {
|
||||
CLAUDE_CHAT_SYNC_PLAN_KIND
|
||||
| CLAUDE_CLI_SYNC_PLAN_KIND
|
||||
| OPENAI_EMBEDDING_SYNC_PLAN_KIND
|
||||
| OPENAI_RERANK_SYNC_PLAN_KIND => {
|
||||
build_standard_sync_plan_from_decision(parts, body_json, payload)?
|
||||
}
|
||||
GEMINI_CHAT_SYNC_PLAN_KIND | GEMINI_CLI_SYNC_PLAN_KIND => {
|
||||
|
||||
@@ -84,6 +84,7 @@ pub(crate) async fn maybe_build_local_same_format_provider_decision_payload_for_
|
||||
json!(super::super::ANTIGRAVITY_ENVELOPE_NAME),
|
||||
);
|
||||
}
|
||||
let provider_api_format = resolved.provider_api_format.clone();
|
||||
let report_context = append_local_failover_policy_to_value(
|
||||
append_execution_contract_fields_to_value(
|
||||
build_local_execution_report_context(LocalExecutionReportContextParts {
|
||||
@@ -100,7 +101,7 @@ pub(crate) async fn maybe_build_local_same_format_provider_decision_payload_for_
|
||||
model_id: Some(&candidate.model_id),
|
||||
global_model_id: Some(&candidate.global_model_id),
|
||||
global_model_name: Some(&candidate.global_model_name),
|
||||
provider_api_format: spec_metadata.api_format,
|
||||
provider_api_format: provider_api_format.as_str(),
|
||||
client_api_format: spec_metadata.api_format,
|
||||
mapped_model: Some(&resolved.mapped_model),
|
||||
candidate_group_id: eligible.orchestration.candidate_group_id.as_deref(),
|
||||
@@ -126,7 +127,7 @@ pub(crate) async fn maybe_build_local_same_format_provider_decision_payload_for_
|
||||
execution_strategy,
|
||||
conversion_mode,
|
||||
spec_metadata.api_format,
|
||||
spec_metadata.api_format,
|
||||
provider_api_format.as_str(),
|
||||
),
|
||||
&resolved.transport,
|
||||
);
|
||||
@@ -136,6 +137,7 @@ pub(crate) async fn maybe_build_local_same_format_provider_decision_payload_for_
|
||||
is_kiro: _,
|
||||
auth_header,
|
||||
auth_value,
|
||||
provider_api_format,
|
||||
mapped_model,
|
||||
report_kind,
|
||||
upstream_is_stream,
|
||||
@@ -161,7 +163,7 @@ pub(crate) async fn maybe_build_local_same_format_provider_decision_payload_for_
|
||||
provider_request_method: None,
|
||||
auth_header,
|
||||
auth_value,
|
||||
provider_api_format: spec_metadata.api_format.to_string(),
|
||||
provider_api_format,
|
||||
client_api_format: spec_metadata.api_format.to_string(),
|
||||
model_name: input.requested_model.clone(),
|
||||
mapped_model,
|
||||
|
||||
@@ -29,12 +29,63 @@ use super::{
|
||||
};
|
||||
use crate::ai_serving::planner::standard::same_format_provider_request_body_failure_extra_data;
|
||||
|
||||
pub(crate) fn resolve_same_format_provider_transport_unsupported_reason_for_trace(
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
provider_api_format: &str,
|
||||
) -> Option<&'static str> {
|
||||
let provider_api_format =
|
||||
match crate::ai_serving::normalize_api_format_alias(provider_api_format).as_str() {
|
||||
"openai:chat" => "openai:chat",
|
||||
"openai:responses" => "openai:responses",
|
||||
"openai:responses:compact" => "openai:responses:compact",
|
||||
"openai:embedding" => "openai:embedding",
|
||||
"openai:rerank" => "openai:rerank",
|
||||
"claude:messages" => "claude:messages",
|
||||
"gemini:generate_content" => "gemini:generate_content",
|
||||
"gemini:embedding" => "gemini:embedding",
|
||||
"jina:embedding" => "jina:embedding",
|
||||
"jina:rerank" => "jina:rerank",
|
||||
"doubao:embedding" => "doubao:embedding",
|
||||
_ => return Some("transport_api_format_unsupported"),
|
||||
};
|
||||
let behavior = policy::classify_same_format_provider_request_behavior(
|
||||
transport,
|
||||
crate::ai_serving::planner::spec_metadata::LocalExecutionSurfaceSpecMetadata {
|
||||
api_format: provider_api_format,
|
||||
require_streaming: false,
|
||||
requested_model_family: None,
|
||||
decision_kind: "trace_candidate_metadata",
|
||||
report_kind: Some("trace_candidate_metadata"),
|
||||
},
|
||||
);
|
||||
if !behavior.is_antigravity
|
||||
&& !behavior.is_claude_code
|
||||
&& !behavior.is_vertex
|
||||
&& !behavior.is_kiro
|
||||
{
|
||||
return None;
|
||||
}
|
||||
|
||||
let family = if provider_api_format.starts_with("gemini:") {
|
||||
crate::ai_serving::LocalSameFormatProviderFamily::Gemini
|
||||
} else {
|
||||
crate::ai_serving::LocalSameFormatProviderFamily::Standard
|
||||
};
|
||||
policy::same_format_provider_transport_unsupported_reason(
|
||||
&behavior,
|
||||
transport,
|
||||
family,
|
||||
provider_api_format,
|
||||
)
|
||||
}
|
||||
|
||||
pub(crate) struct LocalSameFormatProviderCandidatePayloadParts {
|
||||
pub(super) transport: Arc<GatewayProviderTransportSnapshot>,
|
||||
pub(super) is_antigravity: bool,
|
||||
pub(super) is_kiro: bool,
|
||||
pub(super) auth_header: Option<String>,
|
||||
pub(super) auth_value: Option<String>,
|
||||
pub(super) provider_api_format: String,
|
||||
pub(super) mapped_model: String,
|
||||
pub(super) report_kind: &'static str,
|
||||
pub(super) upstream_is_stream: bool,
|
||||
@@ -74,6 +125,7 @@ pub(crate) async fn resolve_local_same_format_provider_candidate_payload_parts(
|
||||
let Some(mut base_provider_request_body) =
|
||||
super::super::request::build_same_format_provider_request_body(
|
||||
body_json,
|
||||
prepared.provider_api_format.as_str(),
|
||||
&prepared.mapped_model,
|
||||
spec,
|
||||
prepared.transport.endpoint.body_rules.as_ref(),
|
||||
@@ -180,6 +232,7 @@ pub(crate) async fn resolve_local_same_format_provider_candidate_payload_parts(
|
||||
parts,
|
||||
&prepared.transport,
|
||||
&prepared.mapped_model,
|
||||
prepared.provider_api_format.as_str(),
|
||||
spec,
|
||||
prepared.upstream_is_stream,
|
||||
prepared.kiro_auth.as_ref(),
|
||||
@@ -248,6 +301,7 @@ pub(crate) async fn resolve_local_same_format_provider_candidate_payload_parts(
|
||||
is_kiro: prepared.is_kiro,
|
||||
auth_header: prepared.auth_header,
|
||||
auth_value: prepared.auth_value,
|
||||
provider_api_format: prepared.provider_api_format,
|
||||
mapped_model: prepared.mapped_model,
|
||||
report_kind: prepared.report_kind,
|
||||
upstream_is_stream: prepared.upstream_is_stream,
|
||||
|
||||
+6
-3
@@ -31,6 +31,7 @@ pub(super) struct PreparedSameFormatProviderCandidate {
|
||||
pub(super) kiro_auth: Option<KiroRequestAuth>,
|
||||
pub(super) auth_header: Option<String>,
|
||||
pub(super) auth_value: Option<String>,
|
||||
pub(super) provider_api_format: String,
|
||||
pub(super) mapped_model: String,
|
||||
pub(super) report_kind: &'static str,
|
||||
pub(super) upstream_is_stream: bool,
|
||||
@@ -49,19 +50,20 @@ pub(super) async fn prepare_local_same_format_provider_candidate(
|
||||
let planner_state = PlannerAppState::new(state);
|
||||
let candidate = &eligible.candidate;
|
||||
let transport = Arc::clone(&eligible.transport);
|
||||
let provider_api_format = eligible.provider_api_format.as_str();
|
||||
let behavior = classify_same_format_provider_request_behavior(&transport, spec_metadata);
|
||||
|
||||
if !same_format_provider_transport_supported(
|
||||
&behavior,
|
||||
&transport,
|
||||
spec.family,
|
||||
spec_metadata.api_format,
|
||||
provider_api_format,
|
||||
) {
|
||||
let skip_reason = same_format_provider_transport_unsupported_reason(
|
||||
&behavior,
|
||||
&transport,
|
||||
spec.family,
|
||||
spec_metadata.api_format,
|
||||
provider_api_format,
|
||||
)
|
||||
.unwrap_or("transport_unsupported");
|
||||
super::super::payload::mark_skipped_local_same_format_provider_candidate(
|
||||
@@ -90,7 +92,7 @@ pub(super) async fn prepare_local_same_format_provider_candidate(
|
||||
&transport,
|
||||
OauthPreparationContext {
|
||||
trace_id,
|
||||
api_format: spec_metadata.api_format,
|
||||
api_format: provider_api_format,
|
||||
operation: "same_format_provider_prepare",
|
||||
},
|
||||
)
|
||||
@@ -168,6 +170,7 @@ pub(super) async fn prepare_local_same_format_provider_candidate(
|
||||
kiro_auth,
|
||||
auth_header,
|
||||
auth_value,
|
||||
provider_api_format: provider_api_format.to_string(),
|
||||
mapped_model,
|
||||
report_kind: behavior.report_kind,
|
||||
upstream_is_stream: behavior.upstream_is_stream,
|
||||
|
||||
@@ -8,6 +8,7 @@ use crate::ai_serving::transport::{
|
||||
|
||||
pub(crate) fn build_same_format_provider_request_body(
|
||||
body_json: &Value,
|
||||
provider_api_format: &str,
|
||||
mapped_model: &str,
|
||||
spec: LocalSameFormatProviderSpec,
|
||||
body_rules: Option<&Value>,
|
||||
@@ -19,7 +20,8 @@ pub(crate) fn build_same_format_provider_request_body(
|
||||
build_same_format_provider_request_body_impl(SameFormatProviderRequestBodyInput {
|
||||
body_json,
|
||||
mapped_model,
|
||||
provider_api_format: spec.api_format,
|
||||
client_api_format: spec.api_format,
|
||||
provider_api_format,
|
||||
source_model: body_json.get("model").and_then(Value::as_str),
|
||||
family: same_format_provider_family(spec.family),
|
||||
body_rules,
|
||||
|
||||
@@ -10,6 +10,7 @@ pub(crate) fn build_same_format_upstream_url(
|
||||
parts: &http::request::Parts,
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
mapped_model: &str,
|
||||
provider_api_format: &str,
|
||||
spec: LocalSameFormatProviderSpec,
|
||||
upstream_is_stream: bool,
|
||||
kiro_auth: Option<&crate::ai_serving::transport::kiro::KiroRequestAuth>,
|
||||
@@ -17,7 +18,7 @@ pub(crate) fn build_same_format_upstream_url(
|
||||
build_same_format_provider_upstream_url_impl(
|
||||
transport,
|
||||
SameFormatProviderUpstreamUrlParams {
|
||||
provider_api_format: spec.api_format,
|
||||
provider_api_format,
|
||||
mapped_model,
|
||||
upstream_is_stream,
|
||||
request_query: parts.uri.query(),
|
||||
|
||||
@@ -28,11 +28,12 @@ pub(crate) use aether_ai_formats::api::{
|
||||
convert_openai_chat_request_to_openai_responses_request,
|
||||
convert_openai_chat_response_to_claude_chat, convert_openai_chat_response_to_gemini_chat,
|
||||
convert_openai_chat_response_to_openai_responses,
|
||||
convert_openai_responses_response_to_openai_chat, convert_standard_chat_response,
|
||||
convert_standard_cli_response, copy_request_number_field, copy_request_number_field_as,
|
||||
core_error_background_report_kind, core_error_default_client_api_format,
|
||||
core_success_background_report_kind, default_model_for_openai_image_operation, encode_done_sse,
|
||||
encode_json_sse, encode_kiro_sse_events, estimate_kiro_tokens, extract_openai_text_content,
|
||||
convert_openai_responses_response_to_openai_chat, convert_request,
|
||||
convert_standard_chat_response, convert_standard_cli_response, copy_request_number_field,
|
||||
copy_request_number_field_as, core_error_background_report_kind,
|
||||
core_error_default_client_api_format, core_success_background_report_kind,
|
||||
default_model_for_openai_image_operation, encode_done_sse, encode_json_sse,
|
||||
encode_kiro_sse_events, estimate_kiro_tokens, extract_openai_text_content,
|
||||
find_kiro_real_thinking_end_tag, find_kiro_real_thinking_end_tag_at_buffer_end,
|
||||
find_kiro_real_thinking_start_tag, force_upstream_streaming_for_provider,
|
||||
implicit_sync_finalize_report_kind, is_core_error_finalize_kind,
|
||||
@@ -75,7 +76,7 @@ pub(crate) use aether_ai_formats::api::{
|
||||
sync_chat_response_conversion_kind, sync_cli_response_conversion_kind,
|
||||
transform_provider_private_stream_line, value_as_u64, AiControlPlanRequest,
|
||||
AiSurfaceFinalizeError, AiSurfaceStreamRewriter, CanonicalStreamFrame, ClaudeClientEmitter,
|
||||
ClaudeProviderState, ExecutionRuntimeAuthContext, FinalizeStreamRewriteMode,
|
||||
ClaudeProviderState, ExecutionRuntimeAuthContext, FinalizeStreamRewriteMode, FormatContext,
|
||||
GeminiClientEmitter, GeminiProviderState, KiroToClaudeCliStreamState, LocalCoreSyncErrorKind,
|
||||
LocalGeminiFilesSpec, LocalOpenAiImageSpec, LocalOpenAiResponsesSpec,
|
||||
LocalSameFormatProviderFamily, LocalSameFormatProviderSpec, LocalStandardSourceFamily,
|
||||
@@ -110,9 +111,10 @@ pub(crate) use aether_ai_formats::api::{
|
||||
KIRO_ENVELOPE_NAME, KIRO_MAX_THINKING_BUFFER, 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_CHAT_SYNC_SUCCESS_REPORT_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,
|
||||
|
||||
@@ -0,0 +1,13 @@
|
||||
pub(crate) fn normalized_signature(api_format: &str) -> Option<&'static str> {
|
||||
match crate::ai_serving::normalize_api_format_alias(api_format).as_str() {
|
||||
"doubao:embedding" => Some("doubao:embedding"),
|
||||
_ => None,
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) fn local_path(api_format: &str) -> Option<&'static str> {
|
||||
match crate::ai_serving::normalize_api_format_alias(api_format).as_str() {
|
||||
"doubao:embedding" => Some("/v1/embeddings"),
|
||||
_ => None,
|
||||
}
|
||||
}
|
||||
@@ -1,6 +1,7 @@
|
||||
pub(crate) fn normalized_signature(api_format: &str) -> Option<&'static str> {
|
||||
match crate::ai_serving::normalize_api_format_alias(api_format).as_str() {
|
||||
"gemini:generate_content" => Some("gemini:generate_content"),
|
||||
"gemini:embedding" => Some("gemini:embedding"),
|
||||
"gemini:video" => Some("gemini:video"),
|
||||
"gemini:files" => Some("gemini:files"),
|
||||
_ => None,
|
||||
@@ -10,6 +11,7 @@ pub(crate) fn normalized_signature(api_format: &str) -> Option<&'static str> {
|
||||
pub(crate) fn local_path(api_format: &str) -> Option<&'static str> {
|
||||
match crate::ai_serving::normalize_api_format_alias(api_format).as_str() {
|
||||
"gemini" | "gemini:generate_content" => Some("/v1beta/models/{model}:{action}"),
|
||||
"gemini:embedding" => Some("/v1/embeddings"),
|
||||
"gemini:video" => Some("/v1beta/models/{model}:predictLongRunning"),
|
||||
"gemini:files" => Some("/v1beta/files"),
|
||||
_ => None,
|
||||
|
||||
@@ -0,0 +1,15 @@
|
||||
pub(crate) fn normalized_signature(api_format: &str) -> Option<&'static str> {
|
||||
match crate::ai_serving::normalize_api_format_alias(api_format).as_str() {
|
||||
"jina:embedding" => Some("jina:embedding"),
|
||||
"jina:rerank" => Some("jina:rerank"),
|
||||
_ => None,
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) fn local_path(api_format: &str) -> Option<&'static str> {
|
||||
match crate::ai_serving::normalize_api_format_alias(api_format).as_str() {
|
||||
"jina:embedding" => Some("/v1/embeddings"),
|
||||
"jina:rerank" => Some("/v1/rerank"),
|
||||
_ => None,
|
||||
}
|
||||
}
|
||||
@@ -1,5 +1,7 @@
|
||||
mod claude;
|
||||
mod doubao;
|
||||
mod gemini;
|
||||
mod jina;
|
||||
mod openai;
|
||||
mod registry;
|
||||
|
||||
|
||||
@@ -1,6 +1,8 @@
|
||||
pub(crate) fn normalized_signature(api_format: &str) -> Option<&'static str> {
|
||||
match crate::ai_serving::normalize_api_format_alias(api_format).as_str() {
|
||||
"openai:chat" => Some("openai:chat"),
|
||||
"openai:embedding" => Some("openai:embedding"),
|
||||
"openai:rerank" => Some("openai:rerank"),
|
||||
"openai:responses" => Some("openai:responses"),
|
||||
"openai:responses:compact" => Some("openai:responses:compact"),
|
||||
"openai:image" => Some("openai:image"),
|
||||
@@ -12,6 +14,8 @@ pub(crate) fn normalized_signature(api_format: &str) -> Option<&'static str> {
|
||||
pub(crate) fn local_path(api_format: &str) -> Option<&'static str> {
|
||||
match crate::ai_serving::normalize_api_format_alias(api_format).as_str() {
|
||||
"openai" | "openai:chat" => Some("/v1/chat/completions"),
|
||||
"openai:embedding" => Some("/v1/embeddings"),
|
||||
"openai:rerank" => Some("/v1/rerank"),
|
||||
"openai:responses" => Some("/v1/responses"),
|
||||
"openai:responses:compact" => Some("/v1/responses/compact"),
|
||||
"openai:image" => Some("/v1/images/generations"),
|
||||
|
||||
@@ -1,7 +1,7 @@
|
||||
use axum::routing::{any, post};
|
||||
use axum::Router;
|
||||
|
||||
use super::{claude, gemini, openai};
|
||||
use super::{claude, doubao, gemini, jina, openai};
|
||||
use crate::{handlers::proxy::proxy_request, state::AppState};
|
||||
|
||||
// Router registration patterns live here so AI public ingress has a single mount registry.
|
||||
@@ -9,6 +9,8 @@ use crate::{handlers::proxy::proxy_request, state::AppState};
|
||||
// which describe operational compatibility surfaces rather than the concrete axum mount list.
|
||||
const AI_POST_ROUTE_PATTERNS: &[&str] = &[
|
||||
"/v1/chat/completions",
|
||||
"/v1/embeddings",
|
||||
"/v1/rerank",
|
||||
"/v1/messages",
|
||||
"/v1/messages/count_tokens",
|
||||
"/v1/responses",
|
||||
@@ -45,6 +47,8 @@ pub(crate) fn public_api_format_local_path(api_format: &str) -> &'static str {
|
||||
openai::local_path(&normalized)
|
||||
.or_else(|| claude::local_path(&normalized))
|
||||
.or_else(|| gemini::local_path(&normalized))
|
||||
.or_else(|| jina::local_path(&normalized))
|
||||
.or_else(|| doubao::local_path(&normalized))
|
||||
.unwrap_or("/")
|
||||
}
|
||||
|
||||
@@ -53,6 +57,8 @@ pub(crate) fn normalize_admin_endpoint_signature(api_format: &str) -> Option<&'s
|
||||
openai::normalized_signature(&normalized)
|
||||
.or_else(|| claude::normalized_signature(&normalized))
|
||||
.or_else(|| gemini::normalized_signature(&normalized))
|
||||
.or_else(|| jina::normalized_signature(&normalized))
|
||||
.or_else(|| doubao::normalized_signature(&normalized))
|
||||
}
|
||||
|
||||
pub(crate) fn admin_endpoint_signature_parts(
|
||||
@@ -71,3 +77,26 @@ pub(crate) fn admin_default_body_rules_for_signature(
|
||||
let _ = provider_type;
|
||||
Some((normalized_api_format, Vec::new()))
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::{admin_endpoint_signature_parts, public_api_format_local_path};
|
||||
|
||||
#[test]
|
||||
fn supports_data_api_endpoint_signatures_and_public_paths() {
|
||||
for (api_format, family, kind, path) in [
|
||||
("openai:embedding", "openai", "embedding", "/v1/embeddings"),
|
||||
("gemini:embedding", "gemini", "embedding", "/v1/embeddings"),
|
||||
("jina:embedding", "jina", "embedding", "/v1/embeddings"),
|
||||
("doubao:embedding", "doubao", "embedding", "/v1/embeddings"),
|
||||
("openai:rerank", "openai", "rerank", "/v1/rerank"),
|
||||
("jina:rerank", "jina", "rerank", "/v1/rerank"),
|
||||
] {
|
||||
assert_eq!(
|
||||
admin_endpoint_signature_parts(api_format),
|
||||
Some((api_format, family, kind))
|
||||
);
|
||||
assert_eq!(public_api_format_local_path(api_format), path);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -109,6 +109,8 @@ pub(crate) const RUST_FRONTDOOR_OWNED_ROUTE_PATTERNS: &[&str] = &[
|
||||
"/v1beta/models",
|
||||
"/v1beta/models/{path...}",
|
||||
"/v1/chat/completions",
|
||||
"/v1/embeddings",
|
||||
"/v1/rerank",
|
||||
"/v1/images/generations",
|
||||
"/v1/images/edits",
|
||||
"/v1/images/variations",
|
||||
|
||||
@@ -16,6 +16,22 @@ pub(super) fn classify_ai_public_route(
|
||||
"openai:chat",
|
||||
true,
|
||||
))
|
||||
} else if method == http::Method::POST && normalized_path == "/v1/embeddings" {
|
||||
Some(classified(
|
||||
"ai_public",
|
||||
"openai",
|
||||
"embedding",
|
||||
"openai:embedding",
|
||||
true,
|
||||
))
|
||||
} else if method == http::Method::POST && normalized_path == "/v1/rerank" {
|
||||
Some(classified(
|
||||
"ai_public",
|
||||
"openai",
|
||||
"rerank",
|
||||
"openai:rerank",
|
||||
true,
|
||||
))
|
||||
} else if method == http::Method::POST
|
||||
&& matches!(normalized_path, "/v1/responses" | "/v1/responses/compact")
|
||||
{
|
||||
|
||||
@@ -20,6 +20,66 @@ fn classifies_claude_count_tokens_as_non_execution_runtime_public_route() {
|
||||
assert!(!decision.is_execution_runtime_candidate());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn classifies_openai_embeddings_as_embedding_not_chat() {
|
||||
let headers = headers(&[("authorization", "Bearer sk-test")]);
|
||||
let uri: Uri = "/v1/embeddings".parse().expect("uri should parse");
|
||||
let decision =
|
||||
classify_control_route(&http::Method::POST, &uri, &headers).expect("route should classify");
|
||||
|
||||
assert_eq!(decision.route_family.as_deref(), Some("openai"));
|
||||
assert_eq!(decision.route_kind.as_deref(), Some("embedding"));
|
||||
assert_ne!(decision.route_kind.as_deref(), Some("chat"));
|
||||
assert_ne!(decision.route_kind.as_deref(), Some("responses"));
|
||||
assert_eq!(
|
||||
decision.auth_endpoint_signature.as_deref(),
|
||||
Some("openai:embedding")
|
||||
);
|
||||
assert!(decision.is_execution_runtime_candidate());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn classifies_openai_rerank_as_rerank_not_chat() {
|
||||
let headers = headers(&[("authorization", "Bearer sk-test")]);
|
||||
let uri: Uri = "/v1/rerank".parse().expect("uri should parse");
|
||||
let decision =
|
||||
classify_control_route(&http::Method::POST, &uri, &headers).expect("route should classify");
|
||||
|
||||
assert_eq!(decision.route_family.as_deref(), Some("openai"));
|
||||
assert_eq!(decision.route_kind.as_deref(), Some("rerank"));
|
||||
assert_ne!(decision.route_kind.as_deref(), Some("chat"));
|
||||
assert_ne!(decision.route_kind.as_deref(), Some("embedding"));
|
||||
assert_eq!(
|
||||
decision.auth_endpoint_signature.as_deref(),
|
||||
Some("openai:rerank")
|
||||
);
|
||||
assert!(decision.is_execution_runtime_candidate());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn classifies_openai_chat_and_responses_separately_from_embedding() {
|
||||
let headers = headers(&[("authorization", "Bearer sk-test")]);
|
||||
let chat_uri: Uri = "/v1/chat/completions".parse().expect("uri should parse");
|
||||
let responses_uri: Uri = "/v1/responses".parse().expect("uri should parse");
|
||||
|
||||
let chat = classify_control_route(&http::Method::POST, &chat_uri, &headers)
|
||||
.expect("chat route should classify");
|
||||
assert_eq!(chat.route_family.as_deref(), Some("openai"));
|
||||
assert_eq!(chat.route_kind.as_deref(), Some("chat"));
|
||||
assert_eq!(chat.auth_endpoint_signature.as_deref(), Some("openai:chat"));
|
||||
assert_ne!(chat.route_kind.as_deref(), Some("embedding"));
|
||||
|
||||
let responses = classify_control_route(&http::Method::POST, &responses_uri, &headers)
|
||||
.expect("responses route should classify");
|
||||
assert_eq!(responses.route_family.as_deref(), Some("openai"));
|
||||
assert_eq!(responses.route_kind.as_deref(), Some("responses"));
|
||||
assert_eq!(
|
||||
responses.auth_endpoint_signature.as_deref(),
|
||||
Some("openai:responses")
|
||||
);
|
||||
assert_ne!(responses.route_kind.as_deref(), Some("embedding"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn classifies_models_list_as_claude_when_headers_match() {
|
||||
let headers = headers(&[
|
||||
|
||||
@@ -37,6 +37,8 @@ pub(crate) fn frontdoor_self_loop_public_ai_path(path: &str) -> bool {
|
||||
"/v1/messages"
|
||||
| "/v1/messages/count_tokens"
|
||||
| "/v1/chat/completions"
|
||||
| "/v1/embeddings"
|
||||
| "/v1/rerank"
|
||||
| "/v1/responses"
|
||||
| "/v1/responses/compact"
|
||||
| "/v1beta/files"
|
||||
|
||||
@@ -8,6 +8,38 @@ use aether_data_contracts::repository::global_models::AdminGlobalModelListQuery;
|
||||
use serde_json::json;
|
||||
use std::collections::BTreeMap;
|
||||
|
||||
const EMBEDDING_API_FORMATS: &[&str] = &[
|
||||
"openai:embedding",
|
||||
"jina:embedding",
|
||||
"gemini:embedding",
|
||||
"doubao:embedding",
|
||||
];
|
||||
|
||||
fn json_value_contains_string(value: &serde_json::Value, expected: &str) -> bool {
|
||||
match value {
|
||||
serde_json::Value::String(value) => value.trim().eq_ignore_ascii_case(expected),
|
||||
serde_json::Value::Array(values) => values
|
||||
.iter()
|
||||
.any(|value| json_value_contains_string(value, expected)),
|
||||
serde_json::Value::Object(object) => object
|
||||
.values()
|
||||
.any(|value| json_value_contains_string(value, expected)),
|
||||
_ => false,
|
||||
}
|
||||
}
|
||||
|
||||
fn json_value_contains_embedding_metadata(value: &serde_json::Value) -> bool {
|
||||
value
|
||||
.as_object()
|
||||
.and_then(|object| object.get("embedding"))
|
||||
.and_then(serde_json::Value::as_bool)
|
||||
.unwrap_or(false)
|
||||
|| json_value_contains_string(value, "embedding")
|
||||
|| EMBEDDING_API_FORMATS
|
||||
.iter()
|
||||
.any(|api_format| json_value_contains_string(value, api_format))
|
||||
}
|
||||
|
||||
pub(crate) async fn build_admin_global_model_providers_payload(
|
||||
state: &AdminAppState<'_>,
|
||||
global_model_id: &str,
|
||||
@@ -53,6 +85,7 @@ pub(crate) async fn build_admin_global_model_providers_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,
|
||||
}))
|
||||
})
|
||||
@@ -123,6 +156,14 @@ pub(crate) async fn build_admin_model_catalog_payload(
|
||||
.and_then(|value| value.get("streaming"))
|
||||
.and_then(serde_json::Value::as_bool)
|
||||
.unwrap_or(false);
|
||||
let mut supports_embedding = global_model
|
||||
.supported_capabilities
|
||||
.as_ref()
|
||||
.is_some_and(json_value_contains_embedding_metadata)
|
||||
|| global_model
|
||||
.config
|
||||
.as_ref()
|
||||
.is_some_and(json_value_contains_embedding_metadata);
|
||||
|
||||
for model in provider_models {
|
||||
let Some(provider) = provider_ids.get(&model.provider_id) else {
|
||||
@@ -143,9 +184,12 @@ pub(crate) async fn build_admin_model_catalog_payload(
|
||||
admin_provider_model_effective_capability(&model, "function_calling");
|
||||
let model_supports_streaming =
|
||||
admin_provider_model_effective_capability(&model, "streaming");
|
||||
let model_supports_embedding =
|
||||
admin_provider_model_effective_capability(&model, "embedding");
|
||||
supports_vision |= model_supports_vision;
|
||||
supports_function_calling |= model_supports_function_calling;
|
||||
supports_streaming |= model_supports_streaming;
|
||||
supports_embedding |= model_supports_embedding;
|
||||
providers.push(json!({
|
||||
"provider_id": provider.id,
|
||||
"provider_name": provider.name,
|
||||
@@ -162,6 +206,7 @@ pub(crate) async fn build_admin_model_catalog_payload(
|
||||
"supports_vision": model_supports_vision,
|
||||
"supports_function_calling": model_supports_function_calling,
|
||||
"supports_streaming": model_supports_streaming,
|
||||
"supports_embedding": model_supports_embedding,
|
||||
"is_active": model.is_active,
|
||||
}));
|
||||
}
|
||||
@@ -189,6 +234,7 @@ pub(crate) async fn build_admin_model_catalog_payload(
|
||||
"supports_vision": supports_vision,
|
||||
"supports_function_calling": supports_function_calling,
|
||||
"supports_streaming": supports_streaming,
|
||||
"supports_embedding": supports_embedding,
|
||||
}),
|
||||
}));
|
||||
}
|
||||
|
||||
@@ -1,6 +1,13 @@
|
||||
use crate::handlers::admin::shared::unix_secs_to_rfc3339;
|
||||
use aether_data_contracts::repository::global_models::StoredAdminProviderModel;
|
||||
|
||||
const EMBEDDING_API_FORMATS: &[&str] = &[
|
||||
"openai:embedding",
|
||||
"jina:embedding",
|
||||
"gemini:embedding",
|
||||
"doubao:embedding",
|
||||
];
|
||||
|
||||
pub(crate) fn model_tiered_pricing_first_tier_value(
|
||||
tiered_pricing: Option<&serde_json::Value>,
|
||||
field_name: &str,
|
||||
@@ -26,6 +33,50 @@ fn model_effective_capability(
|
||||
})
|
||||
}
|
||||
|
||||
fn value_contains_string(value: &serde_json::Value, expected: &str) -> bool {
|
||||
match value {
|
||||
serde_json::Value::String(value) => value.trim().eq_ignore_ascii_case(expected),
|
||||
serde_json::Value::Array(values) => values
|
||||
.iter()
|
||||
.any(|value| value_contains_string(value, expected)),
|
||||
serde_json::Value::Object(object) => object
|
||||
.values()
|
||||
.any(|value| value_contains_string(value, expected)),
|
||||
_ => false,
|
||||
}
|
||||
}
|
||||
|
||||
fn value_has_true_key(value: &serde_json::Value, key: &str) -> bool {
|
||||
value
|
||||
.as_object()
|
||||
.and_then(|object| object.get(key))
|
||||
.and_then(serde_json::Value::as_bool)
|
||||
.unwrap_or(false)
|
||||
}
|
||||
|
||||
fn value_contains_embedding_metadata(value: &serde_json::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)
|
||||
}
|
||||
|
||||
pub(crate) fn timestamp_or_now(value: Option<u64>, now_unix_secs: u64) -> serde_json::Value {
|
||||
unix_secs_to_rfc3339(value.unwrap_or(now_unix_secs))
|
||||
.map(serde_json::Value::String)
|
||||
@@ -110,6 +161,7 @@ pub(crate) fn admin_provider_model_effective_capability(
|
||||
model.global_model_config.as_ref(),
|
||||
"image_generation",
|
||||
),
|
||||
"embedding" => model_effective_embedding_capability(model),
|
||||
_ => false,
|
||||
}
|
||||
}
|
||||
|
||||
@@ -288,6 +288,7 @@ mod tests {
|
||||
" OPENAI:RESPONSES ".to_string(),
|
||||
"claude:messages".to_string(),
|
||||
"gemini:generate_content".to_string(),
|
||||
"jina:rerank".to_string(),
|
||||
"openai:responses".to_string(),
|
||||
]))
|
||||
.expect("formats should normalize"),
|
||||
@@ -295,6 +296,7 @@ mod tests {
|
||||
"openai:responses".to_string(),
|
||||
"claude:messages".to_string(),
|
||||
"gemini:generate_content".to_string(),
|
||||
"jina:rerank".to_string(),
|
||||
])
|
||||
);
|
||||
}
|
||||
|
||||
@@ -33,6 +33,25 @@ const OPENAI_IMAGE_INPUT_FIDELITY_DETAIL: &str = "input_fidelity 仅支持 low
|
||||
const OPENAI_IMAGE_OUTPUT_COMPRESSION_DETAIL: &str = "output_compression 必须是 0-100 的整数";
|
||||
const OPENAI_IMAGE_INVALID_JSON_DETAIL: &str = "图片接口 JSON 请求体无效";
|
||||
const OPENAI_IMAGE_INVALID_MULTIPART_DETAIL: &str = "图片接口 multipart/form-data 请求体无效";
|
||||
const OPENAI_EMBEDDING_CONTENT_TYPE_DETAIL: &str =
|
||||
"Embedding request content-type must be application/json";
|
||||
const OPENAI_EMBEDDING_INVALID_JSON_DETAIL: &str = "Embedding request JSON body is invalid";
|
||||
const OPENAI_EMBEDDING_MODEL_REQUIRED_DETAIL: &str = "Embedding request model is required";
|
||||
const OPENAI_EMBEDDING_INPUT_REQUIRED_DETAIL: &str = "Embedding request input is required";
|
||||
const OPENAI_EMBEDDING_CHAT_PAYLOAD_DETAIL: &str =
|
||||
"Embedding request must use input, not chat messages";
|
||||
const OPENAI_EMBEDDING_STREAM_UNSUPPORTED_DETAIL: &str =
|
||||
"Embedding requests do not support streaming";
|
||||
const OPENAI_RERANK_CONTENT_TYPE_DETAIL: &str =
|
||||
"Rerank request content-type must be application/json";
|
||||
const OPENAI_RERANK_INVALID_JSON_DETAIL: &str = "Rerank request JSON body is invalid";
|
||||
const OPENAI_RERANK_MODEL_REQUIRED_DETAIL: &str = "Rerank request model is required";
|
||||
const OPENAI_RERANK_QUERY_REQUIRED_DETAIL: &str = "Rerank request query is required";
|
||||
const OPENAI_RERANK_DOCUMENTS_REQUIRED_DETAIL: &str = "Rerank request documents are required";
|
||||
const OPENAI_RERANK_TOP_N_DETAIL: &str = "Rerank request top_n must be a positive integer";
|
||||
const OPENAI_RERANK_CHAT_PAYLOAD_DETAIL: &str =
|
||||
"Rerank request must use query/documents, not chat messages";
|
||||
const OPENAI_RERANK_STREAM_UNSUPPORTED_DETAIL: &str = "Rerank requests do not support streaming";
|
||||
|
||||
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
|
||||
enum OpenAiImageOperation {
|
||||
@@ -78,9 +97,15 @@ pub(crate) fn ai_public_local_requires_buffered_body(
|
||||
.as_ref()
|
||||
.is_some_and(|decision| {
|
||||
decision.route_class.as_deref() == Some("ai_public")
|
||||
&& decision.route_family.as_deref() == Some("claude")
|
||||
&& decision.route_kind.as_deref() == Some("count_tokens")
|
||||
&& request_context.request_method == http::Method::POST
|
||||
&& ((decision.route_family.as_deref() == Some("claude")
|
||||
&& decision.route_kind.as_deref() == Some("count_tokens"))
|
||||
|| (decision.route_family.as_deref() == Some("openai")
|
||||
&& decision.route_kind.as_deref() == Some("embedding")
|
||||
&& request_context.request_path == "/v1/embeddings")
|
||||
|| (decision.route_family.as_deref() == Some("openai")
|
||||
&& decision.route_kind.as_deref() == Some("rerank")
|
||||
&& request_context.request_path == "/v1/rerank"))
|
||||
})
|
||||
}
|
||||
|
||||
@@ -124,14 +149,56 @@ fn maybe_build_local_openai_request_validation_response(
|
||||
return None;
|
||||
}
|
||||
|
||||
let request_body = request_body?;
|
||||
|
||||
if decision.route_kind.as_deref() == Some("chat")
|
||||
&& request_context.request_path == "/v1/chat/completions"
|
||||
{
|
||||
return None;
|
||||
}
|
||||
|
||||
if decision.route_kind.as_deref() == Some("embedding")
|
||||
&& request_context.request_path == "/v1/embeddings"
|
||||
{
|
||||
let Some(request_body) = request_body else {
|
||||
return Some(build_ai_public_error_response(
|
||||
http::StatusCode::BAD_REQUEST,
|
||||
OPENAI_EMBEDDING_INVALID_JSON_DETAIL,
|
||||
));
|
||||
};
|
||||
if let Err(detail) = validate_openai_embedding_request(
|
||||
request_context.request_content_type.as_deref(),
|
||||
request_body,
|
||||
) {
|
||||
return Some(build_ai_public_error_response(
|
||||
http::StatusCode::BAD_REQUEST,
|
||||
detail,
|
||||
));
|
||||
}
|
||||
return None;
|
||||
}
|
||||
|
||||
if decision.route_kind.as_deref() == Some("rerank")
|
||||
&& request_context.request_path == "/v1/rerank"
|
||||
{
|
||||
let Some(request_body) = request_body else {
|
||||
return Some(build_ai_public_error_response(
|
||||
http::StatusCode::BAD_REQUEST,
|
||||
OPENAI_RERANK_INVALID_JSON_DETAIL,
|
||||
));
|
||||
};
|
||||
if let Err(detail) = validate_openai_rerank_request(
|
||||
request_context.request_content_type.as_deref(),
|
||||
request_body,
|
||||
) {
|
||||
return Some(build_ai_public_error_response(
|
||||
http::StatusCode::BAD_REQUEST,
|
||||
detail,
|
||||
));
|
||||
}
|
||||
return None;
|
||||
}
|
||||
|
||||
let request_body = request_body?;
|
||||
|
||||
if decision.route_kind.as_deref() != Some("image")
|
||||
|| !matches!(
|
||||
request_context.request_path.as_str(),
|
||||
@@ -293,6 +360,160 @@ fn maybe_build_local_openai_request_validation_response(
|
||||
None
|
||||
}
|
||||
|
||||
fn validate_openai_embedding_request(
|
||||
content_type: Option<&str>,
|
||||
request_body: &Bytes,
|
||||
) -> Result<(), &'static str> {
|
||||
if !content_type
|
||||
.unwrap_or_default()
|
||||
.to_ascii_lowercase()
|
||||
.contains("application/json")
|
||||
{
|
||||
return Err(OPENAI_EMBEDDING_CONTENT_TYPE_DETAIL);
|
||||
}
|
||||
if request_body.is_empty() {
|
||||
return Err(OPENAI_EMBEDDING_INVALID_JSON_DETAIL);
|
||||
}
|
||||
let payload = serde_json::from_slice::<Value>(request_body)
|
||||
.map_err(|_| OPENAI_EMBEDDING_INVALID_JSON_DETAIL)?;
|
||||
let object = payload
|
||||
.as_object()
|
||||
.ok_or(OPENAI_EMBEDDING_INVALID_JSON_DETAIL)?;
|
||||
if object.contains_key("messages") {
|
||||
return Err(OPENAI_EMBEDDING_CHAT_PAYLOAD_DETAIL);
|
||||
}
|
||||
if object
|
||||
.get("stream")
|
||||
.and_then(value_as_bool)
|
||||
.unwrap_or(false)
|
||||
{
|
||||
return Err(OPENAI_EMBEDDING_STREAM_UNSUPPORTED_DETAIL);
|
||||
}
|
||||
if object
|
||||
.get("model")
|
||||
.and_then(Value::as_str)
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
.is_none()
|
||||
{
|
||||
return Err(OPENAI_EMBEDDING_MODEL_REQUIRED_DETAIL);
|
||||
}
|
||||
let Some(input) = object.get("input") else {
|
||||
return Err(OPENAI_EMBEDDING_INPUT_REQUIRED_DETAIL);
|
||||
};
|
||||
if !embedding_input_is_non_empty(input) {
|
||||
return Err(OPENAI_EMBEDDING_INPUT_REQUIRED_DETAIL);
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
fn validate_openai_rerank_request(
|
||||
content_type: Option<&str>,
|
||||
request_body: &Bytes,
|
||||
) -> Result<(), &'static str> {
|
||||
if !content_type
|
||||
.unwrap_or_default()
|
||||
.to_ascii_lowercase()
|
||||
.contains("application/json")
|
||||
{
|
||||
return Err(OPENAI_RERANK_CONTENT_TYPE_DETAIL);
|
||||
}
|
||||
if request_body.is_empty() {
|
||||
return Err(OPENAI_RERANK_INVALID_JSON_DETAIL);
|
||||
}
|
||||
let payload = serde_json::from_slice::<Value>(request_body)
|
||||
.map_err(|_| OPENAI_RERANK_INVALID_JSON_DETAIL)?;
|
||||
let object = payload
|
||||
.as_object()
|
||||
.ok_or(OPENAI_RERANK_INVALID_JSON_DETAIL)?;
|
||||
if object.contains_key("messages") {
|
||||
return Err(OPENAI_RERANK_CHAT_PAYLOAD_DETAIL);
|
||||
}
|
||||
if object
|
||||
.get("stream")
|
||||
.and_then(value_as_bool)
|
||||
.unwrap_or(false)
|
||||
{
|
||||
return Err(OPENAI_RERANK_STREAM_UNSUPPORTED_DETAIL);
|
||||
}
|
||||
if object
|
||||
.get("model")
|
||||
.and_then(Value::as_str)
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
.is_none()
|
||||
{
|
||||
return Err(OPENAI_RERANK_MODEL_REQUIRED_DETAIL);
|
||||
}
|
||||
if object
|
||||
.get("query")
|
||||
.and_then(Value::as_str)
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
.is_none()
|
||||
{
|
||||
return Err(OPENAI_RERANK_QUERY_REQUIRED_DETAIL);
|
||||
}
|
||||
let Some(documents) = object.get("documents").and_then(Value::as_array) else {
|
||||
return Err(OPENAI_RERANK_DOCUMENTS_REQUIRED_DETAIL);
|
||||
};
|
||||
if documents.is_empty() || documents.iter().any(rerank_document_is_empty) {
|
||||
return Err(OPENAI_RERANK_DOCUMENTS_REQUIRED_DETAIL);
|
||||
}
|
||||
if object
|
||||
.get("top_n")
|
||||
.or_else(|| object.get("topN"))
|
||||
.is_some_and(|value| !positive_json_integer(value))
|
||||
{
|
||||
return Err(OPENAI_RERANK_TOP_N_DETAIL);
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
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 positive_json_integer(value: &Value) -> bool {
|
||||
value.as_u64().is_some_and(|number| number > 0)
|
||||
|| value.as_i64().is_some_and(|number| number > 0)
|
||||
|| value
|
||||
.as_str()
|
||||
.and_then(|text| text.trim().parse::<u64>().ok())
|
||||
.is_some_and(|number| number > 0)
|
||||
}
|
||||
|
||||
fn embedding_input_is_non_empty(value: &Value) -> bool {
|
||||
match value {
|
||||
Value::String(text) => !text.trim().is_empty(),
|
||||
Value::Array(items) if !items.is_empty() => embedding_array_input_is_non_empty(items),
|
||||
_ => false,
|
||||
}
|
||||
}
|
||||
|
||||
fn embedding_array_input_is_non_empty(items: &[Value]) -> bool {
|
||||
items
|
||||
.iter()
|
||||
.all(|item| item.as_str().is_some_and(|text| !text.trim().is_empty()))
|
||||
|| embedding_token_array_is_non_empty(items)
|
||||
|| items.iter().all(|item| {
|
||||
item.as_array()
|
||||
.is_some_and(|items| embedding_token_array_is_non_empty(items))
|
||||
})
|
||||
}
|
||||
|
||||
fn embedding_token_array_is_non_empty(items: &[Value]) -> bool {
|
||||
!items.is_empty() && items.iter().all(|item| item.as_u64().is_some())
|
||||
}
|
||||
|
||||
fn image_request_count(value: &Value) -> Option<u64> {
|
||||
value
|
||||
.as_u64()
|
||||
|
||||
@@ -210,6 +210,7 @@ fn serialize_public_catalog_model(model: StoredPublicCatalogModel) -> serde_json
|
||||
"supports_vision": model.supports_vision,
|
||||
"supports_function_calling": model.supports_function_calling,
|
||||
"supports_streaming": model.supports_streaming,
|
||||
"supports_embedding": model.supports_embedding,
|
||||
"is_active": model.is_active,
|
||||
})
|
||||
}
|
||||
|
||||
@@ -17,8 +17,14 @@ pub(crate) fn models_api_format(request_context: &GatewayPublicRequestContext) -
|
||||
"openai:responses" => Some("openai:responses"),
|
||||
"openai:responses:compact" => Some("openai:responses:compact"),
|
||||
"openai:image" => Some("openai:image"),
|
||||
"openai:embedding" => Some("openai:embedding"),
|
||||
"openai:rerank" => Some("openai:rerank"),
|
||||
"claude:messages" => Some("claude:messages"),
|
||||
"gemini:generate_content" => Some("gemini:generate_content"),
|
||||
"gemini:embedding" => Some("gemini:embedding"),
|
||||
"jina:embedding" => Some("jina:embedding"),
|
||||
"jina:rerank" => Some("jina:rerank"),
|
||||
"doubao:embedding" => Some("doubao:embedding"),
|
||||
_ => None,
|
||||
}
|
||||
}
|
||||
@@ -32,6 +38,14 @@ const MODELS_CROSS_FORMAT_QUERY_API_FORMATS: &[&str] = &[
|
||||
"gemini:generate_content",
|
||||
];
|
||||
|
||||
const MODELS_EMBEDDING_QUERY_API_FORMATS: &[&str] = &[
|
||||
"openai:embedding",
|
||||
"jina:embedding",
|
||||
"gemini:embedding",
|
||||
"doubao:embedding",
|
||||
];
|
||||
const MODELS_RERANK_QUERY_API_FORMATS: &[&str] = &["openai:rerank", "jina:rerank"];
|
||||
|
||||
pub(super) fn models_query_api_formats(api_format: &str) -> &'static [&'static str] {
|
||||
match crate::ai_serving::normalize_api_format_alias(api_format).as_str() {
|
||||
"openai:chat"
|
||||
@@ -40,6 +54,10 @@ pub(super) fn models_query_api_formats(api_format: &str) -> &'static [&'static s
|
||||
| "claude:messages"
|
||||
| "gemini:generate_content" => MODELS_CROSS_FORMAT_QUERY_API_FORMATS,
|
||||
"openai:image" => &["openai:image"],
|
||||
"openai:embedding" | "jina:embedding" | "gemini:embedding" | "doubao:embedding" => {
|
||||
MODELS_EMBEDDING_QUERY_API_FORMATS
|
||||
}
|
||||
"openai:rerank" | "jina:rerank" => MODELS_RERANK_QUERY_API_FORMATS,
|
||||
_ => &[],
|
||||
}
|
||||
}
|
||||
|
||||
@@ -362,6 +362,7 @@ pub(super) async fn handle_users_me_providers_get(
|
||||
"supports_vision": model.supports_vision,
|
||||
"supports_function_calling": model.supports_function_calling,
|
||||
"supports_streaming": model.supports_streaming,
|
||||
"supports_embedding": model.supports_embedding,
|
||||
})
|
||||
})
|
||||
.collect::<Vec<_>>(),
|
||||
|
||||
@@ -161,6 +161,7 @@ fn sample_provider_model(
|
||||
Some(global_model_name.to_string()),
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
Some(json!({ "model_mappings": mappings })),
|
||||
)
|
||||
.expect("provider model should build")
|
||||
|
||||
@@ -450,6 +450,128 @@ async fn gateway_handles_admin_model_catalog_locally_with_trusted_admin_principa
|
||||
upstream_handle.abort();
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn admin_global_models_include_embedding_capability() {
|
||||
let upstream_hits = Arc::new(Mutex::new(0usize));
|
||||
let upstream_hits_clone = Arc::clone(&upstream_hits);
|
||||
let upstream = Router::new().route(
|
||||
"/{*path}",
|
||||
any(move |_request: Request| {
|
||||
let upstream_hits_inner = Arc::clone(&upstream_hits_clone);
|
||||
async move {
|
||||
*upstream_hits_inner.lock().expect("mutex should lock") += 1;
|
||||
(StatusCode::OK, Body::from("unexpected upstream hit"))
|
||||
}
|
||||
}),
|
||||
);
|
||||
|
||||
let mut global_model = sample_admin_global_model(
|
||||
"global-embedding-small",
|
||||
"text-embedding-3-small",
|
||||
"Text Embedding 3 Small",
|
||||
);
|
||||
global_model.supported_capabilities = Some(json!(["embedding"]));
|
||||
global_model.config = Some(json!({
|
||||
"api_formats": ["openai:embedding"],
|
||||
"dimensions": 1536,
|
||||
"model_type": "embedding"
|
||||
}));
|
||||
let mut provider_model = sample_admin_provider_model(
|
||||
"model-openai-embedding-small",
|
||||
"provider-openai",
|
||||
"global-embedding-small",
|
||||
"text-embedding-3-small",
|
||||
);
|
||||
provider_model.provider_model_mappings = Some(json!([{
|
||||
"name": "text-embedding-3-small",
|
||||
"priority": 1,
|
||||
"api_formats": ["openai:embedding"]
|
||||
}]));
|
||||
provider_model.config = Some(json!({
|
||||
"api_formats": ["openai:embedding"],
|
||||
"dimensions": 1536,
|
||||
"model_type": "embedding"
|
||||
}));
|
||||
provider_model.supports_streaming = Some(false);
|
||||
provider_model.global_model_name = Some("text-embedding-3-small".to_string());
|
||||
provider_model.global_model_display_name = Some("Text Embedding 3 Small".to_string());
|
||||
provider_model.global_model_supported_capabilities = Some(json!(["embedding"]));
|
||||
provider_model.global_model_config = global_model.config.clone();
|
||||
|
||||
let provider_catalog_repository = Arc::new(InMemoryProviderCatalogReadRepository::seed(
|
||||
vec![sample_provider("provider-openai", "openai", 10)],
|
||||
Vec::new(),
|
||||
Vec::new(),
|
||||
));
|
||||
let global_model_repository = Arc::new(
|
||||
InMemoryGlobalModelReadRepository::seed(Vec::new())
|
||||
.with_admin_global_models(vec![global_model])
|
||||
.with_admin_provider_models(vec![provider_model]),
|
||||
);
|
||||
|
||||
let (upstream_url, upstream_handle) = start_server(upstream).await;
|
||||
let gateway = build_router_with_state(
|
||||
AppState::new()
|
||||
.expect("gateway should build")
|
||||
.with_data_state_for_tests(
|
||||
GatewayDataState::with_provider_catalog_reader_for_tests(
|
||||
provider_catalog_repository,
|
||||
)
|
||||
.with_global_model_repository_for_tests(global_model_repository),
|
||||
),
|
||||
);
|
||||
let (gateway_url, gateway_handle) = start_server(gateway).await;
|
||||
let client = reqwest::Client::new();
|
||||
|
||||
let list_response = client
|
||||
.get(format!("{gateway_url}/api/admin/models/global?limit=20"))
|
||||
.header(crate::constants::GATEWAY_HEADER, "rust-phase3b")
|
||||
.header(TRUSTED_ADMIN_USER_ID_HEADER, "admin-user-123")
|
||||
.header(TRUSTED_ADMIN_USER_ROLE_HEADER, "admin")
|
||||
.header(TRUSTED_ADMIN_SESSION_ID_HEADER, "session-123")
|
||||
.send()
|
||||
.await
|
||||
.expect("request should succeed");
|
||||
assert_eq!(list_response.status(), StatusCode::OK);
|
||||
let list_payload: serde_json::Value =
|
||||
list_response.json().await.expect("json body should parse");
|
||||
assert_eq!(
|
||||
list_payload["models"][0]["supported_capabilities"],
|
||||
json!(["embedding"])
|
||||
);
|
||||
assert_eq!(
|
||||
list_payload["models"][0]["config"]["api_formats"],
|
||||
json!(["openai:embedding"])
|
||||
);
|
||||
|
||||
let catalog_response = client
|
||||
.get(format!("{gateway_url}/api/admin/models/catalog"))
|
||||
.header(crate::constants::GATEWAY_HEADER, "rust-phase3b")
|
||||
.header(TRUSTED_ADMIN_USER_ID_HEADER, "admin-user-123")
|
||||
.header(TRUSTED_ADMIN_USER_ROLE_HEADER, "admin")
|
||||
.header(TRUSTED_ADMIN_SESSION_ID_HEADER, "session-123")
|
||||
.send()
|
||||
.await
|
||||
.expect("request should succeed");
|
||||
assert_eq!(catalog_response.status(), StatusCode::OK);
|
||||
let catalog_payload: serde_json::Value = catalog_response
|
||||
.json()
|
||||
.await
|
||||
.expect("json body should parse");
|
||||
assert_eq!(
|
||||
catalog_payload["models"][0]["capabilities"]["supports_embedding"],
|
||||
true
|
||||
);
|
||||
assert_eq!(
|
||||
catalog_payload["models"][0]["providers"][0]["supports_embedding"],
|
||||
true
|
||||
);
|
||||
assert_eq!(*upstream_hits.lock().expect("mutex should lock"), 0);
|
||||
|
||||
gateway_handle.abort();
|
||||
upstream_handle.abort();
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn gateway_returns_service_unavailable_for_admin_model_catalog_without_required_readers() {
|
||||
let upstream_hits = Arc::new(Mutex::new(0usize));
|
||||
@@ -859,8 +981,8 @@ async fn gateway_creates_admin_global_model_locally_with_trusted_admin_principal
|
||||
"output_price_per_1m": 24.0
|
||||
}]
|
||||
},
|
||||
"supported_capabilities": ["streaming", "vision"],
|
||||
"config": {"streaming": true}
|
||||
"supported_capabilities": ["streaming", "vision", "embedding"],
|
||||
"config": {"streaming": true, "api_formats": ["openai:embedding"], "model_type": "embedding"}
|
||||
}))
|
||||
.send()
|
||||
.await
|
||||
@@ -870,6 +992,14 @@ async fn gateway_creates_admin_global_model_locally_with_trusted_admin_principal
|
||||
let payload: serde_json::Value = response.json().await.expect("json body should parse");
|
||||
assert_eq!(payload["name"], "gpt-5-pro");
|
||||
assert_eq!(payload["display_name"], "GPT 5 Pro");
|
||||
assert_eq!(
|
||||
payload["supported_capabilities"],
|
||||
json!(["streaming", "vision", "embedding"])
|
||||
);
|
||||
assert_eq!(
|
||||
payload["config"]["api_formats"],
|
||||
json!(["openai:embedding"])
|
||||
);
|
||||
assert_eq!(*upstream_hits.lock().expect("mutex should lock"), 0);
|
||||
|
||||
let created = global_model_repository
|
||||
@@ -878,6 +1008,10 @@ async fn gateway_creates_admin_global_model_locally_with_trusted_admin_principal
|
||||
.expect("model lookup should succeed")
|
||||
.expect("model should exist");
|
||||
assert_eq!(created.display_name, "GPT 5 Pro");
|
||||
assert_eq!(
|
||||
created.supported_capabilities,
|
||||
Some(json!(["streaming", "vision", "embedding"]))
|
||||
);
|
||||
|
||||
gateway_handle.abort();
|
||||
upstream_handle.abort();
|
||||
@@ -926,7 +1060,8 @@ async fn gateway_updates_and_deletes_admin_global_model_locally_with_trusted_adm
|
||||
.json(&json!({
|
||||
"display_name": "GPT 5 Updated",
|
||||
"is_active": false,
|
||||
"config": {"streaming": false}
|
||||
"supported_capabilities": ["embedding"],
|
||||
"config": {"streaming": false, "api_formats": ["openai:embedding"], "dimensions": 1536}
|
||||
}))
|
||||
.send()
|
||||
.await
|
||||
@@ -939,6 +1074,14 @@ async fn gateway_updates_and_deletes_admin_global_model_locally_with_trusted_adm
|
||||
.expect("json body should parse");
|
||||
assert_eq!(update_payload["display_name"], "GPT 5 Updated");
|
||||
assert_eq!(update_payload["is_active"], false);
|
||||
assert_eq!(
|
||||
update_payload["supported_capabilities"],
|
||||
json!(["embedding"])
|
||||
);
|
||||
assert_eq!(
|
||||
update_payload["config"]["api_formats"],
|
||||
json!(["openai:embedding"])
|
||||
);
|
||||
|
||||
let delete_response = reqwest::Client::new()
|
||||
.delete(format!(
|
||||
|
||||
@@ -233,7 +233,8 @@ async fn gateway_creates_admin_provider_model_locally_with_trusted_admin_princip
|
||||
"provider_model_name": "gpt-5-upstream",
|
||||
"global_model_id": "global-gpt-5",
|
||||
"supports_vision": true,
|
||||
"config": {"provider_hint": "gpt-5-upstream"}
|
||||
"provider_model_mappings": [{"name": "text-embedding-3-small", "priority": 1, "api_formats": ["openai:embedding"]}],
|
||||
"config": {"provider_hint": "gpt-5-upstream", "api_formats": ["openai:embedding"], "model_type": "embedding"}
|
||||
}))
|
||||
.send()
|
||||
.await
|
||||
@@ -245,6 +246,11 @@ async fn gateway_creates_admin_provider_model_locally_with_trusted_admin_princip
|
||||
assert_eq!(payload["global_model_id"], "global-gpt-5");
|
||||
assert_eq!(payload["provider_model_name"], "gpt-5-upstream");
|
||||
assert_eq!(payload["effective_supports_vision"], true);
|
||||
assert_eq!(payload["effective_supports_embedding"], true);
|
||||
assert_eq!(
|
||||
payload["provider_model_mappings"][0]["api_formats"],
|
||||
json!(["openai:embedding"])
|
||||
);
|
||||
assert_eq!(*upstream_hits.lock().expect("mutex should lock"), 0);
|
||||
|
||||
let created = global_model_repository
|
||||
@@ -258,6 +264,13 @@ async fn gateway_creates_admin_provider_model_locally_with_trusted_admin_princip
|
||||
.expect("models should read");
|
||||
assert_eq!(created.len(), 1);
|
||||
assert_eq!(created[0].provider_model_name, "gpt-5-upstream");
|
||||
assert_eq!(
|
||||
created[0]
|
||||
.config
|
||||
.as_ref()
|
||||
.and_then(|value| value.get("api_formats")),
|
||||
Some(&json!(["openai:embedding"]))
|
||||
);
|
||||
|
||||
gateway_handle.abort();
|
||||
upstream_handle.abort();
|
||||
@@ -320,8 +333,10 @@ async fn gateway_updates_and_deletes_admin_provider_model_locally_with_trusted_a
|
||||
.json(&json!({
|
||||
"provider_model_name": "gpt-5-mini-upstream",
|
||||
"global_model_id": "global-gpt-5-mini",
|
||||
"provider_model_mappings": [{"name": "text-embedding-3-small", "priority": 1, "api_formats": ["openai:embedding"]}],
|
||||
"supports_streaming": false,
|
||||
"is_available": false
|
||||
"is_available": false,
|
||||
"config": {"api_formats": ["openai:embedding"], "model_type": "embedding"}
|
||||
}))
|
||||
.send()
|
||||
.await
|
||||
@@ -334,6 +349,7 @@ async fn gateway_updates_and_deletes_admin_provider_model_locally_with_trusted_a
|
||||
assert_eq!(update_payload["provider_model_name"], "gpt-5-mini-upstream");
|
||||
assert_eq!(update_payload["global_model_id"], "global-gpt-5-mini");
|
||||
assert_eq!(update_payload["is_available"], false);
|
||||
assert_eq!(update_payload["effective_supports_embedding"], true);
|
||||
|
||||
let delete_response = reqwest::Client::new()
|
||||
.delete(format!(
|
||||
@@ -456,27 +472,37 @@ async fn gateway_handles_admin_provider_available_source_models_locally_with_tru
|
||||
Vec::new(),
|
||||
Vec::new(),
|
||||
));
|
||||
let mut global_model = sample_admin_global_model(
|
||||
"global-gpt-5",
|
||||
"text-embedding-3-small",
|
||||
"Text Embedding 3 Small",
|
||||
);
|
||||
global_model.supported_capabilities = Some(json!(["embedding"]));
|
||||
global_model.config = Some(json!({"api_formats": ["openai:embedding"]}));
|
||||
let mut primary_model = sample_admin_provider_model(
|
||||
"model-openai-gpt5",
|
||||
"provider-openai",
|
||||
"global-gpt-5",
|
||||
"text-embedding-3-small",
|
||||
);
|
||||
primary_model.global_model_name = Some("text-embedding-3-small".to_string());
|
||||
primary_model.global_model_display_name = Some("Text Embedding 3 Small".to_string());
|
||||
primary_model.global_model_supported_capabilities = Some(json!(["embedding"]));
|
||||
primary_model.global_model_config = Some(json!({"api_formats": ["openai:embedding"]}));
|
||||
let mut alternate_model = sample_admin_provider_model(
|
||||
"model-openai-gpt5-b",
|
||||
"provider-openai",
|
||||
"global-gpt-5",
|
||||
"gpt-5-alt",
|
||||
);
|
||||
alternate_model.global_model_name = Some("text-embedding-3-small".to_string());
|
||||
alternate_model.global_model_display_name = Some("Text Embedding 3 Small".to_string());
|
||||
alternate_model.global_model_supported_capabilities = Some(json!(["embedding"]));
|
||||
alternate_model.global_model_config = Some(json!({"api_formats": ["openai:embedding"]}));
|
||||
let global_model_repository = Arc::new(
|
||||
InMemoryGlobalModelReadRepository::seed(Vec::new())
|
||||
.with_admin_global_models(vec![sample_admin_global_model(
|
||||
"global-gpt-5",
|
||||
"gpt-5",
|
||||
"GPT 5",
|
||||
)])
|
||||
.with_admin_provider_models(vec![
|
||||
sample_admin_provider_model(
|
||||
"model-openai-gpt5",
|
||||
"provider-openai",
|
||||
"global-gpt-5",
|
||||
"gpt-5-upstream",
|
||||
),
|
||||
sample_admin_provider_model(
|
||||
"model-openai-gpt5-b",
|
||||
"provider-openai",
|
||||
"global-gpt-5",
|
||||
"gpt-5-alt",
|
||||
),
|
||||
]),
|
||||
.with_admin_global_models(vec![global_model])
|
||||
.with_admin_provider_models(vec![primary_model, alternate_model]),
|
||||
);
|
||||
|
||||
let (upstream_url, upstream_handle) = start_server(upstream).await;
|
||||
@@ -507,7 +533,14 @@ async fn gateway_handles_admin_provider_available_source_models_locally_with_tru
|
||||
assert_eq!(response.status(), StatusCode::OK);
|
||||
let payload: serde_json::Value = response.json().await.expect("json body should parse");
|
||||
assert_eq!(payload["total"], 1);
|
||||
assert_eq!(payload["models"][0]["global_model_name"], "gpt-5");
|
||||
assert_eq!(
|
||||
payload["models"][0]["global_model_name"],
|
||||
"text-embedding-3-small"
|
||||
);
|
||||
assert_eq!(
|
||||
payload["models"][0]["capabilities"]["supports_embedding"],
|
||||
true
|
||||
);
|
||||
assert_eq!(*upstream_hits.lock().expect("mutex should lock"), 0);
|
||||
|
||||
gateway_handle.abort();
|
||||
|
||||
@@ -1078,6 +1078,12 @@ async fn gateway_handles_admin_system_api_formats_locally_with_trusted_admin_pri
|
||||
.expect("formats should be an array");
|
||||
assert_eq!(formats[0]["value"], "openai:chat");
|
||||
assert_eq!(formats[0]["default_path"], "/v1/chat/completions");
|
||||
assert!(formats
|
||||
.iter()
|
||||
.any(|item| item["value"] == "openai:embedding"));
|
||||
assert!(formats.iter().any(|item| item["value"] == "openai:rerank"));
|
||||
assert!(formats.iter().any(|item| item["value"] == "jina:embedding"));
|
||||
assert!(formats.iter().any(|item| item["value"] == "jina:rerank"));
|
||||
assert!(formats.iter().any(|item| item["value"] == "gemini:video"));
|
||||
assert_eq!(*upstream_hits.lock().expect("mutex should lock"), 0);
|
||||
|
||||
|
||||
@@ -330,6 +330,7 @@ pub(super) fn sample_admin_provider_model(
|
||||
"output_price_per_1m": 20.0,
|
||||
}]
|
||||
})),
|
||||
Some(json!(["streaming", "vision"])),
|
||||
Some(json!({"streaming": true, "vision": false, "billing": {"currency": "USD"}})),
|
||||
)
|
||||
.expect("admin provider model should build")
|
||||
|
||||
@@ -0,0 +1,423 @@
|
||||
use std::collections::BTreeMap;
|
||||
use std::sync::Arc;
|
||||
|
||||
use aether_contracts::{ExecutionPlan, ExecutionResult, ResponseBody};
|
||||
use aether_crypto::DEVELOPMENT_ENCRYPTION_KEY;
|
||||
use aether_data::repository::candidate_selection::InMemoryMinimalCandidateSelectionReadRepository;
|
||||
use http::StatusCode;
|
||||
use serde_json::json;
|
||||
|
||||
use super::super::{
|
||||
any, build_router_with_state, build_state_with_execution_runtime_override, hash_api_key,
|
||||
sample_currently_usable_auth_snapshot, sample_endpoint, sample_key, sample_provider,
|
||||
start_server, AppState, GatewayDataState, InMemoryAuthApiKeySnapshotRepository,
|
||||
InMemoryProviderCatalogReadRepository, Json, Router,
|
||||
};
|
||||
use crate::constants::{
|
||||
CONTROL_ENDPOINT_SIGNATURE_HEADER, CONTROL_EXECUTION_RUNTIME_HEADER,
|
||||
CONTROL_ROUTE_FAMILY_HEADER, CONTROL_ROUTE_KIND_HEADER, EXECUTION_PATH_HEADER,
|
||||
EXECUTION_PATH_LOCAL_AUTH_DENIED,
|
||||
};
|
||||
use aether_data_contracts::repository::candidate_selection::StoredMinimalCandidateSelectionRow;
|
||||
|
||||
fn embedding_success_state(execution_runtime_url: String) -> AppState {
|
||||
let mut snapshot =
|
||||
sample_currently_usable_auth_snapshot("key-embedding-success", "user-embedding-success");
|
||||
snapshot.user_allowed_providers = None;
|
||||
snapshot.api_key_allowed_providers = None;
|
||||
snapshot.user_allowed_api_formats = Some(vec!["openai:embedding".to_string()]);
|
||||
snapshot.api_key_allowed_api_formats = Some(vec!["openai:embedding".to_string()]);
|
||||
snapshot.user_allowed_models = Some(vec!["text-embedding-3-small".to_string()]);
|
||||
snapshot.api_key_allowed_models = Some(vec!["text-embedding-3-small".to_string()]);
|
||||
let auth_repository = Arc::new(InMemoryAuthApiKeySnapshotRepository::seed(vec![(
|
||||
Some(hash_api_key("sk-embedding-success")),
|
||||
snapshot,
|
||||
)]));
|
||||
let candidate_repository =
|
||||
Arc::new(InMemoryMinimalCandidateSelectionReadRepository::seed(vec![
|
||||
embedding_candidate_row(),
|
||||
]));
|
||||
let provider_catalog_repository = Arc::new(InMemoryProviderCatalogReadRepository::seed(
|
||||
vec![sample_provider(
|
||||
"provider-embedding",
|
||||
"OpenAI Embeddings",
|
||||
1,
|
||||
)],
|
||||
vec![sample_endpoint(
|
||||
"endpoint-embedding",
|
||||
"provider-embedding",
|
||||
"openai:embedding",
|
||||
"https://api.openai.example",
|
||||
)],
|
||||
vec![sample_key(
|
||||
"key-upstream-embedding",
|
||||
"provider-embedding",
|
||||
"openai:embedding",
|
||||
"sk-upstream-embedding",
|
||||
)],
|
||||
));
|
||||
let data_state =
|
||||
GatewayDataState::with_provider_catalog_and_minimal_candidate_selection_for_tests(
|
||||
provider_catalog_repository,
|
||||
candidate_repository,
|
||||
)
|
||||
.with_auth_api_key_reader(auth_repository)
|
||||
.with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY);
|
||||
|
||||
build_state_with_execution_runtime_override(execution_runtime_url)
|
||||
.with_data_state_for_tests(data_state)
|
||||
}
|
||||
|
||||
fn embedding_execution_runtime() -> Router {
|
||||
Router::new().route(
|
||||
"/v1/execute/sync",
|
||||
any(|Json(plan): Json<ExecutionPlan>| async move {
|
||||
assert_embedding_execution_plan(&plan);
|
||||
Json(embedding_execution_result(&plan))
|
||||
}),
|
||||
)
|
||||
}
|
||||
|
||||
fn embedding_candidate_row() -> StoredMinimalCandidateSelectionRow {
|
||||
StoredMinimalCandidateSelectionRow {
|
||||
provider_id: "provider-embedding".to_string(),
|
||||
provider_name: "OpenAI Embeddings".to_string(),
|
||||
provider_type: "custom".to_string(),
|
||||
provider_priority: 1,
|
||||
provider_is_active: true,
|
||||
endpoint_id: "endpoint-embedding".to_string(),
|
||||
endpoint_api_format: "openai:embedding".to_string(),
|
||||
endpoint_api_family: Some("openai".to_string()),
|
||||
endpoint_kind: Some("embedding".to_string()),
|
||||
endpoint_is_active: true,
|
||||
key_id: "key-upstream-embedding".to_string(),
|
||||
key_name: "default".to_string(),
|
||||
key_auth_type: "api_key".to_string(),
|
||||
key_is_active: true,
|
||||
key_api_formats: Some(vec!["openai:embedding".to_string()]),
|
||||
key_allowed_models: None,
|
||||
key_capabilities: None,
|
||||
key_internal_priority: 50,
|
||||
key_global_priority_by_format: None,
|
||||
model_id: "model-embedding-small".to_string(),
|
||||
global_model_id: "global-embedding-small".to_string(),
|
||||
global_model_name: "text-embedding-3-small".to_string(),
|
||||
global_model_mappings: None,
|
||||
global_model_supports_streaming: Some(false),
|
||||
model_provider_model_name: "upstream-embedding".to_string(),
|
||||
model_provider_model_mappings: None,
|
||||
model_supports_streaming: Some(false),
|
||||
model_is_active: true,
|
||||
model_is_available: true,
|
||||
}
|
||||
}
|
||||
|
||||
fn assert_embedding_execution_plan(plan: &ExecutionPlan) {
|
||||
assert_eq!(plan.client_api_format, "openai:embedding");
|
||||
assert_eq!(plan.provider_api_format, "openai:embedding");
|
||||
assert_eq!(plan.method, "POST");
|
||||
assert_eq!(plan.url, "https://api.openai.example/v1/embeddings");
|
||||
assert_eq!(plan.model_name.as_deref(), Some("text-embedding-3-small"));
|
||||
let body = plan.body.json_body.as_ref().expect("json request body");
|
||||
assert_eq!(body["model"], "upstream-embedding");
|
||||
assert!(body.get("input").is_some());
|
||||
}
|
||||
|
||||
fn embedding_execution_result(plan: &ExecutionPlan) -> ExecutionResult {
|
||||
ExecutionResult {
|
||||
request_id: plan.request_id.clone(),
|
||||
candidate_id: plan.candidate_id.clone(),
|
||||
status_code: 200,
|
||||
headers: BTreeMap::from([("content-type".to_string(), "application/json".to_string())]),
|
||||
body: Some(ResponseBody {
|
||||
json_body: Some(json!({
|
||||
"object": "list",
|
||||
"model": "upstream-embedding",
|
||||
"data": [
|
||||
{"object": "embedding", "index": 0, "embedding": [0.1, 0.2, 0.3]}
|
||||
],
|
||||
"usage": {"prompt_tokens": 4, "total_tokens": 4}
|
||||
})),
|
||||
body_bytes_b64: None,
|
||||
}),
|
||||
telemetry: None,
|
||||
error: None,
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn embeddings_route_accepts_openai_payload() {
|
||||
let (execution_runtime_url, execution_runtime_handle) =
|
||||
start_server(embedding_execution_runtime()).await;
|
||||
let gateway = build_router_with_state(embedding_success_state(execution_runtime_url));
|
||||
let (gateway_url, gateway_handle) = start_server(gateway).await;
|
||||
|
||||
let response = reqwest::Client::new()
|
||||
.post(format!("{gateway_url}/v1/embeddings"))
|
||||
.header(http::header::AUTHORIZATION, "Bearer sk-embedding-success")
|
||||
.json(&json!({
|
||||
"model": "text-embedding-3-small",
|
||||
"input": ["hello", "world"]
|
||||
}))
|
||||
.send()
|
||||
.await
|
||||
.expect("request should succeed");
|
||||
|
||||
assert_eq!(response.status(), StatusCode::OK);
|
||||
assert_eq!(
|
||||
response
|
||||
.headers()
|
||||
.get(CONTROL_ROUTE_FAMILY_HEADER)
|
||||
.and_then(|value| value.to_str().ok()),
|
||||
Some("openai")
|
||||
);
|
||||
assert_eq!(
|
||||
response
|
||||
.headers()
|
||||
.get(CONTROL_ROUTE_KIND_HEADER)
|
||||
.and_then(|value| value.to_str().ok()),
|
||||
Some("embedding")
|
||||
);
|
||||
assert_ne!(
|
||||
response
|
||||
.headers()
|
||||
.get(CONTROL_ROUTE_KIND_HEADER)
|
||||
.and_then(|value| value.to_str().ok()),
|
||||
Some("chat")
|
||||
);
|
||||
assert_ne!(
|
||||
response
|
||||
.headers()
|
||||
.get(CONTROL_ROUTE_KIND_HEADER)
|
||||
.and_then(|value| value.to_str().ok()),
|
||||
Some("responses")
|
||||
);
|
||||
assert_eq!(
|
||||
response
|
||||
.headers()
|
||||
.get(CONTROL_ENDPOINT_SIGNATURE_HEADER)
|
||||
.and_then(|value| value.to_str().ok()),
|
||||
Some("openai:embedding")
|
||||
);
|
||||
assert_eq!(
|
||||
response
|
||||
.headers()
|
||||
.get(CONTROL_EXECUTION_RUNTIME_HEADER)
|
||||
.and_then(|value| value.to_str().ok()),
|
||||
Some("true")
|
||||
);
|
||||
let payload: serde_json::Value = response.json().await.expect("body should parse");
|
||||
assert_eq!(payload["object"], "list");
|
||||
assert_eq!(payload["data"][0]["object"], "embedding");
|
||||
assert_eq!(payload["data"][0]["embedding"], json!([0.1, 0.2, 0.3]));
|
||||
|
||||
gateway_handle.abort();
|
||||
execution_runtime_handle.abort();
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn embeddings_route_accepts_all_canonical_input_shapes() {
|
||||
let (execution_runtime_url, execution_runtime_handle) =
|
||||
start_server(embedding_execution_runtime()).await;
|
||||
let gateway = build_router_with_state(embedding_success_state(execution_runtime_url));
|
||||
let (gateway_url, gateway_handle) = start_server(gateway).await;
|
||||
let client = reqwest::Client::new();
|
||||
|
||||
for input in [
|
||||
json!("hello"),
|
||||
json!(["hello", "world"]),
|
||||
json!([1, 2, 3]),
|
||||
json!([[1, 2], [3, 4]]),
|
||||
] {
|
||||
let response = client
|
||||
.post(format!("{gateway_url}/v1/embeddings"))
|
||||
.header(http::header::AUTHORIZATION, "Bearer sk-embedding-success")
|
||||
.json(&json!({
|
||||
"model": "text-embedding-3-small",
|
||||
"input": input
|
||||
}))
|
||||
.send()
|
||||
.await
|
||||
.expect("request should succeed");
|
||||
|
||||
assert_eq!(response.status(), StatusCode::OK);
|
||||
assert_eq!(
|
||||
response
|
||||
.headers()
|
||||
.get(CONTROL_ENDPOINT_SIGNATURE_HEADER)
|
||||
.and_then(|value| value.to_str().ok()),
|
||||
Some("openai:embedding")
|
||||
);
|
||||
let payload: serde_json::Value = response.json().await.expect("body should parse");
|
||||
assert_eq!(payload["data"][0]["embedding"], json!([0.1, 0.2, 0.3]));
|
||||
}
|
||||
|
||||
gateway_handle.abort();
|
||||
execution_runtime_handle.abort();
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn embeddings_route_rejects_invalid_local_payloads() {
|
||||
let gateway = build_router_with_state(AppState::new().expect("gateway should build"));
|
||||
let (gateway_url, gateway_handle) = start_server(gateway).await;
|
||||
let client = reqwest::Client::new();
|
||||
let cases = [
|
||||
("{", "Embedding request JSON body is invalid"),
|
||||
(
|
||||
r#"{"input":"hello"}"#,
|
||||
"Embedding request model is required",
|
||||
),
|
||||
(
|
||||
r#"{"model":"text-embedding-3-small","input":[]}"#,
|
||||
"Embedding request input is required",
|
||||
),
|
||||
(
|
||||
r#"{"model":"text-embedding-3-small","messages":[]}"#,
|
||||
"Embedding request must use input, not chat messages",
|
||||
),
|
||||
(
|
||||
r#"{"model":" ","input":"hello"}"#,
|
||||
"Embedding request model is required",
|
||||
),
|
||||
(
|
||||
r#"{"model":"text-embedding-3-small","input":[[1],[]]}"#,
|
||||
"Embedding request input is required",
|
||||
),
|
||||
(
|
||||
r#"{"model":"text-embedding-3-small","input":"hello","stream":true}"#,
|
||||
"Embedding requests do not support streaming",
|
||||
),
|
||||
];
|
||||
|
||||
for (body, expected_detail) in cases {
|
||||
let response = client
|
||||
.post(format!("{gateway_url}/v1/embeddings"))
|
||||
.header(http::header::CONTENT_TYPE, "application/json")
|
||||
.body(body)
|
||||
.send()
|
||||
.await
|
||||
.expect("request should succeed");
|
||||
|
||||
assert_eq!(response.status(), StatusCode::BAD_REQUEST);
|
||||
assert_eq!(
|
||||
response
|
||||
.headers()
|
||||
.get(CONTROL_ROUTE_KIND_HEADER)
|
||||
.and_then(|value| value.to_str().ok()),
|
||||
Some("embedding")
|
||||
);
|
||||
let payload: serde_json::Value = response.json().await.expect("body should parse");
|
||||
assert_eq!(payload["detail"], expected_detail);
|
||||
}
|
||||
|
||||
gateway_handle.abort();
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn embeddings_route_rejects_non_json_content_type() {
|
||||
let gateway = build_router_with_state(AppState::new().expect("gateway should build"));
|
||||
let (gateway_url, gateway_handle) = start_server(gateway).await;
|
||||
|
||||
let response = reqwest::Client::new()
|
||||
.post(format!("{gateway_url}/v1/embeddings"))
|
||||
.header(http::header::CONTENT_TYPE, "text/plain")
|
||||
.body(r#"{"model":"text-embedding-3-small","input":"hello"}"#)
|
||||
.send()
|
||||
.await
|
||||
.expect("request should succeed");
|
||||
|
||||
assert_eq!(response.status(), StatusCode::BAD_REQUEST);
|
||||
let payload: serde_json::Value = response.json().await.expect("body should parse");
|
||||
assert_eq!(
|
||||
payload["detail"],
|
||||
"Embedding request content-type must be application/json"
|
||||
);
|
||||
|
||||
gateway_handle.abort();
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn embeddings_route_rejects_chat_only_model() {
|
||||
let mut snapshot = sample_currently_usable_auth_snapshot("key-embedding-1", "user-embedding-1");
|
||||
snapshot.user_allowed_api_formats = Some(vec!["openai:embedding".to_string()]);
|
||||
snapshot.api_key_allowed_api_formats = Some(vec!["openai:embedding".to_string()]);
|
||||
snapshot.user_allowed_models = Some(vec!["text-embedding-3-small".to_string()]);
|
||||
snapshot.api_key_allowed_models = Some(vec!["text-embedding-3-small".to_string()]);
|
||||
let repository = Arc::new(InMemoryAuthApiKeySnapshotRepository::seed(vec![(
|
||||
Some(hash_api_key("sk-embedding-model-guard")),
|
||||
snapshot,
|
||||
)]));
|
||||
let gateway = build_router_with_state(
|
||||
AppState::new()
|
||||
.expect("gateway should build")
|
||||
.with_auth_api_key_data_reader_for_tests(repository),
|
||||
);
|
||||
let (gateway_url, gateway_handle) = start_server(gateway).await;
|
||||
|
||||
let response = reqwest::Client::new()
|
||||
.post(format!("{gateway_url}/v1/embeddings"))
|
||||
.header(
|
||||
http::header::AUTHORIZATION,
|
||||
"Bearer sk-embedding-model-guard",
|
||||
)
|
||||
.json(&json!({
|
||||
"model": "gpt-5",
|
||||
"input": "hello"
|
||||
}))
|
||||
.send()
|
||||
.await
|
||||
.expect("request should succeed");
|
||||
|
||||
assert_eq!(response.status(), StatusCode::FORBIDDEN);
|
||||
assert_eq!(
|
||||
response
|
||||
.headers()
|
||||
.get(EXECUTION_PATH_HEADER)
|
||||
.and_then(|value| value.to_str().ok()),
|
||||
Some(EXECUTION_PATH_LOCAL_AUTH_DENIED)
|
||||
);
|
||||
let payload: serde_json::Value = response.json().await.expect("body should parse");
|
||||
assert_eq!(payload["error"]["message"], "当前密钥不允许访问模型 gpt-5");
|
||||
|
||||
gateway_handle.abort();
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn embeddings_route_rejects_chat_only_api_format() {
|
||||
let mut snapshot = sample_currently_usable_auth_snapshot("key-embedding-2", "user-embedding-2");
|
||||
snapshot.user_allowed_api_formats = Some(vec!["openai:chat".to_string()]);
|
||||
snapshot.api_key_allowed_api_formats = Some(vec!["openai:chat".to_string()]);
|
||||
let repository = Arc::new(InMemoryAuthApiKeySnapshotRepository::seed(vec![(
|
||||
Some(hash_api_key("sk-embedding-format-guard")),
|
||||
snapshot,
|
||||
)]));
|
||||
let gateway = build_router_with_state(
|
||||
AppState::new()
|
||||
.expect("gateway should build")
|
||||
.with_auth_api_key_data_reader_for_tests(repository),
|
||||
);
|
||||
let (gateway_url, gateway_handle) = start_server(gateway).await;
|
||||
|
||||
let response = reqwest::Client::new()
|
||||
.post(format!("{gateway_url}/v1/embeddings"))
|
||||
.header(
|
||||
http::header::AUTHORIZATION,
|
||||
"Bearer sk-embedding-format-guard",
|
||||
)
|
||||
.json(&json!({
|
||||
"model": "text-embedding-3-small",
|
||||
"input": "hello"
|
||||
}))
|
||||
.send()
|
||||
.await
|
||||
.expect("request should succeed");
|
||||
|
||||
assert_eq!(response.status(), StatusCode::FORBIDDEN);
|
||||
let payload: serde_json::Value = response.json().await.expect("body should parse");
|
||||
assert_eq!(
|
||||
payload["error"]["message"],
|
||||
"当前密钥不允许访问 openai:embedding 格式"
|
||||
);
|
||||
|
||||
gateway_handle.abort();
|
||||
}
|
||||
@@ -1,2 +1,4 @@
|
||||
mod embeddings;
|
||||
mod local_denials;
|
||||
mod rerank;
|
||||
mod routing;
|
||||
|
||||
@@ -0,0 +1,326 @@
|
||||
use std::collections::BTreeMap;
|
||||
use std::sync::Arc;
|
||||
|
||||
use aether_contracts::{ExecutionPlan, ExecutionResult, ResponseBody};
|
||||
use aether_crypto::DEVELOPMENT_ENCRYPTION_KEY;
|
||||
use aether_data::repository::candidate_selection::InMemoryMinimalCandidateSelectionReadRepository;
|
||||
use aether_data_contracts::repository::candidate_selection::StoredMinimalCandidateSelectionRow;
|
||||
use http::StatusCode;
|
||||
use serde_json::json;
|
||||
|
||||
use super::super::{
|
||||
any, build_router_with_state, build_state_with_execution_runtime_override, hash_api_key,
|
||||
sample_currently_usable_auth_snapshot, sample_endpoint, sample_key, sample_provider,
|
||||
start_server, AppState, GatewayDataState, InMemoryAuthApiKeySnapshotRepository,
|
||||
InMemoryProviderCatalogReadRepository, Json, Router,
|
||||
};
|
||||
use crate::constants::{
|
||||
CONTROL_ENDPOINT_SIGNATURE_HEADER, CONTROL_EXECUTION_RUNTIME_HEADER,
|
||||
CONTROL_ROUTE_FAMILY_HEADER, CONTROL_ROUTE_KIND_HEADER, EXECUTION_PATH_HEADER,
|
||||
EXECUTION_PATH_LOCAL_AUTH_DENIED,
|
||||
};
|
||||
|
||||
fn rerank_success_state(execution_runtime_url: String) -> AppState {
|
||||
let mut snapshot =
|
||||
sample_currently_usable_auth_snapshot("key-rerank-success", "user-rerank-success");
|
||||
snapshot.user_allowed_providers = None;
|
||||
snapshot.api_key_allowed_providers = None;
|
||||
snapshot.user_allowed_api_formats = Some(vec!["openai:rerank".to_string()]);
|
||||
snapshot.api_key_allowed_api_formats = Some(vec!["openai:rerank".to_string()]);
|
||||
snapshot.user_allowed_models = Some(vec!["bge-reranker-base".to_string()]);
|
||||
snapshot.api_key_allowed_models = Some(vec!["bge-reranker-base".to_string()]);
|
||||
let auth_repository = Arc::new(InMemoryAuthApiKeySnapshotRepository::seed(vec![(
|
||||
Some(hash_api_key("sk-rerank-success")),
|
||||
snapshot,
|
||||
)]));
|
||||
let candidate_repository =
|
||||
Arc::new(InMemoryMinimalCandidateSelectionReadRepository::seed(vec![
|
||||
rerank_candidate_row(),
|
||||
]));
|
||||
let provider_catalog_repository = Arc::new(InMemoryProviderCatalogReadRepository::seed(
|
||||
vec![sample_provider("provider-rerank", "OpenAI Rerank", 1)],
|
||||
vec![sample_endpoint(
|
||||
"endpoint-rerank",
|
||||
"provider-rerank",
|
||||
"openai:rerank",
|
||||
"https://api.openai.example",
|
||||
)],
|
||||
vec![sample_key(
|
||||
"key-upstream-rerank",
|
||||
"provider-rerank",
|
||||
"openai:rerank",
|
||||
"sk-upstream-rerank",
|
||||
)],
|
||||
));
|
||||
let data_state =
|
||||
GatewayDataState::with_provider_catalog_and_minimal_candidate_selection_for_tests(
|
||||
provider_catalog_repository,
|
||||
candidate_repository,
|
||||
)
|
||||
.with_auth_api_key_reader(auth_repository)
|
||||
.with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY);
|
||||
|
||||
build_state_with_execution_runtime_override(execution_runtime_url)
|
||||
.with_data_state_for_tests(data_state)
|
||||
}
|
||||
|
||||
fn rerank_execution_runtime() -> Router {
|
||||
Router::new().route(
|
||||
"/v1/execute/sync",
|
||||
any(|Json(plan): Json<ExecutionPlan>| async move {
|
||||
assert_rerank_execution_plan(&plan);
|
||||
Json(rerank_execution_result(&plan))
|
||||
}),
|
||||
)
|
||||
}
|
||||
|
||||
fn rerank_candidate_row() -> StoredMinimalCandidateSelectionRow {
|
||||
StoredMinimalCandidateSelectionRow {
|
||||
provider_id: "provider-rerank".to_string(),
|
||||
provider_name: "OpenAI Rerank".to_string(),
|
||||
provider_type: "custom".to_string(),
|
||||
provider_priority: 1,
|
||||
provider_is_active: true,
|
||||
endpoint_id: "endpoint-rerank".to_string(),
|
||||
endpoint_api_format: "openai:rerank".to_string(),
|
||||
endpoint_api_family: Some("openai".to_string()),
|
||||
endpoint_kind: Some("rerank".to_string()),
|
||||
endpoint_is_active: true,
|
||||
key_id: "key-upstream-rerank".to_string(),
|
||||
key_name: "default".to_string(),
|
||||
key_auth_type: "api_key".to_string(),
|
||||
key_is_active: true,
|
||||
key_api_formats: Some(vec!["openai:rerank".to_string()]),
|
||||
key_allowed_models: None,
|
||||
key_capabilities: None,
|
||||
key_internal_priority: 50,
|
||||
key_global_priority_by_format: None,
|
||||
model_id: "model-rerank-base".to_string(),
|
||||
global_model_id: "global-rerank-base".to_string(),
|
||||
global_model_name: "bge-reranker-base".to_string(),
|
||||
global_model_mappings: None,
|
||||
global_model_supports_streaming: Some(false),
|
||||
model_provider_model_name: "upstream-rerank".to_string(),
|
||||
model_provider_model_mappings: None,
|
||||
model_supports_streaming: Some(false),
|
||||
model_is_active: true,
|
||||
model_is_available: true,
|
||||
}
|
||||
}
|
||||
|
||||
fn assert_rerank_execution_plan(plan: &ExecutionPlan) {
|
||||
assert_eq!(plan.client_api_format, "openai:rerank");
|
||||
assert_eq!(plan.provider_api_format, "openai:rerank");
|
||||
assert_eq!(plan.method, "POST");
|
||||
assert_eq!(plan.url, "https://api.openai.example/v1/rerank");
|
||||
assert_eq!(plan.model_name.as_deref(), Some("bge-reranker-base"));
|
||||
let body = plan.body.json_body.as_ref().expect("json request body");
|
||||
assert_eq!(body["model"], "upstream-rerank");
|
||||
assert_eq!(body["query"], "hello");
|
||||
assert_eq!(body["documents"], json!(["hello world", "goodbye"]));
|
||||
assert_eq!(body["top_n"], 1);
|
||||
}
|
||||
|
||||
fn rerank_execution_result(plan: &ExecutionPlan) -> ExecutionResult {
|
||||
ExecutionResult {
|
||||
request_id: plan.request_id.clone(),
|
||||
candidate_id: plan.candidate_id.clone(),
|
||||
status_code: 200,
|
||||
headers: BTreeMap::from([("content-type".to_string(), "application/json".to_string())]),
|
||||
body: Some(ResponseBody {
|
||||
json_body: Some(json!({
|
||||
"model": "upstream-rerank",
|
||||
"results": [
|
||||
{"index": 0, "relevance_score": 0.98, "document": {"text": "hello world"}}
|
||||
],
|
||||
"usage": {"total_tokens": 8}
|
||||
})),
|
||||
body_bytes_b64: None,
|
||||
}),
|
||||
telemetry: None,
|
||||
error: None,
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn rerank_route_accepts_openai_payload() {
|
||||
let (execution_runtime_url, execution_runtime_handle) =
|
||||
start_server(rerank_execution_runtime()).await;
|
||||
let gateway = build_router_with_state(rerank_success_state(execution_runtime_url));
|
||||
let (gateway_url, gateway_handle) = start_server(gateway).await;
|
||||
|
||||
let response = reqwest::Client::new()
|
||||
.post(format!("{gateway_url}/v1/rerank"))
|
||||
.header(http::header::AUTHORIZATION, "Bearer sk-rerank-success")
|
||||
.json(&json!({
|
||||
"model": "bge-reranker-base",
|
||||
"query": "hello",
|
||||
"documents": ["hello world", "goodbye"],
|
||||
"top_n": 1,
|
||||
"return_documents": true
|
||||
}))
|
||||
.send()
|
||||
.await
|
||||
.expect("request should succeed");
|
||||
|
||||
assert_eq!(response.status(), StatusCode::OK);
|
||||
assert_eq!(
|
||||
response
|
||||
.headers()
|
||||
.get(CONTROL_ROUTE_FAMILY_HEADER)
|
||||
.and_then(|value| value.to_str().ok()),
|
||||
Some("openai")
|
||||
);
|
||||
assert_eq!(
|
||||
response
|
||||
.headers()
|
||||
.get(CONTROL_ROUTE_KIND_HEADER)
|
||||
.and_then(|value| value.to_str().ok()),
|
||||
Some("rerank")
|
||||
);
|
||||
assert_eq!(
|
||||
response
|
||||
.headers()
|
||||
.get(CONTROL_ENDPOINT_SIGNATURE_HEADER)
|
||||
.and_then(|value| value.to_str().ok()),
|
||||
Some("openai:rerank")
|
||||
);
|
||||
assert_eq!(
|
||||
response
|
||||
.headers()
|
||||
.get(CONTROL_EXECUTION_RUNTIME_HEADER)
|
||||
.and_then(|value| value.to_str().ok()),
|
||||
Some("true")
|
||||
);
|
||||
let payload: serde_json::Value = response.json().await.expect("body should parse");
|
||||
assert_eq!(payload["results"][0]["index"], 0);
|
||||
assert_eq!(payload["results"][0]["relevance_score"], 0.98);
|
||||
|
||||
gateway_handle.abort();
|
||||
execution_runtime_handle.abort();
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn rerank_route_rejects_invalid_local_payloads() {
|
||||
let gateway = build_router_with_state(AppState::new().expect("gateway should build"));
|
||||
let (gateway_url, gateway_handle) = start_server(gateway).await;
|
||||
let client = reqwest::Client::new();
|
||||
let cases = [
|
||||
("{", "Rerank request JSON body is invalid"),
|
||||
(
|
||||
r#"{"query":"hello","documents":["doc"]}"#,
|
||||
"Rerank request model is required",
|
||||
),
|
||||
(
|
||||
r#"{"model":"bge-reranker-base","documents":["doc"]}"#,
|
||||
"Rerank request query is required",
|
||||
),
|
||||
(
|
||||
r#"{"model":"bge-reranker-base","query":"hello","documents":[]}"#,
|
||||
"Rerank request documents are required",
|
||||
),
|
||||
(
|
||||
r#"{"model":"bge-reranker-base","query":"hello","messages":[]}"#,
|
||||
"Rerank request must use query/documents, not chat messages",
|
||||
),
|
||||
(
|
||||
r#"{"model":"bge-reranker-base","query":"hello","documents":["doc"],"top_n":0}"#,
|
||||
"Rerank request top_n must be a positive integer",
|
||||
),
|
||||
(
|
||||
r#"{"model":"bge-reranker-base","query":"hello","documents":["doc"],"stream":true}"#,
|
||||
"Rerank requests do not support streaming",
|
||||
),
|
||||
];
|
||||
|
||||
for (body, expected_detail) in cases {
|
||||
let response = client
|
||||
.post(format!("{gateway_url}/v1/rerank"))
|
||||
.header(http::header::CONTENT_TYPE, "application/json")
|
||||
.body(body)
|
||||
.send()
|
||||
.await
|
||||
.expect("request should succeed");
|
||||
|
||||
assert_eq!(response.status(), StatusCode::BAD_REQUEST);
|
||||
assert_eq!(
|
||||
response
|
||||
.headers()
|
||||
.get(CONTROL_ROUTE_KIND_HEADER)
|
||||
.and_then(|value| value.to_str().ok()),
|
||||
Some("rerank")
|
||||
);
|
||||
let payload: serde_json::Value = response.json().await.expect("body should parse");
|
||||
assert_eq!(payload["detail"], expected_detail);
|
||||
}
|
||||
|
||||
gateway_handle.abort();
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn rerank_route_rejects_non_json_content_type() {
|
||||
let gateway = build_router_with_state(AppState::new().expect("gateway should build"));
|
||||
let (gateway_url, gateway_handle) = start_server(gateway).await;
|
||||
|
||||
let response = reqwest::Client::new()
|
||||
.post(format!("{gateway_url}/v1/rerank"))
|
||||
.header(http::header::CONTENT_TYPE, "text/plain")
|
||||
.body(r#"{"model":"bge-reranker-base","query":"hello","documents":["doc"]}"#)
|
||||
.send()
|
||||
.await
|
||||
.expect("request should succeed");
|
||||
|
||||
assert_eq!(response.status(), StatusCode::BAD_REQUEST);
|
||||
let payload: serde_json::Value = response.json().await.expect("body should parse");
|
||||
assert_eq!(
|
||||
payload["detail"],
|
||||
"Rerank request content-type must be application/json"
|
||||
);
|
||||
|
||||
gateway_handle.abort();
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn rerank_route_rejects_chat_only_api_format() {
|
||||
let mut snapshot = sample_currently_usable_auth_snapshot("key-rerank-2", "user-rerank-2");
|
||||
snapshot.user_allowed_api_formats = Some(vec!["openai:chat".to_string()]);
|
||||
snapshot.api_key_allowed_api_formats = Some(vec!["openai:chat".to_string()]);
|
||||
let repository = Arc::new(InMemoryAuthApiKeySnapshotRepository::seed(vec![(
|
||||
Some(hash_api_key("sk-rerank-format-guard")),
|
||||
snapshot,
|
||||
)]));
|
||||
let gateway = build_router_with_state(
|
||||
AppState::new()
|
||||
.expect("gateway should build")
|
||||
.with_auth_api_key_data_reader_for_tests(repository),
|
||||
);
|
||||
let (gateway_url, gateway_handle) = start_server(gateway).await;
|
||||
|
||||
let response = reqwest::Client::new()
|
||||
.post(format!("{gateway_url}/v1/rerank"))
|
||||
.header(http::header::AUTHORIZATION, "Bearer sk-rerank-format-guard")
|
||||
.json(&json!({
|
||||
"model": "bge-reranker-base",
|
||||
"query": "hello",
|
||||
"documents": ["doc"]
|
||||
}))
|
||||
.send()
|
||||
.await
|
||||
.expect("request should succeed");
|
||||
|
||||
assert_eq!(response.status(), StatusCode::FORBIDDEN);
|
||||
assert_eq!(
|
||||
response
|
||||
.headers()
|
||||
.get(EXECUTION_PATH_HEADER)
|
||||
.and_then(|value| value.to_str().ok()),
|
||||
Some(EXECUTION_PATH_LOCAL_AUTH_DENIED)
|
||||
);
|
||||
let payload: serde_json::Value = response.json().await.expect("body should parse");
|
||||
assert_eq!(
|
||||
payload["error"]["message"],
|
||||
"当前密钥不允许访问 openai:rerank 格式"
|
||||
);
|
||||
|
||||
gateway_handle.abort();
|
||||
}
|
||||
@@ -340,6 +340,7 @@ fn sample_public_catalog_model(
|
||||
Some(true),
|
||||
Some(true),
|
||||
Some(true),
|
||||
Some(false),
|
||||
true,
|
||||
)
|
||||
.expect("public catalog model should build")
|
||||
|
||||
@@ -843,6 +843,94 @@ async fn gateway_handles_public_catalog_models_without_proxying_upstream() {
|
||||
upstream_handle.abort();
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn public_catalog_excludes_unsupported_embedding_provider() {
|
||||
let upstream_hits = Arc::new(Mutex::new(0usize));
|
||||
let upstream_hits_clone = Arc::clone(&upstream_hits);
|
||||
let upstream = Router::new().route(
|
||||
"/{*path}",
|
||||
any(move |_request: Request| {
|
||||
let upstream_hits_inner = Arc::clone(&upstream_hits_clone);
|
||||
async move {
|
||||
*upstream_hits_inner.lock().expect("mutex should lock") += 1;
|
||||
(StatusCode::OK, Body::from("proxied"))
|
||||
}
|
||||
}),
|
||||
);
|
||||
|
||||
let mut active_embedding = sample_public_catalog_model(
|
||||
"model-openai-embedding-small",
|
||||
"provider-openai",
|
||||
"openai",
|
||||
"text-embedding-3-small",
|
||||
"text-embedding-3-small",
|
||||
"Text Embedding 3 Small",
|
||||
);
|
||||
active_embedding.supports_embedding = Some(true);
|
||||
active_embedding.supports_streaming = Some(false);
|
||||
let mut inactive_embedding = sample_public_catalog_model(
|
||||
"model-openai-embedding-inactive",
|
||||
"provider-openai",
|
||||
"openai",
|
||||
"text-embedding-3-large",
|
||||
"text-embedding-3-large",
|
||||
"Text Embedding 3 Large",
|
||||
);
|
||||
inactive_embedding.supports_embedding = Some(true);
|
||||
inactive_embedding.is_active = false;
|
||||
let mut unsupported_embedding = sample_public_catalog_model(
|
||||
"model-unsupported-embedding",
|
||||
"provider-openai",
|
||||
"openai",
|
||||
"legacy-embedding",
|
||||
"legacy-embedding",
|
||||
"Legacy Embedding",
|
||||
);
|
||||
unsupported_embedding.supports_embedding = Some(false);
|
||||
unsupported_embedding.is_active = false;
|
||||
|
||||
let global_model_repository = Arc::new(
|
||||
InMemoryGlobalModelReadRepository::seed(Vec::<StoredPublicGlobalModel>::new())
|
||||
.with_public_catalog_models(vec![
|
||||
active_embedding,
|
||||
inactive_embedding,
|
||||
unsupported_embedding,
|
||||
]),
|
||||
);
|
||||
|
||||
let (upstream_url, upstream_handle) = start_server(upstream).await;
|
||||
let gateway = build_router_with_state(
|
||||
AppState::new()
|
||||
.expect("gateway should build")
|
||||
.with_data_state_for_tests(
|
||||
crate::data::GatewayDataState::with_global_model_reader_for_tests(
|
||||
global_model_repository,
|
||||
),
|
||||
),
|
||||
);
|
||||
let (gateway_url, gateway_handle) = start_server(gateway).await;
|
||||
|
||||
let response = reqwest::Client::new()
|
||||
.get(format!(
|
||||
"{gateway_url}/api/public/models?provider_id=provider-openai&limit=10"
|
||||
))
|
||||
.send()
|
||||
.await
|
||||
.expect("request should succeed");
|
||||
|
||||
assert_eq!(response.status(), StatusCode::OK);
|
||||
let payload: serde_json::Value = response.json().await.expect("json body should parse");
|
||||
let models = payload.as_array().expect("models should be an array");
|
||||
assert_eq!(models.len(), 1);
|
||||
assert_eq!(models[0]["id"], "model-openai-embedding-small");
|
||||
assert_eq!(models[0]["supports_embedding"], true);
|
||||
assert_eq!(models[0]["supports_streaming"], false);
|
||||
assert_eq!(*upstream_hits.lock().expect("mutex should lock"), 0);
|
||||
|
||||
gateway_handle.abort();
|
||||
upstream_handle.abort();
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn gateway_handles_public_catalog_search_models_without_proxying_upstream() {
|
||||
let upstream_hits = Arc::new(Mutex::new(0usize));
|
||||
|
||||
@@ -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,
|
||||
})
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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";
|
||||
|
||||
@@ -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::{
|
||||
|
||||
@@ -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!({
|
||||
|
||||
@@ -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()]
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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!(
|
||||
|
||||
@@ -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")
|
||||
|
||||
@@ -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);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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!(
|
||||
|
||||
@@ -8,6 +8,7 @@ pub enum LocalStandardSourceFamily {
|
||||
pub enum LocalStandardSourceMode {
|
||||
Chat,
|
||||
Cli,
|
||||
Embedding,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Copy)]
|
||||
|
||||
@@ -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,
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -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"));
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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 =
|
||||
|
||||
@@ -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::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!("sql.rs"));
|
||||
let supported_capabilities_projection = format!(
|
||||
"{} AS {}",
|
||||
"gm.supported_capabilities", "global_model_supported_capabilities"
|
||||
);
|
||||
assert_eq!(
|
||||
include_str!("sql.rs")
|
||||
.matches(&supported_capabilities_projection)
|
||||
.count(),
|
||||
4
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn repository_constructs_from_lazy_pool() {
|
||||
let factory = PostgresPoolFactory::new(PostgresPoolConfig {
|
||||
|
||||
@@ -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"
|
||||
);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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"
|
||||
));
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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"
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -0,0 +1,144 @@
|
||||
# Embeddings API
|
||||
|
||||
Aether supports OpenAI compatible embedding requests through `POST /v1/embeddings`. Embedding requests are separate from chat and responses requests. They use `input`, never `messages`, and they are always non streaming.
|
||||
|
||||
## Quick Start
|
||||
|
||||
Run this against your Aether gateway URL with a user API key that can access the model and the `openai:embedding` API format.
|
||||
|
||||
```bash
|
||||
curl -sS "http://localhost:8084/v1/embeddings" \
|
||||
-H "Authorization: Bearer sk-your-aether-key" \
|
||||
-H "Content-Type: application/json" \
|
||||
-d '{
|
||||
"model": "text-embedding-3-small",
|
||||
"input": ["hello", "world"],
|
||||
"encoding_format": "float"
|
||||
}'
|
||||
```
|
||||
|
||||
## Public Endpoint
|
||||
|
||||
| Method | Path | Client API format | Route kind |
|
||||
| --- | --- | --- | --- |
|
||||
| `POST` | `/v1/embeddings` | `openai:embedding` | `embedding` |
|
||||
|
||||
The gateway classifies this endpoint as an OpenAI family embedding route with endpoint signature `openai:embedding`. It is not handled as chat or responses.
|
||||
|
||||
## Request Body
|
||||
|
||||
Required fields:
|
||||
|
||||
| Field | Type | Notes |
|
||||
| --- | --- | --- |
|
||||
| `model` | string | Must name a model allowed for the API key and user. Blank strings are rejected. |
|
||||
| `input` | string, string array, integer token array, or nested integer token arrays | Must be non empty. Empty strings, empty arrays, and nested arrays with empty token arrays are rejected. |
|
||||
|
||||
Optional fields that pass through the embedding conversion path when supported by the provider:
|
||||
|
||||
| Field | Notes |
|
||||
| --- | --- |
|
||||
| `encoding_format` | Passed to OpenAI compatible providers. |
|
||||
| `dimensions` | Passed to providers whose embedding request shape supports it. |
|
||||
| `user` | Passed to OpenAI compatible providers. |
|
||||
| `task` | Passed to Jina and OpenAI compatible embedding requests. Jina defaults to `text-matching` when no task is supplied. |
|
||||
|
||||
Accepted `input` shapes:
|
||||
|
||||
```json
|
||||
{ "model": "text-embedding-3-small", "input": "hello" }
|
||||
```
|
||||
|
||||
```json
|
||||
{ "model": "text-embedding-3-small", "input": ["hello", "world"] }
|
||||
```
|
||||
|
||||
```json
|
||||
{ "model": "text-embedding-3-small", "input": [1, 2, 3] }
|
||||
```
|
||||
|
||||
```json
|
||||
{ "model": "text-embedding-3-small", "input": [[1, 2], [3, 4]] }
|
||||
```
|
||||
|
||||
Use string or string array input when routing to Gemini or Doubao embedding providers. Token arrays are accepted by the OpenAI compatible public endpoint, but Gemini and Doubao provider request emitters require text input.
|
||||
|
||||
## Provider Format Mapping
|
||||
|
||||
Embedding routes can select only embedding provider API formats. Chat, responses, image, and generation formats are not valid provider targets for this request type.
|
||||
|
||||
| Provider API format | Upstream path shape | Provider request shape |
|
||||
| --- | --- | --- |
|
||||
| `openai:embedding` | `/v1/embeddings` | OpenAI compatible `{ "model", "input" }` payload. |
|
||||
| `jina:embedding` | `/v1/embeddings` | OpenAI compatible payload with a Jina `task`. Defaults to `text-matching` if omitted. |
|
||||
| `gemini:embedding` | `models/{model}:embedContent` | Single text input uses `content.parts[].text`. Multiple text inputs use `requests[].content.parts[].text`. |
|
||||
| `doubao:embedding` | `/embeddings/multimodal` | Text input is emitted as `input` items like `{ "type": "text", "text": "..." }`. |
|
||||
|
||||
Custom provider endpoint paths are available when the endpoint is configured for an embedding API format. Gemini custom paths can use `{model}` and `{action}`. For `gemini:embedding`, `{action}` expands to `embedContent`.
|
||||
|
||||
## Model And Catalog Requirements
|
||||
|
||||
To use embeddings through the gateway:
|
||||
|
||||
1. The global model should include embedding metadata, for example `supported_capabilities: ["embedding"]`, `config.model_type: "embedding"`, or `config.api_formats` with one of the embedding formats.
|
||||
2. The provider model or mapping must expose an embedding API format, one of `openai:embedding`, `gemini:embedding`, `jina:embedding`, or `doubao:embedding`.
|
||||
3. The user and API key must be allowed to access the model and the `openai:embedding` client API format.
|
||||
4. Public and admin catalog responses expose `supports_embedding` so clients can display embedding capability separately from chat.
|
||||
|
||||
Billing fails closed for embedding global models. A model marked as embedding capable must define either `default_price_per_request` or `default_tiered_pricing.tiers[].input_price_per_1m`. Missing request pricing and missing input token pricing cause the model record to be rejected instead of treated as free.
|
||||
|
||||
No schema migration is needed for embedding metadata. Existing model capability, config, provider mapping, API format, and pricing fields carry the data.
|
||||
|
||||
## Failure Behavior
|
||||
|
||||
The gateway validates deterministic request errors before local execution or provider transport.
|
||||
|
||||
| Case | Example request body or setup | Status | Error detail |
|
||||
| --- | --- | --- | --- |
|
||||
| Invalid JSON | `{` | `400` | `Embedding request JSON body is invalid` |
|
||||
| Missing model | `{ "input": "hello" }` | `400` | `Embedding request model is required` |
|
||||
| Empty input | `{ "model": "text-embedding-3-small", "input": [] }` | `400` | `Embedding request input is required` |
|
||||
| Chat `messages` payload | `{ "model": "text-embedding-3-small", "messages": [] }` | `400` | `Embedding request must use input, not chat messages` |
|
||||
| Streaming requested | `{ "model": "text-embedding-3-small", "input": "hello", "stream": true }` | `400` | `Embedding requests do not support streaming` |
|
||||
| Non JSON content type | `Content-Type: text/plain` with an embedding JSON body | `400` | `Embedding request content-type must be application/json` |
|
||||
| Chat only model | API key allows `text-embedding-3-small`, request uses `gpt-5` | `403` | The key is not allowed to access that model. |
|
||||
| Chat only API format | API key allows `openai:chat` but not `openai:embedding` | `403` | The key is not allowed to access `openai:embedding`. |
|
||||
|
||||
Failure examples:
|
||||
|
||||
```bash
|
||||
curl -sS "http://localhost:8084/v1/embeddings" \
|
||||
-H "Authorization: Bearer sk-your-aether-key" \
|
||||
-H "Content-Type: application/json" \
|
||||
-d '{"model":"text-embedding-3-small","messages":[]}'
|
||||
```
|
||||
|
||||
```bash
|
||||
curl -sS "http://localhost:8084/v1/embeddings" \
|
||||
-H "Authorization: Bearer sk-your-aether-key" \
|
||||
-H "Content-Type: application/json" \
|
||||
-d '{"input":"hello"}'
|
||||
```
|
||||
|
||||
```bash
|
||||
curl -sS "http://localhost:8084/v1/embeddings" \
|
||||
-H "Authorization: Bearer sk-your-aether-key" \
|
||||
-H "Content-Type: application/json" \
|
||||
-d '{"model":"text-embedding-3-small","input":[]}'
|
||||
```
|
||||
|
||||
```bash
|
||||
curl -sS "http://localhost:8084/v1/embeddings" \
|
||||
-H "Authorization: Bearer sk-your-aether-key" \
|
||||
-H "Content-Type: application/json" \
|
||||
-d '{"model":"text-embedding-3-small","input":"hello","stream":true}'
|
||||
```
|
||||
|
||||
```bash
|
||||
curl -sS "http://localhost:8084/v1/embeddings" \
|
||||
-H "Authorization: Bearer sk-your-aether-key" \
|
||||
-H "Content-Type: text/plain" \
|
||||
-d '{"model":"text-embedding-3-small","input":"hello"}'
|
||||
```
|
||||
|
||||
If a valid embedding request passes local validation but no usable provider transport is available, the gateway can return a provider or service availability error. That is different from the deterministic request validation errors above.
|
||||
@@ -0,0 +1,58 @@
|
||||
# Rerank API
|
||||
|
||||
Aether exposes an OpenAI-compatible rerank surface at `POST /v1/rerank` and can route it to providers configured as `openai:rerank` or `jina:rerank`.
|
||||
|
||||
## Request
|
||||
|
||||
```http
|
||||
POST /v1/rerank
|
||||
Authorization: Bearer <aether-api-key>
|
||||
Content-Type: application/json
|
||||
```
|
||||
|
||||
```json
|
||||
{
|
||||
"model": "bge-reranker-base",
|
||||
"query": "What document discusses gateway routing?",
|
||||
"documents": [
|
||||
"Aether routes public AI requests through the Rust gateway.",
|
||||
"This document discusses unrelated content."
|
||||
],
|
||||
"top_n": 1,
|
||||
"return_documents": true
|
||||
}
|
||||
```
|
||||
|
||||
Fields:
|
||||
|
||||
| Field | Required | Notes |
|
||||
| --- | --- | --- |
|
||||
| `model` | Yes | Aether global model name. |
|
||||
| `query` | Yes | Non-empty query string. |
|
||||
| `documents` | Yes | Non-empty array of strings or provider-native document objects. |
|
||||
| `top_n` | No | Positive integer. |
|
||||
| `return_documents` | No | Provider-compatible flag for including matched documents. |
|
||||
|
||||
## Response
|
||||
|
||||
Aether forwards the provider JSON response. OpenAI-compatible and Jina-compatible rerank providers commonly return `results[]`:
|
||||
|
||||
```json
|
||||
{
|
||||
"model": "bge-reranker-base",
|
||||
"results": [
|
||||
{
|
||||
"index": 0,
|
||||
"relevance_score": 0.98,
|
||||
"document": {
|
||||
"text": "Aether routes public AI requests through the Rust gateway."
|
||||
}
|
||||
}
|
||||
],
|
||||
"usage": {
|
||||
"total_tokens": 32
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
Rerank requests must be JSON and do not support `stream` or chat `messages` payloads.
|
||||
@@ -200,6 +200,7 @@ export interface ModelExport {
|
||||
supports_streaming?: boolean | null
|
||||
supports_extended_thinking?: boolean | null
|
||||
supports_image_generation?: boolean | null
|
||||
supports_embedding?: boolean | null
|
||||
is_active: boolean
|
||||
config?: Record<string, unknown>
|
||||
}
|
||||
|
||||
@@ -15,6 +15,30 @@ describe('api format display helpers', () => {
|
||||
expect(normalizeApiFormatAlias('OPENAI_RESPONSES')).toBe(API_FORMATS.OPENAI_RESPONSES)
|
||||
expect(normalizeApiFormatAlias('OPENAI_RESPONSES_COMPACT')).toBe(API_FORMATS.OPENAI_RESPONSES_COMPACT)
|
||||
expect(normalizeApiFormatAlias('GEMINI_GENERATE_CONTENT')).toBe(API_FORMATS.GEMINI_GENERATE_CONTENT)
|
||||
expect(normalizeApiFormatAlias('OPENAI_EMBEDDING')).toBe(API_FORMATS.OPENAI_EMBEDDING)
|
||||
expect(normalizeApiFormatAlias('OPENAI_RERANK')).toBe(API_FORMATS.OPENAI_RERANK)
|
||||
expect(normalizeApiFormatAlias('GEMINI_EMBEDDING')).toBe(API_FORMATS.GEMINI_EMBEDDING)
|
||||
expect(normalizeApiFormatAlias('JINA_EMBEDDING')).toBe(API_FORMATS.JINA_EMBEDDING)
|
||||
expect(normalizeApiFormatAlias('JINA_RERANK')).toBe(API_FORMATS.JINA_RERANK)
|
||||
expect(normalizeApiFormatAlias('DOUBAO_EMBEDDING')).toBe(API_FORMATS.DOUBAO_EMBEDDING)
|
||||
})
|
||||
|
||||
it('formats rerank api format ids distinctly from chat formats', () => {
|
||||
expect(formatApiFormat(API_FORMATS.OPENAI_RERANK)).toBe('OpenAI Rerank')
|
||||
expect(formatApiFormat(API_FORMATS.JINA_RERANK)).toBe('Jina Rerank')
|
||||
expect(formatApiFormatShort(API_FORMATS.OPENAI_RERANK)).toBe('ORR')
|
||||
expect(formatApiFormatShort(API_FORMATS.JINA_RERANK)).toBe('JR')
|
||||
})
|
||||
|
||||
it('formats embedding api format ids distinctly from chat formats', () => {
|
||||
expect(formatApiFormat(API_FORMATS.OPENAI_EMBEDDING)).toBe('OpenAI Embedding')
|
||||
expect(formatApiFormat(API_FORMATS.GEMINI_EMBEDDING)).toBe('Gemini Embedding')
|
||||
expect(formatApiFormat(API_FORMATS.JINA_EMBEDDING)).toBe('Jina Embedding')
|
||||
expect(formatApiFormat(API_FORMATS.DOUBAO_EMBEDDING)).toBe('Doubao Embedding')
|
||||
expect(formatApiFormatShort(API_FORMATS.OPENAI_EMBEDDING)).toBe('OE')
|
||||
expect(formatApiFormatShort(API_FORMATS.GEMINI_EMBEDDING)).toBe('GE')
|
||||
expect(formatApiFormatShort(API_FORMATS.JINA_EMBEDDING)).toBe('JE')
|
||||
expect(formatApiFormatShort(API_FORMATS.DOUBAO_EMBEDDING)).toBe('DE')
|
||||
})
|
||||
|
||||
it('does not remap retired api format ids', () => {
|
||||
@@ -40,15 +64,59 @@ describe('api format display helpers', () => {
|
||||
it('sorts only current canonical formats into known slots', () => {
|
||||
expect(sortApiFormats([
|
||||
'openai:compact',
|
||||
API_FORMATS.DOUBAO_EMBEDDING,
|
||||
API_FORMATS.OPENAI,
|
||||
API_FORMATS.OPENAI_RESPONSES,
|
||||
API_FORMATS.OPENAI_EMBEDDING,
|
||||
API_FORMATS.OPENAI_RERANK,
|
||||
API_FORMATS.GEMINI_EMBEDDING,
|
||||
API_FORMATS.JINA_EMBEDDING,
|
||||
API_FORMATS.JINA_RERANK,
|
||||
])).toEqual([
|
||||
API_FORMATS.OPENAI,
|
||||
API_FORMATS.OPENAI_RESPONSES,
|
||||
API_FORMATS.OPENAI_EMBEDDING,
|
||||
API_FORMATS.OPENAI_RERANK,
|
||||
API_FORMATS.GEMINI_EMBEDDING,
|
||||
API_FORMATS.JINA_EMBEDDING,
|
||||
API_FORMATS.JINA_RERANK,
|
||||
API_FORMATS.DOUBAO_EMBEDDING,
|
||||
'openai:compact',
|
||||
])
|
||||
})
|
||||
|
||||
it('keeps embedding formats after chat/generation formats within each family', () => {
|
||||
expect(sortApiFormats([
|
||||
API_FORMATS.GEMINI_EMBEDDING,
|
||||
API_FORMATS.OPENAI_EMBEDDING,
|
||||
API_FORMATS.OPENAI_RERANK,
|
||||
API_FORMATS.GEMINI_GENERATE_CONTENT,
|
||||
API_FORMATS.OPENAI,
|
||||
])).toEqual([
|
||||
API_FORMATS.OPENAI,
|
||||
API_FORMATS.OPENAI_EMBEDDING,
|
||||
API_FORMATS.OPENAI_RERANK,
|
||||
API_FORMATS.GEMINI_GENERATE_CONTENT,
|
||||
API_FORMATS.GEMINI_EMBEDDING,
|
||||
])
|
||||
})
|
||||
|
||||
it('groups embedding api formats by provider family', () => {
|
||||
expect(groupApiFormats([
|
||||
API_FORMATS.DOUBAO_EMBEDDING,
|
||||
API_FORMATS.JINA_RERANK,
|
||||
API_FORMATS.JINA_EMBEDDING,
|
||||
API_FORMATS.GEMINI_EMBEDDING,
|
||||
API_FORMATS.OPENAI_EMBEDDING,
|
||||
API_FORMATS.OPENAI_RERANK,
|
||||
])).toEqual([
|
||||
{ family: 'openai', label: 'OpenAI', formats: [API_FORMATS.OPENAI_EMBEDDING, API_FORMATS.OPENAI_RERANK] },
|
||||
{ family: 'gemini', label: 'Gemini', formats: [API_FORMATS.GEMINI_EMBEDDING] },
|
||||
{ family: 'jina', label: 'Jina', formats: [API_FORMATS.JINA_EMBEDDING, API_FORMATS.JINA_RERANK] },
|
||||
{ family: 'doubao', label: 'Doubao', formats: [API_FORMATS.DOUBAO_EMBEDDING] },
|
||||
])
|
||||
})
|
||||
|
||||
it('groups retired enum-style aliases as unknown raw families', () => {
|
||||
expect(groupApiFormats(['OPENAI_CLI'])).toEqual([{
|
||||
family: 'openai_cli',
|
||||
|
||||
@@ -8,10 +8,16 @@ export const API_FORMATS = {
|
||||
OPENAI_RESPONSES_COMPACT: 'openai:responses:compact',
|
||||
OPENAI_IMAGE: 'openai:image',
|
||||
OPENAI_VIDEO: 'openai:video',
|
||||
OPENAI_EMBEDDING: 'openai:embedding',
|
||||
OPENAI_RERANK: 'openai:rerank',
|
||||
GEMINI: 'gemini:generate_content',
|
||||
GEMINI_GENERATE_CONTENT: 'gemini:generate_content',
|
||||
GEMINI_VIDEO: 'gemini:video',
|
||||
GEMINI_FILES: 'gemini:files',
|
||||
GEMINI_EMBEDDING: 'gemini:embedding',
|
||||
JINA_EMBEDDING: 'jina:embedding',
|
||||
JINA_RERANK: 'jina:rerank',
|
||||
DOUBAO_EMBEDDING: 'doubao:embedding',
|
||||
} as const
|
||||
|
||||
export type APIFormat = typeof API_FORMATS[keyof typeof API_FORMATS]
|
||||
@@ -24,9 +30,15 @@ export const API_FORMAT_LABELS: Record<string, string> = {
|
||||
[API_FORMATS.OPENAI_RESPONSES_COMPACT]: 'OpenAI Responses Compact',
|
||||
[API_FORMATS.OPENAI_IMAGE]: 'OpenAI Image',
|
||||
[API_FORMATS.OPENAI_VIDEO]: 'OpenAI Video',
|
||||
[API_FORMATS.OPENAI_EMBEDDING]: 'OpenAI Embedding',
|
||||
[API_FORMATS.OPENAI_RERANK]: 'OpenAI Rerank',
|
||||
[API_FORMATS.GEMINI_GENERATE_CONTENT]: 'Gemini Generate Content',
|
||||
[API_FORMATS.GEMINI_VIDEO]: 'Gemini Video',
|
||||
[API_FORMATS.GEMINI_FILES]: 'Gemini Files',
|
||||
[API_FORMATS.GEMINI_EMBEDDING]: 'Gemini Embedding',
|
||||
[API_FORMATS.JINA_EMBEDDING]: 'Jina Embedding',
|
||||
[API_FORMATS.JINA_RERANK]: 'Jina Rerank',
|
||||
[API_FORMATS.DOUBAO_EMBEDDING]: 'Doubao Embedding',
|
||||
CLAUDE: 'Claude Messages',
|
||||
CLAUDE_MESSAGES: 'Claude Messages',
|
||||
OPENAI: 'OpenAI Chat',
|
||||
@@ -34,10 +46,16 @@ export const API_FORMAT_LABELS: Record<string, string> = {
|
||||
OPENAI_RESPONSES_COMPACT: 'OpenAI Responses Compact',
|
||||
OPENAI_IMAGE: 'OpenAI Image',
|
||||
OPENAI_VIDEO: 'OpenAI Video',
|
||||
OPENAI_EMBEDDING: 'OpenAI Embedding',
|
||||
OPENAI_RERANK: 'OpenAI Rerank',
|
||||
GEMINI: 'Gemini Generate Content',
|
||||
GEMINI_GENERATE_CONTENT: 'Gemini Generate Content',
|
||||
GEMINI_VIDEO: 'Gemini Video',
|
||||
GEMINI_FILES: 'Gemini Files',
|
||||
GEMINI_EMBEDDING: 'Gemini Embedding',
|
||||
JINA_EMBEDDING: 'Jina Embedding',
|
||||
JINA_RERANK: 'Jina Rerank',
|
||||
DOUBAO_EMBEDDING: 'Doubao Embedding',
|
||||
}
|
||||
|
||||
// API 格式缩写映射(用于空间紧凑的显示场景)
|
||||
@@ -47,21 +65,33 @@ export const API_FORMAT_SHORT: Record<string, string> = {
|
||||
[API_FORMATS.OPENAI_RESPONSES_COMPACT]: 'ORC',
|
||||
[API_FORMATS.OPENAI_IMAGE]: 'OI',
|
||||
[API_FORMATS.OPENAI_VIDEO]: 'OV',
|
||||
[API_FORMATS.OPENAI_EMBEDDING]: 'OE',
|
||||
[API_FORMATS.OPENAI_RERANK]: 'ORR',
|
||||
[API_FORMATS.CLAUDE_MESSAGES]: 'CM',
|
||||
[API_FORMATS.GEMINI_GENERATE_CONTENT]: 'G',
|
||||
[API_FORMATS.GEMINI_VIDEO]: 'GV',
|
||||
[API_FORMATS.GEMINI_FILES]: 'GF',
|
||||
[API_FORMATS.GEMINI_EMBEDDING]: 'GE',
|
||||
[API_FORMATS.JINA_EMBEDDING]: 'JE',
|
||||
[API_FORMATS.JINA_RERANK]: 'JR',
|
||||
[API_FORMATS.DOUBAO_EMBEDDING]: 'DE',
|
||||
OPENAI: 'O',
|
||||
OPENAI_RESPONSES: 'OR',
|
||||
OPENAI_RESPONSES_COMPACT: 'ORC',
|
||||
OPENAI_IMAGE: 'OI',
|
||||
OPENAI_VIDEO: 'OV',
|
||||
OPENAI_EMBEDDING: 'OE',
|
||||
OPENAI_RERANK: 'ORR',
|
||||
CLAUDE: 'CM',
|
||||
CLAUDE_MESSAGES: 'CM',
|
||||
GEMINI: 'G',
|
||||
GEMINI_GENERATE_CONTENT: 'G',
|
||||
GEMINI_VIDEO: 'GV',
|
||||
GEMINI_FILES: 'GF',
|
||||
GEMINI_EMBEDDING: 'GE',
|
||||
JINA_EMBEDDING: 'JE',
|
||||
JINA_RERANK: 'JR',
|
||||
DOUBAO_EMBEDDING: 'DE',
|
||||
}
|
||||
|
||||
// API 格式排序顺序(统一的显示顺序)
|
||||
@@ -69,12 +99,18 @@ export const API_FORMAT_ORDER: string[] = [
|
||||
API_FORMATS.OPENAI,
|
||||
API_FORMATS.OPENAI_RESPONSES,
|
||||
API_FORMATS.OPENAI_RESPONSES_COMPACT,
|
||||
API_FORMATS.OPENAI_EMBEDDING,
|
||||
API_FORMATS.OPENAI_RERANK,
|
||||
API_FORMATS.OPENAI_IMAGE,
|
||||
API_FORMATS.OPENAI_VIDEO,
|
||||
API_FORMATS.CLAUDE_MESSAGES,
|
||||
API_FORMATS.GEMINI_GENERATE_CONTENT,
|
||||
API_FORMATS.GEMINI_EMBEDDING,
|
||||
API_FORMATS.GEMINI_VIDEO,
|
||||
API_FORMATS.GEMINI_FILES,
|
||||
API_FORMATS.JINA_EMBEDDING,
|
||||
API_FORMATS.JINA_RERANK,
|
||||
API_FORMATS.DOUBAO_EMBEDDING,
|
||||
]
|
||||
|
||||
// Family 显示名称映射
|
||||
@@ -82,6 +118,8 @@ export const API_FORMAT_FAMILY_LABELS: Record<string, string> = {
|
||||
openai: 'OpenAI',
|
||||
claude: 'Claude',
|
||||
gemini: 'Gemini',
|
||||
jina: 'Jina',
|
||||
doubao: 'Doubao',
|
||||
}
|
||||
|
||||
// Kind 显示名称映射
|
||||
@@ -94,10 +132,12 @@ export const API_FORMAT_KIND_LABELS: Record<string, string> = {
|
||||
image: 'Image',
|
||||
video: 'Video',
|
||||
files: 'Files',
|
||||
embedding: 'Embedding',
|
||||
rerank: 'Rerank',
|
||||
}
|
||||
|
||||
// Family 排序顺序
|
||||
const FAMILY_ORDER = ['openai', 'claude', 'gemini']
|
||||
const FAMILY_ORDER = ['openai', 'claude', 'gemini', 'jina', 'doubao']
|
||||
|
||||
// 工具函数:从 API 格式中提取 family 和 kind
|
||||
export function parseApiFormat(format: string): { family: string; kind: string } {
|
||||
@@ -124,6 +164,10 @@ export function normalizeApiFormatAlias(format: string | null | undefined): stri
|
||||
return API_FORMATS.OPENAI_IMAGE
|
||||
case 'OPENAI_VIDEO':
|
||||
return API_FORMATS.OPENAI_VIDEO
|
||||
case 'OPENAI_EMBEDDING':
|
||||
return API_FORMATS.OPENAI_EMBEDDING
|
||||
case 'OPENAI_RERANK':
|
||||
return API_FORMATS.OPENAI_RERANK
|
||||
case 'GEMINI':
|
||||
case 'GEMINI_GENERATE_CONTENT':
|
||||
return API_FORMATS.GEMINI_GENERATE_CONTENT
|
||||
@@ -131,6 +175,14 @@ export function normalizeApiFormatAlias(format: string | null | undefined): stri
|
||||
return API_FORMATS.GEMINI_VIDEO
|
||||
case 'GEMINI_FILES':
|
||||
return API_FORMATS.GEMINI_FILES
|
||||
case 'GEMINI_EMBEDDING':
|
||||
return API_FORMATS.GEMINI_EMBEDDING
|
||||
case 'JINA_EMBEDDING':
|
||||
return API_FORMATS.JINA_EMBEDDING
|
||||
case 'JINA_RERANK':
|
||||
return API_FORMATS.JINA_RERANK
|
||||
case 'DOUBAO_EMBEDDING':
|
||||
return API_FORMATS.DOUBAO_EMBEDDING
|
||||
default:
|
||||
return raw.toLowerCase()
|
||||
}
|
||||
|
||||
@@ -38,6 +38,7 @@ export interface Model {
|
||||
supports_streaming?: boolean | null
|
||||
supports_extended_thinking?: boolean | null
|
||||
supports_image_generation?: boolean | null
|
||||
supports_embedding?: boolean | null
|
||||
// 有效值(合并 Model 和 GlobalModel 默认值后的结果)
|
||||
effective_tiered_pricing?: TieredPricingConfig | null // 有效阶梯计费配置
|
||||
effective_input_price?: number | null
|
||||
@@ -48,6 +49,7 @@ export interface Model {
|
||||
effective_supports_streaming?: boolean | null
|
||||
effective_supports_extended_thinking?: boolean | null
|
||||
effective_supports_image_generation?: boolean | null
|
||||
effective_supports_embedding?: boolean | null
|
||||
is_active: boolean
|
||||
is_available: boolean
|
||||
created_at: string
|
||||
@@ -96,6 +98,7 @@ export interface ModelCapabilities {
|
||||
supports_vision: boolean
|
||||
supports_function_calling: boolean
|
||||
supports_streaming: boolean
|
||||
supports_embedding: boolean
|
||||
[key: string]: boolean
|
||||
}
|
||||
|
||||
@@ -130,6 +133,7 @@ export interface ModelCatalogProviderDetail {
|
||||
supports_vision?: boolean | null
|
||||
supports_function_calling?: boolean | null
|
||||
supports_streaming?: boolean | null
|
||||
supports_embedding?: boolean | null
|
||||
is_active: boolean
|
||||
mapping_id?: string | null
|
||||
}
|
||||
@@ -211,6 +215,7 @@ export interface GlobalModelResponse {
|
||||
default_tiered_pricing: TieredPricingConfig
|
||||
// Key 能力配置 - 模型支持的能力列表
|
||||
supported_capabilities?: string[] | null
|
||||
supports_embedding?: boolean | null
|
||||
// 模型配置(JSON格式)
|
||||
config?: Record<string, unknown> | null
|
||||
// 统计数据
|
||||
|
||||
@@ -340,6 +340,7 @@ export const meApi = {
|
||||
default_price_per_request: number | null
|
||||
default_tiered_pricing: TieredPricingConfig | null
|
||||
supported_capabilities: string[] | null
|
||||
supports_embedding?: boolean | null
|
||||
config: Record<string, unknown> | null
|
||||
usage_count: number
|
||||
}>
|
||||
|
||||
@@ -72,6 +72,7 @@ export interface ModelsDevModelItem {
|
||||
supportsStructuredOutput?: boolean
|
||||
supportsTemperature?: boolean
|
||||
supportsAttachment?: boolean
|
||||
supportsEmbedding?: boolean
|
||||
openWeights?: boolean
|
||||
deprecated?: boolean
|
||||
official?: boolean // 是否来自官方提供商
|
||||
@@ -180,6 +181,9 @@ export async function getModelsDevList(officialOnly: boolean = true): Promise<Mo
|
||||
supportsStructuredOutput: model.structured_output,
|
||||
supportsTemperature: model.temperature,
|
||||
supportsAttachment: model.attachment,
|
||||
supportsEmbedding: model.id.toLowerCase().includes('embedding')
|
||||
|| model.name.toLowerCase().includes('embedding')
|
||||
|| model.family?.toLowerCase().includes('embedding') === true,
|
||||
openWeights: model.open_weights,
|
||||
deprecated: model.deprecated,
|
||||
official: provider.official,
|
||||
|
||||
@@ -15,6 +15,7 @@ export interface PublicGlobalModel {
|
||||
default_price_per_request: number | null // 按次计费价格
|
||||
// Key 能力支持
|
||||
supported_capabilities: string[] | null
|
||||
supports_embedding?: boolean | null
|
||||
// 模型配置(JSON)
|
||||
config: Record<string, unknown> | null
|
||||
// 调用次数
|
||||
|
||||
@@ -142,7 +142,7 @@
|
||||
>描述</Label>
|
||||
<Input
|
||||
id="model-description"
|
||||
:model-value="form.config?.description || ''"
|
||||
:model-value="getConfigInputValue('description')"
|
||||
placeholder="简短描述此模型的特点"
|
||||
@update:model-value="(v) => setConfigField('description', v || undefined)"
|
||||
/>
|
||||
@@ -155,7 +155,7 @@
|
||||
>最大输出 Token</Label>
|
||||
<Input
|
||||
id="model-output-limit"
|
||||
:model-value="form.config?.output_limit ?? ''"
|
||||
:model-value="getConfigInputValue('output_limit')"
|
||||
type="number"
|
||||
min="1"
|
||||
placeholder="如 8192"
|
||||
@@ -169,7 +169,7 @@
|
||||
>上下文窗口</Label>
|
||||
<Input
|
||||
id="model-context-limit"
|
||||
:model-value="form.config?.context_limit ?? ''"
|
||||
:model-value="getConfigInputValue('context_limit')"
|
||||
type="number"
|
||||
min="1"
|
||||
placeholder="如 200000"
|
||||
@@ -177,6 +177,33 @@
|
||||
/>
|
||||
</div>
|
||||
</div>
|
||||
<div class="rounded-lg border border-border/60 bg-muted/20 p-3 space-y-2">
|
||||
<div class="flex items-start gap-2">
|
||||
<Checkbox
|
||||
:model-value="isEmbeddingEnabled"
|
||||
class="mt-0.5"
|
||||
@update:model-value="setEmbeddingEnabled"
|
||||
/>
|
||||
<div class="space-y-1">
|
||||
<div class="text-sm font-medium">
|
||||
Embedding
|
||||
</div>
|
||||
<p class="text-xs text-muted-foreground">
|
||||
标记为 Embeddings 模型,并使用独立的 embedding API 格式,不按 Chat 模型处理。
|
||||
</p>
|
||||
<div
|
||||
v-if="isEmbeddingEnabled"
|
||||
class="flex flex-wrap gap-1.5"
|
||||
>
|
||||
<span
|
||||
v-for="format in embeddingApiFormats"
|
||||
:key="format"
|
||||
class="rounded-md border border-border/60 bg-background px-2 py-0.5 text-[11px] font-mono text-muted-foreground"
|
||||
>{{ format }}</span>
|
||||
</div>
|
||||
</div>
|
||||
</div>
|
||||
</div>
|
||||
</section>
|
||||
|
||||
<!-- 价格配置 -->
|
||||
@@ -333,7 +360,7 @@ import {
|
||||
Loader2, Layers, SquarePen,
|
||||
Search, ChevronRight, Plus, Trash2
|
||||
} from 'lucide-vue-next'
|
||||
import { Dialog, Button, Input, Label } from '@/components/ui'
|
||||
import { Dialog, Button, Input, Label, Checkbox } from '@/components/ui'
|
||||
import { useToast } from '@/composables/useToast'
|
||||
import { useFormDialog } from '@/composables/useFormDialog'
|
||||
import { parseNumberInput, sortResolutionEntries } from '@/utils/form'
|
||||
@@ -349,10 +376,13 @@ import {
|
||||
createGlobalModel,
|
||||
updateGlobalModel,
|
||||
type GlobalModelResponse,
|
||||
type GlobalModelCreate,
|
||||
type GlobalModelUpdate,
|
||||
} from '@/api/global-models'
|
||||
import type { TieredPricingConfig } from '@/api/endpoints/types'
|
||||
import {
|
||||
EMBEDDING_API_FORMATS,
|
||||
buildGlobalModelCreatePayload,
|
||||
buildGlobalModelUpdatePayload,
|
||||
} from './global-model-form-helpers'
|
||||
|
||||
const props = defineProps<{
|
||||
open: boolean
|
||||
@@ -476,6 +506,8 @@ const VIDEO_RESOLUTION_PRICE_PRESETS: Record<
|
||||
],
|
||||
}
|
||||
|
||||
const embeddingApiFormats = [...EMBEDDING_API_FORMATS]
|
||||
|
||||
interface FormData {
|
||||
name: string
|
||||
display_name: string
|
||||
@@ -496,6 +528,12 @@ const defaultForm = (): FormData => ({
|
||||
|
||||
const form = ref<FormData>(defaultForm())
|
||||
|
||||
const isEmbeddingEnabled = computed(() => {
|
||||
return form.value.supported_capabilities?.includes('embedding') === true
|
||||
|| form.value.config?.embedding === true
|
||||
|| form.value.config?.model_type === 'embedding'
|
||||
})
|
||||
|
||||
const KEEP_FALSE_CONFIG_KEYS = new Set(['streaming'])
|
||||
|
||||
// 设置 config 字段
|
||||
@@ -510,6 +548,34 @@ function setConfigField(key: string, value: unknown) {
|
||||
}
|
||||
}
|
||||
|
||||
function getConfigInputValue(key: string): string | number {
|
||||
const value = form.value.config?.[key]
|
||||
return typeof value === 'string' || typeof value === 'number' ? value : ''
|
||||
}
|
||||
|
||||
function setEmbeddingEnabled(enabled: boolean) {
|
||||
const caps = new Set(form.value.supported_capabilities || [])
|
||||
if (enabled) {
|
||||
caps.add('embedding')
|
||||
setConfigField('embedding', true)
|
||||
setConfigField('model_type', 'embedding')
|
||||
setConfigField('streaming', false)
|
||||
form.value.config = {
|
||||
...(form.value.config || {}),
|
||||
api_formats: [...embeddingApiFormats],
|
||||
}
|
||||
} else {
|
||||
caps.delete('embedding')
|
||||
setConfigField('embedding', undefined)
|
||||
if (form.value.config?.model_type === 'embedding') setConfigField('model_type', undefined)
|
||||
if (Array.isArray(form.value.config?.api_formats)
|
||||
&& form.value.config.api_formats.every((format) => embeddingApiFormats.includes(String(format)))) {
|
||||
setConfigField('api_formats', undefined)
|
||||
}
|
||||
}
|
||||
form.value.supported_capabilities = [...caps]
|
||||
}
|
||||
|
||||
function getNested(obj: unknown, path: string): unknown {
|
||||
if (!obj || typeof obj !== 'object') return undefined
|
||||
const parts = path.split('.').filter(Boolean)
|
||||
@@ -670,7 +736,7 @@ function selectModel(model: ModelsDevModelItem) {
|
||||
|
||||
// 构建 config
|
||||
const config: Record<string, unknown> = {
|
||||
streaming: true,
|
||||
streaming: model.supportsEmbedding ? false : true,
|
||||
}
|
||||
if (model.supportsVision) config.vision = true
|
||||
if (model.supportsToolCall) config.function_calling = true
|
||||
@@ -687,6 +753,10 @@ function selectModel(model: ModelsDevModelItem) {
|
||||
if (model.inputModalities?.length) config.input_modalities = model.inputModalities
|
||||
if (model.outputModalities?.length) config.output_modalities = model.outputModalities
|
||||
form.value.config = config
|
||||
form.value.supported_capabilities = model.supportsEmbedding ? ['embedding'] : []
|
||||
if (model.supportsEmbedding) {
|
||||
setEmbeddingEnabled(true)
|
||||
}
|
||||
loadVideoPricingFromConfig()
|
||||
|
||||
if (model.inputPrice !== undefined || model.outputPrice !== undefined) {
|
||||
@@ -796,26 +866,11 @@ async function handleSubmit() {
|
||||
submitting.value = true
|
||||
try {
|
||||
if (isEditMode.value && props.model) {
|
||||
const updateData: GlobalModelUpdate = {
|
||||
display_name: form.value.display_name,
|
||||
config: cleanConfig || null,
|
||||
default_price_per_request: form.value.default_price_per_request ?? null,
|
||||
default_tiered_pricing: finalTieredPricing,
|
||||
supported_capabilities: form.value.supported_capabilities?.length ? form.value.supported_capabilities : null,
|
||||
is_active: form.value.is_active,
|
||||
}
|
||||
const updateData = buildGlobalModelUpdatePayload(form.value, finalTieredPricing)
|
||||
await updateGlobalModel(props.model.id, updateData)
|
||||
success('模型更新成功')
|
||||
} else {
|
||||
const createData: GlobalModelCreate = {
|
||||
name: form.value.name ?? '',
|
||||
display_name: form.value.display_name ?? '',
|
||||
config: cleanConfig,
|
||||
default_price_per_request: form.value.default_price_per_request ?? undefined,
|
||||
default_tiered_pricing: finalTieredPricing,
|
||||
supported_capabilities: form.value.supported_capabilities?.length ? form.value.supported_capabilities : undefined,
|
||||
is_active: form.value.is_active,
|
||||
}
|
||||
const createData = buildGlobalModelCreatePayload(form.value, finalTieredPricing)
|
||||
await createGlobalModel(createData)
|
||||
success('模型创建成功')
|
||||
clearSelection()
|
||||
|
||||
@@ -0,0 +1,62 @@
|
||||
import { describe, expect, it } from 'vitest'
|
||||
|
||||
import {
|
||||
EMBEDDING_API_FORMATS,
|
||||
buildGlobalModelCreatePayload,
|
||||
buildGlobalModelUpdatePayload,
|
||||
} from '../global-model-form-helpers'
|
||||
|
||||
const embeddingPricing = {
|
||||
tiers: [{ up_to: null, input_price_per_1m: 0.02, output_price_per_1m: 0 }],
|
||||
}
|
||||
|
||||
describe('global model form embedding payload helpers', () => {
|
||||
it('preserves embedding metadata in create payloads', () => {
|
||||
const payload = buildGlobalModelCreatePayload({
|
||||
name: 'text-embedding-3-small',
|
||||
display_name: 'text-embedding-3-small',
|
||||
supported_capabilities: ['embedding'],
|
||||
config: {
|
||||
streaming: false,
|
||||
embedding: true,
|
||||
model_type: 'embedding',
|
||||
api_formats: [...EMBEDDING_API_FORMATS],
|
||||
},
|
||||
is_active: true,
|
||||
}, embeddingPricing)
|
||||
|
||||
expect(payload).toMatchObject({
|
||||
name: 'text-embedding-3-small',
|
||||
supported_capabilities: ['embedding'],
|
||||
config: {
|
||||
streaming: false,
|
||||
embedding: true,
|
||||
model_type: 'embedding',
|
||||
api_formats: ['openai:embedding', 'gemini:embedding', 'jina:embedding', 'doubao:embedding'],
|
||||
},
|
||||
})
|
||||
})
|
||||
|
||||
it('preserves embedding metadata in update payloads', () => {
|
||||
const payload = buildGlobalModelUpdatePayload({
|
||||
name: 'unused-on-update',
|
||||
display_name: 'Jina Embeddings v3',
|
||||
supported_capabilities: ['embedding'],
|
||||
config: {
|
||||
streaming: false,
|
||||
embedding: true,
|
||||
model_type: 'embedding',
|
||||
api_formats: ['jina:embedding'],
|
||||
},
|
||||
is_active: true,
|
||||
}, embeddingPricing)
|
||||
|
||||
expect(payload.supported_capabilities).toEqual(['embedding'])
|
||||
expect(payload.config).toEqual({
|
||||
streaming: false,
|
||||
embedding: true,
|
||||
model_type: 'embedding',
|
||||
api_formats: ['jina:embedding'],
|
||||
})
|
||||
})
|
||||
})
|
||||
@@ -0,0 +1,56 @@
|
||||
import type { GlobalModelCreate, GlobalModelUpdate } from '@/api/global-models'
|
||||
import type { TieredPricingConfig } from '@/api/endpoints/types'
|
||||
|
||||
export const EMBEDDING_API_FORMATS = [
|
||||
'openai:embedding',
|
||||
'gemini:embedding',
|
||||
'jina:embedding',
|
||||
'doubao:embedding',
|
||||
] as const
|
||||
|
||||
export const RERANK_API_FORMATS = [
|
||||
'openai:rerank',
|
||||
'jina:rerank',
|
||||
] as const
|
||||
|
||||
export interface GlobalModelFormPayloadState {
|
||||
name: string
|
||||
display_name: string
|
||||
default_price_per_request?: number
|
||||
supported_capabilities?: string[]
|
||||
config?: Record<string, unknown>
|
||||
is_active?: boolean
|
||||
}
|
||||
|
||||
function cleanGlobalModelConfig(form: GlobalModelFormPayloadState): Record<string, unknown> | undefined {
|
||||
return form.config && Object.keys(form.config).length > 0 ? form.config : undefined
|
||||
}
|
||||
|
||||
export function buildGlobalModelCreatePayload(
|
||||
form: GlobalModelFormPayloadState,
|
||||
defaultTieredPricing: TieredPricingConfig,
|
||||
): GlobalModelCreate {
|
||||
return {
|
||||
name: form.name ?? '',
|
||||
display_name: form.display_name ?? '',
|
||||
config: cleanGlobalModelConfig(form),
|
||||
default_price_per_request: form.default_price_per_request ?? undefined,
|
||||
default_tiered_pricing: defaultTieredPricing,
|
||||
supported_capabilities: form.supported_capabilities?.length ? form.supported_capabilities : undefined,
|
||||
is_active: form.is_active,
|
||||
}
|
||||
}
|
||||
|
||||
export function buildGlobalModelUpdatePayload(
|
||||
form: GlobalModelFormPayloadState,
|
||||
defaultTieredPricing: TieredPricingConfig,
|
||||
): GlobalModelUpdate {
|
||||
return {
|
||||
display_name: form.display_name,
|
||||
config: cleanGlobalModelConfig(form) || null,
|
||||
default_price_per_request: form.default_price_per_request ?? null,
|
||||
default_tiered_pricing: defaultTieredPricing,
|
||||
supported_capabilities: form.supported_capabilities?.length ? form.supported_capabilities : null,
|
||||
is_active: form.is_active,
|
||||
}
|
||||
}
|
||||
@@ -41,6 +41,17 @@
|
||||
>
|
||||
所有全局模型已添加到此 Provider
|
||||
</p>
|
||||
<div
|
||||
v-if="selectedGlobalModelSupportsEmbedding"
|
||||
class="rounded-lg border border-border/60 bg-muted/20 px-3 py-2"
|
||||
>
|
||||
<div class="text-sm font-medium">
|
||||
Embedding
|
||||
</div>
|
||||
<p class="text-xs text-muted-foreground">
|
||||
此模型将继承全局模型的 Embeddings 元数据,不按 Chat 能力处理。
|
||||
</p>
|
||||
</div>
|
||||
</div>
|
||||
|
||||
<!-- 编辑模式:显示模型信息 -->
|
||||
@@ -56,6 +67,13 @@
|
||||
<p class="text-sm text-muted-foreground font-mono">
|
||||
{{ editingModel?.provider_model_name }}
|
||||
</p>
|
||||
<Badge
|
||||
v-if="editingModelSupportsEmbedding"
|
||||
variant="secondary"
|
||||
class="mt-2 text-xs"
|
||||
>
|
||||
Embedding
|
||||
</Badge>
|
||||
</div>
|
||||
</div>
|
||||
</div>
|
||||
@@ -214,6 +232,7 @@ import {
|
||||
SelectValue,
|
||||
SelectContent,
|
||||
SelectItem,
|
||||
Badge,
|
||||
} from '@/components/ui'
|
||||
import { useToast } from '@/composables/useToast'
|
||||
import { parseNumberInput, sortResolutionEntries } from '@/utils/form'
|
||||
@@ -221,6 +240,11 @@ import { createModel, updateModel, getProviderModels } from '@/api/endpoints/mod
|
||||
import { listGlobalModels, type GlobalModelResponse } from '@/api/global-models'
|
||||
import TieredPricingEditor from '@/features/models/components/TieredPricingEditor.vue'
|
||||
import type { Model, TieredPricingConfig } from '@/api/endpoints'
|
||||
import {
|
||||
buildProviderModelCreatePayload,
|
||||
buildProviderModelUpdatePayload,
|
||||
modelSupportsEmbedding,
|
||||
} from './provider-model-form-helpers'
|
||||
|
||||
interface Props {
|
||||
open: boolean
|
||||
@@ -245,6 +269,16 @@ const tieredPricingEditorRef = ref<InstanceType<typeof TieredPricingEditor> | nu
|
||||
|
||||
const isEditing = computed(() => !!props.editingModel)
|
||||
|
||||
const selectedGlobalModel = computed(() => {
|
||||
return availableGlobalModels.value.find(model => model.id === form.value.global_model_id) || null
|
||||
})
|
||||
|
||||
const selectedGlobalModelSupportsEmbedding = computed(() => modelSupportsEmbedding(selectedGlobalModel.value))
|
||||
const editingModelSupportsEmbedding = computed(() => {
|
||||
return props.editingModel?.effective_supports_embedding === true
|
||||
|| modelSupportsEmbedding(props.editingModel)
|
||||
})
|
||||
|
||||
// 1h 缓存定价始终显示
|
||||
const showCache1h = true
|
||||
|
||||
@@ -571,35 +605,36 @@ async function handleSubmit() {
|
||||
if (isEditing.value && props.editingModel) {
|
||||
// 编辑模式
|
||||
// 注意:使用 null 而不是 undefined 来显式清空字段(undefined 会被 JSON 序列化忽略)
|
||||
await updateModel(props.providerId, props.editingModel.id, {
|
||||
tiered_pricing: finalTieredPricing,
|
||||
price_per_request: form.value.price_per_request ?? null,
|
||||
config: cleanConfig || null,
|
||||
supports_vision: form.value.supports_vision,
|
||||
supports_function_calling: form.value.supports_function_calling,
|
||||
supports_streaming: form.value.supports_streaming,
|
||||
supports_extended_thinking: form.value.supports_extended_thinking,
|
||||
supports_image_generation: form.value.supports_image_generation,
|
||||
is_active: form.value.is_active
|
||||
})
|
||||
await updateModel(props.providerId, props.editingModel.id, buildProviderModelUpdatePayload({
|
||||
finalTieredPricing,
|
||||
pricePerRequest: form.value.price_per_request,
|
||||
cleanConfig,
|
||||
supportsVision: form.value.supports_vision,
|
||||
supportsFunctionCalling: form.value.supports_function_calling,
|
||||
supportsStreaming: form.value.supports_streaming,
|
||||
supportsExtendedThinking: form.value.supports_extended_thinking,
|
||||
supportsImageGeneration: form.value.supports_image_generation,
|
||||
isActive: form.value.is_active
|
||||
}))
|
||||
showSuccess('模型配置已更新')
|
||||
} else {
|
||||
// 添加模式:只有用户修改了配置才提交 tiered_pricing,否则保持继承关系
|
||||
const selectedModel = availableGlobalModels.value.find(m => m.id === form.value.global_model_id)
|
||||
await createModel(props.providerId, {
|
||||
global_model_id: form.value.global_model_id,
|
||||
provider_model_name: selectedModel?.name || '',
|
||||
// 只有修改了才提交,否则传 undefined 让后端继承 GlobalModel 配置
|
||||
tiered_pricing: tieredPricingModified.value ? finalTieredPricing : undefined,
|
||||
price_per_request: form.value.price_per_request,
|
||||
config: configTouched.value ? cleanConfig : undefined,
|
||||
supports_vision: form.value.supports_vision,
|
||||
supports_function_calling: form.value.supports_function_calling,
|
||||
supports_streaming: form.value.supports_streaming,
|
||||
supports_extended_thinking: form.value.supports_extended_thinking,
|
||||
supports_image_generation: form.value.supports_image_generation,
|
||||
is_active: form.value.is_active
|
||||
})
|
||||
await createModel(props.providerId, buildProviderModelCreatePayload({
|
||||
globalModelId: form.value.global_model_id,
|
||||
providerModelName: selectedModel?.name || '',
|
||||
finalTieredPricing,
|
||||
tieredPricingModified: tieredPricingModified.value,
|
||||
pricePerRequest: form.value.price_per_request,
|
||||
cleanConfig,
|
||||
configTouched: configTouched.value,
|
||||
supportsVision: form.value.supports_vision,
|
||||
supportsFunctionCalling: form.value.supports_function_calling,
|
||||
supportsStreaming: form.value.supports_streaming,
|
||||
supportsExtendedThinking: form.value.supports_extended_thinking,
|
||||
supportsImageGeneration: form.value.supports_image_generation,
|
||||
isActive: form.value.is_active
|
||||
}))
|
||||
showSuccess('模型已添加')
|
||||
}
|
||||
emit('update:open', false)
|
||||
|
||||
+75
@@ -0,0 +1,75 @@
|
||||
import { describe, expect, it } from 'vitest'
|
||||
|
||||
import {
|
||||
buildProviderModelCreatePayload,
|
||||
buildProviderModelUpdatePayload,
|
||||
modelSupportsEmbedding,
|
||||
} from '../provider-model-form-helpers'
|
||||
|
||||
const pricing = {
|
||||
tiers: [{ up_to: null, input_price_per_1m: 0.02, output_price_per_1m: 0 }],
|
||||
}
|
||||
|
||||
describe('provider model form embedding helpers', () => {
|
||||
it.each([
|
||||
{ supported_capabilities: ['embedding'], config: {} },
|
||||
{ supported_capabilities: null, config: { embedding: true } },
|
||||
{ supported_capabilities: null, config: { model_type: 'embedding' } },
|
||||
{ supported_capabilities: null, config: { api_formats: ['doubao:embedding'] } },
|
||||
{ supports_embedding: true, effective_supports_embedding: null, config: {} },
|
||||
{ supports_embedding: null, effective_supports_embedding: true, config: {} },
|
||||
])('detects embedding metadata from %o', (model) => {
|
||||
expect(modelSupportsEmbedding(model)).toBe(true)
|
||||
})
|
||||
|
||||
it('keeps provider create payload inherited from the selected embedding global model', () => {
|
||||
const payload = buildProviderModelCreatePayload({
|
||||
globalModelId: 'gm-embedding',
|
||||
providerModelName: 'text-embedding-3-small',
|
||||
finalTieredPricing: pricing,
|
||||
tieredPricingModified: false,
|
||||
pricePerRequest: undefined,
|
||||
cleanConfig: {
|
||||
embedding: true,
|
||||
model_type: 'embedding',
|
||||
api_formats: ['openai:embedding'],
|
||||
},
|
||||
configTouched: false,
|
||||
supportsStreaming: false,
|
||||
isActive: true,
|
||||
})
|
||||
|
||||
expect(payload).toMatchObject({
|
||||
global_model_id: 'gm-embedding',
|
||||
provider_model_name: 'text-embedding-3-small',
|
||||
tiered_pricing: undefined,
|
||||
config: undefined,
|
||||
supports_streaming: false,
|
||||
})
|
||||
expect('supports_embedding' in payload).toBe(false)
|
||||
})
|
||||
|
||||
it('preserves edited provider embedding config without posting unsupported embedding controls', () => {
|
||||
const payload = buildProviderModelUpdatePayload({
|
||||
finalTieredPricing: pricing,
|
||||
pricePerRequest: undefined,
|
||||
cleanConfig: {
|
||||
streaming: false,
|
||||
embedding: true,
|
||||
model_type: 'embedding',
|
||||
api_formats: ['gemini:embedding'],
|
||||
},
|
||||
supportsStreaming: false,
|
||||
isActive: true,
|
||||
})
|
||||
|
||||
expect(payload.config).toEqual({
|
||||
streaming: false,
|
||||
embedding: true,
|
||||
model_type: 'embedding',
|
||||
api_formats: ['gemini:embedding'],
|
||||
})
|
||||
expect(payload.supports_streaming).toBe(false)
|
||||
expect('supports_embedding' in payload).toBe(false)
|
||||
})
|
||||
})
|
||||
@@ -0,0 +1,79 @@
|
||||
import type { ModelCreate, ModelUpdate, TieredPricingConfig } from '@/api/endpoints'
|
||||
|
||||
interface EmbeddingMetadataCarrier {
|
||||
supported_capabilities?: string[] | null
|
||||
supports_embedding?: boolean | null
|
||||
effective_supports_embedding?: boolean | null
|
||||
config?: Record<string, unknown> | null
|
||||
}
|
||||
|
||||
export interface ProviderModelCreatePayloadInput {
|
||||
globalModelId: string
|
||||
providerModelName: string
|
||||
finalTieredPricing: TieredPricingConfig | null
|
||||
tieredPricingModified: boolean
|
||||
pricePerRequest?: number
|
||||
cleanConfig?: Record<string, unknown>
|
||||
configTouched: boolean
|
||||
supportsVision?: boolean
|
||||
supportsFunctionCalling?: boolean
|
||||
supportsStreaming?: boolean
|
||||
supportsExtendedThinking?: boolean
|
||||
supportsImageGeneration?: boolean
|
||||
isActive: boolean
|
||||
}
|
||||
|
||||
export interface ProviderModelUpdatePayloadInput {
|
||||
finalTieredPricing: TieredPricingConfig | null
|
||||
pricePerRequest?: number
|
||||
cleanConfig?: Record<string, unknown>
|
||||
supportsVision?: boolean
|
||||
supportsFunctionCalling?: boolean
|
||||
supportsStreaming?: boolean
|
||||
supportsExtendedThinking?: boolean
|
||||
supportsImageGeneration?: boolean
|
||||
isActive: boolean
|
||||
}
|
||||
|
||||
export function modelSupportsEmbedding(model: EmbeddingMetadataCarrier | null | undefined): boolean {
|
||||
if (!model) return false
|
||||
if ('effective_supports_embedding' in model && model.effective_supports_embedding === true) return true
|
||||
if ('supports_embedding' in model && model.supports_embedding === true) return true
|
||||
|
||||
const supportedCapabilities = 'supported_capabilities' in model ? model.supported_capabilities : null
|
||||
const config = model.config || {}
|
||||
return supportedCapabilities?.includes('embedding') === true
|
||||
|| config.embedding === true
|
||||
|| config.model_type === 'embedding'
|
||||
|| (Array.isArray(config.api_formats) && config.api_formats.some((format) => String(format).endsWith(':embedding')))
|
||||
}
|
||||
|
||||
export function buildProviderModelCreatePayload(input: ProviderModelCreatePayloadInput): ModelCreate {
|
||||
return {
|
||||
global_model_id: input.globalModelId,
|
||||
provider_model_name: input.providerModelName,
|
||||
tiered_pricing: input.tieredPricingModified && input.finalTieredPricing ? input.finalTieredPricing : undefined,
|
||||
price_per_request: input.pricePerRequest,
|
||||
config: input.configTouched ? input.cleanConfig : undefined,
|
||||
supports_vision: input.supportsVision,
|
||||
supports_function_calling: input.supportsFunctionCalling,
|
||||
supports_streaming: input.supportsStreaming,
|
||||
supports_extended_thinking: input.supportsExtendedThinking,
|
||||
supports_image_generation: input.supportsImageGeneration,
|
||||
is_active: input.isActive,
|
||||
}
|
||||
}
|
||||
|
||||
export function buildProviderModelUpdatePayload(input: ProviderModelUpdatePayloadInput): ModelUpdate {
|
||||
return {
|
||||
tiered_pricing: input.finalTieredPricing,
|
||||
price_per_request: input.pricePerRequest ?? null,
|
||||
config: input.cleanConfig || null,
|
||||
supports_vision: input.supportsVision,
|
||||
supports_function_calling: input.supportsFunctionCalling,
|
||||
supports_streaming: input.supportsStreaming,
|
||||
supports_extended_thinking: input.supportsExtendedThinking,
|
||||
supports_image_generation: input.supportsImageGeneration,
|
||||
is_active: input.isActive,
|
||||
}
|
||||
}
|
||||
@@ -700,7 +700,7 @@ function runMappingTest(testingKey: string, modelName: string) {
|
||||
selectedTestEndpoint.value = activeEndpoints.value[0] ?? null
|
||||
testRequestHeadersResetValue.value = buildDefaultModelTestRequestHeaders()
|
||||
testRequestHeadersDraft.value = testRequestHeadersResetValue.value
|
||||
testRequestBodyResetValue.value = buildDefaultModelTestRequestBody(modelName)
|
||||
testRequestBodyResetValue.value = buildDefaultModelTestRequestBody(modelName, selectedTestEndpoint.value?.api_format)
|
||||
testRequestBodyDraft.value = testRequestBodyResetValue.value
|
||||
}
|
||||
|
||||
|
||||
@@ -543,6 +543,7 @@ async function testModelConnection(model: Model) {
|
||||
testRequestHeadersDraft.value = testRequestHeadersResetValue.value
|
||||
testRequestBodyResetValue.value = buildDefaultModelTestRequestBody(
|
||||
model.global_model_name || model.provider_model_name,
|
||||
selectedTestEndpoint.value?.api_format,
|
||||
)
|
||||
testRequestBodyDraft.value = testRequestBodyResetValue.value
|
||||
modelTest.testResult.value = null
|
||||
|
||||
+46
@@ -0,0 +1,46 @@
|
||||
import { describe, expect, it } from 'vitest'
|
||||
|
||||
import { buildDefaultModelTestRequestBody } from '../model-test-request'
|
||||
|
||||
describe('buildDefaultModelTestRequestBody', () => {
|
||||
it.each([
|
||||
'openai:embedding',
|
||||
'gemini:embedding',
|
||||
'jina:embedding',
|
||||
'doubao:embedding',
|
||||
' OPENAI:EMBEDDING ',
|
||||
])('uses embedding input payloads for %s api formats', (apiFormat) => {
|
||||
const body = JSON.parse(buildDefaultModelTestRequestBody('text-embedding-3-small', apiFormat))
|
||||
|
||||
expect(body).toEqual({
|
||||
model: 'text-embedding-3-small',
|
||||
input: 'This is a test embedding input.',
|
||||
})
|
||||
expect(body.messages).toBeUndefined()
|
||||
expect(body.stream).toBeUndefined()
|
||||
})
|
||||
|
||||
it.each([
|
||||
'openai:rerank',
|
||||
'jina:rerank',
|
||||
' JINA:RERANK ',
|
||||
])('uses rerank query/documents payloads for %s api formats', (apiFormat) => {
|
||||
const body = JSON.parse(buildDefaultModelTestRequestBody('bge-reranker-base', apiFormat))
|
||||
|
||||
expect(body.model).toBe('bge-reranker-base')
|
||||
expect(body.query).toBe('This is a test rerank query.')
|
||||
expect(body.documents).toHaveLength(2)
|
||||
expect(body.top_n).toBe(1)
|
||||
expect(body.return_documents).toBe(true)
|
||||
expect(body.messages).toBeUndefined()
|
||||
expect(body.stream).toBeUndefined()
|
||||
})
|
||||
|
||||
it('keeps chat payloads for chat api formats', () => {
|
||||
const body = JSON.parse(buildDefaultModelTestRequestBody('gpt-5.1', 'openai:chat'))
|
||||
|
||||
expect(body.messages).toEqual([{ role: 'user', content: 'Hello! This is a test message.' }])
|
||||
expect(body.stream).toBe(true)
|
||||
expect(body.input).toBeUndefined()
|
||||
})
|
||||
})
|
||||
@@ -4,7 +4,27 @@ const DEFAULT_MODEL_TEST_MESSAGE = 'Hello! This is a test message.'
|
||||
export const POOL_TEST_CONCURRENCY = 5
|
||||
export const SINGLE_TEST_CONCURRENCY = 1
|
||||
|
||||
export function buildDefaultModelTestRequestBody(modelName: string): string {
|
||||
export function buildDefaultModelTestRequestBody(modelName: string, apiFormat?: string | null): string {
|
||||
if (apiFormat?.trim().toLowerCase().endsWith(':embedding')) {
|
||||
return JSON.stringify({
|
||||
model: modelName,
|
||||
input: 'This is a test embedding input.',
|
||||
}, null, 2)
|
||||
}
|
||||
|
||||
if (apiFormat?.trim().toLowerCase().endsWith(':rerank')) {
|
||||
return JSON.stringify({
|
||||
model: modelName,
|
||||
query: 'This is a test rerank query.',
|
||||
documents: [
|
||||
'This document is relevant to the test query.',
|
||||
'This document is unrelated.',
|
||||
],
|
||||
top_n: 1,
|
||||
return_documents: true,
|
||||
}, null, 2)
|
||||
}
|
||||
|
||||
return JSON.stringify({
|
||||
model: modelName,
|
||||
messages: [
|
||||
|
||||
@@ -8,10 +8,16 @@ const ENDPOINT_SORT_ORDER = [
|
||||
'openai:chat',
|
||||
'openai:responses',
|
||||
'openai:responses:compact',
|
||||
'openai:embedding',
|
||||
'openai:rerank',
|
||||
'gemini:generate_content',
|
||||
'gemini:embedding',
|
||||
'openai:video',
|
||||
'gemini:video',
|
||||
'gemini:files',
|
||||
'jina:embedding',
|
||||
'jina:rerank',
|
||||
'doubao:embedding',
|
||||
]
|
||||
|
||||
/**
|
||||
|
||||
@@ -27,7 +27,13 @@ export function useProviderFilters(
|
||||
{ value: 'openai:chat', label: 'OpenAI Chat' },
|
||||
{ value: 'openai:responses', label: 'OpenAI Responses' },
|
||||
{ value: 'openai:responses:compact', label: 'OpenAI Responses Compact' },
|
||||
{ value: 'openai:embedding', label: 'OpenAI Embedding' },
|
||||
{ value: 'openai:rerank', label: 'OpenAI Rerank' },
|
||||
{ value: 'gemini:generate_content', label: 'Gemini Generate Content' },
|
||||
{ value: 'gemini:embedding', label: 'Gemini Embedding' },
|
||||
{ value: 'jina:embedding', label: 'Jina Embedding' },
|
||||
{ value: 'jina:rerank', label: 'Jina Rerank' },
|
||||
{ value: 'doubao:embedding', label: 'Doubao Embedding' },
|
||||
]
|
||||
|
||||
const modelFilters = computed<FilterOption[]>(() => {
|
||||
|
||||
@@ -0,0 +1,44 @@
|
||||
import { describe, expect, it } from 'vitest'
|
||||
|
||||
import { MOCK_API_FORMATS, MOCK_GLOBAL_MODELS } from '../data'
|
||||
|
||||
describe('embedding mock metadata', () => {
|
||||
it('exposes embedding model metadata to frontend code without chat treatment', () => {
|
||||
const model = MOCK_GLOBAL_MODELS.find(item => item.name === 'text-embedding-3-small')
|
||||
|
||||
expect(model).toMatchObject({
|
||||
supported_capabilities: ['embedding'],
|
||||
supports_embedding: true,
|
||||
config: {
|
||||
streaming: false,
|
||||
embedding: true,
|
||||
model_type: 'embedding',
|
||||
api_formats: ['openai:embedding'],
|
||||
},
|
||||
})
|
||||
})
|
||||
|
||||
it('includes all embedding API formats as distinct catalog formats', () => {
|
||||
const embeddingFormats = MOCK_API_FORMATS.formats
|
||||
.filter(format => format.value.endsWith(':embedding'))
|
||||
.map(format => [format.value, format.label])
|
||||
|
||||
expect(embeddingFormats).toEqual([
|
||||
['openai:embedding', 'OpenAI Embedding'],
|
||||
['gemini:embedding', 'Gemini Embedding'],
|
||||
['jina:embedding', 'Jina Embedding'],
|
||||
['doubao:embedding', 'Doubao Embedding'],
|
||||
])
|
||||
})
|
||||
|
||||
it('includes rerank API formats as distinct catalog formats', () => {
|
||||
const rerankFormats = MOCK_API_FORMATS.formats
|
||||
.filter(format => format.value.endsWith(':rerank'))
|
||||
.map(format => [format.value, format.label])
|
||||
|
||||
expect(rerankFormats).toEqual([
|
||||
['openai:rerank', 'OpenAI Rerank'],
|
||||
['jina:rerank', 'Jina Rerank'],
|
||||
])
|
||||
})
|
||||
})
|
||||
@@ -546,20 +546,21 @@ export const MOCK_PROVIDERS: ProviderWithEndpointsSummary[] = [
|
||||
billing_type: 'pay_as_you_go',
|
||||
monthly_used_usd: 5.29,
|
||||
is_active: true,
|
||||
total_endpoints: 4,
|
||||
active_endpoints: 4,
|
||||
total_endpoints: 5,
|
||||
active_endpoints: 5,
|
||||
total_keys: 11,
|
||||
active_keys: 11,
|
||||
total_models: 8,
|
||||
active_models: 8,
|
||||
total_models: 9,
|
||||
active_models: 9,
|
||||
avg_health_score: 0.863,
|
||||
unhealthy_endpoints: 1,
|
||||
api_formats: ['claude:messages', 'gemini:generate_content', 'openai:chat', 'openai:responses'],
|
||||
api_formats: ['claude:messages', 'gemini:generate_content', 'openai:chat', 'openai:responses', 'openai:embedding'],
|
||||
endpoint_health_details: [
|
||||
{ api_format: 'claude:messages', health_score: 1.0, is_active: true, active_keys: 2 },
|
||||
{ api_format: 'gemini:generate_content', health_score: 1.0, is_active: true, active_keys: 2 },
|
||||
{ api_format: 'openai:chat', health_score: 0.85, is_active: true, active_keys: 2 },
|
||||
{ api_format: 'openai:responses', health_score: 1.0, is_active: true, active_keys: 1 }
|
||||
{ api_format: 'openai:responses', health_score: 1.0, is_active: true, active_keys: 1 },
|
||||
{ api_format: 'openai:embedding', health_score: 0.98, is_active: true, active_keys: 1 }
|
||||
],
|
||||
created_at: '2024-12-07T22:56:09.712806+08:00',
|
||||
updated_at: new Date().toISOString()
|
||||
@@ -805,6 +806,46 @@ export const MOCK_GLOBAL_MODELS: GlobalModelResponse[] = [
|
||||
},
|
||||
provider_count: 2,
|
||||
created_at: '2024-01-01T00:00:00Z'
|
||||
},
|
||||
{
|
||||
id: 'gm-010',
|
||||
name: 'text-embedding-3-small',
|
||||
display_name: 'text-embedding-3-small',
|
||||
is_active: true,
|
||||
default_tiered_pricing: {
|
||||
tiers: [{ up_to: null, input_price_per_1m: 0.02, output_price_per_1m: 0 }]
|
||||
},
|
||||
supported_capabilities: ['embedding'],
|
||||
supports_embedding: true,
|
||||
config: {
|
||||
streaming: false,
|
||||
embedding: true,
|
||||
model_type: 'embedding',
|
||||
api_formats: ['openai:embedding'],
|
||||
dimensions: 1536,
|
||||
description: 'OpenAI 文本向量嵌入模型'
|
||||
},
|
||||
provider_count: 1,
|
||||
created_at: '2024-01-01T00:00:00Z'
|
||||
},
|
||||
{
|
||||
id: 'gm-rerank-001',
|
||||
name: 'bge-reranker-base',
|
||||
display_name: 'bge-reranker-base',
|
||||
is_active: true,
|
||||
default_tiered_pricing: {
|
||||
tiers: [{ up_to: null, input_price_per_1m: 0.05, output_price_per_1m: 0 }]
|
||||
},
|
||||
supported_capabilities: ['rerank'],
|
||||
config: {
|
||||
streaming: false,
|
||||
rerank: true,
|
||||
model_type: 'rerank',
|
||||
api_formats: ['openai:rerank'],
|
||||
description: '文本重排序模型'
|
||||
},
|
||||
provider_count: 1,
|
||||
created_at: '2024-01-01T00:00:00Z'
|
||||
}
|
||||
]
|
||||
|
||||
@@ -878,9 +919,15 @@ export const MOCK_API_FORMATS = {
|
||||
{ value: 'openai:chat', label: 'OpenAI Chat', default_path: '/v1/chat/completions', aliases: [] },
|
||||
{ value: 'openai:responses', label: 'OpenAI Responses', default_path: '/v1/responses', aliases: [] },
|
||||
{ value: 'openai:responses:compact', label: 'OpenAI Responses Compact', default_path: '/v1/responses/compact', aliases: [] },
|
||||
{ value: 'openai:embedding', label: 'OpenAI Embedding', default_path: '/v1/embeddings', aliases: [] },
|
||||
{ value: 'openai:rerank', label: 'OpenAI Rerank', default_path: '/v1/rerank', aliases: [] },
|
||||
{ value: 'openai:image', label: 'OpenAI Image', default_path: '/v1/images/generations', aliases: [] },
|
||||
{ value: 'openai:video', label: 'OpenAI Video', default_path: '/v1/videos', aliases: [] },
|
||||
{ value: 'gemini:generate_content', label: 'Gemini Generate Content', default_path: '/v1beta/models/{model}:{action}', aliases: [] },
|
||||
{ value: 'gemini:video', label: 'Gemini Video', default_path: '/v1beta/models/{model}:predictLongRunning', aliases: [] }
|
||||
{ value: 'gemini:embedding', label: 'Gemini Embedding', default_path: '/v1beta/models/{model}:embedContent', aliases: [] },
|
||||
{ value: 'gemini:video', label: 'Gemini Video', default_path: '/v1beta/models/{model}:predictLongRunning', aliases: [] },
|
||||
{ value: 'jina:embedding', label: 'Jina Embedding', default_path: '/v1/embeddings', aliases: [] },
|
||||
{ value: 'jina:rerank', label: 'Jina Rerank', default_path: '/v1/rerank', aliases: [] },
|
||||
{ value: 'doubao:embedding', label: 'Doubao Embedding', default_path: '/embeddings/multimodal', aliases: [] }
|
||||
]
|
||||
}
|
||||
|
||||
@@ -228,6 +228,19 @@ const MOCK_ENDPOINT_STATUS = {
|
||||
last_event_at: new Date().toISOString(),
|
||||
// 94.0% 成功率:successRate=0.940, failRate=0.043, skipRate=0.017
|
||||
events: generateHealthEvents(100, 0.940, 0.043, 0.017, 800, 600)
|
||||
},
|
||||
{
|
||||
api_format: 'openai:embedding',
|
||||
api_path: '/v1/embeddings',
|
||||
total_attempts: 620,
|
||||
success_count: 612,
|
||||
failed_count: 6,
|
||||
skipped_count: 2,
|
||||
success_rate: 0.987,
|
||||
provider_count: 1,
|
||||
key_count: 1,
|
||||
last_event_at: new Date().toISOString(),
|
||||
events: generateHealthEvents(40, 0.987, 0.01, 0.003, 320, 140)
|
||||
}
|
||||
]
|
||||
}
|
||||
@@ -447,6 +460,12 @@ function getMockEndpointExtras(apiFormat: string) {
|
||||
extras.config = { upstream_stream_policy: 'force_stream' }
|
||||
} else if (normalizedFormat === 'openai:responses') {
|
||||
extras.config = { upstream_stream_policy: 'force_non_stream' }
|
||||
} else if (normalizedFormat === 'openai:embedding') {
|
||||
extras.custom_path = '/v1/embeddings'
|
||||
extras.config = { route_kind: 'embedding' }
|
||||
} else if (normalizedFormat === 'openai:rerank' || normalizedFormat === 'jina:rerank') {
|
||||
extras.custom_path = '/v1/rerank'
|
||||
extras.config = { route_kind: 'rerank' }
|
||||
} else if (normalizedFormat === 'gemini:generate_content') {
|
||||
extras.custom_path = '/v1beta/models/gemini-3-pro-preview:generateContent'
|
||||
extras.body_rules = [
|
||||
@@ -720,6 +739,23 @@ const mockHandlers: Record<string, (config: AxiosRequestConfig) => Promise<Axios
|
||||
})))
|
||||
},
|
||||
|
||||
'GET /api/users/me/available-models': async () => {
|
||||
await delay()
|
||||
const models = MOCK_GLOBAL_MODELS.filter(model => model.is_active).map(model => ({
|
||||
id: model.id,
|
||||
name: model.name,
|
||||
display_name: model.display_name,
|
||||
is_active: model.is_active,
|
||||
default_price_per_request: model.default_price_per_request ?? null,
|
||||
default_tiered_pricing: model.default_tiered_pricing,
|
||||
supported_capabilities: model.supported_capabilities ?? null,
|
||||
supports_embedding: model.supports_embedding ?? null,
|
||||
config: model.config ?? null,
|
||||
usage_count: model.usage_count ?? 0,
|
||||
}))
|
||||
return createMockResponse({ models, total: models.length })
|
||||
},
|
||||
|
||||
'GET /api/users/me/preferences': async () => {
|
||||
await delay()
|
||||
return createMockResponse(getCurrentProfile().preferences || { theme: 'auto', language: 'zh-CN' })
|
||||
@@ -1146,6 +1182,7 @@ const mockHandlers: Record<string, (config: AxiosRequestConfig) => Promise<Axios
|
||||
default_tiered_pricing: m.default_tiered_pricing,
|
||||
default_price_per_request: m.default_price_per_request,
|
||||
supported_capabilities: m.supported_capabilities,
|
||||
supports_embedding: m.supports_embedding,
|
||||
config: m.config
|
||||
})),
|
||||
total: MOCK_GLOBAL_MODELS.length
|
||||
@@ -1419,6 +1456,8 @@ function generateMockModelsForProvider(providerId: string) {
|
||||
const hasClaude = provider.api_formats.some(f => f.includes('claude'))
|
||||
const hasOpenAI = provider.api_formats.some(f => f.includes('openai'))
|
||||
const hasGemini = provider.api_formats.some(f => f.includes('gemini'))
|
||||
const hasEmbedding = provider.api_formats.some(f => f.endsWith(':embedding'))
|
||||
const hasRerank = provider.api_formats.some(f => f.endsWith(':rerank'))
|
||||
|
||||
const models: Record<string, unknown>[] = []
|
||||
const now = new Date().toISOString()
|
||||
@@ -1503,6 +1542,66 @@ function generateMockModelsForProvider(providerId: string) {
|
||||
}
|
||||
)
|
||||
}
|
||||
if (hasEmbedding) {
|
||||
models.push({
|
||||
id: `pm-${providerId}-embedding-1`,
|
||||
provider_id: providerId,
|
||||
global_model_id: 'gm-010',
|
||||
provider_model_name: 'text-embedding-3-small',
|
||||
global_model_name: 'text-embedding-3-small',
|
||||
global_model_display_name: 'text-embedding-3-small',
|
||||
effective_input_price: 0.02,
|
||||
effective_output_price: 0,
|
||||
supports_embedding: true,
|
||||
effective_supports_embedding: true,
|
||||
supports_streaming: false,
|
||||
effective_supports_streaming: false,
|
||||
config: {
|
||||
embedding: true,
|
||||
model_type: 'embedding',
|
||||
api_formats: ['openai:embedding'],
|
||||
},
|
||||
effective_config: {
|
||||
embedding: true,
|
||||
model_type: 'embedding',
|
||||
api_formats: ['openai:embedding'],
|
||||
streaming: false,
|
||||
},
|
||||
is_active: true,
|
||||
is_available: true,
|
||||
created_at: provider.created_at,
|
||||
updated_at: now
|
||||
})
|
||||
}
|
||||
if (hasRerank) {
|
||||
models.push({
|
||||
id: `pm-${providerId}-rerank-1`,
|
||||
provider_id: providerId,
|
||||
global_model_id: 'gm-rerank-001',
|
||||
provider_model_name: 'bge-reranker-base',
|
||||
global_model_name: 'bge-reranker-base',
|
||||
global_model_display_name: 'bge-reranker-base',
|
||||
effective_input_price: 0.05,
|
||||
effective_output_price: 0,
|
||||
supports_streaming: false,
|
||||
effective_supports_streaming: false,
|
||||
config: {
|
||||
rerank: true,
|
||||
model_type: 'rerank',
|
||||
api_formats: ['openai:rerank'],
|
||||
},
|
||||
effective_config: {
|
||||
rerank: true,
|
||||
model_type: 'rerank',
|
||||
api_formats: ['openai:rerank'],
|
||||
streaming: false,
|
||||
},
|
||||
is_active: true,
|
||||
is_available: true,
|
||||
created_at: provider.created_at,
|
||||
updated_at: now
|
||||
})
|
||||
}
|
||||
if (hasGemini) {
|
||||
models.push(
|
||||
{
|
||||
|
||||
@@ -692,6 +692,7 @@ interface ModelProviderDisplay {
|
||||
supports_function_calling?: boolean | null
|
||||
supports_streaming?: boolean | null
|
||||
supports_extended_thinking?: boolean | null
|
||||
supports_embedding?: boolean | null
|
||||
}
|
||||
|
||||
const { success, error: showError } = useToast()
|
||||
@@ -766,6 +767,7 @@ const editingProviderModel = computed<Model | null>(() => {
|
||||
supports_vision: p.supports_vision,
|
||||
supports_function_calling: p.supports_function_calling,
|
||||
supports_extended_thinking: p.supports_extended_thinking,
|
||||
supports_embedding: p.supports_embedding,
|
||||
is_active: p.is_active,
|
||||
global_model_display_name: selectedModel.value?.display_name,
|
||||
} as Model
|
||||
@@ -1180,7 +1182,8 @@ async function loadModelProviders(_globalModelId: string) {
|
||||
// 能力信息
|
||||
supports_vision: p.supports_vision,
|
||||
supports_function_calling: p.supports_function_calling,
|
||||
supports_streaming: p.supports_streaming
|
||||
supports_streaming: p.supports_streaming,
|
||||
supports_embedding: p.supports_embedding
|
||||
}))
|
||||
} catch (err: unknown) {
|
||||
if (requestId !== modelProvidersRequestId) return
|
||||
|
||||
@@ -80,6 +80,14 @@
|
||||
<div>
|
||||
<div class="flex items-center gap-2">
|
||||
<span class="font-medium hover:text-primary transition-colors">{{ model.display_name || model.name }}</span>
|
||||
<Badge
|
||||
v-for="capability in getModelCapabilityLabels(model)"
|
||||
:key="capability"
|
||||
variant="secondary"
|
||||
class="text-[10px] px-1.5 py-0"
|
||||
>
|
||||
{{ capability }}
|
||||
</Badge>
|
||||
</div>
|
||||
<div class="text-xs text-muted-foreground flex items-center gap-1 mt-0.5">
|
||||
<span>{{ model.name }}</span>
|
||||
@@ -150,6 +158,16 @@
|
||||
<div class="flex items-start justify-between gap-3">
|
||||
<div class="flex-1 min-w-0">
|
||||
<span class="font-medium truncate block">{{ model.display_name || model.name }}</span>
|
||||
<div class="flex flex-wrap gap-1 mt-1">
|
||||
<Badge
|
||||
v-for="capability in getModelCapabilityLabels(model)"
|
||||
:key="capability"
|
||||
variant="secondary"
|
||||
class="text-[10px] px-1.5 py-0"
|
||||
>
|
||||
{{ capability }}
|
||||
</Badge>
|
||||
</div>
|
||||
<div class="text-xs text-muted-foreground flex items-center gap-1 mt-0.5">
|
||||
<span class="truncate">{{ model.name }}</span>
|
||||
<button
|
||||
@@ -228,6 +246,7 @@ import UserModelDetailDrawer from './components/UserModelDetailDrawer.vue'
|
||||
import { useRowClick } from '@/composables/useRowClick'
|
||||
import { log } from '@/utils/logger'
|
||||
import { parseApiError } from '@/utils/errorParser'
|
||||
import { getModelCapabilityLabels } from './model-catalog-helpers'
|
||||
|
||||
const { error: showError } = useToast()
|
||||
const { copyToClipboard } = useClipboard()
|
||||
|
||||
@@ -0,0 +1,51 @@
|
||||
import { describe, expect, it } from 'vitest'
|
||||
|
||||
import type { PublicGlobalModel } from '@/api/public-models'
|
||||
import { getModelCapabilityLabels, supportsEmbedding, supportsRerank } from '../model-catalog-helpers'
|
||||
|
||||
function model(overrides: Partial<PublicGlobalModel>): PublicGlobalModel {
|
||||
return {
|
||||
id: 'gm-test',
|
||||
name: 'model-test',
|
||||
display_name: 'Model Test',
|
||||
is_active: true,
|
||||
default_tiered_pricing: null,
|
||||
default_price_per_request: null,
|
||||
supported_capabilities: null,
|
||||
config: null,
|
||||
usage_count: 0,
|
||||
...overrides,
|
||||
}
|
||||
}
|
||||
|
||||
describe('model catalog embedding helpers', () => {
|
||||
it('labels embedding models distinctly from chat models', () => {
|
||||
expect(getModelCapabilityLabels(model({
|
||||
supported_capabilities: ['embedding'],
|
||||
config: { streaming: false, api_formats: ['openai:embedding'] },
|
||||
}))).toEqual(['Embedding'])
|
||||
|
||||
expect(getModelCapabilityLabels(model({
|
||||
config: { streaming: true },
|
||||
}))).toEqual(['Chat'])
|
||||
})
|
||||
|
||||
it('detects embedding metadata from explicit and config-derived frontend fields', () => {
|
||||
expect(supportsEmbedding(model({ supports_embedding: true }))).toBe(true)
|
||||
expect(supportsEmbedding(model({ config: { embedding: true } }))).toBe(true)
|
||||
expect(supportsEmbedding(model({ config: { model_type: 'embedding' } }))).toBe(true)
|
||||
expect(supportsEmbedding(model({ config: { api_formats: ['jina:embedding'] } }))).toBe(true)
|
||||
expect(supportsEmbedding(model({ config: { api_formats: ['openai:chat'] } }))).toBe(false)
|
||||
})
|
||||
|
||||
it('labels rerank models distinctly from chat and embedding models', () => {
|
||||
const rerank = model({
|
||||
supported_capabilities: ['rerank'],
|
||||
config: { streaming: false, api_formats: ['jina:rerank'] },
|
||||
})
|
||||
|
||||
expect(supportsRerank(rerank)).toBe(true)
|
||||
expect(getModelCapabilityLabels(rerank)).toEqual(['Rerank'])
|
||||
expect(supportsRerank(model({ config: { api_formats: ['openai:embedding'] } }))).toBe(false)
|
||||
})
|
||||
})
|
||||
@@ -96,6 +96,23 @@
|
||||
{{ model.config?.image_generation === true ? '支持' : '不支持' }}
|
||||
</Badge>
|
||||
</div>
|
||||
<div class="flex items-center gap-2 p-3 rounded-lg border">
|
||||
<Database class="w-5 h-5 text-muted-foreground" />
|
||||
<div class="flex-1">
|
||||
<p class="text-sm font-medium">
|
||||
Embedding
|
||||
</p>
|
||||
<p class="text-xs text-muted-foreground">
|
||||
向量嵌入
|
||||
</p>
|
||||
</div>
|
||||
<Badge
|
||||
:variant="supportsEmbedding(model) ? 'default' : 'secondary'"
|
||||
class="text-xs"
|
||||
>
|
||||
{{ supportsEmbedding(model) ? '支持' : '不支持' }}
|
||||
</Badge>
|
||||
</div>
|
||||
<div class="flex items-center gap-2 p-3 rounded-lg border">
|
||||
<Eye class="w-5 h-5 text-muted-foreground" />
|
||||
<div class="flex-1">
|
||||
@@ -303,6 +320,7 @@ import {
|
||||
Zap,
|
||||
Copy,
|
||||
Layers,
|
||||
Database,
|
||||
Image as ImageIcon
|
||||
} from 'lucide-vue-next'
|
||||
import { useEscapeKey } from '@/composables/useEscapeKey'
|
||||
@@ -376,6 +394,14 @@ function getFirst1hCachePrice(tieredPricing: TieredPricingConfig | undefined | n
|
||||
return get1hCachePrice(tieredPricing.tiers[0])
|
||||
}
|
||||
|
||||
function supportsEmbedding(model: PublicGlobalModel): boolean {
|
||||
return model.supports_embedding === true
|
||||
|| model.supported_capabilities?.includes('embedding') === true
|
||||
|| model.config?.embedding === true
|
||||
|| model.config?.model_type === 'embedding'
|
||||
|| (Array.isArray(model.config?.api_formats) && model.config.api_formats.some((format) => String(format).endsWith(':embedding')))
|
||||
}
|
||||
|
||||
// 添加 ESC 键监听
|
||||
useEscapeKey(() => {
|
||||
if (props.open) {
|
||||
|
||||
@@ -0,0 +1,41 @@
|
||||
import type { PublicGlobalModel } from '@/api/public-models'
|
||||
|
||||
export function supportsEmbedding(model: PublicGlobalModel): boolean {
|
||||
return model.supports_embedding === true
|
||||
|| model.supported_capabilities?.includes('embedding') === true
|
||||
|| model.config?.embedding === true
|
||||
|| model.config?.model_type === 'embedding'
|
||||
|| (Array.isArray(model.config?.api_formats) && model.config.api_formats.some((format) => String(format).endsWith(':embedding')))
|
||||
}
|
||||
|
||||
export function supportsRerank(model: PublicGlobalModel): boolean {
|
||||
return model.supported_capabilities?.includes('rerank') === true
|
||||
|| model.config?.rerank === true
|
||||
|| model.config?.model_type === 'rerank'
|
||||
|| (Array.isArray(model.config?.api_formats) && model.config.api_formats.some((format) => String(format).endsWith(':rerank')))
|
||||
}
|
||||
|
||||
export function hasVideoPricing(model: PublicGlobalModel): boolean {
|
||||
const billing = model.config?.billing
|
||||
const video = billing && typeof billing === 'object' && !Array.isArray(billing)
|
||||
? (billing as Record<string, unknown>).video
|
||||
: null
|
||||
const priceByResolution = video && typeof video === 'object' && !Array.isArray(video)
|
||||
? (video as Record<string, unknown>).price_per_second_by_resolution
|
||||
: null
|
||||
return !!priceByResolution && typeof priceByResolution === 'object' && Object.keys(priceByResolution).length > 0
|
||||
}
|
||||
|
||||
export function getModelCapabilityLabels(model: PublicGlobalModel): string[] {
|
||||
const labels: string[] = []
|
||||
if (supportsRerank(model)) {
|
||||
labels.push('Rerank')
|
||||
} else if (supportsEmbedding(model)) {
|
||||
labels.push('Embedding')
|
||||
} else {
|
||||
labels.push('Chat')
|
||||
}
|
||||
if (model.config?.image_generation === true) labels.push('Image')
|
||||
if (hasVideoPricing(model)) labels.push('Video')
|
||||
return labels
|
||||
}
|
||||
Reference in New Issue
Block a user