feat: add model directive management

This commit is contained in:
fawney19
2026-05-03 14:48:25 +08:00
parent fe27fb17fb
commit c4ea042eb4
53 changed files with 2655 additions and 182 deletions

View File

@@ -3,6 +3,7 @@ use aether_ai_serving::{
};
use aether_scheduler_core::SchedulerMinimalCandidateSelectionCandidate;
use async_trait::async_trait;
use std::collections::BTreeSet;
use crate::ai_serving::planner::candidate_resolution::SkippedLocalExecutionCandidate;
use crate::ai_serving::{GatewayAuthApiKeySnapshot, PlannerAppState};
@@ -25,6 +26,8 @@ struct GatewayLocalCandidatePreselectionPort<'a> {
auth_snapshot: &'a GatewayAuthApiKeySnapshot,
use_api_format_alias_match: bool,
key_mode: LocalCandidatePreselectionKeyMode,
candidate_api_formats: Vec<String>,
model_directive_enabled_api_formats: BTreeSet<String>,
}
#[async_trait]
@@ -34,13 +37,7 @@ impl AiCandidatePreselectionPort for GatewayLocalCandidatePreselectionPort<'_> {
type Error = GatewayError;
fn candidate_api_formats(&self) -> Vec<String> {
crate::ai_serving::request_candidate_api_formats(
self.client_api_format,
self.require_streaming,
)
.into_iter()
.map(str::to_string)
.collect()
self.candidate_api_formats.clone()
}
fn candidate_api_format_matches_client(&self, candidate_api_format: &str) -> bool {
@@ -84,28 +81,36 @@ impl AiCandidatePreselectionPort for GatewayLocalCandidatePreselectionPort<'_> {
fn candidate_allowed(
&self,
candidate: &Self::Candidate,
_candidate_api_format: &str,
candidate_api_format: &str,
matches_client_format: bool,
) -> bool {
let enable_model_directives = self.model_directive_enabled_api_formats.contains(
&crate::ai_serving::normalize_api_format_alias(candidate_api_format),
);
matches_client_format
|| auth_snapshot_allows_cross_format_candidate(
self.auth_snapshot,
self.requested_model,
candidate,
enable_model_directives,
)
}
fn skipped_candidate_allowed(
&self,
skipped_candidate: &Self::Skipped,
_candidate_api_format: &str,
candidate_api_format: &str,
matches_client_format: bool,
) -> bool {
let enable_model_directives = self.model_directive_enabled_api_formats.contains(
&crate::ai_serving::normalize_api_format_alias(candidate_api_format),
);
matches_client_format
|| auth_snapshot_allows_cross_format_candidate(
self.auth_snapshot,
self.requested_model,
&skipped_candidate.candidate,
enable_model_directives,
)
}
@@ -135,6 +140,24 @@ pub(crate) async fn preselect_local_execution_candidates_with_serving(
>,
GatewayError,
> {
let candidate_api_formats =
crate::ai_serving::request_candidate_api_formats(client_api_format, require_streaming)
.into_iter()
.map(str::to_string)
.collect::<Vec<_>>();
let mut model_directive_enabled_api_formats = BTreeSet::new();
for api_format in &candidate_api_formats {
if crate::system_features::reasoning_model_directive_enabled_for_api_format_and_model(
state.app(),
api_format,
Some(requested_model),
)
.await
{
model_directive_enabled_api_formats
.insert(crate::ai_serving::normalize_api_format_alias(api_format));
}
}
let port = GatewayLocalCandidatePreselectionPort {
state,
client_api_format,
@@ -144,6 +167,8 @@ pub(crate) async fn preselect_local_execution_candidates_with_serving(
auth_snapshot,
use_api_format_alias_match,
key_mode,
candidate_api_formats,
model_directive_enabled_api_formats,
};
run_ai_candidate_preselection(&port).await
@@ -190,6 +215,7 @@ pub(crate) fn auth_snapshot_allows_cross_format_candidate(
auth_snapshot: &GatewayAuthApiKeySnapshot,
requested_model: &str,
candidate: &SchedulerMinimalCandidateSelectionCandidate,
enable_model_directives: bool,
) -> bool {
if let Some(allowed_providers) = auth_snapshot.effective_allowed_providers() {
let provider_allowed = allowed_providers.iter().any(|value| {
@@ -206,9 +232,16 @@ pub(crate) fn auth_snapshot_allows_cross_format_candidate(
}
if let Some(allowed_models) = auth_snapshot.effective_allowed_models() {
let model_allowed = allowed_models
.iter()
.any(|value| value == requested_model || value == &candidate.global_model_name);
let requested_base_model = enable_model_directives
.then(|| crate::ai_serving::model_directive_base_model(requested_model))
.flatten();
let model_allowed = allowed_models.iter().any(|value| {
value == requested_model
|| value == &candidate.global_model_name
|| requested_base_model
.as_ref()
.is_some_and(|base_model| value == base_model)
});
if !model_allowed {
return false;
}

View File

@@ -63,8 +63,15 @@ pub(crate) async fn resolve_local_same_format_provider_candidate_payload_parts(
spec,
)
.await?;
let enable_model_directives =
crate::system_features::reasoning_model_directive_enabled_for_api_format_and_model(
state,
spec.api_format,
Some(&input.requested_model),
)
.await;
let Some(base_provider_request_body) =
let Some(mut base_provider_request_body) =
super::super::request::build_same_format_provider_request_body(
body_json,
&prepared.mapped_model,
@@ -73,6 +80,7 @@ pub(crate) async fn resolve_local_same_format_provider_candidate_payload_parts(
prepared.upstream_is_stream,
prepared.kiro_auth.as_ref(),
prepared.is_claude_code,
enable_model_directives,
)
else {
mark_skipped_local_same_format_provider_candidate_with_extra_data(
@@ -97,6 +105,19 @@ pub(crate) async fn resolve_local_same_format_provider_candidate_payload_parts(
.await;
return None;
};
if let Some(mapping) =
crate::system_features::reasoning_model_directive_mapping_for_api_format_and_model(
state,
spec.api_format,
Some(&input.requested_model),
)
.await
{
crate::ai_serving::apply_model_directive_mapping_patch(
&mut base_provider_request_body,
&mapping,
);
}
let antigravity_auth = if prepared.is_antigravity {
match classify_local_antigravity_request_support(

View File

@@ -14,15 +14,19 @@ pub(crate) fn build_same_format_provider_request_body(
upstream_is_stream: bool,
kiro_auth: Option<&crate::ai_serving::transport::kiro::KiroRequestAuth>,
is_claude_code: bool,
enable_model_directives: bool,
) -> Option<Value> {
build_same_format_provider_request_body_impl(SameFormatProviderRequestBodyInput {
body_json,
mapped_model,
provider_api_format: spec.api_format,
source_model: body_json.get("model").and_then(Value::as_str),
family: same_format_provider_family(spec.family),
body_rules,
upstream_is_stream,
kiro_auth_config: kiro_auth.map(|auth| &auth.auth_config),
is_claude_code,
enable_model_directives,
})
}

View File

@@ -168,8 +168,15 @@ pub(crate) async fn resolve_local_standard_candidate_payload_parts(
transport.provider.provider_type.as_str(),
provider_api_format,
);
let provider_request_body =
match crate::ai_serving::planner::standard::build_standard_request_body(
let enable_model_directives =
crate::system_features::reasoning_model_directive_enabled_for_api_format_and_model(
state,
provider_api_format,
Some(&input.requested_model),
)
.await;
let mut provider_request_body =
match crate::ai_serving::planner::standard::build_standard_request_body_with_model_directives(
body_json,
spec_metadata.api_format,
&prepared_candidate.mapped_model,
@@ -183,6 +190,7 @@ pub(crate) async fn resolve_local_standard_candidate_payload_parts(
transport.endpoint.body_rules.as_ref()
},
Some(input.auth_context.api_key_id.as_str()),
enable_model_directives,
) {
Some(body) => body,
None => {
@@ -204,6 +212,19 @@ pub(crate) async fn resolve_local_standard_candidate_payload_parts(
return None;
}
};
if let Some(mapping) =
crate::system_features::reasoning_model_directive_mapping_for_api_format_and_model(
state,
provider_api_format,
Some(&input.requested_model),
)
.await
{
crate::ai_serving::apply_model_directive_mapping_patch(
&mut provider_request_body,
&mapping,
);
}
if let Some(kiro_auth) = kiro_auth.as_ref() {
return build_kiro_cross_format_payload_parts(

View File

@@ -45,8 +45,8 @@ pub(crate) use crate::ai_serving::{
SyncCliResponseConversionKind,
};
pub(crate) use crate::ai_serving::{
build_standard_request_body, convert_openai_chat_request_to_claude_request,
convert_openai_chat_request_to_gemini_request,
build_standard_request_body, build_standard_request_body_with_model_directives,
convert_openai_chat_request_to_claude_request, convert_openai_chat_request_to_gemini_request,
convert_openai_chat_request_to_openai_responses_request, extract_openai_text_content,
normalize_openai_responses_request_to_openai_chat_request, parse_openai_tool_result_content,
};

View File

@@ -4,8 +4,8 @@ use crate::ai_serving::transport::apply_standard_provider_request_body_rules;
use crate::ai_serving::{
apply_codex_openai_responses_special_body_edits,
apply_openai_responses_compact_special_body_edits,
build_cross_format_openai_chat_request_body as surface_build_cross_format_openai_chat_request_body,
build_local_openai_chat_request_body as surface_build_local_openai_chat_request_body,
build_cross_format_openai_chat_request_body_with_model_directives as surface_build_cross_format_openai_chat_request_body,
build_local_openai_chat_request_body_with_model_directives as surface_build_local_openai_chat_request_body,
GatewayProviderTransportSnapshot,
};
@@ -14,9 +14,14 @@ pub(crate) fn build_local_openai_chat_request_body(
mapped_model: &str,
upstream_is_stream: bool,
body_rules: Option<&Value>,
enable_model_directives: bool,
) -> Option<Value> {
let provider_request_body =
surface_build_local_openai_chat_request_body(body_json, mapped_model, upstream_is_stream)?;
let provider_request_body = surface_build_local_openai_chat_request_body(
body_json,
mapped_model,
upstream_is_stream,
enable_model_directives,
)?;
apply_standard_provider_request_body_rules(provider_request_body, body_rules, body_json)
}
@@ -35,12 +40,14 @@ pub(crate) fn build_cross_format_openai_chat_request_body(
upstream_is_stream: bool,
body_rules: Option<&Value>,
user_api_key_id: Option<&str>,
enable_model_directives: bool,
) -> Option<Value> {
let provider_request_body = surface_build_cross_format_openai_chat_request_body(
body_json,
mapped_model,
provider_api_format,
upstream_is_stream,
enable_model_directives,
)?;
let mut provider_request_body =
apply_standard_provider_request_body_rules(provider_request_body, body_rules, body_json)?;

View File

@@ -4,8 +4,8 @@ use crate::ai_serving::transport::apply_standard_provider_request_body_rules;
use crate::ai_serving::{
apply_codex_openai_responses_special_body_edits,
apply_openai_responses_compact_special_body_edits,
build_cross_format_openai_responses_request_body as surface_build_cross_format_openai_responses_request_body,
build_local_openai_responses_request_body as surface_build_local_openai_responses_request_body,
build_cross_format_openai_responses_request_body_with_model_directives as surface_build_cross_format_openai_responses_request_body,
build_local_openai_responses_request_body_with_model_directives as surface_build_local_openai_responses_request_body,
GatewayProviderTransportSnapshot,
};
@@ -17,11 +17,13 @@ pub(crate) fn build_local_openai_responses_request_body(
provider_api_format: &str,
body_rules: Option<&Value>,
user_api_key_id: Option<&str>,
enable_model_directives: bool,
) -> Option<Value> {
let provider_request_body = surface_build_local_openai_responses_request_body(
body_json,
mapped_model,
require_streaming,
enable_model_directives,
)?;
let mut provider_request_body =
apply_standard_provider_request_body_rules(provider_request_body, body_rules, body_json)?;
@@ -48,6 +50,7 @@ pub(crate) fn build_cross_format_openai_responses_request_body(
provider_type: &str,
body_rules: Option<&Value>,
user_api_key_id: Option<&str>,
enable_model_directives: bool,
) -> Option<Value> {
let provider_request_body = surface_build_cross_format_openai_responses_request_body(
body_json,
@@ -55,6 +58,7 @@ pub(crate) fn build_cross_format_openai_responses_request_body(
client_api_format,
provider_api_format,
upstream_is_stream,
enable_model_directives,
)?;
let mut provider_request_body =
apply_standard_provider_request_body_rules(provider_request_body, body_rules, body_json)?;

View File

@@ -90,6 +90,7 @@ fn builds_openai_chat_cross_format_request_body_from_openai_responses_source() {
"openai",
None,
None,
false,
)
.expect("openai responses to openai chat body should build");
@@ -123,6 +124,7 @@ fn local_openai_responses_wrapper_preserves_body_order_after_edits() {
"openai:responses",
None,
Some("key-123"),
false,
)
.expect("local openai responses body should build");
@@ -160,12 +162,41 @@ fn local_openai_responses_compact_wrapper_strips_store_for_same_format_requests(
"openai:responses:compact",
None,
None,
false,
)
.expect("local openai compact body should build");
assert!(provider_request_body.get("store").is_none());
}
#[test]
fn local_openai_responses_wrapper_applies_model_directive_before_body_rules() {
let body_json = json!({
"model": "gpt-5.4-max",
"input": "hello",
"reasoning": {"effort": "low", "summary": "auto"}
});
let body_rules = json!([
{"action":"set","path":"metadata.override_seen","value":true}
]);
let provider_request_body = build_local_openai_responses_request_body(
&body_json,
"gpt-5.4",
false,
"openai",
"openai:responses",
Some(&body_rules),
None,
true,
)
.expect("local openai responses body should build");
assert_eq!(provider_request_body["reasoning"]["effort"], "xhigh");
assert_eq!(provider_request_body["reasoning"]["summary"], "auto");
assert_eq!(provider_request_body["metadata"]["override_seen"], true);
}
#[test]
fn local_openai_responses_upstream_url_preserves_codex_base_path() {
let request = Request::builder()
@@ -205,6 +236,7 @@ fn strips_metadata_for_codex_openai_responses_requests() {
"codex",
None,
None,
false,
)
.expect("claude cli to codex request should build");
@@ -237,6 +269,7 @@ fn applies_codex_defaults_unless_body_rules_handle_the_field() {
"codex",
Some(&body_rules),
None,
false,
)
.expect("claude cli to codex request should build");
@@ -264,6 +297,7 @@ fn injects_codex_prompt_cache_key_for_openai_responses_cross_format_requests() {
"codex",
None,
Some("key-123"),
false,
)
.expect("claude cli to codex request should build");
@@ -291,6 +325,7 @@ fn injects_codex_prompt_cache_key_for_openai_chat_cross_format_requests() {
false,
None,
Some("key-123"),
false,
)
.expect("openai chat to codex request should build");

View File

@@ -72,6 +72,13 @@ pub(crate) async fn resolve_local_openai_chat_candidate_payload_parts(
let candidate = &eligible.candidate;
let provider_api_format = eligible.provider_api_format.as_str();
let transport = &eligible.transport;
let enable_model_directives =
crate::system_features::reasoning_model_directive_enabled_for_api_format_and_model(
state,
provider_api_format,
Some(&input.requested_model),
)
.await;
if provider_api_format == "openai:chat" {
if let Some(skip_reason) = local_openai_chat_transport_unsupported_reason(transport) {
@@ -122,6 +129,7 @@ pub(crate) async fn resolve_local_openai_chat_candidate_payload_parts(
&prepared_candidate.mapped_model,
upstream_is_stream,
transport.endpoint.body_rules.as_ref(),
enable_model_directives,
) else {
mark_skipped_local_openai_chat_candidate_with_extra_data(
state,
@@ -334,7 +342,7 @@ pub(crate) async fn resolve_local_openai_chat_candidate_payload_parts(
}
};
let Some(provider_request_body) = build_cross_format_openai_chat_request_body(
let Some(mut provider_request_body) = build_cross_format_openai_chat_request_body(
body_json,
&prepared_candidate.mapped_model,
transport.provider.provider_type.as_str(),
@@ -346,6 +354,7 @@ pub(crate) async fn resolve_local_openai_chat_candidate_payload_parts(
transport.endpoint.body_rules.as_ref()
},
Some(input.auth_context.api_key_id.as_str()),
enable_model_directives,
) else {
mark_skipped_local_openai_chat_candidate_with_extra_data(
state,
@@ -364,6 +373,19 @@ pub(crate) async fn resolve_local_openai_chat_candidate_payload_parts(
.await;
return None;
};
if let Some(mapping) =
crate::system_features::reasoning_model_directive_mapping_for_api_format_and_model(
state,
provider_api_format.as_str(),
Some(&input.requested_model),
)
.await
{
crate::ai_serving::apply_model_directive_mapping_patch(
&mut provider_request_body,
&mapping,
);
}
if let Some(kiro_auth) = kiro_auth.as_ref() {
return build_kiro_openai_chat_cross_format_payload_parts(

View File

@@ -215,6 +215,13 @@ pub(crate) async fn resolve_local_openai_responses_candidate_payload_parts(
let auth_header = prepared_candidate.auth_header;
let auth_value = prepared_candidate.auth_value;
let mapped_model = prepared_candidate.mapped_model;
let enable_model_directives =
crate::system_features::reasoning_model_directive_enabled_for_api_format_and_model(
state,
provider_api_format,
Some(&input.requested_model),
)
.await;
let needs_bidirectional_conversion = !same_format && conversion_kind.is_some();
let upstream_is_stream = spec_metadata.require_streaming
@@ -223,7 +230,7 @@ pub(crate) async fn resolve_local_openai_responses_candidate_payload_parts(
transport.provider.provider_type.as_str(),
provider_api_format,
);
let Some(base_provider_request_body) = (if needs_bidirectional_conversion {
let Some(mut base_provider_request_body) = (if needs_bidirectional_conversion {
build_cross_format_openai_responses_request_body(
body_json,
&mapped_model,
@@ -237,6 +244,7 @@ pub(crate) async fn resolve_local_openai_responses_candidate_payload_parts(
transport.endpoint.body_rules.as_ref()
},
Some(input.auth_context.api_key_id.as_str()),
enable_model_directives,
)
} else {
build_local_openai_responses_request_body(
@@ -251,6 +259,7 @@ pub(crate) async fn resolve_local_openai_responses_candidate_payload_parts(
transport.endpoint.body_rules.as_ref()
},
Some(input.auth_context.api_key_id.as_str()),
enable_model_directives,
)
}) else {
mark_skipped_local_openai_responses_candidate_with_extra_data(
@@ -270,6 +279,19 @@ pub(crate) async fn resolve_local_openai_responses_candidate_payload_parts(
.await;
return None;
};
if let Some(mapping) =
crate::system_features::reasoning_model_directive_mapping_for_api_format_and_model(
state,
provider_api_format,
Some(&input.requested_model),
)
.await
{
crate::ai_serving::apply_model_directive_mapping_patch(
&mut base_provider_request_body,
&mapping,
);
}
let antigravity_auth = if is_antigravity {
match classify_local_antigravity_request_support(
transport,

View File

@@ -11,12 +11,15 @@ impl<'a> PlannerAppState<'a> {
requested_model: Option<&str>,
explicit_required_capabilities: Option<&Value>,
) -> Option<Value> {
let enable_model_directives =
crate::system_features::reasoning_model_directive_enabled(self.app()).await;
crate::request_candidate_runtime::resolve_request_candidate_required_capabilities(
self.app(),
user_id,
api_key_id,
requested_model,
explicit_required_capabilities,
enable_model_directives,
)
.await
}

View File

@@ -20,6 +20,13 @@ impl<'a> PlannerAppState<'a> {
auth_snapshot: Option<&GatewayAuthApiKeySnapshot>,
now_unix_secs: u64,
) -> Result<Vec<SchedulerMinimalCandidateSelectionCandidate>, GatewayError> {
let enable_model_directives =
crate::system_features::reasoning_model_directive_enabled_for_api_format_and_model(
self.app(),
api_format,
Some(global_model_name),
)
.await;
crate::scheduler::candidate::list_selectable_candidates(
self.app().data.as_ref(),
self.app(),
@@ -29,6 +36,7 @@ impl<'a> PlannerAppState<'a> {
required_capabilities,
auth_snapshot,
now_unix_secs,
enable_model_directives,
)
.await
}
@@ -52,6 +60,13 @@ impl<'a> PlannerAppState<'a> {
let wait_interval = Duration::from_millis(API_KEY_CONCURRENCY_WAIT_POLL_INTERVAL_MS.max(1));
let wait_deadline = Instant::now() + wait_timeout;
let mut attempt_now_unix_secs = now_unix_secs;
let enable_model_directives =
crate::system_features::reasoning_model_directive_enabled_for_api_format_and_model(
self.app(),
api_format,
Some(global_model_name),
)
.await;
loop {
let result = crate::scheduler::candidate::list_selectable_candidates_with_skip_reasons(
@@ -63,6 +78,7 @@ impl<'a> PlannerAppState<'a> {
required_capabilities,
auth_snapshot,
attempt_now_unix_secs,
enable_model_directives,
)
.await?;