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:
RWDai
2026-04-17 19:45:20 +08:00
committed by GitHub
parent 0ce61bc91c
commit dd4641d618
3 changed files with 1775 additions and 147 deletions

View File

@@ -12,7 +12,7 @@ use super::response::{
}; };
use crate::ai_pipeline::{maybe_build_sync_finalize_outcome, GatewayControlDecision}; use crate::ai_pipeline::{maybe_build_sync_finalize_outcome, GatewayControlDecision};
use crate::execution_runtime; use crate::execution_runtime;
use crate::handlers::admin::request::AdminAppState; use crate::handlers::admin::request::{AdminAppState, AdminGatewayProviderTransportSnapshot};
use crate::model_fetch::ModelFetchRuntimeState; use crate::model_fetch::ModelFetchRuntimeState;
use crate::provider_transport::kiro::{ use crate::provider_transport::kiro::{
build_kiro_generate_assistant_response_url, build_kiro_provider_headers, build_kiro_generate_assistant_response_url, build_kiro_provider_headers,
@@ -33,7 +33,7 @@ use aether_model_fetch::{
}; };
use axum::{ use axum::{
body::{to_bytes, Body}, body::{to_bytes, Body},
http::{HeaderMap, HeaderName, HeaderValue}, http::{self, HeaderMap, HeaderName, HeaderValue},
response::{IntoResponse, Response}, response::{IntoResponse, Response},
Json, Json,
}; };
@@ -242,6 +242,73 @@ fn provider_query_key_supports_endpoint(
.any(|value| value.eq_ignore_ascii_case(endpoint_api_format)) .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( fn provider_query_test_key_sort_key(
provider_type: &str, provider_type: &str,
key: &StoredProviderCatalogKey, key: &StoredProviderCatalogKey,
@@ -319,6 +386,7 @@ async fn provider_query_build_kiro_test_candidates(
state: &AdminAppState<'_>, state: &AdminAppState<'_>,
provider: &StoredProviderCatalogProvider, provider: &StoredProviderCatalogProvider,
payload: &Value, payload: &Value,
requested_model_override: Option<&str>,
) -> Result<Vec<ProviderQueryTestCandidate>, Response<Body>> { ) -> Result<Vec<ProviderQueryTestCandidate>, Response<Body>> {
let provider_ids = vec![provider.id.clone()]; let provider_ids = vec![provider.id.clone()];
let endpoints = state let endpoints = state
@@ -330,29 +398,6 @@ async fn provider_query_build_kiro_test_candidates(
ADMIN_PROVIDER_QUERY_NO_ACTIVE_API_KEY_DETAIL, 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 let all_keys = state
.app() .app()
.list_provider_catalog_keys_by_provider_ids(&provider_ids) .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, ADMIN_PROVIDER_QUERY_NO_ACTIVE_API_KEY_DETAIL,
) )
})?; })?;
let selected_key_id = provider_query_extract_api_key_id(payload); 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() { 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 { let Some(key) = all_keys.iter().find(|key| key.id == api_key_id) else {
return Err(build_admin_provider_query_not_found_response( 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(|| { let requested_model = requested_model_override
build_admin_provider_query_bad_request_response(ADMIN_PROVIDER_QUERY_MODEL_REQUIRED_DETAIL) .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") { let effective_model = if provider_query_test_mode(payload).eq_ignore_ascii_case("direct") {
requested_model.clone() requested_model.clone()
} else { } else {
@@ -661,7 +759,8 @@ async fn provider_query_execute_kiro_test_candidate(
} else { } else {
result.body.as_ref().and_then(|body| body.json_body.clone()) 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) provider_query_extract_error_message(&result)
} else if response_body.is_none() } else if response_body.is_none()
&& provider_query_decode_execution_body(&result) && provider_query_decode_execution_body(&result)
@@ -673,7 +772,7 @@ async fn provider_query_execute_kiro_test_candidate(
}; };
Ok(ProviderQueryExecutionOutcome { Ok(ProviderQueryExecutionOutcome {
status: if error_message.is_some() { status: if did_fail || error_message.is_some() {
"failed" "failed"
} else { } else {
"success" "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( fn provider_query_test_attempt_payload(
candidate_index: usize, candidate_index: usize,
candidate: &ProviderQueryTestCandidate, 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( async fn build_admin_provider_query_kiro_failover_response(
state: &AdminAppState<'_>, state: &AdminAppState<'_>,
payload: &Value, payload: &Value,
@@ -726,14 +1148,6 @@ async fn build_admin_provider_query_kiro_failover_response(
ADMIN_PROVIDER_QUERY_PROVIDER_ID_REQUIRED_DETAIL, 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 let Some(provider) = state
.app() .app()
.read_provider_catalog_providers_by_ids(std::slice::from_ref(&provider_id)) .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, 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( return Ok(build_admin_provider_query_test_model_failover_response(
provider_id, 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) let trace_id = provider_query_extract_request_id(payload)
.unwrap_or_else(|| format!("provider-query-test-{}", Uuid::new_v4().simple())); .unwrap_or_else(|| format!("provider-query-test-{}", Uuid::new_v4().simple()));
let mut attempts = Vec::new(); let mut attempts = Vec::new();
@@ -764,16 +1226,23 @@ async fn build_admin_provider_query_kiro_failover_response(
let mut success_body = None; let mut success_body = None;
for (candidate_index, candidate) in candidates.iter().enumerate() { for (candidate_index, candidate) in candidates.iter().enumerate() {
let execution = provider_query_execute_kiro_test_candidate( let execution = if is_kiro {
state, provider_query_execute_kiro_test_candidate(
&provider, state,
candidate, &provider,
payload, candidate,
route_path, payload,
&trace_id, route_path,
&requested_model, &trace_id,
) &requested_model,
.await?; )
.await?
} else {
provider_query_execute_standard_test_candidate(
state, &provider, candidate, payload, route_path, &trace_id,
)
.await?
};
if execution.status != "skipped" { if execution.status != "skipped" {
total_attempts += 1; total_attempts += 1;
} }
@@ -814,7 +1283,7 @@ async fn build_admin_provider_query_kiro_failover_response(
"total_candidates": candidates.len(), "total_candidates": candidates.len(),
"total_attempts": total_attempts, "total_attempts": total_attempts,
"data": success_body.as_ref().map(|body| json!({ "data": success_body.as_ref().map(|body| json!({
"stream": true, "stream": is_kiro,
"response": body, "response": body,
})), })),
"error": error, "error": error,
@@ -832,6 +1301,9 @@ pub(crate) async fn build_admin_provider_query_test_model_local_response(
"/api/admin/provider-query/test-model", "/api/admin/provider-query/test-model",
) )
.await?; .await?;
if !response.status().is_success() {
return Ok(response);
}
let body = to_bytes(response.into_body(), usize::MAX) let body = to_bytes(response.into_body(), usize::MAX)
.await .await
.map_err(|err| GatewayError::Internal(err.to_string()))?; .map_err(|err| GatewayError::Internal(err.to_string()))?;

View File

@@ -2,9 +2,7 @@ use crate::handlers::admin::provider::query::{
models::{ models::{
build_admin_provider_query_models_response, build_admin_provider_query_models_response,
build_admin_provider_query_test_model_failover_local_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_local_response,
build_admin_provider_query_test_model_response,
}, },
payload::{ payload::{
parse_admin_provider_query_body, provider_query_extract_failover_models, 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?, build_admin_provider_query_models_response(self, &payload).await?,
)), )),
"test_model" => { "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( log_admin_provider_query_validation_failure(
request_context, request_context,
route_kind, route_kind,
@@ -69,7 +67,7 @@ impl<'a> AdminAppState<'a> {
ADMIN_PROVIDER_QUERY_PROVIDER_ID_REQUIRED_DETAIL, 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( log_admin_provider_query_validation_failure(
request_context, request_context,
route_kind, route_kind,
@@ -80,28 +78,12 @@ impl<'a> AdminAppState<'a> {
ADMIN_PROVIDER_QUERY_MODEL_REQUIRED_DETAIL, ADMIN_PROVIDER_QUERY_MODEL_REQUIRED_DETAIL,
))); )));
}; };
let provider_type = self Ok(Some(
.app() build_admin_provider_query_test_model_local_response(self, &payload).await?,
.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,
)))
}
} }
"test_model_failover" => { "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( log_admin_provider_query_validation_failure(
request_context, request_context,
route_kind, route_kind,
@@ -124,29 +106,10 @@ impl<'a> AdminAppState<'a> {
ADMIN_PROVIDER_QUERY_FAILOVER_MODELS_REQUIRED_DETAIL, ADMIN_PROVIDER_QUERY_FAILOVER_MODELS_REQUIRED_DETAIL,
))); )));
} }
let provider_type = self Ok(Some(
.app() build_admin_provider_query_test_model_failover_local_response(self, &payload)
.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,
)
.await?, .await?,
)) ))
} else {
Ok(Some(
build_admin_provider_query_test_model_failover_response(
provider_id,
failover_models,
),
))
}
} }
_ => Ok(Some( _ => Ok(Some(
build_admin_provider_query_models_response(self, &payload).await?, build_admin_provider_query_models_response(self, &payload).await?,

File diff suppressed because it is too large Load Diff