1use roder_api::extension::ExtensionRegistry;
2use roder_api::inference::{
3 InferenceProviderContext, ModelSelection, ReasoningConfig, RuntimeProfile, SpeedPolicyPhase,
4};
5use roder_api::inference_routing::{
6 InferenceRoutingCandidate, InferenceRoutingContext, InferenceRoutingDecision,
7 InferenceRoutingOutcome, InferenceRoutingSignal, InferenceRoutingToolSummary,
8 InferenceRoutingTranscriptSummary,
9};
10use roder_api::tools::ToolSpec;
11use roder_api::transcript::TranscriptItem;
12
13#[derive(Debug, Clone, Default, PartialEq, Eq)]
14pub struct RuntimeInferenceRouterConfig {
15 pub enabled: bool,
16 pub router_id: Option<String>,
17}
18
19impl RuntimeInferenceRouterConfig {
20 pub fn disabled() -> Self {
21 Self::default()
22 }
23
24 pub fn is_active(&self) -> bool {
25 self.enabled && self.router_id.is_some()
26 }
27}
28
29#[derive(Debug, Clone)]
30pub(crate) struct InferenceRoutingRequest<'a> {
31 pub(crate) thread_id: &'a str,
32 pub(crate) turn_id: &'a str,
33 pub(crate) round_index: u32,
34 pub(crate) runtime_profile: RuntimeProfile,
35 pub(crate) phase: SpeedPolicyPhase,
36 pub(crate) profile: Option<&'a str>,
37 pub(crate) default_selection: ModelSelection,
38 pub(crate) transcript: &'a [TranscriptItem],
39 pub(crate) tools: &'a [ToolSpec],
40 pub(crate) candidates: Option<&'a [InferenceRoutingCandidate]>,
41 pub(crate) prior_failures: u32,
42 pub(crate) prior_escalations: u32,
43}
44
45#[derive(Debug, Clone)]
46pub(crate) struct InferenceRoutingSelection {
47 pub(crate) selection: ModelSelection,
48 pub(crate) reasoning: Option<ReasoningConfig>,
49 pub(crate) decision: Option<InferenceRoutingDecision>,
50}
51
52impl InferenceRoutingSelection {
53 fn default(selection: ModelSelection) -> Self {
54 Self {
55 selection,
56 reasoning: None,
57 decision: None,
58 }
59 }
60}
61
62pub(crate) async fn route_inference_selection(
63 registry: &ExtensionRegistry,
64 config: &RuntimeInferenceRouterConfig,
65 request: InferenceRoutingRequest<'_>,
66) -> InferenceRoutingSelection {
67 if !config.is_active() {
68 return InferenceRoutingSelection::default(request.default_selection);
69 }
70
71 let router_id = config.router_id.as_deref().unwrap_or_default();
72 let Some(router) = registry.inference_router(router_id) else {
73 return fallback_selection(
74 request.default_selection,
75 router_id,
76 format!("inference router {router_id:?} is not registered"),
77 );
78 };
79
80 let collected_candidates;
81 let candidates = if let Some(candidates) = request.candidates {
82 candidates
83 } else {
84 collected_candidates = collect_inference_routing_candidates(registry).await;
85 &collected_candidates
86 };
87 let context = InferenceRoutingContext {
88 thread_id: request.thread_id.to_string(),
89 turn_id: request.turn_id.to_string(),
90 round_index: request.round_index,
91 runtime_profile: request.runtime_profile,
92 default_selection: request.default_selection.clone(),
93 requested_selection: None,
94 phase: Some(request.phase),
95 transcript: transcript_summary(request.transcript),
96 tools: tool_summary(request.tools),
97 candidates: candidates.to_vec(),
98 signals: initial_signals(request.phase, request.profile),
99 prior_failures: request.prior_failures,
100 prior_escalations: request.prior_escalations,
101 estimated_input_tokens: approximate_transcript_tokens(request.transcript),
102 };
103
104 let decision = match router.route(context).await {
105 Ok(decision) => decision,
106 Err(err) => {
107 return fallback_selection(
108 request.default_selection,
109 router_id,
110 format!("inference router {router_id:?} failed: {err}"),
111 );
112 }
113 };
114
115 apply_decision(
116 request.default_selection,
117 request.transcript,
118 candidates,
119 decision,
120 )
121}
122
123fn apply_decision(
124 default_selection: ModelSelection,
125 transcript: &[TranscriptItem],
126 candidates: &[InferenceRoutingCandidate],
127 decision: InferenceRoutingDecision,
128) -> InferenceRoutingSelection {
129 match decision.outcome {
130 InferenceRoutingOutcome::Selected | InferenceRoutingOutcome::Escalated => {
131 let Some(selected) = decision.selected.clone() else {
132 return fallback_selection(
133 default_selection,
134 decision.router_id,
135 "router selected no provider/model",
136 );
137 };
138 let Some(candidate) = candidate_for(candidates, &selected) else {
139 return fallback_selection(
140 default_selection,
141 decision.router_id,
142 format!(
143 "router selected unavailable provider/model {}/{}",
144 selected.provider, selected.model
145 ),
146 );
147 };
148 if let Some(reason) = invalid_candidate_reason(candidate, transcript) {
149 return fallback_selection(default_selection, decision.router_id, reason);
150 }
151 let reasoning = decision
152 .reasoning
153 .as_ref()
154 .filter(|reasoning| reasoning_supported(candidate, reasoning))
155 .cloned();
156 InferenceRoutingSelection {
157 selection: selected,
158 reasoning,
159 decision: Some(decision),
160 }
161 }
162 InferenceRoutingOutcome::Abstained | InferenceRoutingOutcome::Fallback => {
163 InferenceRoutingSelection {
164 selection: default_selection,
165 reasoning: None,
166 decision: Some(decision),
167 }
168 }
169 }
170}
171
172fn fallback_selection(
173 default_selection: ModelSelection,
174 router_id: impl Into<String>,
175 reason: impl Into<String>,
176) -> InferenceRoutingSelection {
177 InferenceRoutingSelection {
178 selection: default_selection,
179 reasoning: None,
180 decision: Some(InferenceRoutingDecision::fallback(router_id, reason)),
181 }
182}
183
184pub async fn collect_inference_routing_candidates(
185 registry: &ExtensionRegistry,
186) -> Vec<InferenceRoutingCandidate> {
187 let mut candidates = Vec::new();
188 for engine in ®istry.inference_engines {
189 let provider_id = engine.id();
190 let provider = engine.metadata();
191 let capabilities = engine.capabilities();
192 let auth_available = provider.auth_configured.unwrap_or(true);
193 let models = engine
194 .list_models(InferenceProviderContext {
195 provider_id: &provider_id,
196 })
197 .await
198 .unwrap_or_default();
199 for model in models {
200 let selection = ModelSelection {
201 provider: provider_id.clone(),
202 model: model.id.clone(),
203 };
204 if auth_available {
205 candidates.push(InferenceRoutingCandidate::available(
206 selection,
207 provider.clone(),
208 model,
209 capabilities.clone(),
210 ));
211 } else {
212 candidates.push(InferenceRoutingCandidate::unavailable(
213 selection,
214 provider.clone(),
215 model,
216 capabilities.clone(),
217 "provider authentication is not configured",
218 ));
219 }
220 }
221 }
222 candidates
223}
224
225fn transcript_summary(transcript: &[TranscriptItem]) -> InferenceRoutingTranscriptSummary {
226 InferenceRoutingTranscriptSummary {
227 item_count: transcript.len() as u32,
228 user_message_count: transcript
229 .iter()
230 .filter(|item| matches!(item, TranscriptItem::UserMessage(_)))
231 .count() as u32,
232 assistant_message_count: transcript
233 .iter()
234 .filter(|item| matches!(item, TranscriptItem::AssistantMessage(_)))
235 .count() as u32,
236 tool_result_count: transcript
237 .iter()
238 .filter(|item| matches!(item, TranscriptItem::ToolResult(_)))
239 .count() as u32,
240 has_image_input: transcript_has_images(transcript),
241 latest_user_message_preview: latest_user_message_preview(transcript),
242 recent_tool_names: recent_tool_names(transcript),
243 approximate_tokens: approximate_transcript_tokens(transcript),
244 }
245}
246
247fn tool_summary(tools: &[ToolSpec]) -> InferenceRoutingToolSummary {
248 InferenceRoutingToolSummary {
249 available_count: tools.len() as u32,
250 has_file_tools: tools.iter().any(|tool| {
251 tool.name.contains("file") || tool.name.contains("read") || tool.name.contains("write")
252 }),
253 has_shell_tools: tools
254 .iter()
255 .any(|tool| tool.name.contains("shell") || tool.name.contains("command")),
256 has_network_tools: tools
257 .iter()
258 .any(|tool| tool.name.contains("web") || tool.name.contains("search")),
259 requires_tool_calls: !tools.is_empty(),
260 }
261}
262
263fn initial_signals(phase: SpeedPolicyPhase, profile: Option<&str>) -> Vec<InferenceRoutingSignal> {
264 let mut signals = vec![InferenceRoutingSignal::new("phase", phase.as_str())];
265 if let Some(profile) = profile {
266 signals.push(InferenceRoutingSignal::new("profile", profile));
267 }
268 signals
269}
270
271fn candidate_for<'a>(
272 candidates: &'a [InferenceRoutingCandidate],
273 selection: &ModelSelection,
274) -> Option<&'a InferenceRoutingCandidate> {
275 candidates.iter().find(|candidate| {
276 candidate.selection.provider == selection.provider
277 && candidate.selection.model == selection.model
278 })
279}
280
281fn invalid_candidate_reason(
282 candidate: &InferenceRoutingCandidate,
283 transcript: &[TranscriptItem],
284) -> Option<String> {
285 inference_routing_candidate_unavailable_reason(candidate, transcript_has_images(transcript))
286}
287
288pub fn inference_routing_selection_unavailable_reason(
289 candidates: &[InferenceRoutingCandidate],
290 selection: &ModelSelection,
291 has_image_input: bool,
292) -> Option<String> {
293 let Some(candidate) = candidate_for(candidates, selection) else {
294 return Some(format!(
295 "router selected unavailable provider/model {}/{}",
296 selection.provider, selection.model
297 ));
298 };
299 inference_routing_candidate_unavailable_reason(candidate, has_image_input)
300}
301
302pub fn inference_routing_candidate_unavailable_reason(
303 candidate: &InferenceRoutingCandidate,
304 has_image_input: bool,
305) -> Option<String> {
306 if !candidate.available {
307 return Some(candidate.unavailable_reason.clone().unwrap_or_else(|| {
308 format!(
309 "router selected unavailable provider/model {}/{}",
310 candidate.selection.provider, candidate.selection.model
311 )
312 }));
313 }
314 if has_image_input && !candidate.capabilities.image_input {
315 return Some(format!(
316 "router selected {}/{} but it does not support image input",
317 candidate.selection.provider, candidate.selection.model
318 ));
319 }
320 None
321}
322
323fn reasoning_supported(candidate: &InferenceRoutingCandidate, reasoning: &ReasoningConfig) -> bool {
324 if !reasoning.enabled {
325 return true;
326 }
327 let Some(level) = reasoning.level.as_deref() else {
328 return true;
329 };
330 candidate
331 .model
332 .supported_reasoning
333 .iter()
334 .any(|option| option.effort == level)
335}
336
337fn transcript_has_images(transcript: &[TranscriptItem]) -> bool {
338 transcript.iter().any(|item| {
339 matches!(
340 item,
341 TranscriptItem::UserMessage(message) if !message.images.is_empty()
342 )
343 })
344}
345
346fn approximate_transcript_tokens(transcript: &[TranscriptItem]) -> Option<u32> {
347 let bytes = transcript
348 .iter()
349 .map(transcript_item_text_len)
350 .sum::<usize>();
351 Some((bytes / 4).max(transcript.len()) as u32)
352}
353
354pub(crate) fn transcript_failure_count(transcript: &[TranscriptItem]) -> u32 {
355 transcript
356 .iter()
357 .filter(|item| match item {
358 TranscriptItem::ToolResult(result) => result.is_error,
359 TranscriptItem::Error(_) => true,
360 _ => false,
361 })
362 .count() as u32
363}
364
365pub(crate) fn transcript_failure_count_since(
366 transcript: &[TranscriptItem],
367 start_index: usize,
368) -> u32 {
369 transcript_failure_count(
370 transcript
371 .get(start_index.min(transcript.len())..)
372 .unwrap_or(&[]),
373 )
374}
375
376fn latest_user_message_preview(transcript: &[TranscriptItem]) -> Option<String> {
377 transcript.iter().rev().find_map(|item| {
378 let TranscriptItem::UserMessage(message) = item else {
379 return None;
380 };
381 let text = message.text.trim();
382 (!text.is_empty()).then(|| truncate_chars(text, 600))
383 })
384}
385
386fn recent_tool_names(transcript: &[TranscriptItem]) -> Vec<String> {
387 transcript
388 .iter()
389 .rev()
390 .filter_map(|item| match item {
391 TranscriptItem::ToolCall(call) => Some(call.name.clone()),
392 TranscriptItem::ToolResult(result) => result.name.clone(),
393 _ => None,
394 })
395 .take(12)
396 .collect()
397}
398
399fn truncate_chars(text: &str, max_chars: usize) -> String {
400 text.chars().take(max_chars).collect()
401}
402
403fn transcript_item_text_len(item: &TranscriptItem) -> usize {
404 match item {
405 TranscriptItem::UserMessage(message) => message.text.len(),
406 TranscriptItem::AssistantMessage(message) => message.text.len(),
407 TranscriptItem::ReasoningSummary(summary) => summary.text.len(),
408 TranscriptItem::ToolCall(call) => call.arguments.len(),
409 TranscriptItem::ToolResult(result) => result.result.len(),
410 TranscriptItem::FileChange(change) => change.path.len() + change.change_type.len(),
411 TranscriptItem::ContextCompaction(compaction) => compaction.summary.len(),
412 TranscriptItem::Error(error) => error.message.len(),
413 TranscriptItem::ProviderMetadata(metadata) => metadata.to_string().len(),
414 }
415}
416
417#[cfg(test)]
418mod tests {
419 use roder_api::transcript::{ToolResultRecord, UserMessage};
420
421 use super::*;
422
423 #[test]
424 fn transcript_failure_count_since_ignores_prior_history() {
425 let transcript = vec![
426 TranscriptItem::ToolResult(tool_result("old", true)),
427 TranscriptItem::UserMessage(UserMessage::text("current turn")),
428 TranscriptItem::ToolResult(tool_result("ok", false)),
429 TranscriptItem::ToolResult(tool_result("new", true)),
430 ];
431
432 assert_eq!(transcript_failure_count_since(&transcript, 1), 1);
433 assert_eq!(transcript_failure_count_since(&transcript, 0), 2);
434 }
435
436 fn tool_result(id: &str, is_error: bool) -> ToolResultRecord {
437 ToolResultRecord {
438 id: id.to_string(),
439 name: Some("test".to_string()),
440 result: if is_error { "error" } else { "ok" }.to_string(),
441 display_payload: None,
442 is_error,
443 }
444 }
445}