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 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 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 #[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 #[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 #[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
629pub 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 #[serde(default, skip_serializing_if = "Option::is_none")]
660 pub tokens_before: Option<u32>,
661 #[serde(default, skip_serializing_if = "Option::is_none")]
663 pub tokens_after: Option<u32>,
664 #[serde(default, skip_serializing_if = "Option::is_none")]
666 pub duration_ms: Option<u64>,
667 #[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 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 pub tool_executor: Option<std::sync::Arc<dyn TurnToolExecutor>>,
770}
771
772#[derive(Debug, Clone)]
774pub struct TurnToolOutcome {
775 pub result: String,
776 pub is_error: bool,
777}
778
779#[async_trait::async_trait]
783pub trait TurnToolExecutor: Send + Sync {
784 async fn execute(&self, call: ToolCallCompleted) -> anyhow::Result<TurnToolOutcome>;
785
786 fn register_provider_cleanup(&self, _cleanup: std::sync::Arc<dyn ProviderTurnCleanup>) {}
792}
793
794#[async_trait::async_trait]
797pub trait ProviderTurnCleanup: Send + Sync {
798 fn ownership(&self) -> TurnCleanupOwnership;
800
801 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 fn tool_result_image_input(&self, _model: &str) -> bool {
830 false
831 }
832
833 fn requires_native_compaction(&self) -> bool {
836 false
837 }
838
839 async fn compact_turn(
842 &self,
843 _ctx: InferenceTurnContext<'_>,
844 _request: AgentInferenceRequest,
845 ) -> anyhow::Result<Option<InferenceEventStream>> {
846 Ok(None)
847 }
848
849 async fn stream_turn(
850 &self,
851 ctx: InferenceTurnContext<'_>,
852 request: AgentInferenceRequest,
853 ) -> anyhow::Result<InferenceEventStream>;
854}
855
856#[cfg(test)]
857mod tests {
858 use super::*;
859
860 #[test]
861 fn finish_reason_mapping_normalizes_known_stop_reasons() {
862 assert_eq!(finish_reason_from_stop_reason("end_turn"), "stop");
863 assert_eq!(finish_reason_from_stop_reason("stop"), "stop");
864 assert_eq!(finish_reason_from_stop_reason("stop_sequence"), "stop");
865 assert_eq!(finish_reason_from_stop_reason("max_tokens"), "length");
866 assert_eq!(finish_reason_from_stop_reason("length"), "length");
867 assert_eq!(finish_reason_from_stop_reason("tool_use"), "toolUse");
868 assert_eq!(finish_reason_from_stop_reason("tool_calls"), "toolUse");
869 assert_eq!(
870 finish_reason_from_stop_reason("content_filter"),
871 "contentFilter"
872 );
873 assert_eq!(finish_reason_from_stop_reason("refusal"), "refusal");
874 assert_eq!(finish_reason_from_stop_reason("pause_turn"), "pause_turn");
875 }
876
877 #[test]
878 fn token_usage_accumulates_cache_creation_prompt_tokens() {
879 let mut usage = TokenUsage::new(100, 10, 110)
880 .with_cached_prompt_tokens(80)
881 .with_cache_creation_prompt_tokens(15);
882 usage.add_assign(
883 &TokenUsage::new(50, 5, 55)
884 .with_cached_prompt_tokens(40)
885 .with_cache_creation_prompt_tokens(10),
886 );
887
888 assert_eq!(usage.prompt_tokens, 150);
889 assert_eq!(usage.cached_prompt_tokens, 120);
890 assert_eq!(usage.cache_creation_prompt_tokens, 25);
891 assert!(!usage.is_empty());
892
893 let creation_only = TokenUsage {
894 cache_creation_prompt_tokens: 1,
895 ..TokenUsage::default()
896 };
897 assert!(!creation_only.is_empty());
898 }
899
900 #[test]
901 fn service_tier_fields_default_when_absent_from_older_payloads() {
902 let hints: RuntimeHints = serde_json::from_value(serde_json::json!({
903 "trace_id": null,
904 "prompt_cache_key": null,
905 "auto_compact_token_limit": null
906 }))
907 .unwrap();
908 assert_eq!(hints.service_tier, None);
909 assert!(
910 serde_json::to_value(&hints)
911 .unwrap()
912 .get("service_tier")
913 .is_none()
914 );
915
916 let usage: TokenUsage = serde_json::from_value(serde_json::json!({
917 "prompt_tokens": 1,
918 "completion_tokens": 2,
919 "total_tokens": 3
920 }))
921 .unwrap();
922 assert_eq!(usage.service_tier, None);
923 }
924
925 #[test]
926 fn token_usage_keeps_the_most_recently_reported_service_tier() {
927 let mut usage = TokenUsage::new(10, 1, 11).with_service_tier(Some("priority".into()));
928 usage.add_assign(&TokenUsage::new(10, 1, 11));
929 assert_eq!(usage.service_tier.as_deref(), Some("priority"));
930 usage.add_assign(&TokenUsage::new(10, 1, 11).with_service_tier(Some("default".into())));
931 assert_eq!(usage.service_tier.as_deref(), Some("default"));
932 }
933
934 #[test]
935 fn inference_speed_policy_decision_serializes_runtime_metadata() {
936 let decision = SpeedPolicyDecision {
937 phase: SpeedPolicyPhase::Verification,
938 desired_reasoning: "high".to_string(),
939 applied_reasoning: Some("high".to_string()),
940 supported: true,
941 };
942 let hints = RuntimeHints {
943 speed_policy: Some(decision),
944 ..RuntimeHints::default()
945 };
946
947 let json = serde_json::to_value(hints).unwrap();
948 assert_eq!(
949 json.get("speed_policy")
950 .and_then(|value| value.get("phase"))
951 .and_then(serde_json::Value::as_str),
952 Some("verification")
953 );
954 assert_eq!(
955 json.get("speed_policy")
956 .and_then(|value| value.get("desiredReasoning"))
957 .and_then(serde_json::Value::as_str),
958 Some("high")
959 );
960 assert_eq!(
961 json.get("speed_policy")
962 .and_then(|value| value.get("appliedReasoning"))
963 .and_then(serde_json::Value::as_str),
964 Some("high")
965 );
966 }
967
968 #[test]
969 fn inference_reliability_policy_serializes_runtime_metadata() {
970 let hints = RuntimeHints {
971 reliability: Some(ReliabilityRequestPolicy::default()),
972 ..RuntimeHints::default()
973 };
974
975 let json = serde_json::to_value(hints).unwrap();
976 assert_eq!(
977 json.get("reliability")
978 .and_then(|value| value.get("providerRetryMaxAttempts"))
979 .and_then(serde_json::Value::as_u64),
980 Some(3)
981 );
982 assert_eq!(
983 json.get("reliability")
984 .and_then(|value| value.get("retryEmptyProviderBody"))
985 .and_then(serde_json::Value::as_bool),
986 Some(true)
987 );
988 }
989
990 #[test]
991 fn tool_search_config_serializes_provider_native_request() {
992 let config = ToolSearchConfig {
993 mode: ToolSearchMode::ProviderNative,
994 max_catalog_items: Some(200),
995 include_mcp: true,
996 include_skills: false,
997 fallback_to_explicit_tools: true,
998 provider_variant: ToolSearchProviderVariant::Bm25,
999 };
1000
1001 let value = serde_json::to_value(&config).unwrap();
1002
1003 assert_eq!(value["mode"], "provider_native");
1004 assert_eq!(value["maxCatalogItems"], 200);
1005 assert_eq!(value["includeMcp"], true);
1006 assert_eq!(value["includeSkills"], false);
1007 assert_eq!(value["providerVariant"], "bm25");
1008 assert!(config.is_provider_native_requested());
1009 }
1010
1011 #[test]
1012 fn explicit_tool_search_config_preserves_current_default() {
1013 let config = ToolSearchConfig::default();
1014
1015 assert_eq!(config.mode, ToolSearchMode::Explicit);
1016 assert!(!config.is_provider_native_requested());
1017 assert!(config.fallback_to_explicit_tools);
1018 }
1019
1020 #[test]
1021 fn tool_search_effective_mode_resolution_covers_fallback_matrix() {
1022 let explicit = ToolSearchConfig::explicit();
1023 assert_eq!(
1024 explicit.resolve_effective_mode(true).unwrap(),
1025 EffectiveToolSearchMode::Explicit
1026 );
1027
1028 let auto = ToolSearchConfig {
1029 mode: ToolSearchMode::Auto,
1030 ..ToolSearchConfig::default()
1031 };
1032 assert_eq!(
1033 auto.resolve_effective_mode(true).unwrap(),
1034 EffectiveToolSearchMode::ProviderNative
1035 );
1036 assert_eq!(
1037 auto.resolve_effective_mode(false).unwrap(),
1038 EffectiveToolSearchMode::Explicit
1039 );
1040
1041 let native = ToolSearchConfig::provider_native();
1042 assert_eq!(
1043 native.resolve_effective_mode(true).unwrap(),
1044 EffectiveToolSearchMode::ProviderNative
1045 );
1046 assert_eq!(
1047 native.resolve_effective_mode(false).unwrap(),
1048 EffectiveToolSearchMode::Explicit
1049 );
1050
1051 let strict = ToolSearchConfig {
1052 fallback_to_explicit_tools: false,
1053 ..ToolSearchConfig::provider_native()
1054 };
1055 let error = strict.resolve_effective_mode(false).unwrap_err();
1056 assert_eq!(error, ToolSearchModeError::ProviderNativeUnsupported);
1057 assert!(error.to_string().contains("fallback_to_explicit_tools"));
1058 }
1059
1060 struct BareEngine {
1061 forwards_images_for: Option<&'static str>,
1062 }
1063
1064 #[async_trait::async_trait]
1065 impl InferenceEngine for BareEngine {
1066 fn id(&self) -> InferenceEngineId {
1067 "bare".to_string()
1068 }
1069
1070 fn capabilities(&self) -> InferenceCapabilities {
1071 InferenceCapabilities::coding_agent_default()
1072 }
1073
1074 fn tool_result_image_input(&self, model: &str) -> bool {
1075 self.forwards_images_for == Some(model)
1076 }
1077
1078 async fn list_models(
1079 &self,
1080 _ctx: InferenceProviderContext<'_>,
1081 ) -> anyhow::Result<Vec<ModelDescriptor>> {
1082 Ok(Vec::new())
1083 }
1084
1085 async fn stream_turn(
1086 &self,
1087 _ctx: InferenceTurnContext<'_>,
1088 _request: AgentInferenceRequest,
1089 ) -> anyhow::Result<InferenceEventStream> {
1090 anyhow::bail!("not used")
1091 }
1092 }
1093
1094 struct DefaultEngine;
1097
1098 #[async_trait::async_trait]
1099 impl InferenceEngine for DefaultEngine {
1100 fn id(&self) -> InferenceEngineId {
1101 "default".to_string()
1102 }
1103
1104 fn capabilities(&self) -> InferenceCapabilities {
1105 InferenceCapabilities {
1106 image_input: true,
1107 ..InferenceCapabilities::coding_agent_default()
1108 }
1109 }
1110
1111 async fn list_models(
1112 &self,
1113 _ctx: InferenceProviderContext<'_>,
1114 ) -> anyhow::Result<Vec<ModelDescriptor>> {
1115 Ok(Vec::new())
1116 }
1117
1118 async fn stream_turn(
1119 &self,
1120 _ctx: InferenceTurnContext<'_>,
1121 _request: AgentInferenceRequest,
1122 ) -> anyhow::Result<InferenceEventStream> {
1123 anyhow::bail!("not used")
1124 }
1125 }
1126
1127 #[test]
1128 fn tool_result_images_are_off_unless_the_engine_opts_in_per_model() {
1129 let default = DefaultEngine;
1130 assert!(default.capabilities().image_input);
1131 assert!(!default.tool_result_image_input("any-model"));
1132
1133 let engine: std::sync::Arc<dyn InferenceEngine> = std::sync::Arc::new(BareEngine {
1134 forwards_images_for: Some("vision-model"),
1135 });
1136 assert!(engine.tool_result_image_input("vision-model"));
1137 assert!(!engine.tool_result_image_input("text-model"));
1138 }
1139}