mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-02 01:10:23 +08:00
Fix/provider query non kiro tests (#311)
* fix(admin): enable local provider-query tests for non-kiro providers * test(admin): cover non-kiro provider-query model execution * fix(admin): preserve provider-query test errors and failover retries * fix(admin): handle provider-query test alias and HTTP retry edges * fix(admin): extend provider-query failover coverage and fallback behavior * fix(admin): prefer supported endpoints for provider-query tests * fix(admin): prefer provider-query endpoints with compatible keys * fix(admin): fall back to compatible provider-query endpoints * fix(admin): align provider-query local tests with transport policy --------- Co-authored-by: fawney19 <elky0401@gmail.com>
This commit is contained in:
@@ -12,7 +12,7 @@ use super::response::{
|
||||
};
|
||||
use crate::ai_pipeline::{maybe_build_sync_finalize_outcome, GatewayControlDecision};
|
||||
use crate::execution_runtime;
|
||||
use crate::handlers::admin::request::AdminAppState;
|
||||
use crate::handlers::admin::request::{AdminAppState, AdminGatewayProviderTransportSnapshot};
|
||||
use crate::model_fetch::ModelFetchRuntimeState;
|
||||
use crate::provider_transport::kiro::{
|
||||
build_kiro_generate_assistant_response_url, build_kiro_provider_headers,
|
||||
@@ -33,7 +33,7 @@ use aether_model_fetch::{
|
||||
};
|
||||
use axum::{
|
||||
body::{to_bytes, Body},
|
||||
http::{HeaderMap, HeaderName, HeaderValue},
|
||||
http::{self, HeaderMap, HeaderName, HeaderValue},
|
||||
response::{IntoResponse, Response},
|
||||
Json,
|
||||
};
|
||||
@@ -242,6 +242,73 @@ fn provider_query_key_supports_endpoint(
|
||||
.any(|value| value.eq_ignore_ascii_case(endpoint_api_format))
|
||||
}
|
||||
|
||||
fn provider_query_transport_supports_standard_test_execution(
|
||||
state: &AdminAppState<'_>,
|
||||
transport: &AdminGatewayProviderTransportSnapshot,
|
||||
api_format: &str,
|
||||
) -> bool {
|
||||
match api_format {
|
||||
"openai:chat" => {
|
||||
crate::provider_transport::policy::supports_local_openai_chat_transport(transport)
|
||||
}
|
||||
"claude:chat" => {
|
||||
crate::provider_transport::policy::supports_local_standard_transport_with_network(
|
||||
transport, api_format,
|
||||
)
|
||||
}
|
||||
"gemini:chat" => state.supports_local_gemini_transport_with_network(transport, api_format),
|
||||
_ => false,
|
||||
}
|
||||
}
|
||||
|
||||
async fn provider_query_select_preferred_non_kiro_endpoint(
|
||||
state: &AdminAppState<'_>,
|
||||
provider: &StoredProviderCatalogProvider,
|
||||
endpoints: &[StoredProviderCatalogEndpoint],
|
||||
keys: &[StoredProviderCatalogKey],
|
||||
selected_key_id: Option<&str>,
|
||||
) -> Option<StoredProviderCatalogEndpoint> {
|
||||
for endpoint in endpoints.iter().filter(|endpoint| endpoint.is_active) {
|
||||
if !provider_query_supports_standard_test_api_format(&endpoint.api_format) {
|
||||
continue;
|
||||
}
|
||||
for key in keys {
|
||||
if !key.is_active
|
||||
|| selected_key_id.is_some_and(|value| value != key.id.as_str())
|
||||
|| !provider_query_key_supports_endpoint(key, &endpoint.api_format)
|
||||
{
|
||||
continue;
|
||||
}
|
||||
let Ok(Some(transport)) = state
|
||||
.read_provider_transport_snapshot(&provider.id, &endpoint.id, &key.id)
|
||||
.await
|
||||
else {
|
||||
continue;
|
||||
};
|
||||
if provider_query_transport_supports_standard_test_execution(
|
||||
state,
|
||||
&transport,
|
||||
endpoint.api_format.as_str(),
|
||||
) {
|
||||
return Some(endpoint.clone());
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
endpoints
|
||||
.iter()
|
||||
.find(|endpoint| {
|
||||
endpoint.is_active
|
||||
&& keys.iter().any(|key| {
|
||||
key.is_active
|
||||
&& selected_key_id.is_none_or(|value| value == key.id.as_str())
|
||||
&& provider_query_key_supports_endpoint(key, &endpoint.api_format)
|
||||
})
|
||||
})
|
||||
.or_else(|| endpoints.iter().find(|endpoint| endpoint.is_active))
|
||||
.cloned()
|
||||
}
|
||||
|
||||
fn provider_query_test_key_sort_key(
|
||||
provider_type: &str,
|
||||
key: &StoredProviderCatalogKey,
|
||||
@@ -319,6 +386,7 @@ async fn provider_query_build_kiro_test_candidates(
|
||||
state: &AdminAppState<'_>,
|
||||
provider: &StoredProviderCatalogProvider,
|
||||
payload: &Value,
|
||||
requested_model_override: Option<&str>,
|
||||
) -> Result<Vec<ProviderQueryTestCandidate>, Response<Body>> {
|
||||
let provider_ids = vec![provider.id.clone()];
|
||||
let endpoints = state
|
||||
@@ -330,29 +398,6 @@ async fn provider_query_build_kiro_test_candidates(
|
||||
ADMIN_PROVIDER_QUERY_NO_ACTIVE_API_KEY_DETAIL,
|
||||
)
|
||||
})?;
|
||||
let endpoint = match provider_query_select_kiro_endpoint(
|
||||
&endpoints,
|
||||
provider_query_extract_endpoint_id(payload).as_deref(),
|
||||
provider_query_extract_api_format(payload).as_deref(),
|
||||
) {
|
||||
Ok(Some(endpoint)) => endpoint.clone(),
|
||||
Ok(None) => {
|
||||
return Err(build_admin_provider_query_not_found_response(
|
||||
ADMIN_PROVIDER_QUERY_NO_ACTIVE_API_KEY_DETAIL,
|
||||
));
|
||||
}
|
||||
Err("Endpoint not found") => {
|
||||
return Err(build_admin_provider_query_not_found_response(
|
||||
ADMIN_PROVIDER_QUERY_API_KEY_NOT_FOUND_DETAIL,
|
||||
));
|
||||
}
|
||||
Err(_) => {
|
||||
return Err(build_admin_provider_query_not_found_response(
|
||||
ADMIN_PROVIDER_QUERY_NO_ACTIVE_API_KEY_DETAIL,
|
||||
));
|
||||
}
|
||||
};
|
||||
|
||||
let all_keys = state
|
||||
.app()
|
||||
.list_provider_catalog_keys_by_provider_ids(&provider_ids)
|
||||
@@ -362,8 +407,51 @@ async fn provider_query_build_kiro_test_candidates(
|
||||
ADMIN_PROVIDER_QUERY_NO_ACTIVE_API_KEY_DETAIL,
|
||||
)
|
||||
})?;
|
||||
|
||||
let selected_key_id = provider_query_extract_api_key_id(payload);
|
||||
let requested_endpoint_id = provider_query_extract_endpoint_id(payload);
|
||||
let requested_api_format = provider_query_extract_api_format(payload);
|
||||
let endpoint = if requested_endpoint_id.is_none()
|
||||
&& requested_api_format.is_none()
|
||||
&& !provider.provider_type.trim().eq_ignore_ascii_case("kiro")
|
||||
{
|
||||
provider_query_select_preferred_non_kiro_endpoint(
|
||||
state,
|
||||
provider,
|
||||
&endpoints,
|
||||
&all_keys,
|
||||
selected_key_id.as_deref(),
|
||||
)
|
||||
.await
|
||||
.ok_or_else(|| {
|
||||
build_admin_provider_query_not_found_response(
|
||||
ADMIN_PROVIDER_QUERY_NO_ACTIVE_API_KEY_DETAIL,
|
||||
)
|
||||
})?
|
||||
} else {
|
||||
match provider_query_select_kiro_endpoint(
|
||||
&endpoints,
|
||||
requested_endpoint_id.as_deref(),
|
||||
requested_api_format.as_deref(),
|
||||
) {
|
||||
Ok(Some(endpoint)) => endpoint.clone(),
|
||||
Ok(None) => {
|
||||
return Err(build_admin_provider_query_not_found_response(
|
||||
ADMIN_PROVIDER_QUERY_NO_ACTIVE_API_KEY_DETAIL,
|
||||
));
|
||||
}
|
||||
Err("Endpoint not found") => {
|
||||
return Err(build_admin_provider_query_not_found_response(
|
||||
ADMIN_PROVIDER_QUERY_API_KEY_NOT_FOUND_DETAIL,
|
||||
));
|
||||
}
|
||||
Err(_) => {
|
||||
return Err(build_admin_provider_query_not_found_response(
|
||||
ADMIN_PROVIDER_QUERY_NO_ACTIVE_API_KEY_DETAIL,
|
||||
));
|
||||
}
|
||||
}
|
||||
};
|
||||
|
||||
if let Some(api_key_id) = selected_key_id.as_deref() {
|
||||
let Some(key) = all_keys.iter().find(|key| key.id == api_key_id) else {
|
||||
return Err(build_admin_provider_query_not_found_response(
|
||||
@@ -377,9 +465,19 @@ async fn provider_query_build_kiro_test_candidates(
|
||||
}
|
||||
}
|
||||
|
||||
let requested_model = provider_query_extract_model(payload).ok_or_else(|| {
|
||||
build_admin_provider_query_bad_request_response(ADMIN_PROVIDER_QUERY_MODEL_REQUIRED_DETAIL)
|
||||
})?;
|
||||
let requested_model = requested_model_override
|
||||
.map(ToOwned::to_owned)
|
||||
.or_else(|| provider_query_extract_model(payload))
|
||||
.or_else(|| {
|
||||
super::payload::provider_query_extract_failover_models(payload)
|
||||
.first()
|
||||
.cloned()
|
||||
})
|
||||
.ok_or_else(|| {
|
||||
build_admin_provider_query_bad_request_response(
|
||||
ADMIN_PROVIDER_QUERY_MODEL_REQUIRED_DETAIL,
|
||||
)
|
||||
})?;
|
||||
let effective_model = if provider_query_test_mode(payload).eq_ignore_ascii_case("direct") {
|
||||
requested_model.clone()
|
||||
} else {
|
||||
@@ -661,7 +759,8 @@ async fn provider_query_execute_kiro_test_candidate(
|
||||
} else {
|
||||
result.body.as_ref().and_then(|body| body.json_body.clone())
|
||||
};
|
||||
let error_message = if result.status_code >= 400 {
|
||||
let did_fail = result.status_code >= 400;
|
||||
let error_message = if did_fail {
|
||||
provider_query_extract_error_message(&result)
|
||||
} else if response_body.is_none()
|
||||
&& provider_query_decode_execution_body(&result)
|
||||
@@ -673,7 +772,7 @@ async fn provider_query_execute_kiro_test_candidate(
|
||||
};
|
||||
|
||||
Ok(ProviderQueryExecutionOutcome {
|
||||
status: if error_message.is_some() {
|
||||
status: if did_fail || error_message.is_some() {
|
||||
"failed"
|
||||
} else {
|
||||
"success"
|
||||
@@ -689,6 +788,325 @@ async fn provider_query_execute_kiro_test_candidate(
|
||||
})
|
||||
}
|
||||
|
||||
async fn provider_query_execute_standard_test_candidate(
|
||||
state: &AdminAppState<'_>,
|
||||
provider: &StoredProviderCatalogProvider,
|
||||
candidate: &ProviderQueryTestCandidate,
|
||||
payload: &Value,
|
||||
route_path: &str,
|
||||
trace_id: &str,
|
||||
) -> Result<ProviderQueryExecutionOutcome, GatewayError> {
|
||||
let Some(transport) = state
|
||||
.read_provider_transport_snapshot(&provider.id, &candidate.endpoint.id, &candidate.key.id)
|
||||
.await?
|
||||
else {
|
||||
return Ok(ProviderQueryExecutionOutcome {
|
||||
status: "skipped",
|
||||
error_message: None,
|
||||
status_code: None,
|
||||
latency_ms: None,
|
||||
request_url: String::new(),
|
||||
request_headers: BTreeMap::new(),
|
||||
request_body: Value::Null,
|
||||
response_headers: BTreeMap::new(),
|
||||
response_body: None,
|
||||
});
|
||||
};
|
||||
if !provider_query_transport_supports_standard_test_execution(
|
||||
state,
|
||||
&transport,
|
||||
candidate.endpoint.api_format.as_str(),
|
||||
) {
|
||||
return Ok(ProviderQueryExecutionOutcome {
|
||||
status: "skipped",
|
||||
error_message: None,
|
||||
status_code: None,
|
||||
latency_ms: None,
|
||||
request_url: String::new(),
|
||||
request_headers: BTreeMap::new(),
|
||||
request_body: Value::Null,
|
||||
response_headers: BTreeMap::new(),
|
||||
response_body: None,
|
||||
});
|
||||
}
|
||||
|
||||
let original_request_body =
|
||||
provider_query_build_test_request_body(payload, &candidate.effective_model);
|
||||
let mut request_body = original_request_body.clone();
|
||||
if let Some(object) = request_body.as_object_mut() {
|
||||
object.insert("stream".to_string(), Value::Bool(false));
|
||||
}
|
||||
|
||||
let provider_api_format = candidate.endpoint.api_format.as_str();
|
||||
let provider_request_body = match provider_api_format {
|
||||
"openai:chat" => {
|
||||
let Some(mut provider_request_body) =
|
||||
crate::ai_pipeline::build_local_openai_chat_request_body(
|
||||
&request_body,
|
||||
&candidate.effective_model,
|
||||
false,
|
||||
)
|
||||
else {
|
||||
return Ok(ProviderQueryExecutionOutcome {
|
||||
status: "skipped",
|
||||
error_message: None,
|
||||
status_code: None,
|
||||
latency_ms: None,
|
||||
request_url: String::new(),
|
||||
request_headers: BTreeMap::new(),
|
||||
request_body,
|
||||
response_headers: BTreeMap::new(),
|
||||
response_body: None,
|
||||
});
|
||||
};
|
||||
if !crate::provider_transport::apply_local_body_rules(
|
||||
&mut provider_request_body,
|
||||
transport.endpoint.body_rules.as_ref(),
|
||||
Some(&request_body),
|
||||
) {
|
||||
None
|
||||
} else {
|
||||
Some(provider_request_body)
|
||||
}
|
||||
}
|
||||
"claude:chat" | "gemini:chat" => {
|
||||
let Some(mut provider_request_body) =
|
||||
crate::ai_pipeline::build_cross_format_openai_chat_request_body(
|
||||
&request_body,
|
||||
&candidate.effective_model,
|
||||
provider_api_format,
|
||||
false,
|
||||
)
|
||||
else {
|
||||
return Ok(ProviderQueryExecutionOutcome {
|
||||
status: "skipped",
|
||||
error_message: None,
|
||||
status_code: None,
|
||||
latency_ms: None,
|
||||
request_url: String::new(),
|
||||
request_headers: BTreeMap::new(),
|
||||
request_body,
|
||||
response_headers: BTreeMap::new(),
|
||||
response_body: None,
|
||||
});
|
||||
};
|
||||
if !crate::provider_transport::apply_local_body_rules(
|
||||
&mut provider_request_body,
|
||||
transport.endpoint.body_rules.as_ref(),
|
||||
Some(&request_body),
|
||||
) {
|
||||
None
|
||||
} else {
|
||||
Some(provider_request_body)
|
||||
}
|
||||
}
|
||||
_ => None,
|
||||
};
|
||||
let Some(provider_request_body) = provider_request_body else {
|
||||
return Ok(ProviderQueryExecutionOutcome {
|
||||
status: "skipped",
|
||||
error_message: None,
|
||||
status_code: None,
|
||||
latency_ms: None,
|
||||
request_url: String::new(),
|
||||
request_headers: BTreeMap::new(),
|
||||
request_body,
|
||||
response_headers: BTreeMap::new(),
|
||||
response_body: None,
|
||||
});
|
||||
};
|
||||
|
||||
let oauth_auth = match provider_api_format {
|
||||
"openai:chat" | "claude:chat" => state.resolve_local_oauth_header_auth(&transport).await?,
|
||||
_ => None,
|
||||
};
|
||||
let auth = match provider_api_format {
|
||||
"openai:chat" => {
|
||||
crate::provider_transport::auth::resolve_local_openai_bearer_auth(&transport)
|
||||
.or(oauth_auth)
|
||||
}
|
||||
"claude:chat" => {
|
||||
crate::provider_transport::auth::resolve_local_standard_auth(&transport).or(oauth_auth)
|
||||
}
|
||||
"gemini:chat" => state.resolve_local_gemini_auth(&transport),
|
||||
_ => None,
|
||||
};
|
||||
let Some((auth_header, auth_value)) = auth else {
|
||||
return Ok(ProviderQueryExecutionOutcome {
|
||||
status: "skipped",
|
||||
error_message: None,
|
||||
status_code: None,
|
||||
latency_ms: None,
|
||||
request_url: String::new(),
|
||||
request_headers: BTreeMap::new(),
|
||||
request_body: provider_request_body,
|
||||
response_headers: BTreeMap::new(),
|
||||
response_body: None,
|
||||
});
|
||||
};
|
||||
|
||||
let mut synthetic_request = http::Request::builder()
|
||||
.uri(route_path)
|
||||
.body(())
|
||||
.map_err(|err| GatewayError::Internal(err.to_string()))?;
|
||||
*synthetic_request.headers_mut() = provider_query_extract_request_headers(payload);
|
||||
let (parts, _) = synthetic_request.into_parts();
|
||||
|
||||
let request_url = match provider_api_format {
|
||||
"openai:chat" => {
|
||||
let custom_path = transport
|
||||
.endpoint
|
||||
.custom_path
|
||||
.as_deref()
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty());
|
||||
match custom_path {
|
||||
Some(path) => state.build_passthrough_path_url(
|
||||
&transport.endpoint.base_url,
|
||||
path,
|
||||
parts.uri.query(),
|
||||
&[],
|
||||
),
|
||||
None => Some(
|
||||
state.build_openai_chat_url(&transport.endpoint.base_url, parts.uri.query()),
|
||||
),
|
||||
}
|
||||
}
|
||||
"claude:chat" => {
|
||||
let custom_path = transport
|
||||
.endpoint
|
||||
.custom_path
|
||||
.as_deref()
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty());
|
||||
match custom_path {
|
||||
Some(path) => state.build_passthrough_path_url(
|
||||
&transport.endpoint.base_url,
|
||||
path,
|
||||
parts.uri.query(),
|
||||
&[],
|
||||
),
|
||||
None => Some(
|
||||
state
|
||||
.build_claude_messages_url(&transport.endpoint.base_url, parts.uri.query()),
|
||||
),
|
||||
}
|
||||
}
|
||||
"gemini:chat" => {
|
||||
let custom_path = transport
|
||||
.endpoint
|
||||
.custom_path
|
||||
.as_deref()
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty());
|
||||
match custom_path {
|
||||
Some(path) => state.build_passthrough_path_url(
|
||||
&transport.endpoint.base_url,
|
||||
path,
|
||||
parts.uri.query(),
|
||||
&["key"],
|
||||
),
|
||||
None => state.build_gemini_content_url(
|
||||
&transport.endpoint.base_url,
|
||||
&candidate.effective_model,
|
||||
false,
|
||||
parts.uri.query(),
|
||||
),
|
||||
}
|
||||
}
|
||||
_ => None,
|
||||
}
|
||||
.ok_or_else(|| GatewayError::Internal("provider request url is unavailable".to_string()))?;
|
||||
|
||||
let mut request_headers = match provider_api_format {
|
||||
"claude:chat" => crate::provider_transport::auth::build_claude_passthrough_headers(
|
||||
&parts.headers,
|
||||
&auth_header,
|
||||
&auth_value,
|
||||
&BTreeMap::new(),
|
||||
Some("application/json"),
|
||||
),
|
||||
_ => crate::provider_transport::auth::build_openai_passthrough_headers(
|
||||
&parts.headers,
|
||||
&auth_header,
|
||||
&auth_value,
|
||||
&BTreeMap::new(),
|
||||
Some("application/json"),
|
||||
),
|
||||
};
|
||||
if !state.apply_local_header_rules(
|
||||
&mut request_headers,
|
||||
transport.endpoint.header_rules.as_ref(),
|
||||
&[auth_header.as_str(), "content-type"],
|
||||
&provider_request_body,
|
||||
Some(&request_body),
|
||||
) {
|
||||
return Ok(ProviderQueryExecutionOutcome {
|
||||
status: "failed",
|
||||
error_message: Some("provider request headers build failed".to_string()),
|
||||
status_code: None,
|
||||
latency_ms: None,
|
||||
request_url,
|
||||
request_headers,
|
||||
request_body: provider_request_body,
|
||||
response_headers: BTreeMap::new(),
|
||||
response_body: None,
|
||||
});
|
||||
}
|
||||
crate::provider_transport::ensure_upstream_auth_header(
|
||||
&mut request_headers,
|
||||
&auth_header,
|
||||
&auth_value,
|
||||
);
|
||||
|
||||
let plan = ExecutionPlan {
|
||||
request_id: trace_id.to_string(),
|
||||
candidate_id: Some(format!("provider-query-{}", candidate.key.id)),
|
||||
provider_name: Some(provider.name.clone()),
|
||||
provider_id: provider.id.clone(),
|
||||
endpoint_id: candidate.endpoint.id.clone(),
|
||||
key_id: candidate.key.id.clone(),
|
||||
method: "POST".to_string(),
|
||||
url: request_url.clone(),
|
||||
headers: request_headers.clone(),
|
||||
content_type: Some("application/json".to_string()),
|
||||
content_encoding: None,
|
||||
body: RequestBody::from_json(provider_request_body.clone()),
|
||||
stream: false,
|
||||
client_api_format: "openai:chat".to_string(),
|
||||
provider_api_format: candidate.endpoint.api_format.clone(),
|
||||
model_name: Some(candidate.effective_model.clone()),
|
||||
proxy: state
|
||||
.resolve_transport_proxy_snapshot_with_tunnel_affinity(&transport)
|
||||
.await,
|
||||
tls_profile: state.resolve_transport_tls_profile(&transport),
|
||||
timeouts: state.resolve_transport_execution_timeouts(&transport),
|
||||
};
|
||||
|
||||
let result = state
|
||||
.execute_execution_runtime_sync_plan(Some(trace_id), &plan)
|
||||
.await?;
|
||||
let response_body = result.body.as_ref().and_then(|body| body.json_body.clone());
|
||||
let did_fail = result.status_code >= 400;
|
||||
let error_message = if did_fail {
|
||||
provider_query_extract_error_message(&result)
|
||||
} else {
|
||||
None
|
||||
};
|
||||
|
||||
Ok(ProviderQueryExecutionOutcome {
|
||||
status: if did_fail { "failed" } else { "success" },
|
||||
error_message,
|
||||
status_code: Some(result.status_code),
|
||||
latency_ms: result.telemetry.as_ref().and_then(|value| value.elapsed_ms),
|
||||
request_url,
|
||||
request_headers,
|
||||
request_body: provider_request_body,
|
||||
response_headers: result.headers,
|
||||
response_body,
|
||||
})
|
||||
}
|
||||
|
||||
fn provider_query_test_attempt_payload(
|
||||
candidate_index: usize,
|
||||
candidate: &ProviderQueryTestCandidate,
|
||||
@@ -716,6 +1134,10 @@ fn provider_query_test_attempt_payload(
|
||||
})
|
||||
}
|
||||
|
||||
fn provider_query_supports_standard_test_api_format(api_format: &str) -> bool {
|
||||
matches!(api_format, "openai:chat" | "claude:chat" | "gemini:chat")
|
||||
}
|
||||
|
||||
async fn build_admin_provider_query_kiro_failover_response(
|
||||
state: &AdminAppState<'_>,
|
||||
payload: &Value,
|
||||
@@ -726,14 +1148,6 @@ async fn build_admin_provider_query_kiro_failover_response(
|
||||
ADMIN_PROVIDER_QUERY_PROVIDER_ID_REQUIRED_DETAIL,
|
||||
));
|
||||
};
|
||||
let requested_model = match provider_query_extract_model(payload) {
|
||||
Some(model) => model,
|
||||
None => {
|
||||
return Ok(build_admin_provider_query_bad_request_response(
|
||||
ADMIN_PROVIDER_QUERY_MODEL_REQUIRED_DETAIL,
|
||||
));
|
||||
}
|
||||
};
|
||||
let Some(provider) = state
|
||||
.app()
|
||||
.read_provider_catalog_providers_by_ids(std::slice::from_ref(&provider_id))
|
||||
@@ -745,18 +1159,66 @@ async fn build_admin_provider_query_kiro_failover_response(
|
||||
ADMIN_PROVIDER_QUERY_PROVIDER_NOT_FOUND_DETAIL,
|
||||
));
|
||||
};
|
||||
if !provider.provider_type.trim().eq_ignore_ascii_case("kiro") {
|
||||
let failover_models = super::payload::provider_query_extract_failover_models(payload);
|
||||
let is_kiro = provider.provider_type.trim().eq_ignore_ascii_case("kiro");
|
||||
let Some(requested_model) =
|
||||
provider_query_extract_model(payload).or_else(|| failover_models.first().cloned())
|
||||
else {
|
||||
return Ok(build_admin_provider_query_bad_request_response(
|
||||
ADMIN_PROVIDER_QUERY_MODEL_REQUIRED_DETAIL,
|
||||
));
|
||||
};
|
||||
|
||||
let requested_models = if is_kiro {
|
||||
vec![requested_model.clone()]
|
||||
} else if failover_models.is_empty() {
|
||||
vec![requested_model.clone()]
|
||||
} else {
|
||||
failover_models.clone()
|
||||
};
|
||||
let mut candidates = Vec::new();
|
||||
for requested_failover_model in &requested_models {
|
||||
match provider_query_build_kiro_test_candidates(
|
||||
state,
|
||||
&provider,
|
||||
payload,
|
||||
Some(requested_failover_model.as_str()),
|
||||
)
|
||||
.await
|
||||
{
|
||||
Ok(mut built_candidates) => candidates.append(&mut built_candidates),
|
||||
Err(response) => return Ok(response),
|
||||
}
|
||||
}
|
||||
if !is_kiro {
|
||||
let mut supported_candidates = Vec::new();
|
||||
for candidate in candidates {
|
||||
let Some(transport) = state
|
||||
.read_provider_transport_snapshot(
|
||||
&provider.id,
|
||||
&candidate.endpoint.id,
|
||||
&candidate.key.id,
|
||||
)
|
||||
.await?
|
||||
else {
|
||||
continue;
|
||||
};
|
||||
if provider_query_transport_supports_standard_test_execution(
|
||||
state,
|
||||
&transport,
|
||||
candidate.endpoint.api_format.as_str(),
|
||||
) {
|
||||
supported_candidates.push(candidate);
|
||||
}
|
||||
}
|
||||
candidates = supported_candidates;
|
||||
}
|
||||
if !is_kiro && candidates.is_empty() {
|
||||
return Ok(build_admin_provider_query_test_model_failover_response(
|
||||
provider_id,
|
||||
super::payload::provider_query_extract_failover_models(payload),
|
||||
requested_models,
|
||||
));
|
||||
}
|
||||
|
||||
let candidates =
|
||||
match provider_query_build_kiro_test_candidates(state, &provider, payload).await {
|
||||
Ok(candidates) => candidates,
|
||||
Err(response) => return Ok(response),
|
||||
};
|
||||
let trace_id = provider_query_extract_request_id(payload)
|
||||
.unwrap_or_else(|| format!("provider-query-test-{}", Uuid::new_v4().simple()));
|
||||
let mut attempts = Vec::new();
|
||||
@@ -764,16 +1226,23 @@ async fn build_admin_provider_query_kiro_failover_response(
|
||||
let mut success_body = None;
|
||||
|
||||
for (candidate_index, candidate) in candidates.iter().enumerate() {
|
||||
let execution = provider_query_execute_kiro_test_candidate(
|
||||
state,
|
||||
&provider,
|
||||
candidate,
|
||||
payload,
|
||||
route_path,
|
||||
&trace_id,
|
||||
&requested_model,
|
||||
)
|
||||
.await?;
|
||||
let execution = if is_kiro {
|
||||
provider_query_execute_kiro_test_candidate(
|
||||
state,
|
||||
&provider,
|
||||
candidate,
|
||||
payload,
|
||||
route_path,
|
||||
&trace_id,
|
||||
&requested_model,
|
||||
)
|
||||
.await?
|
||||
} else {
|
||||
provider_query_execute_standard_test_candidate(
|
||||
state, &provider, candidate, payload, route_path, &trace_id,
|
||||
)
|
||||
.await?
|
||||
};
|
||||
if execution.status != "skipped" {
|
||||
total_attempts += 1;
|
||||
}
|
||||
@@ -814,7 +1283,7 @@ async fn build_admin_provider_query_kiro_failover_response(
|
||||
"total_candidates": candidates.len(),
|
||||
"total_attempts": total_attempts,
|
||||
"data": success_body.as_ref().map(|body| json!({
|
||||
"stream": true,
|
||||
"stream": is_kiro,
|
||||
"response": body,
|
||||
})),
|
||||
"error": error,
|
||||
@@ -832,6 +1301,9 @@ pub(crate) async fn build_admin_provider_query_test_model_local_response(
|
||||
"/api/admin/provider-query/test-model",
|
||||
)
|
||||
.await?;
|
||||
if !response.status().is_success() {
|
||||
return Ok(response);
|
||||
}
|
||||
let body = to_bytes(response.into_body(), usize::MAX)
|
||||
.await
|
||||
.map_err(|err| GatewayError::Internal(err.to_string()))?;
|
||||
|
||||
@@ -2,9 +2,7 @@ use crate::handlers::admin::provider::query::{
|
||||
models::{
|
||||
build_admin_provider_query_models_response,
|
||||
build_admin_provider_query_test_model_failover_local_response,
|
||||
build_admin_provider_query_test_model_failover_response,
|
||||
build_admin_provider_query_test_model_local_response,
|
||||
build_admin_provider_query_test_model_response,
|
||||
},
|
||||
payload::{
|
||||
parse_admin_provider_query_body, provider_query_extract_failover_models,
|
||||
@@ -58,7 +56,7 @@ impl<'a> AdminAppState<'a> {
|
||||
build_admin_provider_query_models_response(self, &payload).await?,
|
||||
)),
|
||||
"test_model" => {
|
||||
let Some(provider_id) = provider_query_extract_provider_id(&payload) else {
|
||||
let Some(_provider_id) = provider_query_extract_provider_id(&payload) else {
|
||||
log_admin_provider_query_validation_failure(
|
||||
request_context,
|
||||
route_kind,
|
||||
@@ -69,7 +67,7 @@ impl<'a> AdminAppState<'a> {
|
||||
ADMIN_PROVIDER_QUERY_PROVIDER_ID_REQUIRED_DETAIL,
|
||||
)));
|
||||
};
|
||||
let Some(model) = provider_query_extract_model(&payload) else {
|
||||
let Some(_model) = provider_query_extract_model(&payload) else {
|
||||
log_admin_provider_query_validation_failure(
|
||||
request_context,
|
||||
route_kind,
|
||||
@@ -80,28 +78,12 @@ impl<'a> AdminAppState<'a> {
|
||||
ADMIN_PROVIDER_QUERY_MODEL_REQUIRED_DETAIL,
|
||||
)));
|
||||
};
|
||||
let provider_type = self
|
||||
.app()
|
||||
.read_provider_catalog_providers_by_ids(std::slice::from_ref(&provider_id))
|
||||
.await?
|
||||
.into_iter()
|
||||
.find(|provider| provider.id == provider_id)
|
||||
.map(|provider| provider.provider_type)
|
||||
.unwrap_or_default();
|
||||
if provider_type.trim().eq_ignore_ascii_case("kiro") {
|
||||
Ok(Some(
|
||||
build_admin_provider_query_test_model_local_response(self, &payload)
|
||||
.await?,
|
||||
))
|
||||
} else {
|
||||
Ok(Some(build_admin_provider_query_test_model_response(
|
||||
provider_id,
|
||||
model,
|
||||
)))
|
||||
}
|
||||
Ok(Some(
|
||||
build_admin_provider_query_test_model_local_response(self, &payload).await?,
|
||||
))
|
||||
}
|
||||
"test_model_failover" => {
|
||||
let Some(provider_id) = provider_query_extract_provider_id(&payload) else {
|
||||
let Some(_provider_id) = provider_query_extract_provider_id(&payload) else {
|
||||
log_admin_provider_query_validation_failure(
|
||||
request_context,
|
||||
route_kind,
|
||||
@@ -124,29 +106,10 @@ impl<'a> AdminAppState<'a> {
|
||||
ADMIN_PROVIDER_QUERY_FAILOVER_MODELS_REQUIRED_DETAIL,
|
||||
)));
|
||||
}
|
||||
let provider_type = self
|
||||
.app()
|
||||
.read_provider_catalog_providers_by_ids(std::slice::from_ref(&provider_id))
|
||||
.await?
|
||||
.into_iter()
|
||||
.find(|provider| provider.id == provider_id)
|
||||
.map(|provider| provider.provider_type)
|
||||
.unwrap_or_default();
|
||||
if provider_type.trim().eq_ignore_ascii_case("kiro") {
|
||||
Ok(Some(
|
||||
build_admin_provider_query_test_model_failover_local_response(
|
||||
self, &payload,
|
||||
)
|
||||
Ok(Some(
|
||||
build_admin_provider_query_test_model_failover_local_response(self, &payload)
|
||||
.await?,
|
||||
))
|
||||
} else {
|
||||
Ok(Some(
|
||||
build_admin_provider_query_test_model_failover_response(
|
||||
provider_id,
|
||||
failover_models,
|
||||
),
|
||||
))
|
||||
}
|
||||
))
|
||||
}
|
||||
_ => Ok(Some(
|
||||
build_admin_provider_query_models_response(self, &payload).await?,
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
Reference in New Issue
Block a user