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     * Provider service tier requested for this call (OpenAI `service_tier`,
464     * e.g. `"priority"` for Fast mode, `"flex"`, `"default"`, `"auto"`).
465     * `None` leaves the provider default. Providers that have no tier
466     * concept ignore it; the tier actually served is reported on
467     * `TokenUsage::service_tier`.
468     */
469    #[serde(default, skip_serializing_if = "Option::is_none")]
470    pub service_tier: Option<String>,
471}
472
473#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
474pub struct AgentInferenceRequest {
475    pub model: ModelSelection,
476    pub instructions: InstructionBundle,
477    pub transcript: Vec<TranscriptItem>,
478    pub tools: Vec<ToolSpec>,
479    pub tool_choice: ToolChoice,
480    pub reasoning: ReasoningConfig,
481    pub output: OutputConfig,
482    pub runtime: RuntimeHints,
483    pub metadata: serde_json::Value,
484}
485
486#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
487pub struct MessageDelta {
488    pub text: String,
489    #[serde(default, skip_serializing_if = "Option::is_none")]
490    pub phase: Option<String>,
491}
492
493#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
494pub struct ReasoningDelta {
495    pub text: String,
496}
497
498#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
499pub struct ToolCallStarted {
500    pub id: String,
501    pub name: String,
502}
503
504#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
505pub struct ToolCallDelta {
506    pub id: String,
507    pub arguments_delta: String,
508}
509
510#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
511pub struct ToolCallCompleted {
512    pub id: String,
513    pub name: String,
514    pub arguments: String,
515}
516
517#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
518pub struct HostedToolCallStarted {
519    pub id: String,
520    pub name: String,
521}
522
523#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
524pub struct HostedToolCallCompleted {
525    pub id: String,
526    pub name: String,
527    pub arguments: String,
528}
529
530#[derive(Debug, Clone, Default, Serialize, Deserialize, PartialEq)]
531pub struct TokenUsage {
532    pub prompt_tokens: u32,
533    pub completion_tokens: u32,
534    pub total_tokens: u32,
535    #[serde(default)]
536    pub cached_prompt_tokens: u32,
537    /**
538     * Prompt tokens written to the provider's prompt cache this step. Like
539     * `cached_prompt_tokens`, this is a subset of `prompt_tokens`, not an
540     * additional count; hosts use it to bill cache writes at the provider's
541     * cache-write rate.
542     */
543    #[serde(default)]
544    pub cache_creation_prompt_tokens: u32,
545    #[serde(default, skip_serializing_if = "Option::is_none")]
546    pub cache_hit_rate: Option<f64>,
547    /**
548     * Service tier the provider reports it actually served this step with
549     * (OpenAI `response.service_tier`). A request for a faster tier may be
550     * downgraded under load and reported here as `"default"`, which is what
551     * the step is billed at. `None` when the provider did not report one.
552     * When usage is aggregated with `add_assign`, the most recently reported
553     * tier wins; consumers needing per-step precision read per-step usage.
554     */
555    #[serde(default, skip_serializing_if = "Option::is_none")]
556    pub service_tier: Option<String>,
557}
558
559impl TokenUsage {
560    pub fn new(prompt_tokens: u32, completion_tokens: u32, total_tokens: u32) -> Self {
561        Self {
562            prompt_tokens,
563            completion_tokens,
564            total_tokens,
565            cached_prompt_tokens: 0,
566            cache_creation_prompt_tokens: 0,
567            cache_hit_rate: cache_hit_rate(prompt_tokens, 0),
568            service_tier: None,
569        }
570    }
571
572    pub fn with_service_tier(mut self, service_tier: Option<String>) -> Self {
573        self.service_tier = service_tier;
574        self
575    }
576
577    pub fn with_cached_prompt_tokens(mut self, cached_prompt_tokens: u32) -> Self {
578        self.cached_prompt_tokens = cached_prompt_tokens.min(self.prompt_tokens);
579        self.cache_hit_rate = cache_hit_rate(self.prompt_tokens, self.cached_prompt_tokens);
580        self
581    }
582
583    pub fn with_cache_creation_prompt_tokens(mut self, cache_creation_prompt_tokens: u32) -> Self {
584        self.cache_creation_prompt_tokens = cache_creation_prompt_tokens.min(self.prompt_tokens);
585        self
586    }
587
588    pub fn add_assign(&mut self, usage: &TokenUsage) {
589        self.prompt_tokens = self.prompt_tokens.saturating_add(usage.prompt_tokens);
590        self.completion_tokens = self
591            .completion_tokens
592            .saturating_add(usage.completion_tokens);
593        self.total_tokens = self.total_tokens.saturating_add(usage.total_tokens);
594        self.cached_prompt_tokens = self
595            .cached_prompt_tokens
596            .saturating_add(usage.cached_prompt_tokens);
597        self.cache_creation_prompt_tokens = self
598            .cache_creation_prompt_tokens
599            .saturating_add(usage.cache_creation_prompt_tokens);
600        self.cache_hit_rate = cache_hit_rate(self.prompt_tokens, self.cached_prompt_tokens);
601        if usage.service_tier.is_some() {
602            self.service_tier.clone_from(&usage.service_tier);
603        }
604    }
605
606    pub fn is_empty(&self) -> bool {
607        self.prompt_tokens == 0
608            && self.completion_tokens == 0
609            && self.total_tokens == 0
610            && self.cached_prompt_tokens == 0
611            && self.cache_creation_prompt_tokens == 0
612    }
613}
614
615pub fn cache_hit_rate(prompt_tokens: u32, cached_prompt_tokens: u32) -> Option<f64> {
616    if prompt_tokens == 0 {
617        None
618    } else {
619        Some(f64::from(cached_prompt_tokens.min(prompt_tokens)) / f64::from(prompt_tokens))
620    }
621}
622
623#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
624pub struct CompletionMetadata {
625    pub stop_reason: Option<String>,
626    pub provider_response_id: Option<String>,
627}
628
629/**
630 * Canonical mapping from provider-native stop reasons to the finish reason
631 * surfaced as `finishReason` on `turn/completed`. Only the terminal inference
632 * step's stop reason reaches the turn surface, so `toolUse` appears only when
633 * a turn genuinely ends on a tool-use step (e.g. tool rounds exhausted).
634 * Unknown stop reasons pass through unchanged.
635 */
636pub fn finish_reason_from_stop_reason(stop_reason: &str) -> String {
637    match stop_reason {
638        "end_turn" | "stop" | "stop_sequence" => "stop",
639        "max_tokens" | "length" => "length",
640        "tool_use" | "tool_calls" => "toolUse",
641        "content_filter" => "contentFilter",
642        "refusal" => "refusal",
643        other => other,
644    }
645    .to_string()
646}
647
648#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
649pub struct InferenceFailure {
650    pub message: String,
651}
652
653#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
654pub struct CompactionProgress {
655    pub status: String,
656    #[serde(default, skip_serializing_if = "Option::is_none")]
657    pub item_id: Option<String>,
658    /// Estimated prompt tokens before this compaction pass, when known.
659    #[serde(default, skip_serializing_if = "Option::is_none")]
660    pub tokens_before: Option<u32>,
661    /// Estimated prompt tokens after this compaction pass, when known.
662    #[serde(default, skip_serializing_if = "Option::is_none")]
663    pub tokens_after: Option<u32>,
664    /// Wall-clock duration of the compaction pass, when measured.
665    #[serde(default, skip_serializing_if = "Option::is_none")]
666    pub duration_ms: Option<u64>,
667    /// Opaque provider compaction output item (e.g. OpenAI `type: "compaction"`).
668    /// Runtime persists this as a transcript boundary even if the stream dies
669    /// before the full ProviderMetadata/response.completed frame arrives.
670    #[serde(default, skip_serializing_if = "Option::is_none")]
671    pub item: Option<serde_json::Value>,
672}
673
674#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
675pub enum InferenceEvent {
676    MessageDelta(MessageDelta),
677    ReasoningDelta(ReasoningDelta),
678    ToolCallStarted(ToolCallStarted),
679    ToolCallDelta(ToolCallDelta),
680    ToolCallCompleted(ToolCallCompleted),
681    HostedToolCallStarted(HostedToolCallStarted),
682    HostedToolCallCompleted(HostedToolCallCompleted),
683    Compaction(CompactionProgress),
684    Usage(TokenUsage),
685    Completed(CompletionMetadata),
686    Failed(InferenceFailure),
687    ProviderMetadata(serde_json::Value),
688}
689
690pub type InferenceEventStream =
691    Pin<Box<dyn Stream<Item = anyhow::Result<InferenceEvent>> + Send + 'static>>;
692
693#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
694pub struct InferenceCapabilities {
695    pub streaming: bool,
696    pub tool_calls: bool,
697    pub parallel_tool_calls: bool,
698    pub reasoning_summaries: bool,
699    pub structured_output: bool,
700    pub image_input: bool,
701    pub prompt_cache: bool,
702    pub provider_metadata: bool,
703    pub tool_search: bool,
704}
705
706impl InferenceCapabilities {
707    pub fn text_only() -> Self {
708        Self {
709            streaming: true,
710            tool_calls: false,
711            parallel_tool_calls: false,
712            reasoning_summaries: false,
713            structured_output: false,
714            image_input: false,
715            prompt_cache: false,
716            provider_metadata: false,
717            tool_search: false,
718        }
719    }
720
721    pub fn coding_agent_default() -> Self {
722        Self {
723            streaming: true,
724            tool_calls: true,
725            parallel_tool_calls: true,
726            reasoning_summaries: false,
727            structured_output: false,
728            image_input: false,
729            prompt_cache: false,
730            provider_metadata: true,
731            tool_search: false,
732        }
733    }
734}
735
736#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
737pub struct ModelDescriptor {
738    pub id: String,
739    pub name: String,
740    pub context_window: Option<u32>,
741    #[serde(default, skip_serializing_if = "Option::is_none")]
742    pub default_reasoning: Option<String>,
743    #[serde(default, skip_serializing_if = "Vec::is_empty")]
744    pub supported_reasoning: Vec<ReasoningEffortDescriptor>,
745}
746
747#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
748pub struct ReasoningEffortDescriptor {
749    pub effort: String,
750    pub description: String,
751}
752
753pub struct InferenceProviderContext<'a> {
754    pub provider_id: &'a str,
755}
756
757pub struct InferenceTurnContext<'a> {
758    pub thread_id: &'a str,
759    pub turn_id: &'a str,
760    /// Optional callback that executes a single tool call through Roder's tool
761    /// registry and policy, returning its result. Provided by the runtime for
762    /// providers that drive their own in-stream agent loop (e.g. the Cursor
763    /// bidi agent-runtime client, which must execute read/write/shell exec
764    /// requests mid-stream rather than ending the turn). Most providers ignore
765    /// it and surface tool calls as `ToolCallCompleted` events instead.
766    pub tool_executor: Option<std::sync::Arc<dyn TurnToolExecutor>>,
767}
768
769/// Result of executing one tool call via [`TurnToolExecutor`].
770#[derive(Debug, Clone)]
771pub struct TurnToolOutcome {
772    pub result: String,
773    pub is_error: bool,
774}
775
776/// Executes a single tool call through the runtime's registry + policy.
777/// Implemented by the runtime; used by providers that run their own in-stream
778/// agent loop.
779#[async_trait::async_trait]
780pub trait TurnToolExecutor: Send + Sync {
781    async fn execute(&self, call: ToolCallCompleted) -> anyhow::Result<TurnToolOutcome>;
782
783    /// Registers provider-owned cleanup for the current turn. Most providers
784    /// do not own a local child process and therefore never call this. The
785    /// runtime uses a registered handle only after interrupting the turn, when
786    /// it can distinguish its own task ending from provider cleanup being
787    /// acknowledged.
788    fn register_provider_cleanup(&self, _cleanup: std::sync::Arc<dyn ProviderTurnCleanup>) {}
789}
790
791/// A provider-owned cleanup acknowledgement for one turn. Implementations must
792/// not expose OS handles or command lines through this public boundary.
793#[async_trait::async_trait]
794pub trait ProviderTurnCleanup: Send + Sync {
795    /// Returns the conservative ownership state before cleanup completes.
796    fn ownership(&self) -> TurnCleanupOwnership;
797
798    /// Resolves only after the provider has completed its owned cleanup path.
799    async fn wait_for_cleanup(&self) -> anyhow::Result<()>;
800}
801
802#[async_trait::async_trait]
803pub trait InferenceEngine: Send + Sync + 'static {
804    fn id(&self) -> InferenceEngineId;
805    fn capabilities(&self) -> InferenceCapabilities;
806
807    fn metadata(&self) -> InferenceProviderMetadata {
808        InferenceProviderMetadata::local(self.id())
809    }
810
811    async fn list_models(
812        &self,
813        ctx: InferenceProviderContext<'_>,
814    ) -> anyhow::Result<Vec<ModelDescriptor>>;
815
816    async fn stream_turn(
817        &self,
818        ctx: InferenceTurnContext<'_>,
819        request: AgentInferenceRequest,
820    ) -> anyhow::Result<InferenceEventStream>;
821}
822
823#[cfg(test)]
824mod tests {
825    use super::*;
826
827    #[test]
828    fn finish_reason_mapping_normalizes_known_stop_reasons() {
829        assert_eq!(finish_reason_from_stop_reason("end_turn"), "stop");
830        assert_eq!(finish_reason_from_stop_reason("stop"), "stop");
831        assert_eq!(finish_reason_from_stop_reason("stop_sequence"), "stop");
832        assert_eq!(finish_reason_from_stop_reason("max_tokens"), "length");
833        assert_eq!(finish_reason_from_stop_reason("length"), "length");
834        assert_eq!(finish_reason_from_stop_reason("tool_use"), "toolUse");
835        assert_eq!(finish_reason_from_stop_reason("tool_calls"), "toolUse");
836        assert_eq!(
837            finish_reason_from_stop_reason("content_filter"),
838            "contentFilter"
839        );
840        assert_eq!(finish_reason_from_stop_reason("refusal"), "refusal");
841        assert_eq!(finish_reason_from_stop_reason("pause_turn"), "pause_turn");
842    }
843
844    #[test]
845    fn token_usage_accumulates_cache_creation_prompt_tokens() {
846        let mut usage = TokenUsage::new(100, 10, 110)
847            .with_cached_prompt_tokens(80)
848            .with_cache_creation_prompt_tokens(15);
849        usage.add_assign(
850            &TokenUsage::new(50, 5, 55)
851                .with_cached_prompt_tokens(40)
852                .with_cache_creation_prompt_tokens(10),
853        );
854
855        assert_eq!(usage.prompt_tokens, 150);
856        assert_eq!(usage.cached_prompt_tokens, 120);
857        assert_eq!(usage.cache_creation_prompt_tokens, 25);
858        assert!(!usage.is_empty());
859
860        let creation_only = TokenUsage {
861            cache_creation_prompt_tokens: 1,
862            ..TokenUsage::default()
863        };
864        assert!(!creation_only.is_empty());
865    }
866
867    #[test]
868    fn service_tier_fields_default_when_absent_from_older_payloads() {
869        let hints: RuntimeHints = serde_json::from_value(serde_json::json!({
870            "trace_id": null,
871            "prompt_cache_key": null,
872            "auto_compact_token_limit": null
873        }))
874        .unwrap();
875        assert_eq!(hints.service_tier, None);
876        assert!(
877            serde_json::to_value(&hints)
878                .unwrap()
879                .get("service_tier")
880                .is_none()
881        );
882
883        let usage: TokenUsage = serde_json::from_value(serde_json::json!({
884            "prompt_tokens": 1,
885            "completion_tokens": 2,
886            "total_tokens": 3
887        }))
888        .unwrap();
889        assert_eq!(usage.service_tier, None);
890    }
891
892    #[test]
893    fn token_usage_keeps_the_most_recently_reported_service_tier() {
894        let mut usage = TokenUsage::new(10, 1, 11).with_service_tier(Some("priority".into()));
895        usage.add_assign(&TokenUsage::new(10, 1, 11));
896        assert_eq!(usage.service_tier.as_deref(), Some("priority"));
897        usage.add_assign(&TokenUsage::new(10, 1, 11).with_service_tier(Some("default".into())));
898        assert_eq!(usage.service_tier.as_deref(), Some("default"));
899    }
900
901    #[test]
902    fn inference_speed_policy_decision_serializes_runtime_metadata() {
903        let decision = SpeedPolicyDecision {
904            phase: SpeedPolicyPhase::Verification,
905            desired_reasoning: "high".to_string(),
906            applied_reasoning: Some("high".to_string()),
907            supported: true,
908        };
909        let hints = RuntimeHints {
910            speed_policy: Some(decision),
911            ..RuntimeHints::default()
912        };
913
914        let json = serde_json::to_value(hints).unwrap();
915        assert_eq!(
916            json.get("speed_policy")
917                .and_then(|value| value.get("phase"))
918                .and_then(serde_json::Value::as_str),
919            Some("verification")
920        );
921        assert_eq!(
922            json.get("speed_policy")
923                .and_then(|value| value.get("desiredReasoning"))
924                .and_then(serde_json::Value::as_str),
925            Some("high")
926        );
927        assert_eq!(
928            json.get("speed_policy")
929                .and_then(|value| value.get("appliedReasoning"))
930                .and_then(serde_json::Value::as_str),
931            Some("high")
932        );
933    }
934
935    #[test]
936    fn inference_reliability_policy_serializes_runtime_metadata() {
937        let hints = RuntimeHints {
938            reliability: Some(ReliabilityRequestPolicy::default()),
939            ..RuntimeHints::default()
940        };
941
942        let json = serde_json::to_value(hints).unwrap();
943        assert_eq!(
944            json.get("reliability")
945                .and_then(|value| value.get("providerRetryMaxAttempts"))
946                .and_then(serde_json::Value::as_u64),
947            Some(3)
948        );
949        assert_eq!(
950            json.get("reliability")
951                .and_then(|value| value.get("retryEmptyProviderBody"))
952                .and_then(serde_json::Value::as_bool),
953            Some(true)
954        );
955    }
956
957    #[test]
958    fn tool_search_config_serializes_provider_native_request() {
959        let config = ToolSearchConfig {
960            mode: ToolSearchMode::ProviderNative,
961            max_catalog_items: Some(200),
962            include_mcp: true,
963            include_skills: false,
964            fallback_to_explicit_tools: true,
965            provider_variant: ToolSearchProviderVariant::Bm25,
966        };
967
968        let value = serde_json::to_value(&config).unwrap();
969
970        assert_eq!(value["mode"], "provider_native");
971        assert_eq!(value["maxCatalogItems"], 200);
972        assert_eq!(value["includeMcp"], true);
973        assert_eq!(value["includeSkills"], false);
974        assert_eq!(value["providerVariant"], "bm25");
975        assert!(config.is_provider_native_requested());
976    }
977
978    #[test]
979    fn explicit_tool_search_config_preserves_current_default() {
980        let config = ToolSearchConfig::default();
981
982        assert_eq!(config.mode, ToolSearchMode::Explicit);
983        assert!(!config.is_provider_native_requested());
984        assert!(config.fallback_to_explicit_tools);
985    }
986
987    #[test]
988    fn tool_search_effective_mode_resolution_covers_fallback_matrix() {
989        let explicit = ToolSearchConfig::explicit();
990        assert_eq!(
991            explicit.resolve_effective_mode(true).unwrap(),
992            EffectiveToolSearchMode::Explicit
993        );
994
995        let auto = ToolSearchConfig {
996            mode: ToolSearchMode::Auto,
997            ..ToolSearchConfig::default()
998        };
999        assert_eq!(
1000            auto.resolve_effective_mode(true).unwrap(),
1001            EffectiveToolSearchMode::ProviderNative
1002        );
1003        assert_eq!(
1004            auto.resolve_effective_mode(false).unwrap(),
1005            EffectiveToolSearchMode::Explicit
1006        );
1007
1008        let native = ToolSearchConfig::provider_native();
1009        assert_eq!(
1010            native.resolve_effective_mode(true).unwrap(),
1011            EffectiveToolSearchMode::ProviderNative
1012        );
1013        assert_eq!(
1014            native.resolve_effective_mode(false).unwrap(),
1015            EffectiveToolSearchMode::Explicit
1016        );
1017
1018        let strict = ToolSearchConfig {
1019            fallback_to_explicit_tools: false,
1020            ..ToolSearchConfig::provider_native()
1021        };
1022        let error = strict.resolve_effective_mode(false).unwrap_err();
1023        assert_eq!(error, ToolSearchModeError::ProviderNativeUnsupported);
1024        assert!(error.to_string().contains("fallback_to_explicit_tools"));
1025    }
1026}