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 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 pub tool_executor: Option<std::sync::Arc<dyn TurnToolExecutor>>,
767}
768
769#[derive(Debug, Clone)]
771pub struct TurnToolOutcome {
772 pub result: String,
773 pub is_error: bool,
774}
775
776#[async_trait::async_trait]
780pub trait TurnToolExecutor: Send + Sync {
781 async fn execute(&self, call: ToolCallCompleted) -> anyhow::Result<TurnToolOutcome>;
782
783 fn register_provider_cleanup(&self, _cleanup: std::sync::Arc<dyn ProviderTurnCleanup>) {}
789}
790
791#[async_trait::async_trait]
794pub trait ProviderTurnCleanup: Send + Sync {
795 fn ownership(&self) -> TurnCleanupOwnership;
797
798 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}