1use std::sync::Arc;
21
22use bamboo_domain::reasoning::ReasoningEffort;
23use bamboo_domain::ProviderModelRef;
24use bamboo_llm::{Config, ProviderRegistry, ResolvedModel};
25
26use crate::model_config_helper::{
27 resolve_background_model, resolve_fast_model, resolve_subagent_model,
28 resolve_task_summary_model, resolve_vision_model,
29};
30
31pub struct GlobalAreaModels {
38 pub fast: Option<ResolvedModel>,
40 pub fast_ref: Option<ProviderModelRef>,
41 pub background: Option<ResolvedModel>,
43 pub background_ref: Option<ProviderModelRef>,
44 pub summarization: Option<ResolvedModel>,
46 pub summarization_ref: Option<ProviderModelRef>,
47}
48
49pub fn resolve_global_area_models(
58 config: &Config,
59 provider_name: &str,
60 provider_registry: &Arc<ProviderRegistry>,
61) -> GlobalAreaModels {
62 let defaults = config.defaults.as_ref();
63 GlobalAreaModels {
64 fast: resolve_fast_model(config, provider_name, provider_registry),
65 fast_ref: defaults.and_then(|d| d.fast.clone()),
66 background: resolve_background_model(config, provider_name, provider_registry),
67 background_ref: defaults.and_then(|d| d.memory_background.clone()),
68 summarization: resolve_task_summary_model(config, provider_name, provider_registry),
69 summarization_ref: defaults.and_then(|d| d.task_summary.clone()),
70 }
71}
72
73pub fn resolve_global_vision_model(
77 config: &Config,
78 provider_name: &str,
79 provider_registry: &Arc<ProviderRegistry>,
80) -> Option<ResolvedModel> {
81 resolve_vision_model(config, provider_name, provider_registry)
82}
83
84pub fn resolve_global_subagent_model(
89 config: &Config,
90 provider_name: &str,
91 provider_registry: &Arc<ProviderRegistry>,
92 subagent_type: &str,
93) -> Option<ResolvedModel> {
94 resolve_subagent_model(config, provider_name, provider_registry, subagent_type)
95}
96
97#[derive(Debug, Clone, Copy, PartialEq, Eq)]
100pub enum ReasoningEffortSource {
101 Session,
102 Request,
103 ProviderDefault,
104 None,
105}
106
107impl ReasoningEffortSource {
108 pub fn as_str(self) -> &'static str {
109 match self {
110 Self::Session => "session",
111 Self::Request => "request",
112 Self::ProviderDefault => "provider_default",
113 Self::None => "none",
114 }
115 }
116}
117
118pub fn resolve_effective_reasoning_effort(
126 session_effort: Option<ReasoningEffort>,
127 request_effort: Option<ReasoningEffort>,
128 provider_default: Option<ReasoningEffort>,
129) -> (Option<ReasoningEffort>, ReasoningEffortSource) {
130 if let Some(effort) = session_effort {
131 (Some(effort), ReasoningEffortSource::Session)
132 } else if let Some(effort) = request_effort {
133 (Some(effort), ReasoningEffortSource::Request)
134 } else if let Some(effort) = provider_default {
135 (Some(effort), ReasoningEffortSource::ProviderDefault)
136 } else {
137 (None, ReasoningEffortSource::None)
138 }
139}
140
141#[cfg(test)]
142mod tests {
143 use super::*;
144
145 macro_rules! test_config {
146 (@assign $config:ident, providers, $value:expr) => { *$config.providers_mut() = $value; };
147 (@assign $config:ident, memory, $value:expr) => { *$config.memory_mut() = $value; };
148 (@assign $config:ident, subagents, $value:expr) => { *$config.subagents_mut() = $value; };
149 (@assign $config:ident, $field:ident, $value:expr) => { $config.$field = $value; };
150 ($($field:ident: $value:expr),* $(,)?) => {{
151 let mut config = Config::default();
152 $(test_config!(@assign config, $field, $value);)*
153 config
154 }};
155 }
156 use bamboo_agent_core::tools::ToolSchema;
157 use bamboo_agent_core::Message;
158 use bamboo_config::{DefaultsConfig, FeatureFlags};
159 use bamboo_config::{OpenAIConfig, ProviderConfigs};
160 use bamboo_domain::{Session, DEFAULT_REASONING_EFFORT};
161 use bamboo_llm::{LLMError, LLMProvider, LLMStream};
162 use std::collections::HashMap;
163
164 struct NoopProvider;
165
166 #[async_trait::async_trait]
167 impl LLMProvider for NoopProvider {
168 async fn chat_stream(
169 &self,
170 _messages: &[Message],
171 _tools: &[ToolSchema],
172 _max_output_tokens: Option<u32>,
173 _model: &str,
174 ) -> Result<LLMStream, LLMError> {
175 Err(LLMError::Api("noop".to_string()))
176 }
177 }
178
179 fn test_registry() -> Arc<ProviderRegistry> {
180 let mut providers: HashMap<String, Arc<dyn LLMProvider>> = HashMap::new();
181 providers.insert("openai".to_string(), Arc::new(NoopProvider));
182 Arc::new(ProviderRegistry::new(providers, "openai".to_string()))
183 }
184
185 fn defaults_with_all_areas() -> DefaultsConfig {
186 DefaultsConfig {
187 chat: ProviderModelRef::new("openai", "gpt-chat"),
188 fast: Some(ProviderModelRef::new("openai", "gpt-fast")),
189 task_summary: Some(ProviderModelRef::new("openai", "gpt-summary")),
190 vision: Some(ProviderModelRef::new("openai", "gpt-vision")),
191 memory_background: Some(ProviderModelRef::new("openai", "gpt-memory")),
192 planning: None,
193 search: None,
194 code_review: None,
195 sub_agent: Some(ProviderModelRef::new("openai", "gpt-sub")),
196 subagent_models: HashMap::new(),
197 }
198 }
199
200 fn config_with_defaults(defaults: DefaultsConfig) -> Config {
201 test_config! {
202 provider: "openai".to_string(),
203 features: FeatureFlags {
204 provider_model_ref: true,
205 ..Default::default()
206 },
207 defaults: Some(defaults),
208 }
209 }
210
211 #[test]
214 fn global_area_models_read_each_area_from_its_own_default() {
215 let config = config_with_defaults(defaults_with_all_areas());
216 let areas = resolve_global_area_models(&config, "openai", &test_registry());
217
218 assert_eq!(
219 areas.fast.as_ref().map(|m| m.model_name.as_str()),
220 Some("gpt-fast")
221 );
222 assert_eq!(
223 areas.summarization.as_ref().map(|m| m.model_name.as_str()),
224 Some("gpt-summary")
225 );
226 assert_eq!(
227 areas.background.as_ref().map(|m| m.model_name.as_str()),
228 Some("gpt-memory")
229 );
230 assert_eq!(
232 areas.fast_ref,
233 Some(ProviderModelRef::new("openai", "gpt-fast"))
234 );
235 assert_eq!(
236 areas.summarization_ref,
237 Some(ProviderModelRef::new("openai", "gpt-summary"))
238 );
239 assert_eq!(
240 areas.background_ref,
241 Some(ProviderModelRef::new("openai", "gpt-memory"))
242 );
243 }
244
245 #[test]
252 fn global_area_models_are_independent_of_any_session() {
253 let config = config_with_defaults(defaults_with_all_areas());
254 let registry = test_registry();
255
256 let before = resolve_global_area_models(&config, "openai", ®istry);
257
258 let mut session = Session::new("s1", "some-exotic-session-model");
260 session.model_ref = Some(ProviderModelRef::new("openai", "some-exotic-session-model"));
261 session.reasoning_effort = Some(ReasoningEffort::Max);
262 let _ = &session; let after = resolve_global_area_models(&config, "openai", ®istry);
265
266 assert_eq!(
267 before.fast.as_ref().map(|m| m.model_name.clone()),
268 after.fast.as_ref().map(|m| m.model_name.clone())
269 );
270 assert_eq!(
271 before.background.as_ref().map(|m| m.model_name.clone()),
272 after.background.as_ref().map(|m| m.model_name.clone())
273 );
274 assert_eq!(
275 before.summarization.as_ref().map(|m| m.model_name.clone()),
276 after.summarization.as_ref().map(|m| m.model_name.clone())
277 );
278 assert_ne!(
280 after.fast.as_ref().map(|m| m.model_name.as_str()),
281 Some("some-exotic-session-model")
282 );
283 }
284
285 #[test]
286 fn vision_model_is_global_from_defaults() {
287 let config = config_with_defaults(defaults_with_all_areas());
288 let vision = resolve_global_vision_model(&config, "openai", &test_registry());
289 assert_eq!(
290 vision.as_ref().map(|m| m.model_name.as_str()),
291 Some("gpt-vision")
292 );
293 }
294
295 #[test]
296 fn subagent_model_is_global_from_defaults() {
297 let config = config_with_defaults(defaults_with_all_areas());
298 let sub = resolve_global_subagent_model(&config, "openai", &test_registry(), "coder");
300 assert_eq!(sub.as_ref().map(|m| m.model_name.as_str()), Some("gpt-sub"));
301 }
302
303 #[test]
304 fn background_falls_back_to_fast_when_memory_background_unset() {
305 let mut defaults = defaults_with_all_areas();
306 defaults.memory_background = None;
307 let config = config_with_defaults(defaults);
308
309 let areas = resolve_global_area_models(&config, "openai", &test_registry());
310 assert_eq!(
312 areas.background.as_ref().map(|m| m.model_name.as_str()),
313 Some("gpt-fast")
314 );
315 }
316
317 #[test]
318 fn legacy_mode_resolves_fast_from_provider_config() {
319 let config = test_config! {
321 provider: "openai".to_string(),
322 features: FeatureFlags {
323 provider_model_ref: false,
324 ..Default::default()
325 },
326 defaults: None,
327 providers: ProviderConfigs {
328 openai: Some(OpenAIConfig {
329 api_key: "test".to_string(),
330 api_key_from_env: false,
331 api_key_encrypted: None,
332 credential_ref: None,
333 base_url: None,
334 model: Some("gpt-4o".to_string()),
335 fast_model: Some("gpt-4o-mini".to_string()),
336 vision_model: None,
337 reasoning_effort: None,
338 responses_only_models: vec![],
339 request_overrides: None,
340 extra: Default::default(),
341 }),
342 ..ProviderConfigs::default()
343 },
344 };
345
346 let areas = resolve_global_area_models(&config, "openai", &test_registry());
347 assert_eq!(
348 areas.fast.as_ref().map(|m| m.model_name.as_str()),
349 Some("gpt-4o-mini")
350 );
351 }
352
353 #[test]
356 fn reasoning_prefers_session_then_request_then_provider() {
357 assert_eq!(
358 resolve_effective_reasoning_effort(
359 Some(ReasoningEffort::Max),
360 Some(ReasoningEffort::High),
361 Some(ReasoningEffort::Low),
362 ),
363 (Some(ReasoningEffort::Max), ReasoningEffortSource::Session)
364 );
365 assert_eq!(
366 resolve_effective_reasoning_effort(
367 None,
368 Some(ReasoningEffort::High),
369 Some(ReasoningEffort::Low),
370 ),
371 (Some(ReasoningEffort::High), ReasoningEffortSource::Request)
372 );
373 assert_eq!(
374 resolve_effective_reasoning_effort(None, None, Some(ReasoningEffort::Low)),
375 (
376 Some(ReasoningEffort::Low),
377 ReasoningEffortSource::ProviderDefault
378 )
379 );
380 }
381
382 #[test]
383 fn reasoning_none_when_nothing_configured() {
384 let (effort, source) = resolve_effective_reasoning_effort(None, None, None);
385 assert_eq!(effort, None);
386 assert_eq!(source, ReasoningEffortSource::None);
387 }
388
389 #[test]
390 fn canonical_default_is_medium_and_used_as_terminal() {
391 assert_eq!(DEFAULT_REASONING_EFFORT, ReasoningEffort::Medium);
394 let (effort, _) = resolve_effective_reasoning_effort(None, None, None);
395 assert_eq!(
396 effort.unwrap_or(DEFAULT_REASONING_EFFORT),
397 ReasoningEffort::Medium
398 );
399 }
400}