use serde::{Deserialize, Serialize};
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub struct ModelSwitchPlan {
pub source_model: String,
pub target_model: String,
pub source_provider: Option<String>,
pub target_provider: Option<String>,
pub context_adaptation: ContextAdaptationPlan,
pub capability_diff: CapabilityDiff,
pub compaction_triggered: bool,
pub estimated_tokens_after: usize,
pub target_window: Option<usize>,
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub struct ContextAdaptationPlan {
pub needs_compaction: bool,
pub tail_messages: usize,
pub tail_fits: bool,
pub retained_tokens: usize,
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub struct CapabilityDiff {
pub hidden_tools: Vec<String>,
pub unsupported_modalities: Vec<String>,
pub window_shrink: Option<usize>,
pub unsupported_features: Vec<String>,
pub tools_supported: bool,
pub streaming_supported: bool,
}
impl Default for CapabilityDiff {
fn default() -> Self {
Self {
hidden_tools: Vec::new(),
unsupported_modalities: Vec::new(),
window_shrink: None,
unsupported_features: Vec::new(),
tools_supported: true,
streaming_supported: true,
}
}
}
#[derive(Debug, Clone, PartialEq)]
pub enum ModelSwitchRequest {
Managed {
provider_slug: String,
model: String,
api_key: Option<String>,
},
Unmanaged {
model: String,
provider_name: String,
},
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub struct ModelSwitchRecord {
pub from_model: String,
pub to_model: String,
pub from_provider: Option<String>,
pub to_provider: Option<String>,
pub adapted: bool,
pub capability_diff: CapabilityDiff,
pub timestamp: chrono::DateTime<chrono::Utc>,
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub struct CompactionInfo {
pub tokens_before: u32,
pub tokens_after: u32,
pub method: String,
}
#[derive(Debug, Clone)]
pub struct PendingModelSwitch {
pub request: ModelSwitchRequest,
pub requested_at: chrono::DateTime<chrono::Utc>,
}
impl ModelSwitchPlan {
pub fn new(source_model: impl Into<String>, target_model: impl Into<String>) -> Self {
Self {
source_model: source_model.into(),
target_model: target_model.into(),
source_provider: None,
target_provider: None,
context_adaptation: ContextAdaptationPlan {
needs_compaction: false,
tail_messages: 20,
tail_fits: true,
retained_tokens: 0,
},
capability_diff: CapabilityDiff::default(),
compaction_triggered: false,
estimated_tokens_after: 0,
target_window: None,
}
}
pub fn requires_adaptation(&self) -> bool {
self.context_adaptation.needs_compaction
|| !self.capability_diff.hidden_tools.is_empty()
|| !self.capability_diff.unsupported_modalities.is_empty()
|| !self.capability_diff.unsupported_features.is_empty()
}
}
impl ContextAdaptationPlan {
pub fn fits(&self) -> bool {
self.tail_fits && !self.needs_compaction
}
}
pub struct ModelSwitchPlanner;
impl ModelSwitchPlanner {
pub fn create_plan(
source_model: &str,
target_model: &str,
target_window: Option<usize>,
current_tokens: usize,
) -> ModelSwitchPlan {
Self::create_plan_with_capabilities(
source_model,
target_model,
target_window,
current_tokens,
None,
None,
)
}
pub fn create_plan_with_capabilities(
source_model: &str,
target_model: &str,
target_window: Option<usize>,
current_tokens: usize,
source_capabilities: Option<&ModelCapabilityMetadata>,
target_capabilities: Option<&ModelCapabilityMetadata>,
) -> ModelSwitchPlan {
let mut plan = ModelSwitchPlan::new(source_model, target_model);
plan.target_window = target_window;
if let (Some(source), Some(target)) = (source_capabilities, target_capabilities) {
plan.capability_diff = compare_capabilities(source, target);
}
if let Some(window) = target_window {
let _estimated_tail_tokens = Self::estimate_tail_tokens(current_tokens);
if current_tokens > window {
plan.context_adaptation.needs_compaction = true;
plan.compaction_triggered = true;
let avg_tokens_per_message = if current_tokens > 0 {
current_tokens / 20 } else {
100
};
let max_tail_messages = (window / avg_tokens_per_message.max(1)).max(5);
plan.context_adaptation.tail_messages = max_tail_messages.min(20);
plan.context_adaptation.retained_tokens =
plan.context_adaptation.tail_messages * avg_tokens_per_message;
plan.context_adaptation.tail_fits =
plan.context_adaptation.retained_tokens <= window;
plan.estimated_tokens_after = plan.context_adaptation.retained_tokens;
} else {
plan.context_adaptation.needs_compaction = false;
plan.context_adaptation.tail_messages = 20;
plan.context_adaptation.tail_fits = true;
plan.context_adaptation.retained_tokens = current_tokens;
plan.estimated_tokens_after = current_tokens;
}
} else {
plan.context_adaptation.needs_compaction = false;
plan.context_adaptation.tail_fits = true;
plan.context_adaptation.retained_tokens = current_tokens;
plan.estimated_tokens_after = current_tokens;
}
plan
}
fn estimate_tail_tokens(total_tokens: usize) -> usize {
(total_tokens as f64 * 0.6) as usize
}
pub fn estimate_session_tokens(
uncompacted_tokens: usize,
compressed_blocks: &[crate::context::models::CompressedBlock],
) -> usize {
let compressed_estimate: usize = compressed_blocks
.iter()
.map(|block| block.summary.len() / 4) .sum();
uncompacted_tokens + compressed_estimate
}
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub struct ModelCapabilityMetadata {
pub model: String,
pub provider: String,
pub context_window: usize,
pub supports_tools: bool,
pub supports_streaming: bool,
#[serde(default)]
pub supports_reasoning_effort: bool,
#[serde(default)]
pub reasoning_effort_values: Vec<String>,
pub supported_modalities: Vec<String>,
pub unsupported_tools: Vec<String>,
}
pub fn compare_capabilities(
source: &ModelCapabilityMetadata,
target: &ModelCapabilityMetadata,
) -> CapabilityDiff {
let mut diff = CapabilityDiff::default();
if target.context_window < source.context_window {
diff.window_shrink = Some(source.context_window - target.context_window);
}
if !target.supports_tools {
diff.tools_supported = false;
}
if !target.supports_streaming {
diff.streaming_supported = false;
}
if source.supports_reasoning_effort && !target.supports_reasoning_effort {
diff.unsupported_features.push("reasoning_effort".into());
}
for modality in &source.supported_modalities {
if !target.supported_modalities.contains(modality) {
diff.unsupported_modalities.push(modality.clone());
}
}
for tool in &target.unsupported_tools {
diff.hidden_tools.push(tool.clone());
}
diff
}
#[derive(Debug, Clone, Default)]
pub struct ModelCapabilityRegistry {
models: std::collections::HashMap<(String, String), ModelCapabilityMetadata>,
}
impl ModelCapabilityRegistry {
pub fn new() -> Self {
Self::default()
}
pub fn register(&mut self, metadata: ModelCapabilityMetadata) {
let key = (metadata.provider.clone(), metadata.model.clone());
self.models.insert(key, metadata);
}
pub fn get(&self, provider: &str, model: &str) -> Option<&ModelCapabilityMetadata> {
self.models.get(&(provider.to_string(), model.to_string()))
}
pub fn compare(
&self,
source_provider: &str,
source_model: &str,
target_provider: &str,
target_model: &str,
) -> Option<CapabilityDiff> {
let source = self.get(source_provider, source_model)?;
let target = self.get(target_provider, target_model)?;
Some(compare_capabilities(source, target))
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_model_switch_plan_requires_adaptation() {
let plan = ModelSwitchPlan {
source_model: "model-a".into(),
target_model: "model-b".into(),
source_provider: None,
target_provider: None,
context_adaptation: ContextAdaptationPlan {
needs_compaction: true,
tail_messages: 10,
tail_fits: true,
retained_tokens: 5000,
},
capability_diff: CapabilityDiff::default(),
compaction_triggered: true,
estimated_tokens_after: 5000,
target_window: Some(10000),
};
assert!(plan.requires_adaptation());
}
#[test]
fn test_model_switch_plan_no_adaptation_needed() {
let plan = ModelSwitchPlan::new("model-a", "model-b");
assert!(!plan.requires_adaptation());
assert!(plan.context_adaptation.fits());
}
#[test]
fn test_model_switch_planner_context_fits() {
let plan = ModelSwitchPlanner::create_plan("model-a", "model-b", Some(100000), 50000);
assert!(!plan.context_adaptation.needs_compaction);
assert!(plan.context_adaptation.fits());
assert_eq!(plan.context_adaptation.tail_messages, 20);
}
#[test]
fn test_model_switch_planner_context_exceeds_window() {
let plan = ModelSwitchPlanner::create_plan("model-a", "model-b", Some(10000), 50000);
assert!(plan.context_adaptation.needs_compaction);
assert!(plan.compaction_triggered);
assert!(!plan.context_adaptation.fits());
assert!(plan.context_adaptation.tail_messages < 20);
}
#[test]
fn test_model_switch_planner_unknown_window() {
let plan = ModelSwitchPlanner::create_plan("model-a", "model-b", None, 50000);
assert!(!plan.context_adaptation.needs_compaction);
assert!(plan.context_adaptation.fits());
}
#[test]
fn test_estimate_session_tokens() {
let blocks = vec![crate::context::models::CompressedBlock::new(
"c0001",
"Test topic",
"m0001-m0010",
"a".repeat(400),
)];
let tokens = ModelSwitchPlanner::estimate_session_tokens(1000, &blocks);
assert_eq!(tokens, 1100); }
#[test]
fn test_capability_diff_default() {
let diff = CapabilityDiff::default();
assert!(diff.hidden_tools.is_empty());
assert!(diff.unsupported_modalities.is_empty());
assert!(diff.unsupported_features.is_empty());
assert!(diff.tools_supported);
assert!(diff.streaming_supported);
assert_eq!(diff.window_shrink, None);
}
#[test]
fn test_compare_capabilities_window_shrink() {
let source = ModelCapabilityMetadata {
model: "model-a".into(),
provider: "provider-a".into(),
context_window: 100000,
supports_tools: true,
supports_streaming: true,
supports_reasoning_effort: true,
reasoning_effort_values: vec!["low".into(), "medium".into(), "high".into()],
supported_modalities: vec!["text".into()],
unsupported_tools: vec![],
};
let target = ModelCapabilityMetadata {
model: "model-b".into(),
provider: "provider-b".into(),
context_window: 50000,
supports_tools: true,
supports_streaming: true,
supports_reasoning_effort: true,
reasoning_effort_values: vec!["low".into(), "medium".into(), "high".into()],
supported_modalities: vec!["text".into()],
unsupported_tools: vec![],
};
let diff = compare_capabilities(&source, &target);
assert_eq!(diff.window_shrink, Some(50000));
assert!(diff.tools_supported);
}
#[test]
fn test_compare_capabilities_tool_loss() {
let source = ModelCapabilityMetadata {
model: "model-a".into(),
provider: "provider-a".into(),
context_window: 100000,
supports_tools: true,
supports_streaming: true,
supports_reasoning_effort: true,
reasoning_effort_values: vec!["low".into(), "medium".into(), "high".into()],
supported_modalities: vec!["text".into()],
unsupported_tools: vec![],
};
let target = ModelCapabilityMetadata {
model: "model-b".into(),
provider: "provider-b".into(),
context_window: 100000,
supports_tools: false,
supports_streaming: true,
supports_reasoning_effort: false,
reasoning_effort_values: vec![],
supported_modalities: vec!["text".into()],
unsupported_tools: vec!["tool1".into()],
};
let diff = compare_capabilities(&source, &target);
assert!(!diff.tools_supported);
assert_eq!(diff.hidden_tools, vec!["tool1"]);
assert_eq!(diff.unsupported_features, vec!["reasoning_effort"]);
}
#[test]
fn test_model_capability_metadata_exposes_reasoning_effort() {
let meta = ModelCapabilityMetadata {
model: "o3".into(),
provider: "openai".into(),
context_window: 200000,
supports_tools: true,
supports_streaming: true,
supports_reasoning_effort: true,
reasoning_effort_values: vec!["low".into(), "medium".into(), "high".into()],
supported_modalities: vec!["text".into()],
unsupported_tools: vec![],
};
assert!(meta.supports_reasoning_effort);
assert_eq!(meta.reasoning_effort_values, vec!["low", "medium", "high"]);
}
#[test]
fn test_model_capability_registry() {
let mut registry = ModelCapabilityRegistry::new();
let meta = ModelCapabilityMetadata {
model: "gpt-4".into(),
provider: "openai".into(),
context_window: 8192,
supports_tools: true,
supports_streaming: true,
supports_reasoning_effort: false,
reasoning_effort_values: vec![],
supported_modalities: vec!["text".into()],
unsupported_tools: vec![],
};
registry.register(meta);
assert!(registry.get("openai", "gpt-4").is_some());
assert!(registry.get("openai", "gpt-3").is_none());
let diff = registry.compare("openai", "gpt-4", "openai", "gpt-4");
assert!(diff.is_some());
let diff = diff.unwrap();
assert_eq!(diff.window_shrink, None);
assert!(diff.tools_supported);
}
}