1use std::fmt::Write;
2
3use anyhow::Result;
4use async_trait::async_trait;
5use futures::StreamExt;
6
7use crate::provider::LlmProvider;
8use crate::streaming::StreamBox;
9use agent_sdk_foundation::llm::{ChatOutcome, ChatRequest, ChatResponse, Message, Role};
10
11#[derive(Debug, Clone, Copy, PartialEq, Eq)]
17pub enum ModelTier {
18 Fast,
20 Capable,
22 Advanced,
24}
25
26#[derive(Debug, Clone, Copy, PartialEq, Eq)]
33pub enum TaskComplexity {
34 Simple,
36 Moderate,
38 Complex,
40}
41
42impl TaskComplexity {
43 #[must_use]
44 pub const fn recommended_tier(self) -> ModelTier {
45 match self {
46 Self::Simple => ModelTier::Fast,
47 Self::Moderate => ModelTier::Capable,
48 Self::Complex => ModelTier::Advanced,
49 }
50 }
51}
52
53pub struct ModelRouter<C, S, A> {
74 classifier: C,
75 fast: S,
76 capable: S,
77 advanced: A,
78}
79
80impl<C, S, A> ModelRouter<C, S, A>
81where
82 C: LlmProvider,
83 S: LlmProvider,
84 A: LlmProvider,
85{
86 pub const fn new(classifier: C, fast: S, capable: S, advanced: A) -> Self {
87 Self {
88 classifier,
89 fast,
90 capable,
91 advanced,
92 }
93 }
94
95 pub async fn classify(&self, request: &ChatRequest) -> Result<TaskComplexity> {
98 let classification_prompt = build_classification_prompt(request);
99
100 let classification_request = ChatRequest {
101 system: CLASSIFICATION_SYSTEM.to_owned(),
102 messages: vec![Message::user(classification_prompt)],
103 tools: None,
104 max_tokens: 50,
105 max_tokens_explicit: true,
106 session_id: None,
107 cached_content: None,
108 thinking: None,
109 tool_choice: None,
110 response_format: None,
111 cache: None,
112 };
113
114 match self.classifier.chat(classification_request).await? {
115 ChatOutcome::Success(response) => {
116 let complexity = parse_complexity(&response);
117 log::debug!(
118 "Model router classified request as {:?} using {}",
119 complexity,
120 self.classifier.model()
121 );
122 Ok(complexity)
123 }
124 ChatOutcome::RateLimited(_) => {
125 log::warn!("Classifier rate limited, defaulting to Complex");
126 Ok(TaskComplexity::Complex)
127 }
128 ChatOutcome::InvalidRequest(e) => {
129 log::error!("Classifier invalid request: {e}, defaulting to Complex");
130 Ok(TaskComplexity::Complex)
131 }
132 ChatOutcome::ServerError(e) => {
133 log::error!("Classifier server error: {e}, defaulting to Complex");
134 Ok(TaskComplexity::Complex)
135 }
136 _ => {
139 log::error!("Classifier returned unrecognized outcome, defaulting to Complex");
140 Ok(TaskComplexity::Complex)
141 }
142 }
143 }
144
145 pub async fn route(&self, request: ChatRequest) -> Result<ChatOutcome> {
148 let complexity = self.classify(&request).await?;
149 let tier = complexity.recommended_tier();
150
151 log::info!("Routing request to {tier:?} tier (complexity: {complexity:?})");
152
153 match tier {
154 ModelTier::Fast => self.fast.chat(request).await,
155 ModelTier::Capable => self.capable.chat(request).await,
156 ModelTier::Advanced => self.advanced.chat(request).await,
157 }
158 }
159
160 pub async fn route_with_tier(
163 &self,
164 request: ChatRequest,
165 tier: ModelTier,
166 ) -> Result<ChatOutcome> {
167 match tier {
168 ModelTier::Fast => self.fast.chat(request).await,
169 ModelTier::Capable => self.capable.chat(request).await,
170 ModelTier::Advanced => self.advanced.chat(request).await,
171 }
172 }
173
174 #[must_use]
175 pub const fn fast_provider(&self) -> &S {
176 &self.fast
177 }
178
179 #[must_use]
180 pub const fn capable_provider(&self) -> &S {
181 &self.capable
182 }
183
184 #[must_use]
185 pub const fn advanced_provider(&self) -> &A {
186 &self.advanced
187 }
188}
189
190#[async_trait]
191impl<C, S, A> LlmProvider for ModelRouter<C, S, A>
192where
193 C: LlmProvider,
194 S: LlmProvider,
195 A: LlmProvider,
196{
197 async fn chat(&self, request: ChatRequest) -> Result<ChatOutcome> {
198 self.route(request).await
199 }
200
201 fn chat_stream(&self, request: ChatRequest) -> StreamBox<'_> {
202 Box::pin(async_stream::stream! {
203 let tier = match self.classify(&request).await {
204 Ok(complexity) => complexity.recommended_tier(),
205 Err(error) => {
206 yield Err(error);
207 return;
208 }
209 };
210 log::info!("Streaming request to {tier:?} tier");
211 let mut stream = match tier {
212 ModelTier::Fast => self.fast.chat_stream(request),
213 ModelTier::Capable => self.capable.chat_stream(request),
214 ModelTier::Advanced => self.advanced.chat_stream(request),
215 };
216 while let Some(item) = stream.next().await {
217 yield item;
218 }
219 })
220 }
221
222 fn model(&self) -> &str {
225 self.capable.model()
226 }
227
228 fn provider(&self) -> &'static str {
231 self.capable.provider()
232 }
233
234 fn route(&self) -> &str {
240 self.capable.route()
241 }
242}
243
244const CLASSIFICATION_SYSTEM: &str = r"You are a task complexity classifier. Analyze the user's request and classify it as one of: SIMPLE, MODERATE, or COMPLEX.
245
246SIMPLE tasks:
247- Basic questions with factual answers
248- Simple calculations
249- Direct lookups or retrievals
250- Yes/no questions
251- Single-step operations
252
253MODERATE tasks:
254- Multi-step reasoning
255- Summarization
256- Basic analysis
257- Comparisons
258- Standard tool usage
259
260COMPLEX tasks:
261- Creative writing or content generation
262- Multi-step planning
263- Complex analysis or synthesis
264- Nuanced decisions
265- Tasks requiring deep domain knowledge
266- Financial advice or calculations
267- Multi-tool orchestration
268
269Respond with ONLY one word: SIMPLE, MODERATE, or COMPLEX.";
270
271fn build_classification_prompt(request: &ChatRequest) -> String {
272 let mut prompt = String::new();
273
274 prompt.push_str("Classify this task:\n\n");
275
276 if !request.system.is_empty() {
277 prompt.push_str("System context: ");
278 let truncated = truncate_on_char_boundary(&request.system, 200);
279 prompt.push_str(truncated);
280 if truncated.len() < request.system.len() {
281 prompt.push_str("...");
282 }
283 prompt.push_str("\n\n");
284 }
285
286 if let Some(last_user_message) = request.messages.iter().rev().find(|m| m.role == Role::User)
287 && let Some(text) = last_user_message.content.first_text()
288 {
289 prompt.push_str("User request: ");
290 let truncated = truncate_on_char_boundary(text, 500);
291 prompt.push_str(truncated);
292 if truncated.len() < text.len() {
293 prompt.push_str("...");
294 }
295 }
296
297 if let Some(tools) = &request.tools {
298 let _ = write!(prompt, "\n\nAvailable tools: {}", tools.len());
299 }
300
301 prompt
302}
303
304fn truncate_on_char_boundary(s: &str, max_bytes: usize) -> &str {
308 if s.len() <= max_bytes {
309 return s;
310 }
311 let mut end = max_bytes;
312 while end > 0 && !s.is_char_boundary(end) {
313 end -= 1;
314 }
315 &s[..end]
316}
317
318fn parse_complexity(response: &ChatResponse) -> TaskComplexity {
319 let text = response.first_text().unwrap_or("").to_uppercase();
320
321 if text.contains("SIMPLE") {
322 TaskComplexity::Simple
323 } else if text.contains("MODERATE") {
324 TaskComplexity::Moderate
325 } else {
326 TaskComplexity::Complex
327 }
328}
329
330#[cfg(test)]
331mod tests {
332 use super::*;
333 use crate::streaming::StreamDelta;
334 use agent_sdk_foundation::llm::{ContentBlock, StopReason, Usage};
335 use anyhow::Result;
336 use async_trait::async_trait;
337 use futures::StreamExt;
338
339 struct StaticProvider {
343 name: &'static str,
344 reply: &'static str,
345 }
346
347 #[async_trait]
348 impl LlmProvider for StaticProvider {
349 async fn chat(&self, _request: ChatRequest) -> Result<ChatOutcome> {
350 Ok(ChatOutcome::Success(ChatResponse {
351 id: "r".to_owned(),
352 content: vec![ContentBlock::Text {
353 text: self.reply.to_owned(),
354 }],
355 model: self.name.to_owned(),
356 stop_reason: Some(StopReason::EndTurn),
357 usage: Usage {
358 served_speed: None,
359 input_tokens: 1,
360 output_tokens: 1,
361 cached_input_tokens: 0,
362 cache_creation_input_tokens: 0,
363 },
364 }))
365 }
366
367 fn model(&self) -> &str {
368 self.name
369 }
370
371 fn provider(&self) -> &'static str {
372 self.name
373 }
374 }
375
376 #[tokio::test]
380 async fn streamed_dispatch_attributes_the_tier_that_served() -> Result<()> {
381 let router = ModelRouter::new(
382 StaticProvider {
383 name: "classifier",
384 reply: "SIMPLE",
385 },
386 StaticProvider {
387 name: "fast-tier",
388 reply: "quick answer",
389 },
390 StaticProvider {
391 name: "capable-tier",
392 reply: "unused",
393 },
394 StaticProvider {
395 name: "advanced-tier",
396 reply: "unused",
397 },
398 );
399 assert_eq!(LlmProvider::route(&router), "capable-tier");
400
401 let request = ChatRequest::new("system", vec![Message::user("2+2?")]);
402 let mut stream = router.chat_stream(request);
403 let mut served = None;
404 while let Some(item) = stream.next().await {
405 if let StreamDelta::Done { served_route, .. } = item? {
406 served = served_route;
407 }
408 }
409 assert_eq!(served.as_deref(), Some("fast-tier"));
410 Ok(())
411 }
412
413 #[test]
414 fn complexity_to_tier() {
415 assert_eq!(TaskComplexity::Simple.recommended_tier(), ModelTier::Fast);
416 assert_eq!(
417 TaskComplexity::Moderate.recommended_tier(),
418 ModelTier::Capable
419 );
420 assert_eq!(
421 TaskComplexity::Complex.recommended_tier(),
422 ModelTier::Advanced
423 );
424 }
425
426 #[test]
427 fn truncate_on_char_boundary_never_splits_multibyte_char() {
428 let s = "😀😀😀";
432 for max in 0..=s.len() {
433 let truncated = truncate_on_char_boundary(s, max);
434 assert!(s.starts_with(truncated));
436 assert!(truncated.len() <= max);
437 }
438 assert_eq!(truncate_on_char_boundary(s, 4), "😀");
439 assert_eq!(truncate_on_char_boundary(s, 5), "😀");
440 assert_eq!(truncate_on_char_boundary(s, 100), s);
441 }
442
443 #[test]
444 fn build_classification_prompt_handles_multibyte_at_limit() {
445 let system = "é".repeat(150); let request = ChatRequest::new(system, vec![Message::user("日本語".repeat(300))]);
449 let prompt = build_classification_prompt(&request);
451 assert!(prompt.contains("System context:"));
452 assert!(prompt.ends_with("..."));
453 }
454}