mirror of
https://github.com/fabro-sh/fabro.git
synced 2026-09-24 00:51:19 +00:00
feat(llm): validate model request controls
Propagate run-level model controls into workflow LLM requests, type speed at the request boundary, and reject unsupported speed or reasoning controls before provider dispatch.
This commit is contained in:
parent
cfb4ea91a3
commit
b969b02026
11 changed files with 422 additions and 42 deletions
|
|
@ -2,7 +2,7 @@ use std::collections::HashMap;
|
|||
use std::sync::Arc;
|
||||
use std::time::Duration;
|
||||
|
||||
use fabro_llm::types::ReasoningEffort;
|
||||
use fabro_llm::types::{ReasoningEffort, Speed};
|
||||
use fabro_mcp::config::McpServerSettings;
|
||||
|
||||
/// Callback invoked before each tool execution. Return `Ok(())` to allow,
|
||||
|
|
@ -65,7 +65,7 @@ pub struct SessionOptions {
|
|||
pub default_command_timeout_ms: u64,
|
||||
pub max_command_timeout_ms: u64,
|
||||
pub reasoning_effort: Option<ReasoningEffort>,
|
||||
pub speed: Option<String>,
|
||||
pub speed: Option<Speed>,
|
||||
pub tool_output_limits: HashMap<String, usize>,
|
||||
pub tool_line_limits: HashMap<String, usize>,
|
||||
/// Override the provider's default max_tokens when set.
|
||||
|
|
|
|||
|
|
@ -966,7 +966,7 @@ impl Session {
|
|||
self.config.reasoning_effort = effort;
|
||||
}
|
||||
|
||||
pub fn set_speed(&mut self, speed: Option<String>) {
|
||||
pub fn set_speed(&mut self, speed: Option<Speed>) {
|
||||
self.config.speed = speed;
|
||||
}
|
||||
|
||||
|
|
@ -1344,11 +1344,6 @@ impl Session {
|
|||
});
|
||||
|
||||
// Emit AssistantMessage with enriched data from the response
|
||||
let speed = self
|
||||
.config
|
||||
.speed
|
||||
.as_deref()
|
||||
.and_then(|value| value.parse::<Speed>().ok());
|
||||
let model = ModelRef {
|
||||
provider: self.provider_profile.provider_id(),
|
||||
model_id: if response.model.is_empty() {
|
||||
|
|
@ -1356,7 +1351,7 @@ impl Session {
|
|||
} else {
|
||||
response.model.clone()
|
||||
},
|
||||
speed,
|
||||
speed: self.config.speed,
|
||||
};
|
||||
self.event_emitter
|
||||
.emit(self.id.clone(), AgentEvent::AssistantMessage {
|
||||
|
|
@ -1565,7 +1560,7 @@ impl Session {
|
|||
}),
|
||||
stop_sequences: None,
|
||||
reasoning_effort: self.config.reasoning_effort,
|
||||
speed: self.config.speed.clone(),
|
||||
speed: self.config.speed,
|
||||
metadata: None,
|
||||
provider_options: None,
|
||||
}
|
||||
|
|
|
|||
|
|
@ -10,7 +10,7 @@ use crate::adapter_registry::{AdapterConfig, factory_for};
|
|||
use crate::error::Error;
|
||||
use crate::middleware::{Middleware, NextFn, NextStreamFn};
|
||||
use crate::provider::{ProviderAdapter, StreamEventStream};
|
||||
use crate::types::{Request, Response};
|
||||
use crate::types::{Request, Response, Speed};
|
||||
|
||||
/// The core client that routes requests to provider adapters (Section 2.2, 3).
|
||||
#[derive(Clone)]
|
||||
|
|
@ -199,6 +199,44 @@ impl Client {
|
|||
})
|
||||
}
|
||||
|
||||
fn validate_request_controls(&self, request: &Request) -> Result<(), Error> {
|
||||
let Some(catalog) = &self.catalog else {
|
||||
return Ok(());
|
||||
};
|
||||
let Some(settings) = catalog.model_settings(&request.model) else {
|
||||
return Ok(());
|
||||
};
|
||||
let model_id = catalog
|
||||
.get(&request.model)
|
||||
.map_or(request.model.as_str(), |model| model.id.as_str());
|
||||
|
||||
if let Some(effort) = request.reasoning_effort {
|
||||
if !settings.controls.reasoning_effort.contains(&effort) {
|
||||
return Err(Error::Configuration {
|
||||
message: format!(
|
||||
"model '{model_id}' does not support reasoning_effort '{effort}'; allowed values: {}",
|
||||
format_control_values(&settings.controls.reasoning_effort),
|
||||
),
|
||||
source: None,
|
||||
});
|
||||
}
|
||||
}
|
||||
|
||||
if let Some(speed) = request.speed {
|
||||
if speed != Speed::Standard && !settings.controls.speed.contains(&speed) {
|
||||
return Err(Error::Configuration {
|
||||
message: format!(
|
||||
"model '{model_id}' does not support speed '{speed}'; allowed values: standard{}",
|
||||
format_additional_speeds(&settings.controls.speed),
|
||||
),
|
||||
source: None,
|
||||
});
|
||||
}
|
||||
}
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// Send a blocking request (Section 4.1).
|
||||
///
|
||||
/// # Errors
|
||||
|
|
@ -207,6 +245,7 @@ impl Client {
|
|||
/// registered, or any provider/middleware error encountered during the
|
||||
/// request.
|
||||
pub async fn complete(&self, request: &Request) -> Result<Response, Error> {
|
||||
self.validate_request_controls(request)?;
|
||||
let provider = self.resolve_provider(request)?;
|
||||
|
||||
if self.middleware.is_empty() {
|
||||
|
|
@ -240,6 +279,7 @@ impl Client {
|
|||
/// registered, or any provider/middleware error encountered during the
|
||||
/// request.
|
||||
pub async fn stream(&self, request: &Request) -> Result<StreamEventStream, Error> {
|
||||
self.validate_request_controls(request)?;
|
||||
let provider = self.resolve_provider(request)?;
|
||||
|
||||
if self.middleware.is_empty() {
|
||||
|
|
@ -304,6 +344,26 @@ impl Client {
|
|||
}
|
||||
}
|
||||
|
||||
fn format_control_values<T: ToString>(values: &[T]) -> String {
|
||||
if values.is_empty() {
|
||||
"none".to_string()
|
||||
} else {
|
||||
values
|
||||
.iter()
|
||||
.map(ToString::to_string)
|
||||
.collect::<Vec<_>>()
|
||||
.join(", ")
|
||||
}
|
||||
}
|
||||
|
||||
fn format_additional_speeds(values: &[Speed]) -> String {
|
||||
if values.is_empty() {
|
||||
String::new()
|
||||
} else {
|
||||
format!(", {}", format_control_values(values))
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) fn auth_value(auth_header: &ApiKeyHeader) -> String {
|
||||
match auth_header {
|
||||
ApiKeyHeader::Bearer(value) | ApiKeyHeader::Custom { value, .. } => value.clone(),
|
||||
|
|
@ -483,6 +543,132 @@ mod tests {
|
|||
assert!(matches!(result.unwrap_err(), Error::Configuration { .. }));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn complete_rejects_unsupported_reasoning_effort_before_dispatch() {
|
||||
let catalog =
|
||||
Arc::new(Catalog::from_builtin_with_overrides(&LlmCatalogSettings::default()).unwrap());
|
||||
let mut client = Client::new(HashMap::new(), None, vec![]);
|
||||
client.catalog = Some(Arc::clone(&catalog));
|
||||
client
|
||||
.register_provider(Arc::new(MockProvider::new("kimi", "should not dispatch")))
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
let mut request = test_request();
|
||||
request.model = "kimi-k2.5".to_string();
|
||||
request.provider = Some("kimi".to_string());
|
||||
request.reasoning_effort = Some(ReasoningEffort::High);
|
||||
|
||||
let err = client.complete(&request).await.unwrap_err();
|
||||
|
||||
assert!(matches!(
|
||||
err,
|
||||
Error::Configuration {
|
||||
ref message,
|
||||
..
|
||||
} if message.contains("model 'kimi-k2.5' does not support reasoning_effort 'high'")
|
||||
));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn complete_rejects_unsupported_speed_before_dispatch() {
|
||||
let catalog =
|
||||
Arc::new(Catalog::from_builtin_with_overrides(&LlmCatalogSettings::default()).unwrap());
|
||||
let mut client = Client::new(HashMap::new(), None, vec![]);
|
||||
client.catalog = Some(Arc::clone(&catalog));
|
||||
client
|
||||
.register_provider(Arc::new(MockProvider::new("openai", "should not dispatch")))
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
let mut request = test_request();
|
||||
request.model = "gpt-5.4".to_string();
|
||||
request.provider = Some("openai".to_string());
|
||||
request.speed = Some(Speed::Fast);
|
||||
|
||||
let err = client.complete(&request).await.unwrap_err();
|
||||
|
||||
assert!(matches!(
|
||||
err,
|
||||
Error::Configuration {
|
||||
ref message,
|
||||
..
|
||||
} if message.contains("model 'gpt-5.4' does not support speed 'fast'")
|
||||
));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn complete_accepts_standard_speed_without_catalog_declaration() {
|
||||
let catalog =
|
||||
Arc::new(Catalog::from_builtin_with_overrides(&LlmCatalogSettings::default()).unwrap());
|
||||
let mut client = Client::new(HashMap::new(), None, vec![]);
|
||||
client.catalog = Some(Arc::clone(&catalog));
|
||||
client
|
||||
.register_provider(Arc::new(MockProvider::new("openai", "standard")))
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
let mut request = test_request();
|
||||
request.model = "gpt-5.4".to_string();
|
||||
request.provider = Some("openai".to_string());
|
||||
request.speed = Some(Speed::Standard);
|
||||
|
||||
let response = client.complete(&request).await.unwrap();
|
||||
|
||||
assert_eq!(response.text(), "standard");
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn complete_skips_control_validation_for_unknown_model_passthrough() {
|
||||
let catalog =
|
||||
Arc::new(Catalog::from_builtin_with_overrides(&LlmCatalogSettings::default()).unwrap());
|
||||
let mut client = Client::new(HashMap::new(), None, vec![]);
|
||||
client.catalog = Some(Arc::clone(&catalog));
|
||||
client
|
||||
.register_provider(Arc::new(MockProvider::new("openai", "passthrough")))
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
let mut request = test_request();
|
||||
request.model = "custom-model".to_string();
|
||||
request.provider = Some("openai".to_string());
|
||||
request.reasoning_effort = Some(ReasoningEffort::High);
|
||||
request.speed = Some(Speed::Fast);
|
||||
|
||||
let response = client.complete(&request).await.unwrap();
|
||||
|
||||
assert_eq!(response.text(), "passthrough");
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn stream_rejects_unsupported_speed_before_dispatch() {
|
||||
let catalog =
|
||||
Arc::new(Catalog::from_builtin_with_overrides(&LlmCatalogSettings::default()).unwrap());
|
||||
let mut client = Client::new(HashMap::new(), None, vec![]);
|
||||
client.catalog = Some(Arc::clone(&catalog));
|
||||
client
|
||||
.register_provider(Arc::new(MockProvider::new("openai", "should not dispatch")))
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
let mut request = test_request();
|
||||
request.model = "gpt-5.4".to_string();
|
||||
request.provider = Some("openai".to_string());
|
||||
request.speed = Some(Speed::Fast);
|
||||
|
||||
let Err(err) = client.stream(&request).await else {
|
||||
panic!("unsupported speed should fail before stream dispatch");
|
||||
};
|
||||
|
||||
assert!(matches!(
|
||||
err,
|
||||
Error::Configuration {
|
||||
ref message,
|
||||
..
|
||||
} if message.contains("model 'gpt-5.4' does not support speed 'fast'")
|
||||
));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn from_credentials_registers_multiple_providers() {
|
||||
let client = Client::from_credentials(vec![
|
||||
|
|
|
|||
|
|
@ -17,8 +17,8 @@ use crate::retry::retry;
|
|||
use crate::tools::{RepairToolCallFn, Tool, execute_all_tools_with_repair};
|
||||
use crate::types::{
|
||||
FinishReason, GenerateResult, Message, ObjectStreamEvent, ReasoningEffort, Request, Response,
|
||||
ResponseFormat, ResponseFormatType, RetryPolicy, StepResult, StreamEvent, TimeoutOptions,
|
||||
TokenCounts, ToolCall, ToolChoice, ToolDefinition,
|
||||
ResponseFormat, ResponseFormatType, RetryPolicy, Speed, StepResult, StreamEvent,
|
||||
TimeoutOptions, TokenCounts, ToolCall, ToolChoice, ToolDefinition,
|
||||
};
|
||||
|
||||
fn build_initial_messages(params: &GenerateParams) -> Result<Vec<Message>, Error> {
|
||||
|
|
@ -57,7 +57,7 @@ fn build_request(
|
|||
max_tokens: params.max_tokens,
|
||||
stop_sequences: params.stop_sequences.clone(),
|
||||
reasoning_effort: params.reasoning_effort,
|
||||
speed: params.speed.clone(),
|
||||
speed: params.speed,
|
||||
metadata: params.metadata.clone(),
|
||||
provider_options: params.provider_options.clone(),
|
||||
}
|
||||
|
|
@ -280,7 +280,7 @@ pub struct GenerateParams {
|
|||
pub max_tokens: Option<i64>,
|
||||
pub stop_sequences: Option<Vec<String>>,
|
||||
pub reasoning_effort: Option<ReasoningEffort>,
|
||||
pub speed: Option<String>,
|
||||
pub speed: Option<Speed>,
|
||||
pub provider: Option<String>,
|
||||
pub provider_options: Option<serde_json::Value>,
|
||||
pub metadata: Option<std::collections::HashMap<String, String>>,
|
||||
|
|
@ -402,6 +402,12 @@ impl GenerateParams {
|
|||
self
|
||||
}
|
||||
|
||||
#[must_use]
|
||||
pub const fn speed(mut self, speed: Speed) -> Self {
|
||||
self.speed = Some(speed);
|
||||
self
|
||||
}
|
||||
|
||||
#[must_use]
|
||||
pub fn provider_options(mut self, provider_options: serde_json::Value) -> Self {
|
||||
self.provider_options = Some(provider_options);
|
||||
|
|
@ -1480,6 +1486,7 @@ mod tests {
|
|||
.max_tokens(100)
|
||||
.stop_sequences(vec!["STOP".to_string()])
|
||||
.reasoning_effort(ReasoningEffort::High)
|
||||
.speed(Speed::Fast)
|
||||
.provider("anthropic")
|
||||
.provider_options(serde_json::json!({"key": "value"}))
|
||||
.max_retries(5)
|
||||
|
|
@ -1499,6 +1506,7 @@ mod tests {
|
|||
assert_eq!(params.max_tokens, Some(100));
|
||||
assert_eq!(params.stop_sequences, Some(vec!["STOP".to_string()]));
|
||||
assert_eq!(params.reasoning_effort, Some(ReasoningEffort::High));
|
||||
assert_eq!(params.speed, Some(Speed::Fast));
|
||||
assert_eq!(params.provider.as_deref(), Some("anthropic"));
|
||||
assert!(params.provider_options.is_some());
|
||||
assert_eq!(params.max_retries, 5);
|
||||
|
|
|
|||
|
|
@ -13,7 +13,7 @@ use crate::providers::common::{
|
|||
};
|
||||
use crate::types::{
|
||||
AdapterTimeout, ContentPart, FinishReason, Message, RateLimitInfo, ReasoningEffort, Request,
|
||||
Response, ResponseFormatType, Role, StreamEvent, ThinkingData, TokenCounts, ToolCall,
|
||||
Response, ResponseFormatType, Role, Speed, StreamEvent, ThinkingData, TokenCounts, ToolCall,
|
||||
ToolChoice, ToolDefinition,
|
||||
};
|
||||
|
||||
|
|
@ -1209,7 +1209,7 @@ async fn build_api_request(
|
|||
output_config = None;
|
||||
}
|
||||
|
||||
let is_fast = request.speed.as_deref() == Some("fast");
|
||||
let is_fast = request.speed == Some(Speed::Fast);
|
||||
|
||||
let api_request = ApiRequest {
|
||||
model: common::api_model_id(adapter.catalog.as_deref(), &request.model),
|
||||
|
|
@ -1223,7 +1223,11 @@ async fn build_api_request(
|
|||
tool_choice: tool_choice_json,
|
||||
thinking,
|
||||
output_config,
|
||||
speed: request.speed.clone(),
|
||||
speed: request
|
||||
.speed
|
||||
.filter(|speed| *speed != Speed::Standard)
|
||||
.map(<&'static str>::from)
|
||||
.map(str::to_string),
|
||||
metadata: request.metadata.clone(),
|
||||
stream,
|
||||
};
|
||||
|
|
@ -2412,7 +2416,7 @@ mod tests {
|
|||
async fn build_api_request_sets_speed() {
|
||||
let adapter = Adapter::new("test-key");
|
||||
let request = Request {
|
||||
speed: Some("fast".to_string()),
|
||||
speed: Some(Speed::Fast),
|
||||
..make_base_request()
|
||||
};
|
||||
|
||||
|
|
@ -2424,7 +2428,7 @@ mod tests {
|
|||
async fn build_api_request_injects_fast_mode_beta_header() {
|
||||
let adapter = Adapter::new("test-key");
|
||||
let request = Request {
|
||||
speed: Some("fast".to_string()),
|
||||
speed: Some(Speed::Fast),
|
||||
..make_base_request()
|
||||
};
|
||||
|
||||
|
|
|
|||
|
|
@ -368,7 +368,7 @@ impl<'de> Deserialize<'de> for FinishReason {
|
|||
|
||||
// --- 3.9 TokenCounts ---
|
||||
|
||||
pub use fabro_model::TokenCounts;
|
||||
pub use fabro_model::{Speed, TokenCounts};
|
||||
|
||||
// --- 3.10 ResponseFormat ---
|
||||
|
||||
|
|
@ -430,7 +430,7 @@ pub struct Request {
|
|||
pub max_tokens: Option<i64>,
|
||||
pub stop_sequences: Option<Vec<String>>,
|
||||
pub reasoning_effort: Option<ReasoningEffort>,
|
||||
pub speed: Option<String>,
|
||||
pub speed: Option<Speed>,
|
||||
pub metadata: Option<HashMap<String, String>>,
|
||||
pub provider_options: Option<serde_json::Value>,
|
||||
}
|
||||
|
|
|
|||
|
|
@ -9,12 +9,13 @@ use fabro_agent::{
|
|||
ToolEnvProvider, Turn,
|
||||
};
|
||||
use fabro_auth::{CredentialSource, EnvCredentialSource};
|
||||
use fabro_graphviz::graph::Node;
|
||||
use fabro_graphviz::graph::{AttrValue, Node};
|
||||
use fabro_llm::client::Client;
|
||||
use fabro_llm::types::{Message, Request, TokenCounts};
|
||||
use fabro_llm::types::{Message, ReasoningEffort, Request, Speed, TokenCounts};
|
||||
use fabro_mcp::config::McpServerSettings;
|
||||
use fabro_model::catalog::LlmCatalogSettings;
|
||||
use fabro_model::{AgentProfileKind, Catalog, FallbackTarget, Provider, ProviderId, adapter};
|
||||
use fabro_types::settings::run::RunModelControls;
|
||||
use fabro_types::{SessionCapability, StageId};
|
||||
use tokio::sync::Mutex as TokioMutex;
|
||||
use tokio::task::JoinHandle;
|
||||
|
|
@ -106,6 +107,12 @@ struct ProviderContext {
|
|||
profile_kind: AgentProfileKind,
|
||||
}
|
||||
|
||||
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
|
||||
struct EffectiveRequestControls {
|
||||
reasoning_effort: Option<ReasoningEffort>,
|
||||
speed: Option<Speed>,
|
||||
}
|
||||
|
||||
fn classify_agent_error(err: fabro_agent::Error, allow_failover: bool) -> AgentApiErrorDisposition {
|
||||
match err {
|
||||
fabro_agent::Error::Interrupted(fabro_agent::InterruptReason::Cancelled) => {
|
||||
|
|
@ -196,6 +203,72 @@ fn default_profile_kind(provider: Provider) -> AgentProfileKind {
|
|||
}
|
||||
}
|
||||
|
||||
fn effective_request_controls(
|
||||
catalog: &Catalog,
|
||||
run_model_controls: &RunModelControls,
|
||||
model: &str,
|
||||
node: &Node,
|
||||
) -> Result<EffectiveRequestControls, Error> {
|
||||
let reasoning_effort = match control_attr(node, "reasoning_effort")
|
||||
.or(run_model_controls.reasoning_effort.as_deref())
|
||||
{
|
||||
Some(value) => Some(parse_reasoning_effort(node, value)?),
|
||||
None => legacy_reasoning_effort_default(catalog, model),
|
||||
};
|
||||
let speed = control_attr(node, "speed")
|
||||
.or(run_model_controls.speed.as_deref())
|
||||
.map(|value| parse_speed(node, value))
|
||||
.transpose()?;
|
||||
|
||||
Ok(EffectiveRequestControls {
|
||||
reasoning_effort,
|
||||
speed,
|
||||
})
|
||||
}
|
||||
|
||||
fn control_attr<'a>(node: &'a Node, key: &str) -> Option<&'a str> {
|
||||
node.attrs.get(key).and_then(AttrValue::as_str)
|
||||
}
|
||||
|
||||
fn parse_reasoning_effort(node: &Node, value: &str) -> Result<ReasoningEffort, Error> {
|
||||
value.parse::<ReasoningEffort>().map_err(|source| {
|
||||
Error::handler_with_source(
|
||||
format!(
|
||||
"Invalid reasoning_effort \"{value}\" for node \"{}\"; expected one of: low, medium, high, xhigh, max",
|
||||
node.id
|
||||
),
|
||||
&source,
|
||||
)
|
||||
})
|
||||
}
|
||||
|
||||
fn parse_speed(node: &Node, value: &str) -> Result<Speed, Error> {
|
||||
value.parse::<Speed>().map_err(|source| {
|
||||
Error::handler_with_source(
|
||||
format!(
|
||||
"Invalid speed \"{value}\" for node \"{}\"; expected one of: standard, fast",
|
||||
node.id
|
||||
),
|
||||
&source,
|
||||
)
|
||||
})
|
||||
}
|
||||
|
||||
fn legacy_reasoning_effort_default(catalog: &Catalog, model: &str) -> Option<ReasoningEffort> {
|
||||
match catalog.model_settings(model) {
|
||||
Some(settings)
|
||||
if settings
|
||||
.controls
|
||||
.reasoning_effort
|
||||
.contains(&ReasoningEffort::High) =>
|
||||
{
|
||||
Some(ReasoningEffort::High)
|
||||
}
|
||||
Some(_) => None,
|
||||
None => Some(ReasoningEffort::High),
|
||||
}
|
||||
}
|
||||
|
||||
fn profile_provider_for_catalog_provider(
|
||||
provider_id: &ProviderId,
|
||||
profile_kind: AgentProfileKind,
|
||||
|
|
@ -293,17 +366,18 @@ fn spawn_event_forwarder(
|
|||
/// For `full` fidelity nodes sharing a thread key, sessions are cached
|
||||
/// and reused so the LLM sees the full conversation history.
|
||||
pub struct AgentApiBackend {
|
||||
model: String,
|
||||
provider: Provider,
|
||||
provider_id: ProviderId,
|
||||
profile_kind: AgentProfileKind,
|
||||
fallback_chain: Vec<FallbackTarget>,
|
||||
sessions: Mutex<HashMap<String, Session>>,
|
||||
tool_env: Option<Arc<dyn ToolEnvProvider>>,
|
||||
mcp_servers: Vec<McpServerSettings>,
|
||||
source: Arc<dyn CredentialSource>,
|
||||
steering_hub: Arc<SteeringHub>,
|
||||
catalog: Arc<Catalog>,
|
||||
model: String,
|
||||
provider: Provider,
|
||||
provider_id: ProviderId,
|
||||
profile_kind: AgentProfileKind,
|
||||
fallback_chain: Vec<FallbackTarget>,
|
||||
sessions: Mutex<HashMap<String, Session>>,
|
||||
tool_env: Option<Arc<dyn ToolEnvProvider>>,
|
||||
mcp_servers: Vec<McpServerSettings>,
|
||||
run_model_controls: RunModelControls,
|
||||
source: Arc<dyn CredentialSource>,
|
||||
steering_hub: Arc<SteeringHub>,
|
||||
catalog: Arc<Catalog>,
|
||||
}
|
||||
|
||||
impl AgentApiBackend {
|
||||
|
|
@ -351,6 +425,7 @@ impl AgentApiBackend {
|
|||
sessions: Mutex::new(HashMap::new()),
|
||||
tool_env: None,
|
||||
mcp_servers: Vec::new(),
|
||||
run_model_controls: RunModelControls::default(),
|
||||
source,
|
||||
steering_hub,
|
||||
catalog,
|
||||
|
|
@ -391,6 +466,20 @@ impl AgentApiBackend {
|
|||
self
|
||||
}
|
||||
|
||||
#[must_use]
|
||||
pub fn with_run_model_controls(mut self, controls: RunModelControls) -> Self {
|
||||
self.run_model_controls = controls;
|
||||
self
|
||||
}
|
||||
|
||||
fn effective_request_controls(
|
||||
&self,
|
||||
model: &str,
|
||||
node: &Node,
|
||||
) -> Result<EffectiveRequestControls, Error> {
|
||||
effective_request_controls(self.catalog.as_ref(), &self.run_model_controls, model, node)
|
||||
}
|
||||
|
||||
fn resolve_provider_context(
|
||||
&self,
|
||||
model: &str,
|
||||
|
|
@ -451,6 +540,7 @@ impl AgentApiBackend {
|
|||
sandbox,
|
||||
self.source.as_ref(),
|
||||
Arc::clone(&self.catalog),
|
||||
&self.run_model_controls,
|
||||
self.tool_env.as_ref(),
|
||||
tool_hooks,
|
||||
self.mcp_servers.clone(),
|
||||
|
|
@ -465,11 +555,14 @@ impl AgentApiBackend {
|
|||
sandbox: &Arc<dyn Sandbox>,
|
||||
source: &dyn CredentialSource,
|
||||
catalog: Arc<Catalog>,
|
||||
run_model_controls: &RunModelControls,
|
||||
tool_env: Option<&Arc<dyn ToolEnvProvider>>,
|
||||
tool_hooks: Option<Arc<dyn fabro_agent::ToolHookCallback>>,
|
||||
mcp_servers: Vec<McpServerSettings>,
|
||||
) -> Result<Session, Error> {
|
||||
let client = Client::from_source_with_catalog(source, catalog)
|
||||
let controls =
|
||||
effective_request_controls(catalog.as_ref(), run_model_controls, model, node)?;
|
||||
let client = Client::from_source_with_catalog(source, Arc::clone(&catalog))
|
||||
.await
|
||||
.map_err(|e| Error::handler_with_source("Failed to create LLM client", &e))?;
|
||||
|
||||
|
|
@ -482,8 +575,8 @@ impl AgentApiBackend {
|
|||
|
||||
let config = SessionOptions {
|
||||
max_tokens: node.max_tokens(),
|
||||
reasoning_effort: node.reasoning_effort().parse().ok(),
|
||||
speed: node.speed().map(String::from),
|
||||
reasoning_effort: controls.reasoning_effort,
|
||||
speed: controls.speed,
|
||||
tool_hooks,
|
||||
mcp_servers,
|
||||
..SessionOptions::default()
|
||||
|
|
@ -511,7 +604,11 @@ impl AgentApiBackend {
|
|||
factory_client.clone(),
|
||||
child_profile,
|
||||
Arc::clone(&factory_env),
|
||||
SessionOptions::default(),
|
||||
SessionOptions {
|
||||
reasoning_effort: controls.reasoning_effort,
|
||||
speed: controls.speed,
|
||||
..SessionOptions::default()
|
||||
},
|
||||
None,
|
||||
);
|
||||
if let Some(provider) = &factory_tool_env {
|
||||
|
|
@ -614,6 +711,7 @@ impl CodergenBackend for AgentApiBackend {
|
|||
let model = node.model().unwrap_or(&self.model);
|
||||
let provider = self.resolve_provider_context(model, node.provider())?;
|
||||
let provider_id = provider.provider_id.to_string();
|
||||
let controls = self.effective_request_controls(model, node)?;
|
||||
|
||||
let max_tokens = node
|
||||
.max_tokens()
|
||||
|
|
@ -629,8 +727,8 @@ impl CodergenBackend for AgentApiBackend {
|
|||
model: model.to_string(),
|
||||
messages,
|
||||
provider: Some(provider_id),
|
||||
reasoning_effort: node.reasoning_effort().parse().ok(),
|
||||
speed: node.speed().map(String::from),
|
||||
reasoning_effort: controls.reasoning_effort,
|
||||
speed: controls.speed,
|
||||
tools: None,
|
||||
tool_choice: None,
|
||||
response_format: None,
|
||||
|
|
@ -692,11 +790,14 @@ impl CodergenBackend for AgentApiBackend {
|
|||
.get(&target.model)
|
||||
.and_then(|m| m.limits.max_output)
|
||||
});
|
||||
let fallback_controls = self.effective_request_controls(&target.model, node)?;
|
||||
|
||||
let fallback_request = Request {
|
||||
model: target.model.clone(),
|
||||
provider: Some(target.provider.clone()),
|
||||
max_tokens,
|
||||
reasoning_effort: fallback_controls.reasoning_effort,
|
||||
speed: fallback_controls.speed,
|
||||
..request.clone()
|
||||
};
|
||||
|
||||
|
|
@ -920,6 +1021,7 @@ impl CodergenBackend for AgentApiBackend {
|
|||
sandbox,
|
||||
self.source.as_ref(),
|
||||
Arc::clone(&self.catalog),
|
||||
&self.run_model_controls,
|
||||
self.tool_env.as_ref(),
|
||||
tool_hooks.clone(),
|
||||
self.mcp_servers.clone(),
|
||||
|
|
@ -1386,6 +1488,75 @@ effort = false
|
|||
assert_eq!(provider.provider, Provider::OpenAiCompatible);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn run_model_controls_apply_when_node_omits_controls() {
|
||||
let backend = AgentApiBackend::new_from_env(
|
||||
"gpt-5.4".to_string(),
|
||||
Provider::OpenAi,
|
||||
Vec::new(),
|
||||
SteeringHub::for_tests(),
|
||||
)
|
||||
.with_run_model_controls(fabro_types::settings::run::RunModelControls {
|
||||
reasoning_effort: Some("low".to_string()),
|
||||
speed: Some("fast".to_string()),
|
||||
});
|
||||
let node = Node::new("work");
|
||||
|
||||
let controls = backend
|
||||
.effective_request_controls("gpt-5.4", &node)
|
||||
.unwrap();
|
||||
|
||||
assert_eq!(controls.reasoning_effort, Some(ReasoningEffort::Low));
|
||||
assert_eq!(controls.speed, Some(Speed::Fast));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn node_controls_override_run_model_controls() {
|
||||
let backend = AgentApiBackend::new_from_env(
|
||||
"gpt-5.4".to_string(),
|
||||
Provider::OpenAi,
|
||||
Vec::new(),
|
||||
SteeringHub::for_tests(),
|
||||
)
|
||||
.with_run_model_controls(fabro_types::settings::run::RunModelControls {
|
||||
reasoning_effort: Some("low".to_string()),
|
||||
speed: Some("fast".to_string()),
|
||||
});
|
||||
let mut node = Node::new("work");
|
||||
node.attrs.insert(
|
||||
"reasoning_effort".to_string(),
|
||||
fabro_graphviz::graph::AttrValue::String("high".to_string()),
|
||||
);
|
||||
node.attrs.insert(
|
||||
"speed".to_string(),
|
||||
fabro_graphviz::graph::AttrValue::String("standard".to_string()),
|
||||
);
|
||||
|
||||
let controls = backend
|
||||
.effective_request_controls("gpt-5.4", &node)
|
||||
.unwrap();
|
||||
|
||||
assert_eq!(controls.reasoning_effort, Some(ReasoningEffort::High));
|
||||
assert_eq!(controls.speed, Some(Speed::Standard));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn known_model_without_effort_omits_legacy_high_default() {
|
||||
let backend = AgentApiBackend::new_from_env(
|
||||
"kimi-k2.5".to_string(),
|
||||
Provider::Kimi,
|
||||
Vec::new(),
|
||||
SteeringHub::for_tests(),
|
||||
);
|
||||
let node = Node::new("work");
|
||||
|
||||
let controls = backend
|
||||
.effective_request_controls("kimi-k2.5", &node)
|
||||
.unwrap();
|
||||
|
||||
assert_eq!(controls.reasoning_effort, None);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn api_backend_uses_source_credentials() {
|
||||
let dir = tempfile::tempdir().unwrap();
|
||||
|
|
|
|||
|
|
@ -441,6 +441,7 @@ impl RunSession {
|
|||
profile_kind,
|
||||
fallback_chain,
|
||||
mcp_servers,
|
||||
model_controls: resolved.model.controls.clone(),
|
||||
dry_run: resolved.execution.mode == RunMode::DryRun,
|
||||
},
|
||||
interviewer,
|
||||
|
|
|
|||
|
|
@ -17,6 +17,7 @@ use fabro_hooks::HookSettings;
|
|||
use fabro_interview::AutoApproveInterviewer;
|
||||
use fabro_sandbox::SandboxSpec;
|
||||
use fabro_store::Database;
|
||||
use fabro_types::settings::run::RunModelControls;
|
||||
use fabro_types::{Principal, RunId, SystemActorKind, WorkflowSettings, fixtures, format_blob_ref};
|
||||
use object_store::memory::InMemory;
|
||||
|
||||
|
|
@ -252,6 +253,7 @@ async fn execute_test_run_with_options(
|
|||
profile_kind: fabro_model::AgentProfileKind::Anthropic,
|
||||
fallback_chain: Vec::new(),
|
||||
mcp_servers: Vec::new(),
|
||||
model_controls: RunModelControls::default(),
|
||||
dry_run: true,
|
||||
},
|
||||
interviewer: Arc::new(AutoApproveInterviewer::engine()),
|
||||
|
|
@ -316,6 +318,7 @@ async fn execute_runs_start_to_exit_and_returns_final_context() {
|
|||
profile_kind: fabro_model::AgentProfileKind::Anthropic,
|
||||
fallback_chain: Vec::new(),
|
||||
mcp_servers: Vec::new(),
|
||||
model_controls: RunModelControls::default(),
|
||||
dry_run: true,
|
||||
},
|
||||
interviewer: Arc::new(AutoApproveInterviewer::engine()),
|
||||
|
|
@ -395,6 +398,7 @@ async fn run_with_lifecycle(
|
|||
profile_kind: fabro_model::AgentProfileKind::Anthropic,
|
||||
fallback_chain: Vec::new(),
|
||||
mcp_servers: Vec::new(),
|
||||
model_controls: RunModelControls::default(),
|
||||
dry_run: true,
|
||||
},
|
||||
interviewer: Arc::new(AutoApproveInterviewer::engine()),
|
||||
|
|
|
|||
|
|
@ -156,6 +156,7 @@ async fn build_registry(
|
|||
let profile_kind = spec.profile_kind;
|
||||
let fallback_chain = spec.fallback_chain.clone();
|
||||
let mcp_servers = spec.mcp_servers.clone();
|
||||
let model_controls = spec.model_controls.clone();
|
||||
let llm_source_for_api = Arc::clone(&llm_source);
|
||||
let catalog_for_api = Arc::clone(&catalog);
|
||||
let steering_hub_for_api = Arc::clone(&steering_hub);
|
||||
|
|
@ -172,6 +173,7 @@ async fn build_registry(
|
|||
Arc::clone(&steering_hub_for_api),
|
||||
Arc::clone(&catalog_for_api),
|
||||
)
|
||||
.with_run_model_controls(model_controls.clone())
|
||||
.with_tool_env_provider(tool_env_provider.clone())
|
||||
.with_mcp_servers(mcp_servers.clone());
|
||||
let cli = cli_resolver
|
||||
|
|
@ -717,6 +719,7 @@ mod tests {
|
|||
use fabro_model::catalog::LlmCatalogSettings;
|
||||
use fabro_sandbox::SandboxSpec;
|
||||
use fabro_store::Database;
|
||||
use fabro_types::settings::run::RunModelControls;
|
||||
use fabro_types::{EventBody, RunEvent, RunId, WorkflowSettings, fixtures};
|
||||
use fabro_vault::{SecretType, Vault};
|
||||
use object_store::memory::InMemory;
|
||||
|
|
@ -886,6 +889,7 @@ mod tests {
|
|||
profile_kind: fabro_model::AgentProfileKind::Anthropic,
|
||||
fallback_chain: Vec::new(),
|
||||
mcp_servers: Vec::new(),
|
||||
model_controls: RunModelControls::default(),
|
||||
dry_run: true,
|
||||
},
|
||||
interviewer: Arc::new(AutoApproveInterviewer::engine()),
|
||||
|
|
@ -948,6 +952,7 @@ mod tests {
|
|||
profile_kind: fabro_model::AgentProfileKind::Anthropic,
|
||||
fallback_chain: Vec::new(),
|
||||
mcp_servers: Vec::new(),
|
||||
model_controls: RunModelControls::default(),
|
||||
dry_run: true,
|
||||
},
|
||||
interviewer: Arc::new(AutoApproveInterviewer::engine()),
|
||||
|
|
@ -1044,6 +1049,7 @@ mod tests {
|
|||
profile_kind: fabro_model::AgentProfileKind::Anthropic,
|
||||
fallback_chain: Vec::new(),
|
||||
mcp_servers: Vec::new(),
|
||||
model_controls: RunModelControls::default(),
|
||||
dry_run: false,
|
||||
},
|
||||
Arc::new(AutoApproveInterviewer::engine()),
|
||||
|
|
@ -1162,6 +1168,7 @@ mod tests {
|
|||
profile_kind: fabro_model::AgentProfileKind::OpenAi,
|
||||
fallback_chain: Vec::new(),
|
||||
mcp_servers: Vec::new(),
|
||||
model_controls: RunModelControls::default(),
|
||||
dry_run: false,
|
||||
},
|
||||
interviewer: Arc::new(AutoApproveInterviewer::engine()),
|
||||
|
|
@ -1261,6 +1268,7 @@ mod tests {
|
|||
profile_kind: fabro_model::AgentProfileKind::Anthropic,
|
||||
fallback_chain: Vec::new(),
|
||||
mcp_servers: Vec::new(),
|
||||
model_controls: RunModelControls::default(),
|
||||
dry_run: true,
|
||||
},
|
||||
interviewer: Arc::new(AutoApproveInterviewer::engine()),
|
||||
|
|
@ -1378,6 +1386,7 @@ mod tests {
|
|||
profile_kind: fabro_model::AgentProfileKind::Anthropic,
|
||||
fallback_chain: Vec::new(),
|
||||
mcp_servers: Vec::new(),
|
||||
model_controls: RunModelControls::default(),
|
||||
dry_run: true,
|
||||
},
|
||||
interviewer: Arc::new(AutoApproveInterviewer::engine()),
|
||||
|
|
@ -1444,6 +1453,7 @@ mod tests {
|
|||
profile_kind: fabro_model::AgentProfileKind::Anthropic,
|
||||
fallback_chain: Vec::new(),
|
||||
mcp_servers: Vec::new(),
|
||||
model_controls: RunModelControls::default(),
|
||||
dry_run: true,
|
||||
},
|
||||
interviewer: Arc::new(AutoApproveInterviewer::engine()),
|
||||
|
|
|
|||
|
|
@ -9,7 +9,7 @@ use fabro_mcp::config::McpServerSettings;
|
|||
use fabro_model::{AgentProfileKind, Catalog, FallbackTarget, ProviderId};
|
||||
use fabro_sandbox::SandboxSpec;
|
||||
use fabro_types::RunId;
|
||||
use fabro_types::settings::run::PullRequestSettings;
|
||||
use fabro_types::settings::run::{PullRequestSettings, RunModelControls};
|
||||
use fabro_validate::{Diagnostic, Severity};
|
||||
use fabro_vault::Vault;
|
||||
use tokio::sync::RwLock as AsyncRwLock;
|
||||
|
|
@ -224,6 +224,7 @@ pub struct LlmSpec {
|
|||
pub profile_kind: AgentProfileKind,
|
||||
pub fallback_chain: Vec<FallbackTarget>,
|
||||
pub mcp_servers: Vec<McpServerSettings>,
|
||||
pub model_controls: RunModelControls,
|
||||
pub dry_run: bool,
|
||||
}
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue