Skip to main content

roder_api/
inference.rs

1use std::pin::Pin;
2
3use futures::Stream;
4use serde::{Deserialize, Serialize};
5
6use crate::extension::InferenceEngineId;
7use crate::lifecycle::TurnCleanupOwnership;
8use crate::reliability::ReliabilityRequestPolicy;
9use crate::tools::{ToolChoice, ToolSpec};
10use crate::transcript::TranscriptItem;
11
12#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
13pub struct ModelSelection {
14    pub provider: String,
15    pub model: String,
16}
17
18#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
19#[serde(rename_all = "snake_case")]
20pub enum ProviderAuthType {
21    None,
22    ApiKey,
23    OAuth,
24}
25
26#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
27pub struct InferenceProviderMetadata {
28    pub name: String,
29    pub description: Option<String>,
30    pub auth_type: ProviderAuthType,
31    pub auth_label: Option<String>,
32    pub auth_configured: Option<bool>,
33    pub recommended: bool,
34    pub sort_order: i32,
35}
36
37impl InferenceProviderMetadata {
38    pub fn local(name: impl Into<String>) -> Self {
39        Self {
40            name: name.into(),
41            description: None,
42            auth_type: ProviderAuthType::None,
43            auth_label: None,
44            auth_configured: Some(true),
45            recommended: false,
46            sort_order: 100,
47        }
48    }
49}
50
51#[derive(Debug, Clone, Copy, Default, Serialize, Deserialize, PartialEq, Eq)]
52#[serde(rename_all = "snake_case")]
53pub enum ToolSearchMode {
54    #[default]
55    Explicit,
56    Auto,
57    ProviderNative,
58}
59
60impl ToolSearchMode {
61    pub fn allows_provider_native(self) -> bool {
62        matches!(self, Self::Auto | Self::ProviderNative)
63    }
64}
65
66#[derive(Debug, Clone, Copy, Default, Serialize, Deserialize, PartialEq, Eq)]
67#[serde(rename_all = "snake_case")]
68pub enum ToolSearchProviderVariant {
69    #[default]
70    Default,
71    Regex,
72    Bm25,
73}
74
75#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
76#[serde(rename_all = "camelCase")]
77pub struct ToolSearchConfig {
78    #[serde(default)]
79    pub mode: ToolSearchMode,
80    #[serde(default, skip_serializing_if = "Option::is_none")]
81    pub max_catalog_items: Option<u32>,
82    #[serde(default)]
83    pub include_mcp: bool,
84    #[serde(default)]
85    pub include_skills: bool,
86    #[serde(default)]
87    pub fallback_to_explicit_tools: bool,
88    #[serde(default)]
89    pub provider_variant: ToolSearchProviderVariant,
90}
91
92impl Default for ToolSearchConfig {
93    fn default() -> Self {
94        Self {
95            mode: ToolSearchMode::Explicit,
96            max_catalog_items: None,
97            include_mcp: true,
98            include_skills: true,
99            fallback_to_explicit_tools: true,
100            provider_variant: ToolSearchProviderVariant::Default,
101        }
102    }
103}
104
105impl ToolSearchConfig {
106    pub fn explicit() -> Self {
107        Self {
108            mode: ToolSearchMode::Explicit,
109            ..Self::default()
110        }
111    }
112
113    pub fn provider_native() -> Self {
114        Self {
115            mode: ToolSearchMode::ProviderNative,
116            ..Self::default()
117        }
118    }
119
120    pub fn is_provider_native_requested(&self) -> bool {
121        self.mode.allows_provider_native()
122    }
123
124    /**
125     * Resolve the effective tool-search mode for one provider/model turn.
126     *
127     * `Auto` silently falls back to explicit tools when the provider/model
128     * does not support native tool search. An explicit `ProviderNative`
129     * request only falls back when `fallback_to_explicit_tools` allows it;
130     * otherwise the turn must fail closed with the returned diagnostic.
131     */
132    pub fn resolve_effective_mode(
133        &self,
134        provider_native_supported: bool,
135    ) -> Result<EffectiveToolSearchMode, ToolSearchModeError> {
136        match self.mode {
137            ToolSearchMode::Explicit => Ok(EffectiveToolSearchMode::Explicit),
138            ToolSearchMode::Auto => {
139                if provider_native_supported {
140                    Ok(EffectiveToolSearchMode::ProviderNative)
141                } else {
142                    Ok(EffectiveToolSearchMode::Explicit)
143                }
144            }
145            ToolSearchMode::ProviderNative => {
146                if provider_native_supported {
147                    Ok(EffectiveToolSearchMode::ProviderNative)
148                } else if self.fallback_to_explicit_tools {
149                    Ok(EffectiveToolSearchMode::Explicit)
150                } else {
151                    Err(ToolSearchModeError::ProviderNativeUnsupported)
152                }
153            }
154        }
155    }
156}
157
158#[derive(Debug, Clone, Copy, Serialize, Deserialize, PartialEq, Eq)]
159#[serde(rename_all = "snake_case")]
160pub enum EffectiveToolSearchMode {
161    Explicit,
162    ProviderNative,
163}
164
165#[derive(Debug, Clone, Copy, PartialEq, Eq)]
166pub enum ToolSearchModeError {
167    ProviderNativeUnsupported,
168}
169
170impl std::fmt::Display for ToolSearchModeError {
171    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
172        match self {
173            Self::ProviderNativeUnsupported => write!(
174                f,
175                "provider-native tool search was requested but the selected provider/model does \
176                 not support it and fallback_to_explicit_tools is disabled; enable fallback or \
177                 pick a supported model"
178            ),
179        }
180    }
181}
182
183impl std::error::Error for ToolSearchModeError {}
184
185#[derive(Debug, Clone, Default, Serialize, Deserialize, PartialEq, Eq)]
186#[serde(rename_all = "camelCase")]
187pub struct ToolSearchConfigOverlay {
188    #[serde(default, skip_serializing_if = "Option::is_none")]
189    pub mode: Option<ToolSearchMode>,
190    #[serde(default, skip_serializing_if = "Option::is_none")]
191    pub max_catalog_items: Option<u32>,
192    #[serde(default, skip_serializing_if = "Option::is_none")]
193    pub include_mcp: Option<bool>,
194    #[serde(default, skip_serializing_if = "Option::is_none")]
195    pub include_skills: Option<bool>,
196    #[serde(default, skip_serializing_if = "Option::is_none")]
197    pub fallback_to_explicit_tools: Option<bool>,
198    #[serde(default, skip_serializing_if = "Option::is_none")]
199    pub provider_variant: Option<ToolSearchProviderVariant>,
200}
201
202impl ToolSearchConfigOverlay {
203    pub fn overlay(&mut self, other: &Self) {
204        if other.mode.is_some() {
205            self.mode = other.mode;
206        }
207        if other.max_catalog_items.is_some() {
208            self.max_catalog_items = other.max_catalog_items;
209        }
210        if other.include_mcp.is_some() {
211            self.include_mcp = other.include_mcp;
212        }
213        if other.include_skills.is_some() {
214            self.include_skills = other.include_skills;
215        }
216        if other.fallback_to_explicit_tools.is_some() {
217            self.fallback_to_explicit_tools = other.fallback_to_explicit_tools;
218        }
219        if other.provider_variant.is_some() {
220            self.provider_variant = other.provider_variant;
221        }
222    }
223
224    pub fn apply_to(&self, config: &mut ToolSearchConfig) {
225        if let Some(mode) = self.mode {
226            config.mode = mode;
227        }
228        if let Some(max_catalog_items) = self.max_catalog_items {
229            config.max_catalog_items = Some(max_catalog_items);
230        }
231        if let Some(include_mcp) = self.include_mcp {
232            config.include_mcp = include_mcp;
233        }
234        if let Some(include_skills) = self.include_skills {
235            config.include_skills = include_skills;
236        }
237        if let Some(fallback_to_explicit_tools) = self.fallback_to_explicit_tools {
238            config.fallback_to_explicit_tools = fallback_to_explicit_tools;
239        }
240        if let Some(provider_variant) = self.provider_variant {
241            config.provider_variant = provider_variant;
242        }
243    }
244}
245
246#[derive(Debug, Clone, Default, Serialize, Deserialize, PartialEq, Eq)]
247pub struct InstructionBundle {
248    pub system: Option<String>,
249    pub developer: Option<String>,
250    /**
251     * Per-turn developer-authority context supplied on turn/start. Volatile:
252     * providers must render it after `system` and `developer` so prompt-cache
253     * breakpoints on the stable prefix survive per-turn changes. Never
254     * persisted to thread state.
255     */
256    pub developer_context: Option<String>,
257}
258
259#[derive(Debug, Clone, Copy, Default, Serialize, Deserialize, PartialEq, Eq)]
260#[serde(rename_all = "snake_case")]
261pub enum RuntimeProfile {
262    #[default]
263    Interactive,
264    NonInteractive,
265    Eval,
266}
267
268impl RuntimeProfile {
269    pub fn as_str(self) -> &'static str {
270        match self {
271            Self::Interactive => "interactive",
272            Self::NonInteractive => "non_interactive",
273            Self::Eval => "eval",
274        }
275    }
276
277    pub fn is_non_interactive(self) -> bool {
278        matches!(self, Self::NonInteractive | Self::Eval)
279    }
280}
281
282impl std::str::FromStr for RuntimeProfile {
283    type Err = anyhow::Error;
284
285    fn from_str(value: &str) -> Result<Self, Self::Err> {
286        match value.trim().to_ascii_lowercase().as_str() {
287            "interactive" => Ok(Self::Interactive),
288            "non_interactive" | "non-interactive" | "headless" => Ok(Self::NonInteractive),
289            "eval" => Ok(Self::Eval),
290            other => anyhow::bail!(
291                "unsupported runtime profile {other:?}; expected interactive, non_interactive, or eval"
292            ),
293        }
294    }
295}
296
297#[derive(Debug, Clone, Default, Serialize, Deserialize, PartialEq, Eq)]
298pub struct ReasoningConfig {
299    pub enabled: bool,
300    pub level: Option<String>,
301}
302
303#[derive(Debug, Clone, Copy, Default, Serialize, Deserialize, PartialEq, Eq)]
304#[serde(rename_all = "snake_case")]
305pub enum ProviderFamily {
306    #[default]
307    Mock,
308    OpenAi,
309    Anthropic,
310    Gemini,
311    Xai,
312    Opencode,
313    Poolside,
314    Cursor,
315}
316
317#[derive(Debug, Clone, Copy, Default, Serialize, Deserialize, PartialEq, Eq)]
318#[serde(rename_all = "snake_case")]
319pub enum ModelSchemaPolicy {
320    #[default]
321    StandardRequiredFirst,
322    RequiredFirstFlat,
323}
324
325#[derive(Debug, Clone, Copy, Default, Serialize, Deserialize, PartialEq, Eq)]
326#[serde(rename_all = "snake_case")]
327pub enum ModelInstructionOverlay {
328    #[default]
329    Standard,
330    LiteralToolOutputs,
331    IntuitiveContext,
332}
333
334#[derive(Debug, Clone, Default, Serialize, Deserialize, PartialEq, Eq)]
335#[serde(rename_all = "camelCase")]
336pub struct ModelProfileReasoning {
337    #[serde(default, skip_serializing_if = "Option::is_none")]
338    pub orientation: Option<String>,
339    #[serde(default, skip_serializing_if = "Option::is_none")]
340    pub execution: Option<String>,
341    #[serde(default, skip_serializing_if = "Option::is_none")]
342    pub verification: Option<String>,
343    #[serde(default, skip_serializing_if = "Option::is_none")]
344    pub recovery: Option<String>,
345}
346
347#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
348#[serde(rename_all = "camelCase")]
349pub struct ModelHarnessProfile {
350    pub model: String,
351    pub provider: String,
352    pub provider_family: ProviderFamily,
353    #[serde(default, skip_serializing_if = "Option::is_none")]
354    pub edit_tool: Option<String>,
355    #[serde(default)]
356    pub schema_policy: ModelSchemaPolicy,
357    #[serde(default)]
358    pub instruction_overlay: ModelInstructionOverlay,
359    #[serde(default)]
360    pub reasoning: ModelProfileReasoning,
361    #[serde(default, skip_serializing_if = "Option::is_none")]
362    pub parallel_tool_calls: Option<bool>,
363    #[serde(default, skip_serializing_if = "Option::is_none")]
364    pub auto_compact_token_limit: Option<u32>,
365}
366
367#[derive(Debug, Clone, Copy, Default, Serialize, Deserialize, PartialEq, Eq)]
368#[serde(rename_all = "snake_case")]
369pub enum SpeedPolicyPhase {
370    #[default]
371    Orientation,
372    Execution,
373    Verification,
374    Recovery,
375}
376
377impl SpeedPolicyPhase {
378    pub fn as_str(self) -> &'static str {
379        match self {
380            Self::Orientation => "orientation",
381            Self::Execution => "execution",
382            Self::Verification => "verification",
383            Self::Recovery => "recovery",
384        }
385    }
386}
387
388#[derive(Debug, Clone, Default, Serialize, Deserialize, PartialEq, Eq)]
389#[serde(rename_all = "camelCase")]
390pub struct SpeedPolicyDecision {
391    pub phase: SpeedPolicyPhase,
392    pub desired_reasoning: String,
393    pub applied_reasoning: Option<String>,
394    pub supported: bool,
395}
396
397#[derive(Debug, Clone, Default, Serialize, Deserialize, PartialEq)]
398pub struct OutputConfig {
399    pub max_tokens: Option<u32>,
400    pub temperature: Option<f32>,
401    pub top_p: Option<f32>,
402    pub response_format: Option<serde_json::Value>,
403}
404
405#[derive(Debug, Clone, Copy, Default, Serialize, Deserialize, PartialEq, Eq)]
406#[serde(rename_all = "snake_case")]
407pub enum HostedWebSearchMode {
408    #[default]
409    Disabled,
410    Cached,
411    Live,
412}
413
414#[derive(Debug, Clone, Default, Serialize, Deserialize, PartialEq, Eq)]
415pub struct HostedWebSearchConfig {
416    pub mode: HostedWebSearchMode,
417}
418
419impl HostedWebSearchConfig {
420    pub fn disabled() -> Self {
421        Self {
422            mode: HostedWebSearchMode::Disabled,
423        }
424    }
425
426    pub fn cached() -> Self {
427        Self {
428            mode: HostedWebSearchMode::Cached,
429        }
430    }
431
432    pub fn live() -> Self {
433        Self {
434            mode: HostedWebSearchMode::Live,
435        }
436    }
437
438    pub fn is_enabled(&self) -> bool {
439        self.mode != HostedWebSearchMode::Disabled
440    }
441}
442
443#[derive(Debug, Clone, Default, Serialize, Deserialize, PartialEq, Eq)]
444pub struct RuntimeHints {
445    pub trace_id: Option<String>,
446    pub prompt_cache_key: Option<String>,
447    pub auto_compact_token_limit: Option<u32>,
448    #[serde(default)]
449    pub profile: RuntimeProfile,
450    #[serde(default, skip_serializing_if = "Option::is_none")]
451    pub parallel_tool_calls: Option<bool>,
452    #[serde(default)]
453    pub hosted_web_search: HostedWebSearchConfig,
454    #[serde(default)]
455    pub tool_search: ToolSearchConfig,
456    #[serde(default, skip_serializing_if = "Option::is_none")]
457    pub speed_policy: Option<SpeedPolicyDecision>,
458    #[serde(default, skip_serializing_if = "Option::is_none")]
459    pub deadline_remaining_seconds: Option<u64>,
460    #[serde(default, skip_serializing_if = "Option::is_none")]
461    pub reliability: Option<ReliabilityRequestPolicy>,
462}
463
464#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
465pub struct AgentInferenceRequest {
466    pub model: ModelSelection,
467    pub instructions: InstructionBundle,
468    pub transcript: Vec<TranscriptItem>,
469    pub tools: Vec<ToolSpec>,
470    pub tool_choice: ToolChoice,
471    pub reasoning: ReasoningConfig,
472    pub output: OutputConfig,
473    pub runtime: RuntimeHints,
474    pub metadata: serde_json::Value,
475}
476
477#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
478pub struct MessageDelta {
479    pub text: String,
480    #[serde(default, skip_serializing_if = "Option::is_none")]
481    pub phase: Option<String>,
482}
483
484#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
485pub struct ReasoningDelta {
486    pub text: String,
487}
488
489#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
490pub struct ToolCallStarted {
491    pub id: String,
492    pub name: String,
493}
494
495#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
496pub struct ToolCallDelta {
497    pub id: String,
498    pub arguments_delta: String,
499}
500
501#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
502pub struct ToolCallCompleted {
503    pub id: String,
504    pub name: String,
505    pub arguments: String,
506}
507
508#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
509pub struct HostedToolCallStarted {
510    pub id: String,
511    pub name: String,
512}
513
514#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
515pub struct HostedToolCallCompleted {
516    pub id: String,
517    pub name: String,
518    pub arguments: String,
519}
520
521#[derive(Debug, Clone, Default, Serialize, Deserialize, PartialEq)]
522pub struct TokenUsage {
523    pub prompt_tokens: u32,
524    pub completion_tokens: u32,
525    pub total_tokens: u32,
526    #[serde(default)]
527    pub cached_prompt_tokens: u32,
528    /**
529     * Prompt tokens written to the provider's prompt cache this step. Like
530     * `cached_prompt_tokens`, this is a subset of `prompt_tokens`, not an
531     * additional count; hosts use it to bill cache writes at the provider's
532     * cache-write rate.
533     */
534    #[serde(default)]
535    pub cache_creation_prompt_tokens: u32,
536    #[serde(default, skip_serializing_if = "Option::is_none")]
537    pub cache_hit_rate: Option<f64>,
538}
539
540impl TokenUsage {
541    pub fn new(prompt_tokens: u32, completion_tokens: u32, total_tokens: u32) -> Self {
542        Self {
543            prompt_tokens,
544            completion_tokens,
545            total_tokens,
546            cached_prompt_tokens: 0,
547            cache_creation_prompt_tokens: 0,
548            cache_hit_rate: cache_hit_rate(prompt_tokens, 0),
549        }
550    }
551
552    pub fn with_cached_prompt_tokens(mut self, cached_prompt_tokens: u32) -> Self {
553        self.cached_prompt_tokens = cached_prompt_tokens.min(self.prompt_tokens);
554        self.cache_hit_rate = cache_hit_rate(self.prompt_tokens, self.cached_prompt_tokens);
555        self
556    }
557
558    pub fn with_cache_creation_prompt_tokens(mut self, cache_creation_prompt_tokens: u32) -> Self {
559        self.cache_creation_prompt_tokens = cache_creation_prompt_tokens.min(self.prompt_tokens);
560        self
561    }
562
563    pub fn add_assign(&mut self, usage: &TokenUsage) {
564        self.prompt_tokens = self.prompt_tokens.saturating_add(usage.prompt_tokens);
565        self.completion_tokens = self
566            .completion_tokens
567            .saturating_add(usage.completion_tokens);
568        self.total_tokens = self.total_tokens.saturating_add(usage.total_tokens);
569        self.cached_prompt_tokens = self
570            .cached_prompt_tokens
571            .saturating_add(usage.cached_prompt_tokens);
572        self.cache_creation_prompt_tokens = self
573            .cache_creation_prompt_tokens
574            .saturating_add(usage.cache_creation_prompt_tokens);
575        self.cache_hit_rate = cache_hit_rate(self.prompt_tokens, self.cached_prompt_tokens);
576    }
577
578    pub fn is_empty(&self) -> bool {
579        self.prompt_tokens == 0
580            && self.completion_tokens == 0
581            && self.total_tokens == 0
582            && self.cached_prompt_tokens == 0
583            && self.cache_creation_prompt_tokens == 0
584    }
585}
586
587pub fn cache_hit_rate(prompt_tokens: u32, cached_prompt_tokens: u32) -> Option<f64> {
588    if prompt_tokens == 0 {
589        None
590    } else {
591        Some(f64::from(cached_prompt_tokens.min(prompt_tokens)) / f64::from(prompt_tokens))
592    }
593}
594
595#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
596pub struct CompletionMetadata {
597    pub stop_reason: Option<String>,
598    pub provider_response_id: Option<String>,
599}
600
601/**
602 * Canonical mapping from provider-native stop reasons to the finish reason
603 * surfaced as `finishReason` on `turn/completed`. Only the terminal inference
604 * step's stop reason reaches the turn surface, so `toolUse` appears only when
605 * a turn genuinely ends on a tool-use step (e.g. tool rounds exhausted).
606 * Unknown stop reasons pass through unchanged.
607 */
608pub fn finish_reason_from_stop_reason(stop_reason: &str) -> String {
609    match stop_reason {
610        "end_turn" | "stop" | "stop_sequence" => "stop",
611        "max_tokens" | "length" => "length",
612        "tool_use" | "tool_calls" => "toolUse",
613        "content_filter" => "contentFilter",
614        "refusal" => "refusal",
615        other => other,
616    }
617    .to_string()
618}
619
620#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
621pub struct InferenceFailure {
622    pub message: String,
623}
624
625#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
626pub struct CompactionProgress {
627    pub status: String,
628    #[serde(default, skip_serializing_if = "Option::is_none")]
629    pub item_id: Option<String>,
630    /// Estimated prompt tokens before this compaction pass, when known.
631    #[serde(default, skip_serializing_if = "Option::is_none")]
632    pub tokens_before: Option<u32>,
633    /// Estimated prompt tokens after this compaction pass, when known.
634    #[serde(default, skip_serializing_if = "Option::is_none")]
635    pub tokens_after: Option<u32>,
636    /// Wall-clock duration of the compaction pass, when measured.
637    #[serde(default, skip_serializing_if = "Option::is_none")]
638    pub duration_ms: Option<u64>,
639    /// Opaque provider compaction output item (e.g. OpenAI `type: "compaction"`).
640    /// Runtime persists this as a transcript boundary even if the stream dies
641    /// before the full ProviderMetadata/response.completed frame arrives.
642    #[serde(default, skip_serializing_if = "Option::is_none")]
643    pub item: Option<serde_json::Value>,
644}
645
646#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
647pub enum InferenceEvent {
648    MessageDelta(MessageDelta),
649    ReasoningDelta(ReasoningDelta),
650    ToolCallStarted(ToolCallStarted),
651    ToolCallDelta(ToolCallDelta),
652    ToolCallCompleted(ToolCallCompleted),
653    HostedToolCallStarted(HostedToolCallStarted),
654    HostedToolCallCompleted(HostedToolCallCompleted),
655    Compaction(CompactionProgress),
656    Usage(TokenUsage),
657    Completed(CompletionMetadata),
658    Failed(InferenceFailure),
659    ProviderMetadata(serde_json::Value),
660}
661
662pub type InferenceEventStream =
663    Pin<Box<dyn Stream<Item = anyhow::Result<InferenceEvent>> + Send + 'static>>;
664
665#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
666pub struct InferenceCapabilities {
667    pub streaming: bool,
668    pub tool_calls: bool,
669    pub parallel_tool_calls: bool,
670    pub reasoning_summaries: bool,
671    pub structured_output: bool,
672    pub image_input: bool,
673    pub prompt_cache: bool,
674    pub provider_metadata: bool,
675    pub tool_search: bool,
676}
677
678impl InferenceCapabilities {
679    pub fn text_only() -> Self {
680        Self {
681            streaming: true,
682            tool_calls: false,
683            parallel_tool_calls: false,
684            reasoning_summaries: false,
685            structured_output: false,
686            image_input: false,
687            prompt_cache: false,
688            provider_metadata: false,
689            tool_search: false,
690        }
691    }
692
693    pub fn coding_agent_default() -> Self {
694        Self {
695            streaming: true,
696            tool_calls: true,
697            parallel_tool_calls: true,
698            reasoning_summaries: false,
699            structured_output: false,
700            image_input: false,
701            prompt_cache: false,
702            provider_metadata: true,
703            tool_search: false,
704        }
705    }
706}
707
708#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
709pub struct ModelDescriptor {
710    pub id: String,
711    pub name: String,
712    pub context_window: Option<u32>,
713    #[serde(default, skip_serializing_if = "Option::is_none")]
714    pub default_reasoning: Option<String>,
715    #[serde(default, skip_serializing_if = "Vec::is_empty")]
716    pub supported_reasoning: Vec<ReasoningEffortDescriptor>,
717}
718
719#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
720pub struct ReasoningEffortDescriptor {
721    pub effort: String,
722    pub description: String,
723}
724
725pub struct InferenceProviderContext<'a> {
726    pub provider_id: &'a str,
727}
728
729pub struct InferenceTurnContext<'a> {
730    pub thread_id: &'a str,
731    pub turn_id: &'a str,
732    /// Optional callback that executes a single tool call through Roder's tool
733    /// registry and policy, returning its result. Provided by the runtime for
734    /// providers that drive their own in-stream agent loop (e.g. the Cursor
735    /// bidi agent-runtime client, which must execute read/write/shell exec
736    /// requests mid-stream rather than ending the turn). Most providers ignore
737    /// it and surface tool calls as `ToolCallCompleted` events instead.
738    pub tool_executor: Option<std::sync::Arc<dyn TurnToolExecutor>>,
739}
740
741/// Result of executing one tool call via [`TurnToolExecutor`].
742#[derive(Debug, Clone)]
743pub struct TurnToolOutcome {
744    pub result: String,
745    pub is_error: bool,
746}
747
748/// Executes a single tool call through the runtime's registry + policy.
749/// Implemented by the runtime; used by providers that run their own in-stream
750/// agent loop.
751#[async_trait::async_trait]
752pub trait TurnToolExecutor: Send + Sync {
753    async fn execute(&self, call: ToolCallCompleted) -> anyhow::Result<TurnToolOutcome>;
754
755    /// Registers provider-owned cleanup for the current turn. Most providers
756    /// do not own a local child process and therefore never call this. The
757    /// runtime uses a registered handle only after interrupting the turn, when
758    /// it can distinguish its own task ending from provider cleanup being
759    /// acknowledged.
760    fn register_provider_cleanup(&self, _cleanup: std::sync::Arc<dyn ProviderTurnCleanup>) {}
761}
762
763/// A provider-owned cleanup acknowledgement for one turn. Implementations must
764/// not expose OS handles or command lines through this public boundary.
765#[async_trait::async_trait]
766pub trait ProviderTurnCleanup: Send + Sync {
767    /// Returns the conservative ownership state before cleanup completes.
768    fn ownership(&self) -> TurnCleanupOwnership;
769
770    /// Resolves only after the provider has completed its owned cleanup path.
771    async fn wait_for_cleanup(&self) -> anyhow::Result<()>;
772}
773
774#[async_trait::async_trait]
775pub trait InferenceEngine: Send + Sync + 'static {
776    fn id(&self) -> InferenceEngineId;
777    fn capabilities(&self) -> InferenceCapabilities;
778
779    fn metadata(&self) -> InferenceProviderMetadata {
780        InferenceProviderMetadata::local(self.id())
781    }
782
783    async fn list_models(
784        &self,
785        ctx: InferenceProviderContext<'_>,
786    ) -> anyhow::Result<Vec<ModelDescriptor>>;
787
788    async fn stream_turn(
789        &self,
790        ctx: InferenceTurnContext<'_>,
791        request: AgentInferenceRequest,
792    ) -> anyhow::Result<InferenceEventStream>;
793}
794
795#[cfg(test)]
796mod tests {
797    use super::*;
798
799    #[test]
800    fn finish_reason_mapping_normalizes_known_stop_reasons() {
801        assert_eq!(finish_reason_from_stop_reason("end_turn"), "stop");
802        assert_eq!(finish_reason_from_stop_reason("stop"), "stop");
803        assert_eq!(finish_reason_from_stop_reason("stop_sequence"), "stop");
804        assert_eq!(finish_reason_from_stop_reason("max_tokens"), "length");
805        assert_eq!(finish_reason_from_stop_reason("length"), "length");
806        assert_eq!(finish_reason_from_stop_reason("tool_use"), "toolUse");
807        assert_eq!(finish_reason_from_stop_reason("tool_calls"), "toolUse");
808        assert_eq!(
809            finish_reason_from_stop_reason("content_filter"),
810            "contentFilter"
811        );
812        assert_eq!(finish_reason_from_stop_reason("refusal"), "refusal");
813        assert_eq!(finish_reason_from_stop_reason("pause_turn"), "pause_turn");
814    }
815
816    #[test]
817    fn token_usage_accumulates_cache_creation_prompt_tokens() {
818        let mut usage = TokenUsage::new(100, 10, 110)
819            .with_cached_prompt_tokens(80)
820            .with_cache_creation_prompt_tokens(15);
821        usage.add_assign(
822            &TokenUsage::new(50, 5, 55)
823                .with_cached_prompt_tokens(40)
824                .with_cache_creation_prompt_tokens(10),
825        );
826
827        assert_eq!(usage.prompt_tokens, 150);
828        assert_eq!(usage.cached_prompt_tokens, 120);
829        assert_eq!(usage.cache_creation_prompt_tokens, 25);
830        assert!(!usage.is_empty());
831
832        let creation_only = TokenUsage {
833            cache_creation_prompt_tokens: 1,
834            ..TokenUsage::default()
835        };
836        assert!(!creation_only.is_empty());
837    }
838
839    #[test]
840    fn inference_speed_policy_decision_serializes_runtime_metadata() {
841        let decision = SpeedPolicyDecision {
842            phase: SpeedPolicyPhase::Verification,
843            desired_reasoning: "high".to_string(),
844            applied_reasoning: Some("high".to_string()),
845            supported: true,
846        };
847        let hints = RuntimeHints {
848            speed_policy: Some(decision),
849            ..RuntimeHints::default()
850        };
851
852        let json = serde_json::to_value(hints).unwrap();
853        assert_eq!(
854            json.get("speed_policy")
855                .and_then(|value| value.get("phase"))
856                .and_then(serde_json::Value::as_str),
857            Some("verification")
858        );
859        assert_eq!(
860            json.get("speed_policy")
861                .and_then(|value| value.get("desiredReasoning"))
862                .and_then(serde_json::Value::as_str),
863            Some("high")
864        );
865        assert_eq!(
866            json.get("speed_policy")
867                .and_then(|value| value.get("appliedReasoning"))
868                .and_then(serde_json::Value::as_str),
869            Some("high")
870        );
871    }
872
873    #[test]
874    fn inference_reliability_policy_serializes_runtime_metadata() {
875        let hints = RuntimeHints {
876            reliability: Some(ReliabilityRequestPolicy::default()),
877            ..RuntimeHints::default()
878        };
879
880        let json = serde_json::to_value(hints).unwrap();
881        assert_eq!(
882            json.get("reliability")
883                .and_then(|value| value.get("providerRetryMaxAttempts"))
884                .and_then(serde_json::Value::as_u64),
885            Some(3)
886        );
887        assert_eq!(
888            json.get("reliability")
889                .and_then(|value| value.get("retryEmptyProviderBody"))
890                .and_then(serde_json::Value::as_bool),
891            Some(true)
892        );
893    }
894
895    #[test]
896    fn tool_search_config_serializes_provider_native_request() {
897        let config = ToolSearchConfig {
898            mode: ToolSearchMode::ProviderNative,
899            max_catalog_items: Some(200),
900            include_mcp: true,
901            include_skills: false,
902            fallback_to_explicit_tools: true,
903            provider_variant: ToolSearchProviderVariant::Bm25,
904        };
905
906        let value = serde_json::to_value(&config).unwrap();
907
908        assert_eq!(value["mode"], "provider_native");
909        assert_eq!(value["maxCatalogItems"], 200);
910        assert_eq!(value["includeMcp"], true);
911        assert_eq!(value["includeSkills"], false);
912        assert_eq!(value["providerVariant"], "bm25");
913        assert!(config.is_provider_native_requested());
914    }
915
916    #[test]
917    fn explicit_tool_search_config_preserves_current_default() {
918        let config = ToolSearchConfig::default();
919
920        assert_eq!(config.mode, ToolSearchMode::Explicit);
921        assert!(!config.is_provider_native_requested());
922        assert!(config.fallback_to_explicit_tools);
923    }
924
925    #[test]
926    fn tool_search_effective_mode_resolution_covers_fallback_matrix() {
927        let explicit = ToolSearchConfig::explicit();
928        assert_eq!(
929            explicit.resolve_effective_mode(true).unwrap(),
930            EffectiveToolSearchMode::Explicit
931        );
932
933        let auto = ToolSearchConfig {
934            mode: ToolSearchMode::Auto,
935            ..ToolSearchConfig::default()
936        };
937        assert_eq!(
938            auto.resolve_effective_mode(true).unwrap(),
939            EffectiveToolSearchMode::ProviderNative
940        );
941        assert_eq!(
942            auto.resolve_effective_mode(false).unwrap(),
943            EffectiveToolSearchMode::Explicit
944        );
945
946        let native = ToolSearchConfig::provider_native();
947        assert_eq!(
948            native.resolve_effective_mode(true).unwrap(),
949            EffectiveToolSearchMode::ProviderNative
950        );
951        assert_eq!(
952            native.resolve_effective_mode(false).unwrap(),
953            EffectiveToolSearchMode::Explicit
954        );
955
956        let strict = ToolSearchConfig {
957            fallback_to_explicit_tools: false,
958            ..ToolSearchConfig::provider_native()
959        };
960        let error = strict.resolve_effective_mode(false).unwrap_err();
961        assert_eq!(error, ToolSearchModeError::ProviderNativeUnsupported);
962        assert!(error.to_string().contains("fallback_to_explicit_tools"));
963    }
964}