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}
463
464#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
465pub struct AgentInferenceRequest {
466 pub model: ModelSelection,
467 pub instructions: InstructionBundle,
468 pub transcript: Vec<TranscriptItem>,
469 pub tools: Vec<ToolSpec>,
470 pub tool_choice: ToolChoice,
471 pub reasoning: ReasoningConfig,
472 pub output: OutputConfig,
473 pub runtime: RuntimeHints,
474 pub metadata: serde_json::Value,
475}
476
477#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
478pub struct MessageDelta {
479 pub text: String,
480 #[serde(default, skip_serializing_if = "Option::is_none")]
481 pub phase: Option<String>,
482}
483
484#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
485pub struct ReasoningDelta {
486 pub text: String,
487}
488
489#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
490pub struct ToolCallStarted {
491 pub id: String,
492 pub name: String,
493}
494
495#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
496pub struct ToolCallDelta {
497 pub id: String,
498 pub arguments_delta: String,
499}
500
501#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
502pub struct ToolCallCompleted {
503 pub id: String,
504 pub name: String,
505 pub arguments: String,
506}
507
508#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
509pub struct HostedToolCallStarted {
510 pub id: String,
511 pub name: String,
512}
513
514#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
515pub struct HostedToolCallCompleted {
516 pub id: String,
517 pub name: String,
518 pub arguments: String,
519}
520
521#[derive(Debug, Clone, Default, Serialize, Deserialize, PartialEq)]
522pub struct TokenUsage {
523 pub prompt_tokens: u32,
524 pub completion_tokens: u32,
525 pub total_tokens: u32,
526 #[serde(default)]
527 pub cached_prompt_tokens: u32,
528 #[serde(default)]
535 pub cache_creation_prompt_tokens: u32,
536 #[serde(default, skip_serializing_if = "Option::is_none")]
537 pub cache_hit_rate: Option<f64>,
538}
539
540impl TokenUsage {
541 pub fn new(prompt_tokens: u32, completion_tokens: u32, total_tokens: u32) -> Self {
542 Self {
543 prompt_tokens,
544 completion_tokens,
545 total_tokens,
546 cached_prompt_tokens: 0,
547 cache_creation_prompt_tokens: 0,
548 cache_hit_rate: cache_hit_rate(prompt_tokens, 0),
549 }
550 }
551
552 pub fn with_cached_prompt_tokens(mut self, cached_prompt_tokens: u32) -> Self {
553 self.cached_prompt_tokens = cached_prompt_tokens.min(self.prompt_tokens);
554 self.cache_hit_rate = cache_hit_rate(self.prompt_tokens, self.cached_prompt_tokens);
555 self
556 }
557
558 pub fn with_cache_creation_prompt_tokens(mut self, cache_creation_prompt_tokens: u32) -> Self {
559 self.cache_creation_prompt_tokens = cache_creation_prompt_tokens.min(self.prompt_tokens);
560 self
561 }
562
563 pub fn add_assign(&mut self, usage: &TokenUsage) {
564 self.prompt_tokens = self.prompt_tokens.saturating_add(usage.prompt_tokens);
565 self.completion_tokens = self
566 .completion_tokens
567 .saturating_add(usage.completion_tokens);
568 self.total_tokens = self.total_tokens.saturating_add(usage.total_tokens);
569 self.cached_prompt_tokens = self
570 .cached_prompt_tokens
571 .saturating_add(usage.cached_prompt_tokens);
572 self.cache_creation_prompt_tokens = self
573 .cache_creation_prompt_tokens
574 .saturating_add(usage.cache_creation_prompt_tokens);
575 self.cache_hit_rate = cache_hit_rate(self.prompt_tokens, self.cached_prompt_tokens);
576 }
577
578 pub fn is_empty(&self) -> bool {
579 self.prompt_tokens == 0
580 && self.completion_tokens == 0
581 && self.total_tokens == 0
582 && self.cached_prompt_tokens == 0
583 && self.cache_creation_prompt_tokens == 0
584 }
585}
586
587pub fn cache_hit_rate(prompt_tokens: u32, cached_prompt_tokens: u32) -> Option<f64> {
588 if prompt_tokens == 0 {
589 None
590 } else {
591 Some(f64::from(cached_prompt_tokens.min(prompt_tokens)) / f64::from(prompt_tokens))
592 }
593}
594
595#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
596pub struct CompletionMetadata {
597 pub stop_reason: Option<String>,
598 pub provider_response_id: Option<String>,
599}
600
601pub fn finish_reason_from_stop_reason(stop_reason: &str) -> String {
609 match stop_reason {
610 "end_turn" | "stop" | "stop_sequence" => "stop",
611 "max_tokens" | "length" => "length",
612 "tool_use" | "tool_calls" => "toolUse",
613 "content_filter" => "contentFilter",
614 "refusal" => "refusal",
615 other => other,
616 }
617 .to_string()
618}
619
620#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
621pub struct InferenceFailure {
622 pub message: String,
623}
624
625#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
626pub struct CompactionProgress {
627 pub status: String,
628 #[serde(default, skip_serializing_if = "Option::is_none")]
629 pub item_id: Option<String>,
630 #[serde(default, skip_serializing_if = "Option::is_none")]
632 pub tokens_before: Option<u32>,
633 #[serde(default, skip_serializing_if = "Option::is_none")]
635 pub tokens_after: Option<u32>,
636 #[serde(default, skip_serializing_if = "Option::is_none")]
638 pub duration_ms: Option<u64>,
639 #[serde(default, skip_serializing_if = "Option::is_none")]
643 pub item: Option<serde_json::Value>,
644}
645
646#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
647pub enum InferenceEvent {
648 MessageDelta(MessageDelta),
649 ReasoningDelta(ReasoningDelta),
650 ToolCallStarted(ToolCallStarted),
651 ToolCallDelta(ToolCallDelta),
652 ToolCallCompleted(ToolCallCompleted),
653 HostedToolCallStarted(HostedToolCallStarted),
654 HostedToolCallCompleted(HostedToolCallCompleted),
655 Compaction(CompactionProgress),
656 Usage(TokenUsage),
657 Completed(CompletionMetadata),
658 Failed(InferenceFailure),
659 ProviderMetadata(serde_json::Value),
660}
661
662pub type InferenceEventStream =
663 Pin<Box<dyn Stream<Item = anyhow::Result<InferenceEvent>> + Send + 'static>>;
664
665#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
666pub struct InferenceCapabilities {
667 pub streaming: bool,
668 pub tool_calls: bool,
669 pub parallel_tool_calls: bool,
670 pub reasoning_summaries: bool,
671 pub structured_output: bool,
672 pub image_input: bool,
673 pub prompt_cache: bool,
674 pub provider_metadata: bool,
675 pub tool_search: bool,
676}
677
678impl InferenceCapabilities {
679 pub fn text_only() -> Self {
680 Self {
681 streaming: true,
682 tool_calls: false,
683 parallel_tool_calls: false,
684 reasoning_summaries: false,
685 structured_output: false,
686 image_input: false,
687 prompt_cache: false,
688 provider_metadata: false,
689 tool_search: false,
690 }
691 }
692
693 pub fn coding_agent_default() -> Self {
694 Self {
695 streaming: true,
696 tool_calls: true,
697 parallel_tool_calls: true,
698 reasoning_summaries: false,
699 structured_output: false,
700 image_input: false,
701 prompt_cache: false,
702 provider_metadata: true,
703 tool_search: false,
704 }
705 }
706}
707
708#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
709pub struct ModelDescriptor {
710 pub id: String,
711 pub name: String,
712 pub context_window: Option<u32>,
713 #[serde(default, skip_serializing_if = "Option::is_none")]
714 pub default_reasoning: Option<String>,
715 #[serde(default, skip_serializing_if = "Vec::is_empty")]
716 pub supported_reasoning: Vec<ReasoningEffortDescriptor>,
717}
718
719#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
720pub struct ReasoningEffortDescriptor {
721 pub effort: String,
722 pub description: String,
723}
724
725pub struct InferenceProviderContext<'a> {
726 pub provider_id: &'a str,
727}
728
729pub struct InferenceTurnContext<'a> {
730 pub thread_id: &'a str,
731 pub turn_id: &'a str,
732 pub tool_executor: Option<std::sync::Arc<dyn TurnToolExecutor>>,
739}
740
741#[derive(Debug, Clone)]
743pub struct TurnToolOutcome {
744 pub result: String,
745 pub is_error: bool,
746}
747
748#[async_trait::async_trait]
752pub trait TurnToolExecutor: Send + Sync {
753 async fn execute(&self, call: ToolCallCompleted) -> anyhow::Result<TurnToolOutcome>;
754
755 fn register_provider_cleanup(&self, _cleanup: std::sync::Arc<dyn ProviderTurnCleanup>) {}
761}
762
763#[async_trait::async_trait]
766pub trait ProviderTurnCleanup: Send + Sync {
767 fn ownership(&self) -> TurnCleanupOwnership;
769
770 async fn wait_for_cleanup(&self) -> anyhow::Result<()>;
772}
773
774#[async_trait::async_trait]
775pub trait InferenceEngine: Send + Sync + 'static {
776 fn id(&self) -> InferenceEngineId;
777 fn capabilities(&self) -> InferenceCapabilities;
778
779 fn metadata(&self) -> InferenceProviderMetadata {
780 InferenceProviderMetadata::local(self.id())
781 }
782
783 async fn list_models(
784 &self,
785 ctx: InferenceProviderContext<'_>,
786 ) -> anyhow::Result<Vec<ModelDescriptor>>;
787
788 async fn stream_turn(
789 &self,
790 ctx: InferenceTurnContext<'_>,
791 request: AgentInferenceRequest,
792 ) -> anyhow::Result<InferenceEventStream>;
793}
794
795#[cfg(test)]
796mod tests {
797 use super::*;
798
799 #[test]
800 fn finish_reason_mapping_normalizes_known_stop_reasons() {
801 assert_eq!(finish_reason_from_stop_reason("end_turn"), "stop");
802 assert_eq!(finish_reason_from_stop_reason("stop"), "stop");
803 assert_eq!(finish_reason_from_stop_reason("stop_sequence"), "stop");
804 assert_eq!(finish_reason_from_stop_reason("max_tokens"), "length");
805 assert_eq!(finish_reason_from_stop_reason("length"), "length");
806 assert_eq!(finish_reason_from_stop_reason("tool_use"), "toolUse");
807 assert_eq!(finish_reason_from_stop_reason("tool_calls"), "toolUse");
808 assert_eq!(
809 finish_reason_from_stop_reason("content_filter"),
810 "contentFilter"
811 );
812 assert_eq!(finish_reason_from_stop_reason("refusal"), "refusal");
813 assert_eq!(finish_reason_from_stop_reason("pause_turn"), "pause_turn");
814 }
815
816 #[test]
817 fn token_usage_accumulates_cache_creation_prompt_tokens() {
818 let mut usage = TokenUsage::new(100, 10, 110)
819 .with_cached_prompt_tokens(80)
820 .with_cache_creation_prompt_tokens(15);
821 usage.add_assign(
822 &TokenUsage::new(50, 5, 55)
823 .with_cached_prompt_tokens(40)
824 .with_cache_creation_prompt_tokens(10),
825 );
826
827 assert_eq!(usage.prompt_tokens, 150);
828 assert_eq!(usage.cached_prompt_tokens, 120);
829 assert_eq!(usage.cache_creation_prompt_tokens, 25);
830 assert!(!usage.is_empty());
831
832 let creation_only = TokenUsage {
833 cache_creation_prompt_tokens: 1,
834 ..TokenUsage::default()
835 };
836 assert!(!creation_only.is_empty());
837 }
838
839 #[test]
840 fn inference_speed_policy_decision_serializes_runtime_metadata() {
841 let decision = SpeedPolicyDecision {
842 phase: SpeedPolicyPhase::Verification,
843 desired_reasoning: "high".to_string(),
844 applied_reasoning: Some("high".to_string()),
845 supported: true,
846 };
847 let hints = RuntimeHints {
848 speed_policy: Some(decision),
849 ..RuntimeHints::default()
850 };
851
852 let json = serde_json::to_value(hints).unwrap();
853 assert_eq!(
854 json.get("speed_policy")
855 .and_then(|value| value.get("phase"))
856 .and_then(serde_json::Value::as_str),
857 Some("verification")
858 );
859 assert_eq!(
860 json.get("speed_policy")
861 .and_then(|value| value.get("desiredReasoning"))
862 .and_then(serde_json::Value::as_str),
863 Some("high")
864 );
865 assert_eq!(
866 json.get("speed_policy")
867 .and_then(|value| value.get("appliedReasoning"))
868 .and_then(serde_json::Value::as_str),
869 Some("high")
870 );
871 }
872
873 #[test]
874 fn inference_reliability_policy_serializes_runtime_metadata() {
875 let hints = RuntimeHints {
876 reliability: Some(ReliabilityRequestPolicy::default()),
877 ..RuntimeHints::default()
878 };
879
880 let json = serde_json::to_value(hints).unwrap();
881 assert_eq!(
882 json.get("reliability")
883 .and_then(|value| value.get("providerRetryMaxAttempts"))
884 .and_then(serde_json::Value::as_u64),
885 Some(3)
886 );
887 assert_eq!(
888 json.get("reliability")
889 .and_then(|value| value.get("retryEmptyProviderBody"))
890 .and_then(serde_json::Value::as_bool),
891 Some(true)
892 );
893 }
894
895 #[test]
896 fn tool_search_config_serializes_provider_native_request() {
897 let config = ToolSearchConfig {
898 mode: ToolSearchMode::ProviderNative,
899 max_catalog_items: Some(200),
900 include_mcp: true,
901 include_skills: false,
902 fallback_to_explicit_tools: true,
903 provider_variant: ToolSearchProviderVariant::Bm25,
904 };
905
906 let value = serde_json::to_value(&config).unwrap();
907
908 assert_eq!(value["mode"], "provider_native");
909 assert_eq!(value["maxCatalogItems"], 200);
910 assert_eq!(value["includeMcp"], true);
911 assert_eq!(value["includeSkills"], false);
912 assert_eq!(value["providerVariant"], "bm25");
913 assert!(config.is_provider_native_requested());
914 }
915
916 #[test]
917 fn explicit_tool_search_config_preserves_current_default() {
918 let config = ToolSearchConfig::default();
919
920 assert_eq!(config.mode, ToolSearchMode::Explicit);
921 assert!(!config.is_provider_native_requested());
922 assert!(config.fallback_to_explicit_tools);
923 }
924
925 #[test]
926 fn tool_search_effective_mode_resolution_covers_fallback_matrix() {
927 let explicit = ToolSearchConfig::explicit();
928 assert_eq!(
929 explicit.resolve_effective_mode(true).unwrap(),
930 EffectiveToolSearchMode::Explicit
931 );
932
933 let auto = ToolSearchConfig {
934 mode: ToolSearchMode::Auto,
935 ..ToolSearchConfig::default()
936 };
937 assert_eq!(
938 auto.resolve_effective_mode(true).unwrap(),
939 EffectiveToolSearchMode::ProviderNative
940 );
941 assert_eq!(
942 auto.resolve_effective_mode(false).unwrap(),
943 EffectiveToolSearchMode::Explicit
944 );
945
946 let native = ToolSearchConfig::provider_native();
947 assert_eq!(
948 native.resolve_effective_mode(true).unwrap(),
949 EffectiveToolSearchMode::ProviderNative
950 );
951 assert_eq!(
952 native.resolve_effective_mode(false).unwrap(),
953 EffectiveToolSearchMode::Explicit
954 );
955
956 let strict = ToolSearchConfig {
957 fallback_to_explicit_tools: false,
958 ..ToolSearchConfig::provider_native()
959 };
960 let error = strict.resolve_effective_mode(false).unwrap_err();
961 assert_eq!(error, ToolSearchModeError::ProviderNativeUnsupported);
962 assert!(error.to_string().contains("fallback_to_explicit_tools"));
963 }
964}