fix(summon): re-apply canonical limits when delegate overrides model (#9183)
Signed-off-by: Kyle De Freitas <kdefreitas@squareup.com> Signed-off-by: Douwe Osinga <douwe@squareup.com> Co-authored-by: Douwe Osinga <douwe@squareup.com>
This commit is contained in:
@@ -1281,6 +1281,61 @@ impl SummonClient {
|
|||||||
Ok(task_config)
|
Ok(task_config)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
fn resolve_model_config(
|
||||||
|
&self,
|
||||||
|
params: &DelegateParams,
|
||||||
|
recipe: &Recipe,
|
||||||
|
session: &crate::session::Session,
|
||||||
|
provider_name: &str,
|
||||||
|
) -> Result<crate::model::ModelConfig, anyhow::Error> {
|
||||||
|
let mut model_config = session.model_config.clone().map(Ok).unwrap_or_else(|| {
|
||||||
|
crate::model::ModelConfig::new("default")
|
||||||
|
.map(|c| c.with_canonical_limits(provider_name))
|
||||||
|
})?;
|
||||||
|
|
||||||
|
let override_model = params
|
||||||
|
.model
|
||||||
|
.clone()
|
||||||
|
.or_else(|| recipe.settings.as_ref().and_then(|s| s.goose_model.clone()))
|
||||||
|
.or_else(|| {
|
||||||
|
Config::global()
|
||||||
|
.get_param::<String>("GOOSE_SUBAGENT_MODEL")
|
||||||
|
.ok()
|
||||||
|
});
|
||||||
|
|
||||||
|
if let Some(model) = override_model {
|
||||||
|
if model != model_config.model_name {
|
||||||
|
// Build the new config from scratch so canonical fields
|
||||||
|
// (context_limit, max_tokens, reasoning) and env-derived
|
||||||
|
// overrides (GOOSE_CONTEXT_LIMIT, GOOSE_MAX_TOKENS) match the
|
||||||
|
// overridden model, then preserve session-level state that is
|
||||||
|
// not model-specific from the parent.
|
||||||
|
let parent = model_config;
|
||||||
|
let mut cfg =
|
||||||
|
crate::model::ModelConfig::new(&model)?.with_canonical_limits(provider_name);
|
||||||
|
cfg.toolshim = parent.toolshim;
|
||||||
|
cfg.toolshim_model = parent.toolshim_model;
|
||||||
|
cfg.fast_model_config = parent.fast_model_config;
|
||||||
|
cfg.temperature = cfg.temperature.or(parent.temperature);
|
||||||
|
if let Some(parent_params) = parent.request_params {
|
||||||
|
let merged = cfg.request_params.get_or_insert_with(Default::default);
|
||||||
|
for (k, v) in parent_params {
|
||||||
|
merged.insert(k, v);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
model_config = cfg;
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
if let Some(temp) = params.temperature {
|
||||||
|
model_config = model_config.with_temperature(Some(temp));
|
||||||
|
} else if let Some(temp) = recipe.settings.as_ref().and_then(|s| s.temperature) {
|
||||||
|
model_config = model_config.with_temperature(Some(temp));
|
||||||
|
}
|
||||||
|
|
||||||
|
Ok(model_config)
|
||||||
|
}
|
||||||
|
|
||||||
async fn resolve_provider(
|
async fn resolve_provider(
|
||||||
&self,
|
&self,
|
||||||
params: &DelegateParams,
|
params: &DelegateParams,
|
||||||
@@ -1304,29 +1359,7 @@ impl SummonClient {
|
|||||||
.or_else(|| session.provider_name.clone())
|
.or_else(|| session.provider_name.clone())
|
||||||
.ok_or_else(|| anyhow::anyhow!("No provider configured"))?;
|
.ok_or_else(|| anyhow::anyhow!("No provider configured"))?;
|
||||||
|
|
||||||
let mut model_config = session.model_config.clone().map(Ok).unwrap_or_else(|| {
|
let model_config = self.resolve_model_config(params, recipe, session, &provider_name)?;
|
||||||
crate::model::ModelConfig::new("default")
|
|
||||||
.map(|c| c.with_canonical_limits(&provider_name))
|
|
||||||
})?;
|
|
||||||
|
|
||||||
if let Some(model) = ¶ms.model {
|
|
||||||
model_config.model_name = model.clone();
|
|
||||||
} else if let Some(model) = recipe
|
|
||||||
.settings
|
|
||||||
.as_ref()
|
|
||||||
.and_then(|s| s.goose_model.as_ref())
|
|
||||||
{
|
|
||||||
model_config.model_name = model.clone();
|
|
||||||
} else if let Ok(model) = Config::global().get_param::<String>("GOOSE_SUBAGENT_MODEL") {
|
|
||||||
model_config.model_name = model;
|
|
||||||
}
|
|
||||||
|
|
||||||
if let Some(temp) = params.temperature {
|
|
||||||
model_config = model_config.with_temperature(Some(temp));
|
|
||||||
} else if let Some(temp) = recipe.settings.as_ref().and_then(|s| s.temperature) {
|
|
||||||
model_config = model_config.with_temperature(Some(temp));
|
|
||||||
}
|
|
||||||
|
|
||||||
providers::create(&provider_name, model_config, Vec::new()).await
|
providers::create(&provider_name, model_config, Vec::new()).await
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -1672,7 +1705,7 @@ impl McpClientTrait for SummonClient {
|
|||||||
mod tests {
|
mod tests {
|
||||||
use super::*;
|
use super::*;
|
||||||
use serial_test::serial;
|
use serial_test::serial;
|
||||||
use std::collections::HashSet;
|
use std::collections::{HashMap, HashSet};
|
||||||
use std::fs;
|
use std::fs;
|
||||||
use std::sync::Arc;
|
use std::sync::Arc;
|
||||||
use tempfile::TempDir;
|
use tempfile::TempDir;
|
||||||
@@ -2041,6 +2074,106 @@ You review code."#;
|
|||||||
);
|
);
|
||||||
}
|
}
|
||||||
|
|
||||||
|
fn empty_recipe() -> crate::recipe::Recipe {
|
||||||
|
crate::recipe::Recipe {
|
||||||
|
version: "1.0.0".to_string(),
|
||||||
|
title: String::new(),
|
||||||
|
description: String::new(),
|
||||||
|
instructions: None,
|
||||||
|
prompt: None,
|
||||||
|
extensions: None,
|
||||||
|
settings: None,
|
||||||
|
activities: None,
|
||||||
|
author: None,
|
||||||
|
parameters: None,
|
||||||
|
response: None,
|
||||||
|
sub_recipes: None,
|
||||||
|
retry: None,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
const PARENT_MODEL: &str = "claude-3-5-sonnet-20241022";
|
||||||
|
const OVERRIDE_MODEL: &str = "claude-opus-4-6";
|
||||||
|
const PROVIDER: &str = "anthropic";
|
||||||
|
|
||||||
|
fn session_with(parent: crate::model::ModelConfig) -> crate::session::Session {
|
||||||
|
crate::session::Session {
|
||||||
|
provider_name: Some(PROVIDER.to_string()),
|
||||||
|
model_config: Some(parent),
|
||||||
|
..Default::default()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
fn resolve_with_override(
|
||||||
|
model: Option<&str>,
|
||||||
|
parent: crate::model::ModelConfig,
|
||||||
|
) -> crate::model::ModelConfig {
|
||||||
|
let client = SummonClient::new(create_test_context()).unwrap();
|
||||||
|
let params = DelegateParams {
|
||||||
|
model: model.map(String::from),
|
||||||
|
..Default::default()
|
||||||
|
};
|
||||||
|
client
|
||||||
|
.resolve_model_config(¶ms, &empty_recipe(), &session_with(parent), PROVIDER)
|
||||||
|
.expect("resolve_model_config")
|
||||||
|
}
|
||||||
|
|
||||||
|
fn parent_config() -> crate::model::ModelConfig {
|
||||||
|
crate::model::ModelConfig::new(PARENT_MODEL)
|
||||||
|
.unwrap()
|
||||||
|
.with_canonical_limits(PROVIDER)
|
||||||
|
}
|
||||||
|
|
||||||
|
#[tokio::test]
|
||||||
|
#[serial]
|
||||||
|
async fn test_resolve_model_config_applies_canonical_limits_to_overridden_model() {
|
||||||
|
let _env = env_lock::lock_env([
|
||||||
|
("GOOSE_CONTEXT_LIMIT", None::<&str>),
|
||||||
|
("GOOSE_MAX_TOKENS", None::<&str>),
|
||||||
|
("GOOSE_SUBAGENT_MODEL", None::<&str>),
|
||||||
|
]);
|
||||||
|
|
||||||
|
let parent = parent_config();
|
||||||
|
let overridden = crate::model::ModelConfig::new(OVERRIDE_MODEL)
|
||||||
|
.unwrap()
|
||||||
|
.with_canonical_limits(PROVIDER);
|
||||||
|
assert_ne!(parent.context_limit, overridden.context_limit);
|
||||||
|
assert_ne!(parent.reasoning, overridden.reasoning);
|
||||||
|
|
||||||
|
let resolved = resolve_with_override(Some(OVERRIDE_MODEL), parent);
|
||||||
|
|
||||||
|
assert_eq!(resolved.model_name, OVERRIDE_MODEL);
|
||||||
|
assert_eq!(resolved.context_limit, overridden.context_limit);
|
||||||
|
assert_eq!(resolved.max_tokens, overridden.max_tokens);
|
||||||
|
assert_eq!(resolved.reasoning, overridden.reasoning);
|
||||||
|
}
|
||||||
|
|
||||||
|
#[tokio::test]
|
||||||
|
#[serial]
|
||||||
|
async fn test_resolve_model_config_preserves_parent_request_params_on_override() {
|
||||||
|
let _env = env_lock::lock_env([
|
||||||
|
("GOOSE_CONTEXT_LIMIT", None::<&str>),
|
||||||
|
("GOOSE_MAX_TOKENS", None::<&str>),
|
||||||
|
("GOOSE_SUBAGENT_MODEL", None::<&str>),
|
||||||
|
]);
|
||||||
|
|
||||||
|
let mut parent = parent_config();
|
||||||
|
parent.request_params = Some(HashMap::from([(
|
||||||
|
"anthropic_beta".to_string(),
|
||||||
|
serde_json::json!("custom-beta-header"),
|
||||||
|
)]));
|
||||||
|
|
||||||
|
let resolved = resolve_with_override(Some(OVERRIDE_MODEL), parent);
|
||||||
|
|
||||||
|
assert_eq!(
|
||||||
|
resolved
|
||||||
|
.request_params
|
||||||
|
.as_ref()
|
||||||
|
.and_then(|p| p.get("anthropic_beta")),
|
||||||
|
Some(&serde_json::json!("custom-beta-header")),
|
||||||
|
);
|
||||||
|
}
|
||||||
|
|
||||||
fn extract_text(content: &Content) -> &str {
|
fn extract_text(content: &Content) -> &str {
|
||||||
use rmcp::model::RawContent;
|
use rmcp::model::RawContent;
|
||||||
match &content.raw {
|
match &content.raw {
|
||||||
|
|||||||
Reference in New Issue
Block a user