Skip to main content

roder_core/
inference_routing.rs

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 &registry.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}