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    /// A complete provider output item. Persist immediately; unfinished deltas
682    /// are presentation only and must not enter replay history on interruption.
683    OutputItemCompleted(serde_json::Value),
684    HostedToolCallStarted(HostedToolCallStarted),
685    HostedToolCallCompleted(HostedToolCallCompleted),
686    Compaction(CompactionProgress),
687    Usage(TokenUsage),
688    Completed(CompletionMetadata),
689    Failed(InferenceFailure),
690    ProviderMetadata(serde_json::Value),
691}
692
693pub type InferenceEventStream =
694    Pin<Box<dyn Stream<Item = anyhow::Result<InferenceEvent>> + Send + 'static>>;
695
696#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
697pub struct InferenceCapabilities {
698    pub streaming: bool,
699    pub tool_calls: bool,
700    pub parallel_tool_calls: bool,
701    pub reasoning_summaries: bool,
702    pub structured_output: bool,
703    pub image_input: bool,
704    pub prompt_cache: bool,
705    pub provider_metadata: bool,
706    pub tool_search: bool,
707}
708
709impl InferenceCapabilities {
710    pub fn text_only() -> Self {
711        Self {
712            streaming: true,
713            tool_calls: false,
714            parallel_tool_calls: false,
715            reasoning_summaries: false,
716            structured_output: false,
717            image_input: false,
718            prompt_cache: false,
719            provider_metadata: false,
720            tool_search: false,
721        }
722    }
723
724    pub fn coding_agent_default() -> Self {
725        Self {
726            streaming: true,
727            tool_calls: true,
728            parallel_tool_calls: true,
729            reasoning_summaries: false,
730            structured_output: false,
731            image_input: false,
732            prompt_cache: false,
733            provider_metadata: true,
734            tool_search: false,
735        }
736    }
737}
738
739#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
740pub struct ModelDescriptor {
741    pub id: String,
742    pub name: String,
743    pub context_window: Option<u32>,
744    #[serde(default, skip_serializing_if = "Option::is_none")]
745    pub default_reasoning: Option<String>,
746    #[serde(default, skip_serializing_if = "Vec::is_empty")]
747    pub supported_reasoning: Vec<ReasoningEffortDescriptor>,
748}
749
750#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
751pub struct ReasoningEffortDescriptor {
752    pub effort: String,
753    pub description: String,
754}
755
756pub struct InferenceProviderContext<'a> {
757    pub provider_id: &'a str,
758}
759
760pub struct InferenceTurnContext<'a> {
761    pub thread_id: &'a str,
762    pub turn_id: &'a str,
763    /// Optional callback that executes a single tool call through Roder's tool
764    /// registry and policy, returning its result. Provided by the runtime for
765    /// providers that drive their own in-stream agent loop (e.g. the Cursor
766    /// bidi agent-runtime client, which must execute read/write/shell exec
767    /// requests mid-stream rather than ending the turn). Most providers ignore
768    /// it and surface tool calls as `ToolCallCompleted` events instead.
769    pub tool_executor: Option<std::sync::Arc<dyn TurnToolExecutor>>,
770}
771
772/// Result of executing one tool call via [`TurnToolExecutor`].
773#[derive(Debug, Clone)]
774pub struct TurnToolOutcome {
775    pub result: String,
776    pub is_error: bool,
777}
778
779/// Executes a single tool call through the runtime's registry + policy.
780/// Implemented by the runtime; used by providers that run their own in-stream
781/// agent loop.
782#[async_trait::async_trait]
783pub trait TurnToolExecutor: Send + Sync {
784    async fn execute(&self, call: ToolCallCompleted) -> anyhow::Result<TurnToolOutcome>;
785
786    /// Registers provider-owned cleanup for the current turn. Most providers
787    /// do not own a local child process and therefore never call this. The
788    /// runtime uses a registered handle only after interrupting the turn, when
789    /// it can distinguish its own task ending from provider cleanup being
790    /// acknowledged.
791    fn register_provider_cleanup(&self, _cleanup: std::sync::Arc<dyn ProviderTurnCleanup>) {}
792}
793
794/// A provider-owned cleanup acknowledgement for one turn. Implementations must
795/// not expose OS handles or command lines through this public boundary.
796#[async_trait::async_trait]
797pub trait ProviderTurnCleanup: Send + Sync {
798    /// Returns the conservative ownership state before cleanup completes.
799    fn ownership(&self) -> TurnCleanupOwnership;
800
801    /// Resolves only after the provider has completed its owned cleanup path.
802    async fn wait_for_cleanup(&self) -> anyhow::Result<()>;
803}
804
805#[async_trait::async_trait]
806pub trait InferenceEngine: Send + Sync + 'static {
807    fn id(&self) -> InferenceEngineId;
808    fn capabilities(&self) -> InferenceCapabilities;
809
810    fn metadata(&self) -> InferenceProviderMetadata {
811        InferenceProviderMetadata::local(self.id())
812    }
813
814    async fn list_models(
815        &self,
816        ctx: InferenceProviderContext<'_>,
817    ) -> anyhow::Result<Vec<ModelDescriptor>>;
818
819    /// Native, opaque provider compaction. Unsupported engines return None.
820    /// The stream must publish a completed boundary and one terminal completion.
821    async fn compact_turn(
822        &self,
823        _ctx: InferenceTurnContext<'_>,
824        _request: AgentInferenceRequest,
825    ) -> anyhow::Result<Option<InferenceEventStream>> {
826        Ok(None)
827    }
828
829    async fn stream_turn(
830        &self,
831        ctx: InferenceTurnContext<'_>,
832        request: AgentInferenceRequest,
833    ) -> anyhow::Result<InferenceEventStream>;
834}
835
836#[cfg(test)]
837mod tests {
838    use super::*;
839
840    #[test]
841    fn finish_reason_mapping_normalizes_known_stop_reasons() {
842        assert_eq!(finish_reason_from_stop_reason("end_turn"), "stop");
843        assert_eq!(finish_reason_from_stop_reason("stop"), "stop");
844        assert_eq!(finish_reason_from_stop_reason("stop_sequence"), "stop");
845        assert_eq!(finish_reason_from_stop_reason("max_tokens"), "length");
846        assert_eq!(finish_reason_from_stop_reason("length"), "length");
847        assert_eq!(finish_reason_from_stop_reason("tool_use"), "toolUse");
848        assert_eq!(finish_reason_from_stop_reason("tool_calls"), "toolUse");
849        assert_eq!(
850            finish_reason_from_stop_reason("content_filter"),
851            "contentFilter"
852        );
853        assert_eq!(finish_reason_from_stop_reason("refusal"), "refusal");
854        assert_eq!(finish_reason_from_stop_reason("pause_turn"), "pause_turn");
855    }
856
857    #[test]
858    fn token_usage_accumulates_cache_creation_prompt_tokens() {
859        let mut usage = TokenUsage::new(100, 10, 110)
860            .with_cached_prompt_tokens(80)
861            .with_cache_creation_prompt_tokens(15);
862        usage.add_assign(
863            &TokenUsage::new(50, 5, 55)
864                .with_cached_prompt_tokens(40)
865                .with_cache_creation_prompt_tokens(10),
866        );
867
868        assert_eq!(usage.prompt_tokens, 150);
869        assert_eq!(usage.cached_prompt_tokens, 120);
870        assert_eq!(usage.cache_creation_prompt_tokens, 25);
871        assert!(!usage.is_empty());
872
873        let creation_only = TokenUsage {
874            cache_creation_prompt_tokens: 1,
875            ..TokenUsage::default()
876        };
877        assert!(!creation_only.is_empty());
878    }
879
880    #[test]
881    fn service_tier_fields_default_when_absent_from_older_payloads() {
882        let hints: RuntimeHints = serde_json::from_value(serde_json::json!({
883            "trace_id": null,
884            "prompt_cache_key": null,
885            "auto_compact_token_limit": null
886        }))
887        .unwrap();
888        assert_eq!(hints.service_tier, None);
889        assert!(
890            serde_json::to_value(&hints)
891                .unwrap()
892                .get("service_tier")
893                .is_none()
894        );
895
896        let usage: TokenUsage = serde_json::from_value(serde_json::json!({
897            "prompt_tokens": 1,
898            "completion_tokens": 2,
899            "total_tokens": 3
900        }))
901        .unwrap();
902        assert_eq!(usage.service_tier, None);
903    }
904
905    #[test]
906    fn token_usage_keeps_the_most_recently_reported_service_tier() {
907        let mut usage = TokenUsage::new(10, 1, 11).with_service_tier(Some("priority".into()));
908        usage.add_assign(&TokenUsage::new(10, 1, 11));
909        assert_eq!(usage.service_tier.as_deref(), Some("priority"));
910        usage.add_assign(&TokenUsage::new(10, 1, 11).with_service_tier(Some("default".into())));
911        assert_eq!(usage.service_tier.as_deref(), Some("default"));
912    }
913
914    #[test]
915    fn inference_speed_policy_decision_serializes_runtime_metadata() {
916        let decision = SpeedPolicyDecision {
917            phase: SpeedPolicyPhase::Verification,
918            desired_reasoning: "high".to_string(),
919            applied_reasoning: Some("high".to_string()),
920            supported: true,
921        };
922        let hints = RuntimeHints {
923            speed_policy: Some(decision),
924            ..RuntimeHints::default()
925        };
926
927        let json = serde_json::to_value(hints).unwrap();
928        assert_eq!(
929            json.get("speed_policy")
930                .and_then(|value| value.get("phase"))
931                .and_then(serde_json::Value::as_str),
932            Some("verification")
933        );
934        assert_eq!(
935            json.get("speed_policy")
936                .and_then(|value| value.get("desiredReasoning"))
937                .and_then(serde_json::Value::as_str),
938            Some("high")
939        );
940        assert_eq!(
941            json.get("speed_policy")
942                .and_then(|value| value.get("appliedReasoning"))
943                .and_then(serde_json::Value::as_str),
944            Some("high")
945        );
946    }
947
948    #[test]
949    fn inference_reliability_policy_serializes_runtime_metadata() {
950        let hints = RuntimeHints {
951            reliability: Some(ReliabilityRequestPolicy::default()),
952            ..RuntimeHints::default()
953        };
954
955        let json = serde_json::to_value(hints).unwrap();
956        assert_eq!(
957            json.get("reliability")
958                .and_then(|value| value.get("providerRetryMaxAttempts"))
959                .and_then(serde_json::Value::as_u64),
960            Some(3)
961        );
962        assert_eq!(
963            json.get("reliability")
964                .and_then(|value| value.get("retryEmptyProviderBody"))
965                .and_then(serde_json::Value::as_bool),
966            Some(true)
967        );
968    }
969
970    #[test]
971    fn tool_search_config_serializes_provider_native_request() {
972        let config = ToolSearchConfig {
973            mode: ToolSearchMode::ProviderNative,
974            max_catalog_items: Some(200),
975            include_mcp: true,
976            include_skills: false,
977            fallback_to_explicit_tools: true,
978            provider_variant: ToolSearchProviderVariant::Bm25,
979        };
980
981        let value = serde_json::to_value(&config).unwrap();
982
983        assert_eq!(value["mode"], "provider_native");
984        assert_eq!(value["maxCatalogItems"], 200);
985        assert_eq!(value["includeMcp"], true);
986        assert_eq!(value["includeSkills"], false);
987        assert_eq!(value["providerVariant"], "bm25");
988        assert!(config.is_provider_native_requested());
989    }
990
991    #[test]
992    fn explicit_tool_search_config_preserves_current_default() {
993        let config = ToolSearchConfig::default();
994
995        assert_eq!(config.mode, ToolSearchMode::Explicit);
996        assert!(!config.is_provider_native_requested());
997        assert!(config.fallback_to_explicit_tools);
998    }
999
1000    #[test]
1001    fn tool_search_effective_mode_resolution_covers_fallback_matrix() {
1002        let explicit = ToolSearchConfig::explicit();
1003        assert_eq!(
1004            explicit.resolve_effective_mode(true).unwrap(),
1005            EffectiveToolSearchMode::Explicit
1006        );
1007
1008        let auto = ToolSearchConfig {
1009            mode: ToolSearchMode::Auto,
1010            ..ToolSearchConfig::default()
1011        };
1012        assert_eq!(
1013            auto.resolve_effective_mode(true).unwrap(),
1014            EffectiveToolSearchMode::ProviderNative
1015        );
1016        assert_eq!(
1017            auto.resolve_effective_mode(false).unwrap(),
1018            EffectiveToolSearchMode::Explicit
1019        );
1020
1021        let native = ToolSearchConfig::provider_native();
1022        assert_eq!(
1023            native.resolve_effective_mode(true).unwrap(),
1024            EffectiveToolSearchMode::ProviderNative
1025        );
1026        assert_eq!(
1027            native.resolve_effective_mode(false).unwrap(),
1028            EffectiveToolSearchMode::Explicit
1029        );
1030
1031        let strict = ToolSearchConfig {
1032            fallback_to_explicit_tools: false,
1033            ..ToolSearchConfig::provider_native()
1034        };
1035        let error = strict.resolve_effective_mode(false).unwrap_err();
1036        assert_eq!(error, ToolSearchModeError::ProviderNativeUnsupported);
1037        assert!(error.to_string().contains("fallback_to_explicit_tools"));
1038    }
1039}